diff --git a/litellm-rust/crates/cache/src/semantic.rs b/litellm-rust/crates/cache/src/semantic.rs index b27ab269edd..990f5e47aa8 100644 --- a/litellm-rust/crates/cache/src/semantic.rs +++ b/litellm-rust/crates/cache/src/semantic.rs @@ -4,7 +4,7 @@ //! `RedisSemanticCache._get_prompt_from_kwargs` (inherited by Valkey) adds Responses API //! `input`. Qdrant reads messages only. Each backend picks one of the two extractors here. -use std::{future::Future, io}; +use std::{collections::HashMap, future::Future, io}; use serde::Serialize; use serde_json::{ @@ -84,11 +84,35 @@ impl Embedder for PreparedEmbedding { } /// `get_str_from_messages_with_tools`: every message's content text, tool calls and tool results, -/// then its OpenAI `tool_calls`, then its search results. +/// then its OpenAI `tool_calls`, then its search results. Each tool result is tagged with the +/// position of the call it answers. pub fn str_from_messages(messages: &[Value]) -> String { + let messages: Vec<_> = messages.iter().filter_map(Value::as_object).collect(); + let call_ordinals = tool_call_ordinals(messages.iter().flat_map(|message| { + let block_ids = message + .get("content") + .and_then(Value::as_array) + .into_iter() + .flatten() + .filter(|block| block.get("type").and_then(Value::as_str) == Some("tool_use")) + .filter_map(|block| block.get("id")); + let tool_call_ids = message + .get("tool_calls") + .and_then(Value::as_array) + .into_iter() + .flatten() + .filter_map(|tool_call| tool_call.get("id")); + block_ids.chain(tool_call_ids) + })); let mut text = String::new(); - for message in messages.iter().filter_map(Value::as_object) { - push_content_text(&mut text, message.get("content")); + for message in messages { + if message.get("role").and_then(Value::as_str) == Some("tool") { + text.push_str(&tool_result_tag( + message.get("tool_call_id"), + &call_ordinals, + )); + } + push_content_text(&mut text, message.get("content"), &call_ordinals); if let Some(Value::Array(tool_calls)) = message.get("tool_calls") { for tool_call in tool_calls.iter().filter_map(Value::as_object) { let function = tool_call.get("function"); @@ -104,7 +128,11 @@ pub fn str_from_messages(messages: &[Value]) -> String { } /// `_content_str_with_tools`: text parts, Anthropic `tool_use` blocks and `tool_result` content. -fn push_content_text(text: &mut String, content: Option<&Value>) { +fn push_content_text( + text: &mut String, + content: Option<&Value>, + call_ordinals: &HashMap<&str, usize>, +) { match content { Some(Value::String(content)) => text.push_str(content), Some(Value::Array(blocks)) => { @@ -113,7 +141,10 @@ fn push_content_text(text: &mut String, content: Option<&Value>) { Some("tool_use") => { text.push_str(&tool_call_json(block.get("name"), block.get("input"))); } - Some("tool_result") => push_content_text(text, block.get("content")), + Some("tool_result") => { + text.push_str(&tool_result_tag(block.get("tool_use_id"), call_ordinals)); + push_content_text(text, block.get("content"), call_ordinals); + } _ => { if let Some(block_text) = block.get("text").and_then(Value::as_str) { text.push_str(block_text); @@ -126,6 +157,26 @@ fn push_content_text(text: &mut String, content: Option<&Value>) { } } +/// `tool_call_ordinals`: the 1-based position of each distinct string call id, first seen first. +fn tool_call_ordinals<'a>(call_ids: impl Iterator) -> HashMap<&'a str, usize> { + let mut ordinals = HashMap::new(); + for call_id in call_ids.filter_map(Value::as_str) { + let next = ordinals.len() + 1; + ordinals.entry(call_id).or_insert(next); + } + ordinals +} + +/// `tool_result_str`: `{"result_of_call":N}` for a result answering a known call, else empty. +fn tool_result_tag(call_id: Option<&Value>, call_ordinals: &HashMap<&str, usize>) -> String { + call_id + .and_then(Value::as_str) + .and_then(|call_id| call_ordinals.get(call_id)) + .map_or_else(String::new, |ordinal| { + format!("{{\"result_of_call\":{ordinal}}}") + }) +} + /// `tool_call_str`: the compact `{"name":...,"arguments":...}` a tool call contributes. fn tool_call_json(name: Option<&Value>, arguments: Option<&Value>) -> String { compact_json(&json!({ @@ -149,8 +200,16 @@ pub fn prompt_from_context(context: &SemanticCacheContext) -> Option { return Some(str_from_messages(messages)); } let input = context.input.as_ref()?; + let call_ordinals = tool_call_ordinals( + input + .as_array() + .into_iter() + .flatten() + .filter(|item| item.get("type").and_then(Value::as_str) == Some("function_call")) + .filter_map(|item| item.get("call_id")), + ); let mut parts = Vec::new(); - collect_input_text(input, &mut parts); + collect_input_text(input, &mut parts, &call_ordinals); let prompt = python_strip(&parts.join("\n")).to_owned(); (!prompt.is_empty()).then_some(prompt) } @@ -179,14 +238,18 @@ fn push_search_results_text(text: &mut String, search_results: Option<&Value>) { } } -fn collect_input_text(value: &Value, parts: &mut Vec) { +fn collect_input_text( + value: &Value, + parts: &mut Vec, + call_ordinals: &HashMap<&str, usize>, +) { match value { Value::String(text) => { push_trimmed(text, parts); } Value::Array(items) => { for item in items { - collect_input_text(item, parts); + collect_input_text(item, parts, call_ordinals); } } Value::Object(map) => { @@ -194,14 +257,24 @@ fn collect_input_text(value: &Value, parts: &mut Vec) { parts.push(tool_call_json(map.get("name"), map.get("arguments"))); return; } + if map.get("type").and_then(Value::as_str) == Some("function_call_output") { + let tag = tool_result_tag(map.get("call_id"), call_ordinals); + if !tag.is_empty() { + parts.push(tag); + if let Some(output) = map.get("output") { + collect_input_text(output, parts, call_ordinals); + } + return; + } + } if let Some(content) = map.get("content").filter(|content| !content.is_null()) { - collect_input_text(content, parts); + collect_input_text(content, parts, call_ordinals); return; } for key in ["text", "output", "input_text", "output_text"] { match map.get(key) { Some(nested @ Value::Array(_)) => { - collect_input_text(nested, parts); + collect_input_text(nested, parts, call_ordinals); return; } Some(Value::String(text)) if push_trimmed(text, parts) => return, diff --git a/litellm-rust/crates/cache/tests/semantic.rs b/litellm-rust/crates/cache/tests/semantic.rs index a29844bc241..7ede4948044 100644 --- a/litellm-rust/crates/cache/tests/semantic.rs +++ b/litellm-rust/crates/cache/tests/semantic.rs @@ -133,7 +133,40 @@ fn context(messages: Option, input: Option) -> SemanticCacheContex ]}, {"role": "tool", "tool_call_id": "c1", "content": "ok"}, ]), - r#"writing{"name":"write","arguments":"{\"path\": \"a\"}"}{"name":"write","arguments":"{\"path\": \"b\"}"}ok"#, + r#"writing{"name":"write","arguments":"{\"path\": \"a\"}"}{"name":"write","arguments":"{\"path\": \"b\"}"}{"result_of_call":1}ok"#, +)] +#[case::anthropic_parallel_tool_results_tagged_with_their_call( + json!([ + {"role": "assistant", "content": [ + {"type": "tool_use", "id": "t1", "name": "Read", "input": {"path": "a"}}, + {"type": "tool_use", "id": "t2", "name": "Read", "input": {"path": "b"}}, + ]}, + {"role": "user", "content": [ + {"type": "tool_result", "tool_use_id": "t2", "content": "B"}, + {"type": "tool_result", "tool_use_id": "t1", "content": "A"}, + ]}, + ]), + r#"{"name":"Read","arguments":{"path":"a"}}{"name":"Read","arguments":{"path":"b"}}{"result_of_call":2}B{"result_of_call":1}A"#, +)] +#[case::reused_call_ids_keep_first_position( + json!([ + {"role": "assistant", "content": null, "tool_calls": [ + {"id": "call_0", "type": "function", "function": {"name": "a", "arguments": "{}"}}, + ]}, + {"role": "assistant", "content": null, "tool_calls": [ + {"id": "call_0", "type": "function", "function": {"name": "b", "arguments": "{}"}}, + {"id": "call_1", "type": "function", "function": {"name": "c", "arguments": "{}"}}, + ]}, + {"role": "tool", "tool_call_id": "call_1", "content": "C"}, + ]), + r#"{"name":"a","arguments":"{}"}{"name":"b","arguments":"{}"}{"name":"c","arguments":"{}"}{"result_of_call":2}C"#, +)] +#[case::unknown_call_ids_left_untagged( + json!([ + {"role": "tool", "tool_call_id": "c9", "content": "ok"}, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "t9", "content": "done"}]}, + ]), + "okdone", )] #[case::malformed_tool_call_entries( json!([{"role": "assistant", "content": null, "tool_calls": [{"id": "c1", "type": "function"}, "junk"]}]), @@ -238,7 +271,7 @@ fn prompt_from_messages_reads_messages_only( {"type": "function_call", "call_id": "c1", "name": "write_file", "arguments": "{\"path\":\"a.yaml\"}"}, {"type": "function_call_output", "call_id": "c1", "output": "ok"}, ])), - Some("update the config\n{\"name\":\"write_file\",\"arguments\":\"{\\\"path\\\":\\\"a.yaml\\\"}\"}\nok"), + Some("update the config\n{\"name\":\"write_file\",\"arguments\":\"{\\\"path\\\":\\\"a.yaml\\\"}\"}\n{\"result_of_call\":1}\nok"), )] #[case::responses_structured_function_call_output( None, @@ -246,7 +279,22 @@ fn prompt_from_messages_reads_messages_only( {"type": "function_call", "call_id": "c1", "name": "write_file", "arguments": "{}"}, {"type": "function_call_output", "call_id": "c1", "output": [{"type": "input_text", "text": "denied"}]}, ])), - Some("{\"name\":\"write_file\",\"arguments\":\"{}\"}\ndenied"), + Some("{\"name\":\"write_file\",\"arguments\":\"{}\"}\n{\"result_of_call\":1}\ndenied"), +)] +#[case::responses_parallel_outputs_tagged_with_their_call( + None, + Some(json!([ + {"type": "function_call", "call_id": "c1", "name": "read", "arguments": "a"}, + {"type": "function_call", "call_id": "c2", "name": "read", "arguments": "b"}, + {"type": "function_call_output", "call_id": "c2", "output": "B"}, + {"type": "function_call_output", "call_id": "c1", "output": "A"}, + ])), + Some("{\"name\":\"read\",\"arguments\":\"a\"}\n{\"name\":\"read\",\"arguments\":\"b\"}\n{\"result_of_call\":2}\nB\n{\"result_of_call\":1}\nA"), +)] +#[case::responses_unknown_call_id_output_untagged( + None, + Some(json!([{"type": "function_call_output", "call_id": "c9", "output": "orphan"}])), + Some("orphan"), )] fn prompt_from_context_matches_python( #[case] messages: Option, diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index 9968297f01d..88f5acf026e 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -22,7 +22,9 @@ from litellm.constants import SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS from litellm.litellm_core_utils.asyncify import asyncify from litellm.litellm_core_utils.prompt_templates.common_utils import ( get_str_from_messages_with_tools, + tool_call_ordinals, tool_call_str, + tool_result_str, ) from litellm.types.utils import EmbeddingResponse @@ -269,14 +271,17 @@ class RedisSemanticCache(BaseCache): if "input" not in kwargs: return None + responses_input: Final = kwargs.get("input") prompt_parts: Final[list[str]] = [] - cls._collect_responses_input_text(kwargs.get("input"), prompt_parts) + cls._collect_responses_input_text(responses_input, prompt_parts, cls._responses_call_ordinals(responses_input)) prompt: Final = "\n".join(prompt_parts).strip() return prompt or None @classmethod - def _collect_responses_input_text(cls, value: object, prompt_parts: list[str]) -> None: - value = cls._function_call_as_prompt(cls._coerce_response_input_value(value)) + def _collect_responses_input_text( + cls, value: object, prompt_parts: list[str], call_ordinals: Mapping[str, int] + ) -> None: + value = cls._function_call_as_prompt(cls._coerce_response_input_value(value), call_ordinals) if value is None: return @@ -288,21 +293,21 @@ class RedisSemanticCache(BaseCache): if isinstance(value, (list, tuple)): for item in value: - cls._collect_responses_input_text(item, prompt_parts) + cls._collect_responses_input_text(item, prompt_parts, call_ordinals) return if isinstance(value, dict): content = value.get("content") if content is not None: - cls._collect_responses_input_text(content, prompt_parts) + cls._collect_responses_input_text(content, prompt_parts, call_ordinals) return - cls._collect_responses_text_fields(value, prompt_parts) + cls._collect_responses_text_fields(value, prompt_parts, call_ordinals) return content = getattr(value, "content", None) if content is not None: - cls._collect_responses_input_text(content, prompt_parts) + cls._collect_responses_input_text(content, prompt_parts, call_ordinals) return for text_key in ("text", "output", "input_text", "output_text"): @@ -314,21 +319,38 @@ class RedisSemanticCache(BaseCache): return @classmethod - def _collect_responses_text_fields(cls, value: dict, prompt_parts: list[str]) -> None: + def _collect_responses_text_fields( + cls, value: dict, prompt_parts: list[str], call_ordinals: Mapping[str, int] + ) -> None: for text_key in ("text", "output", "input_text", "output_text"): text_value = value.get(text_key) if isinstance(text_value, (list, tuple)): - cls._collect_responses_input_text(text_value, prompt_parts) + cls._collect_responses_input_text(text_value, prompt_parts, call_ordinals) return if isinstance(text_value, str) and (stripped_text := text_value.strip()): prompt_parts.append(stripped_text) return + @classmethod + def _responses_call_ordinals(cls, responses_input: object) -> Mapping[str, int]: + items: Final = responses_input if isinstance(responses_input, (list, tuple)) else () + dumped_items: Final = (cls._coerce_response_input_value(item) for item in items) + return tool_call_ordinals( + item.get("call_id") + for item in dumped_items + if isinstance(item, dict) and item.get("type") == "function_call" + ) + @staticmethod - def _function_call_as_prompt(value: object) -> object: - if isinstance(value, dict) and value.get("type") == "function_call": + def _function_call_as_prompt(value: object, call_ordinals: Mapping[str, int]) -> object: + if not isinstance(value, dict): + return value + if value.get("type") == "function_call": return tool_call_str(value.get("name"), value.get("arguments")) - return value + result_tag: Final = ( + tool_result_str(value.get("call_id"), call_ordinals) if value.get("type") == "function_call_output" else "" + ) + return (result_tag, value.get("output")) if result_tag else value @staticmethod def _coerce_response_input_value(value: object) -> object: diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index d7fe1849670..dddb9f3b468 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -199,42 +199,68 @@ def get_str_from_messages_with_tools(messages: object) -> str: """ ``get_str_from_messages`` that also keeps each conversation's tool calls and tool results, so agent turns that differ only in their tool exchange (Anthropic ``tool_use`` / ``tool_result``, OpenAI ``tool_calls``) - produce different text + produce different text. Each result is tagged with the position of the call it answers, since call ids are + random per session """ - return "".join(_message_str_with_tools(message) for message in _str_mappings(messages)) + message_mappings: Final = tuple(_str_mappings(messages)) + call_ordinals: Final = tool_call_ordinals(_message_tool_call_ids(message_mappings)) + return "".join(_message_str_with_tools(message, call_ordinals) for message in message_mappings) def tool_call_str(name: object, arguments: object) -> str: return f'{{"name":{_compact_json(name)},"arguments":{_compact_json(arguments)}}}' +def tool_result_str(call_id: object, call_ordinals: Mapping[str, int]) -> str: + ordinal: Final = call_ordinals.get(call_id) if isinstance(call_id, str) else None + return "" if ordinal is None else f'{{"result_of_call":{ordinal}}}' + + +def tool_call_ordinals(call_ids: Iterable[object]) -> Mapping[str, int]: + string_ids: Final = (call_id for call_id in call_ids if isinstance(call_id, str)) + return MappingProxyType({call_id: ordinal for ordinal, call_id in enumerate(dict.fromkeys(string_ids), start=1)}) + + def _compact_json(value: object) -> str: return json.dumps(value, separators=(",", ":"), default=str) -def _message_str_with_tools(message: Mapping[str, object]) -> str: +def _message_tool_call_ids(messages: Iterable[Mapping[str, object]]) -> Iterator[object]: + for message in messages: + yield from ( + block.get("id") for block in _str_mappings(message.get("content")) if block.get("type") == "tool_use" + ) + yield from (tool_call.get("id") for tool_call in _str_mappings(message.get("tool_calls"))) + + +def _message_str_with_tools(message: Mapping[str, object], call_ordinals: Mapping[str, int]) -> str: + result_tag: Final = ( + tool_result_str(message.get("tool_call_id"), call_ordinals) if message.get("role") == "tool" else "" + ) return ( - _content_str_with_tools(message.get("content")) + result_tag + + _content_str_with_tools(message.get("content"), call_ordinals) + "".join(_openai_tool_call_str(tool_call) for tool_call in _str_mappings(message.get("tool_calls"))) + extract_search_results_text(message.get("search_results")) ) -def _content_str_with_tools(content: object) -> str: +def _content_str_with_tools(content: object, call_ordinals: Mapping[str, int]) -> str: if isinstance(content, str): return content - return "".join(_block_str_with_tools(block) for block in _str_mappings(content)) + return "".join(_block_str_with_tools(block, call_ordinals) for block in _str_mappings(content)) -def _block_str_with_tools(block: Mapping[str, object]) -> str: - match block.get("type"): - case "tool_use": - return tool_call_str(block.get("name"), block.get("input")) - case "tool_result": - return _content_str_with_tools(block.get("content")) - case _: - text: Final = block.get("text") - return text if isinstance(text, str) else "" +def _block_str_with_tools(block: Mapping[str, object], call_ordinals: Mapping[str, int]) -> str: + block_type: Final = block.get("type") + if block_type == "tool_use": + return tool_call_str(block.get("name"), block.get("input")) + if block_type == "tool_result": + return tool_result_str(block.get("tool_use_id"), call_ordinals) + _content_str_with_tools( + block.get("content"), call_ordinals + ) + text: Final = block.get("text") + return text if isinstance(text, str) else "" def _openai_tool_call_str(tool_call: Mapping[str, object]) -> str: diff --git a/tests/integration/caching/test_semantic_cache_tool_turns.py b/tests/integration/caching/test_semantic_cache_tool_turns.py index 6c9a0b060fc..5d6e6a956a1 100644 --- a/tests/integration/caching/test_semantic_cache_tool_turns.py +++ b/tests/integration/caching/test_semantic_cache_tool_turns.py @@ -241,10 +241,13 @@ def test_claude_code_tool_turns_on_messages_are_not_served_the_first_turn_answer }, ] - answers: Final = _send_turns(semantic_proxy, "/v1/messages", _CLAUDE_MODEL, ([task], list_files, read_file)) + answers: Final = _send_turns( + semantic_proxy, "/v1/messages", _CLAUDE_MODEL, ([task], list_files, read_file, list_files) + ) - assert len(set(answers)) == 3, f"a later tool turn replayed an earlier cached answer: {answers}" - assert answers == semantic_proxy.peer.answered[-3:], semantic_proxy.peer.answered + assert len(set(answers[:3])) == 3, f"a later tool turn replayed an earlier cached answer: {answers}" + assert answers[:3] == semantic_proxy.peer.answered[-3:], semantic_proxy.peer.answered + assert answers[3] == answers[1], f"a repeated tool turn missed the cache: {answers}" def test_openai_agent_loop_tool_calls_on_chat_completions_are_not_served_a_cached_answer( @@ -270,8 +273,9 @@ def test_openai_agent_loop_tool_calls_on_chat_completions_are_not_served_a_cache ] answers: Final = _send_turns( - semantic_proxy, "/v1/chat/completions", _CHAT_MODEL, (wrote("a.yaml"), wrote("b.yaml")) + semantic_proxy, "/v1/chat/completions", _CHAT_MODEL, (wrote("a.yaml"), wrote("b.yaml"), wrote("a.yaml")) ) - assert len(set(answers)) == 2, f"a different tool call replayed an earlier cached answer: {answers}" - assert answers == semantic_proxy.peer.answered[-2:], semantic_proxy.peer.answered + assert len(set(answers[:2])) == 2, f"a different tool call replayed an earlier cached answer: {answers}" + assert answers[:2] == semantic_proxy.peer.answered[-2:], semantic_proxy.peer.answered + assert answers[2] == answers[0], f"a repeated tool call missed the cache: {answers}" diff --git a/tests/unit/caching/test_redis_semantic_cache.py b/tests/unit/caching/test_redis_semantic_cache.py index 38700dce62c..6ca4481648b 100644 --- a/tests/unit/caching/test_redis_semantic_cache.py +++ b/tests/unit/caching/test_redis_semantic_cache.py @@ -593,10 +593,10 @@ def test_redis_semantic_cache_prompt_extraction_keeps_tool_turns_distinct(): ] assert RedisSemanticCache._get_prompt_from_kwargs(messages=turn("ls")) == ( - 'fix the failing test{"name":"Bash","arguments":{"cmd":"ls"}}ok' + 'fix the failing test{"name":"Bash","arguments":{"cmd":"ls"}}{"result_of_call":1}ok' ) assert RedisSemanticCache._get_prompt_from_kwargs(messages=turn("pwd")) == ( - 'fix the failing test{"name":"Bash","arguments":{"cmd":"pwd"}}ok' + 'fix the failing test{"name":"Bash","arguments":{"cmd":"pwd"}}{"result_of_call":1}ok' ) @@ -611,7 +611,9 @@ def test_redis_semantic_cache_prompt_extraction_keeps_responses_function_calls() ] ) - assert prompt == 'update the config\n{"name":"write_file","arguments":"{\\"path\\":\\"a.yaml\\"}"}\nok' + assert prompt == ( + 'update the config\n{"name":"write_file","arguments":"{\\"path\\":\\"a.yaml\\"}"}\n{"result_of_call":1}\nok' + ) def test_redis_semantic_cache_prompt_extraction_keeps_structured_function_call_outputs(): @@ -631,10 +633,30 @@ def test_redis_semantic_cache_prompt_extraction_keeps_structured_function_call_o ) expected_call = '{"name":"write_file","arguments":"{\\"path\\":\\"a.txt\\"}"}' - assert prompt_for("wrote 5 bytes") == f"write hello\n{expected_call}\nwrote 5 bytes" + assert prompt_for("wrote 5 bytes") == f'write hello\n{expected_call}\n{{"result_of_call":1}}\nwrote 5 bytes' assert prompt_for("wrote 5 bytes") != prompt_for("PermissionError") +def test_redis_semantic_cache_prompt_extraction_tells_apart_parallel_outputs_answering_different_calls(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + def prompt_for(first_output_call_id: str, second_output_call_id: str) -> str | None: + return RedisSemanticCache._get_prompt_from_kwargs( + input=[ + {"type": "function_call", "call_id": "c1", "name": "read", "arguments": "a"}, + {"type": "function_call", "call_id": "c2", "name": "read", "arguments": "b"}, + {"type": "function_call_output", "call_id": first_output_call_id, "output": "empty"}, + {"type": "function_call_output", "call_id": second_output_call_id, "output": "secret"}, + ] + ) + + assert prompt_for("c2", "c1") == ( + '{"name":"read","arguments":"a"}\n{"name":"read","arguments":"b"}\n' + '{"result_of_call":2}\nempty\n{"result_of_call":1}\nsecret' + ) + assert prompt_for("c1", "c2") != prompt_for("c2", "c1") + + def test_redis_semantic_cache_prompt_extraction_handles_model_objects(): from litellm.caching.redis_semantic_cache import RedisSemanticCache diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py index 4c1af95f4fe..2bd593bea0d 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py @@ -7,8 +7,6 @@ from typing import Final import pytest -from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message - from litellm.litellm_core_utils.prompt_templates.common_utils import ( ENCRYPTED_REASONING_SIGNATURE_PREFIX, TOOL_RESULT_IMAGE_BOUNDARY, @@ -32,6 +30,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( system_messages_first, update_messages_with_model_file_ids, ) +from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message _ARTIFACT_FIELD_PATTERN: Final = r'^(?!__.*__$)[^\p{Cc}\p{Cf}\p{Zl}\p{Zp}"\\./[\]]{1,200}$' @@ -2063,9 +2062,51 @@ _TASK: Final = {"role": "user", "content": "fix the failing test"} {"role": "tool", "tool_call_id": "c1", "content": "ok"}, ], 'fix the failing testwriting{"name":"write","arguments":"{\\"path\\": \\"a\\"}"}' - '{"name":"write","arguments":"{\\"path\\": \\"b\\"}"}ok', + '{"name":"write","arguments":"{\\"path\\": \\"b\\"}"}{"result_of_call":1}ok', id="openai-tool-calls-in-order-before-tool-result", ), + pytest.param( + [ + { + "role": "assistant", + "content": [ + {"type": "tool_use", "id": "t1", "name": "Read", "input": {"path": "a"}}, + {"type": "tool_use", "id": "t2", "name": "Read", "input": {"path": "b"}}, + ], + }, + { + "role": "user", + "content": [ + {"type": "tool_result", "tool_use_id": "t2", "content": "B"}, + {"type": "tool_result", "tool_use_id": "t1", "content": "A"}, + ], + }, + ], + '{"name":"Read","arguments":{"path":"a"}}{"name":"Read","arguments":{"path":"b"}}' + '{"result_of_call":2}B{"result_of_call":1}A', + id="anthropic-parallel-tool-results-tagged-with-their-call", + ), + pytest.param( + [ + { + "role": "assistant", + "content": None, + "tool_calls": [{"id": "call_0", "type": "function", "function": {"name": "a", "arguments": "{}"}}], + }, + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_0", "type": "function", "function": {"name": "b", "arguments": "{}"}}, + {"id": "call_1", "type": "function", "function": {"name": "c", "arguments": "{}"}}, + ], + }, + {"role": "tool", "tool_call_id": "call_1", "content": "C"}, + ], + '{"name":"a","arguments":"{}"}{"name":"b","arguments":"{}"}{"name":"c","arguments":"{}"}' + '{"result_of_call":2}C', + id="reused-call-ids-keep-first-position", + ), pytest.param( [ Message( @@ -2120,3 +2161,35 @@ def test_get_str_from_messages_with_tools_keeps_tool_exchange(messages: list[obj ) def test_get_str_from_messages_with_tools_matches_get_str_from_messages_without_tools(messages: list[object]) -> None: assert get_str_from_messages_with_tools(messages) == get_str_from_messages(messages) # pyright: ignore[reportArgumentType] # untyped fixtures + + +def _parallel_reads(result_for_a: str, result_for_b: str, *, call_id_prefix: str = "c") -> list[object]: + return [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": f"{call_id_prefix}1", "type": "function", "function": {"name": "read", "arguments": '"a"'}}, + {"id": f"{call_id_prefix}2", "type": "function", "function": {"name": "read", "arguments": '"b"'}}, + ], + }, + {"role": "tool", "tool_call_id": f"{call_id_prefix}1", "content": result_for_a}, + {"role": "tool", "tool_call_id": f"{call_id_prefix}2", "content": result_for_b}, + ] + + +def _results_in_swapped_order(result_for_a: str, result_for_b: str) -> list[object]: + call, answer_a, answer_b = _parallel_reads(result_for_a, result_for_b) + return [call, answer_b, answer_a] + + +def test_get_str_from_messages_with_tools_tells_apart_parallel_results_answering_different_calls() -> None: + assert get_str_from_messages_with_tools(_parallel_reads("empty", "secret")) != get_str_from_messages_with_tools( + _results_in_swapped_order("secret", "empty") + ) + + +def test_get_str_from_messages_with_tools_ignores_call_ids_that_differ_between_sessions() -> None: + assert get_str_from_messages_with_tools( + _parallel_reads("A", "B", call_id_prefix="toolu_") + ) == get_str_from_messages_with_tools(_parallel_reads("A", "B"))