diff --git a/crates/tinyinference-llm/src/prompt_tools/mod.rs b/crates/tinyinference-llm/src/prompt_tools/mod.rs index 5eaca90..358af54 100644 --- a/crates/tinyinference-llm/src/prompt_tools/mod.rs +++ b/crates/tinyinference-llm/src/prompt_tools/mod.rs @@ -229,6 +229,41 @@ pub fn ensure_resolvable_user_turn(messages: &[Message]) -> Vec { out } +/// Keep the active request at the end of a prompt-guided tool continuation. +/// Some chat templates treat the synthetic tool-result user turn as transport +/// data and resolve an earlier user message as the query. Repeating the latest +/// real text request after a terminal result gives those templates the correct +/// query without changing the durable transcript. +#[must_use] +pub fn anchor_user_request_after_tool_result(messages: &[Message]) -> Vec { + if !messages.last().is_some_and(|message| { + matches!(message, Message::User(_)) + && message + .text() + .trim_start() + .starts_with(TOOL_RESULTS_PREFIX.trim_end()) + }) { + return messages.to_vec(); + } + let Some(request) = messages + .iter() + .rev() + .skip(1) + .find(|message| is_resolvable_user_query(message)) + else { + return messages.to_vec(); + }; + let request = request.text(); + if request.trim().is_empty() { + return messages.to_vec(); + } + let mut out = messages.to_vec(); + out.push(Message::user(format!( + "Continue the latest user request using the tool result above. Latest user request:\n{request}" + ))); + out +} + /// Converts a recovered call into this crate's [`ToolCall`], minting a /// process-unique id. /// diff --git a/crates/tinyinference-llm/src/prompt_tools/test.rs b/crates/tinyinference-llm/src/prompt_tools/test.rs index e09ff87..1b9609a 100644 --- a/crates/tinyinference-llm/src/prompt_tools/test.rs +++ b/crates/tinyinference-llm/src/prompt_tools/test.rs @@ -166,6 +166,44 @@ fn user_turn_normalization_does_not_count_folded_tool_results() { assert_eq!(out[1].text(), CONTINUATION_USER_TURN); } +#[test] +fn tool_continuation_anchors_the_latest_request_after_search_results() { + let messages = coalesce_tool_results(&[ + Message::system("system"), + Message::user("hey"), + Message::assistant("Hey! What's up?"), + Message::user("fetch my latest email"), + Message::assistant(""), + Message::tool("search-1", "GMAIL_FETCH_EMAILS schema"), + ]); + let anchored = anchor_user_request_after_tool_result(&messages); + assert_eq!(anchored.len(), messages.len() + 1); + assert!( + anchored[anchored.len() - 2] + .text() + .contains("GMAIL_FETCH_EMAILS") + ); + assert!( + anchored + .last() + .unwrap() + .text() + .contains("fetch my latest email") + ); + assert!(!anchored.last().unwrap().text().contains("Hey! What's up?")); + assert_eq!(anchor_user_request_after_tool_result(&anchored), anchored); +} + +#[test] +fn tool_continuation_without_a_real_user_request_stays_unchanged() { + let messages = coalesce_tool_results(&[ + Message::system("system"), + Message::assistant("calling"), + Message::tool("call-1", "result"), + ]); + assert_eq!(anchor_user_request_after_tool_result(&messages), messages); +} + #[test] fn user_turn_normalization_ignores_blank_and_accepts_non_text_turns() { let out = ensure_resolvable_user_turn(&[Message::system("system"), Message::user(" ")]); diff --git a/crates/tinyinference-llm/src/providers/openai/transport.rs b/crates/tinyinference-llm/src/providers/openai/transport.rs index cf17789..3b68192 100644 --- a/crates/tinyinference-llm/src/providers/openai/transport.rs +++ b/crates/tinyinference-llm/src/providers/openai/transport.rs @@ -1141,8 +1141,9 @@ impl OpenAiModel { // into the text protocol before the protocol block is added. let coalesced = crate::prompt_tools::coalesce_tool_results(&request.messages); let resolvable = crate::prompt_tools::ensure_resolvable_user_turn(&coalesced); + let anchored = crate::prompt_tools::anchor_user_request_after_tool_result(&resolvable); instructed_messages = crate::prompt_tools::with_tool_instructions( - &resolvable, + &anchored, &prompt_tool_schemas, &request.tool_choice, );