From f462fbeb2e59b6743aea8af5a5b92d2a8407faa8 Mon Sep 17 00:00:00 2001 From: kerry Date: Wed, 30 Sep 2026 16:43:47 +0000 Subject: [PATCH] fix(caching): keep tool calls and tool results in semantic cache prompts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/crates/cache/src/semantic.rs | 59 +++- litellm-rust/crates/cache/tests/semantic.rs | 48 +++ litellm/caching/qdrant_semantic_cache.py | 10 +- litellm/caching/redis_semantic_cache.py | 9 +- .../prompt_templates/common_utils.py | 63 ++++ tests/integration/_support/manifest.py | 1 + .../caching/test_semantic_cache_tool_turns.py | 277 ++++++++++++++++++ tests/integration/run.py | 2 +- .../unit/caching/test_redis_semantic_cache.py | 35 +++ ...ore_utils_prompt_templates_common_utils.py | 116 ++++++++ 10 files changed, 600 insertions(+), 20 deletions(-) create mode 100644 tests/integration/caching/test_semantic_cache_tool_turns.py diff --git a/litellm-rust/crates/cache/src/semantic.rs b/litellm-rust/crates/cache/src/semantic.rs index c88706a213b..5b619151285 100644 --- a/litellm-rust/crates/cache/src/semantic.rs +++ b/litellm-rust/crates/cache/src/semantic.rs @@ -1,6 +1,6 @@ //! The embedding and prompt contract every semantic backend shares. //! -//! Python's semantic caches all read their prompt through `get_str_from_messages`, and +//! Python's semantic caches all read their prompt through `get_str_from_messages_with_tools`, and //! `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. @@ -8,7 +8,7 @@ use std::{future::Future, io}; use serde::Serialize; use serde_json::{ - Value, + Value, json, ser::{CharEscape, Formatter, Serializer}, }; @@ -83,26 +83,57 @@ impl Embedder for PreparedEmbedding { } } -/// `get_str_from_messages`: every message's text content followed by its search results. +/// `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. pub fn str_from_messages(messages: &[Value]) -> String { let mut text = String::new(); for message in messages.iter().filter_map(Value::as_object) { - match message.get("content") { - Some(Value::String(content)) => text.push_str(content), - Some(Value::Array(parts)) => { - for part in parts { - if let Some(part_text) = part.get("text").and_then(Value::as_str) { - text.push_str(part_text); - } - } + push_content_text(&mut text, message.get("content")); + 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"); + text.push_str(&tool_call_json( + function.and_then(|function| function.get("name")), + function.and_then(|function| function.get("arguments")), + )); } - _ => {} } push_search_results_text(&mut text, message.get("search_results")); } text } +/// `_content_str_with_tools`: text parts, Anthropic `tool_use` blocks and `tool_result` content. +fn push_content_text(text: &mut String, content: Option<&Value>) { + match content { + Some(Value::String(content)) => text.push_str(content), + Some(Value::Array(blocks)) => { + for block in blocks.iter().filter_map(Value::as_object) { + match block.get("type").and_then(Value::as_str) { + 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")), + _ => { + if let Some(block_text) = block.get("text").and_then(Value::as_str) { + text.push_str(block_text); + } + } + } + } + } + _ => {} + } +} + +/// `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!({ + "name": name.unwrap_or(&Value::Null), + "arguments": arguments.unwrap_or(&Value::Null), + })) +} + /// The messages prompt Qdrant embeds: `None` when the request carries no messages. pub fn prompt_from_messages(context: &SemanticCacheContext) -> Option { let messages = context.messages.as_ref()?.as_array()?; @@ -159,6 +190,10 @@ fn collect_input_text(value: &Value, parts: &mut Vec) { } } Value::Object(map) => { + if map.get("type").and_then(Value::as_str) == Some("function_call") { + parts.push(tool_call_json(map.get("name"), map.get("arguments"))); + return; + } if let Some(content) = map.get("content").filter(|content| !content.is_null()) { collect_input_text(content, parts); return; diff --git a/litellm-rust/crates/cache/tests/semantic.rs b/litellm-rust/crates/cache/tests/semantic.rs index 97a552a8010..ab6b904af1a 100644 --- a/litellm-rust/crates/cache/tests/semantic.rs +++ b/litellm-rust/crates/cache/tests/semantic.rs @@ -100,6 +100,45 @@ fn context(messages: Option, input: Option) -> SemanticCacheContex json!([{"role": "tool", "search_results": [{"citations": false}, {"citations": 3}]}]), "false3", )] +#[case::anthropic_tool_use_name_and_input_without_id( + json!([ + {"role": "user", "content": "fix the failing test"}, + {"role": "assistant", "content": [ + {"type": "tool_use", "id": "t1", "name": "Bash", "input": {"cmd": "ls"}}, + ]}, + ]), + r#"fix the failing test{"name":"Bash","arguments":{"cmd":"ls"}}"#, +)] +#[case::anthropic_string_tool_result( + json!([{"role": "user", "content": [ + {"type": "tool_result", "tool_use_id": "t1", "content": "calc.py"}, + ]}]), + "calc.py", +)] +#[case::anthropic_nested_text_tool_result_then_text( + json!([{"role": "user", "content": [ + {"type": "tool_result", "tool_use_id": "t1", "content": [ + {"type": "text", "text": "a"}, + {"type": "text", "text": "b"}, + ]}, + {"type": "text", "text": "next"}, + ]}]), + "abnext", +)] +#[case::openai_tool_calls_in_order_before_tool_result( + json!([ + {"role": "assistant", "content": "writing", "tool_calls": [ + {"id": "c1", "type": "function", "function": {"name": "write", "arguments": "{\"path\": \"a\"}"}}, + {"id": "c2", "type": "function", "function": {"name": "write", "arguments": "{\"path\": \"b\"}"}}, + ]}, + {"role": "tool", "tool_call_id": "c1", "content": "ok"}, + ]), + r#"writing{"name":"write","arguments":"{\"path\": \"a\"}"}{"name":"write","arguments":"{\"path\": \"b\"}"}ok"#, +)] +#[case::malformed_tool_call_entries( + json!([{"role": "assistant", "content": null, "tool_calls": [{"id": "c1", "type": "function"}, "junk"]}]), + r#"{"name":null,"arguments":null}"#, +)] fn str_from_messages_matches_python(#[case] messages: Value, #[case] expected: &str) { assert_eq!(str_from_messages(messages.as_array().unwrap()), expected); } @@ -192,6 +231,15 @@ fn prompt_from_messages_reads_messages_only( )] #[case::nested_lists(None, Some(json!([["a", [" b "]], "", "c"])), Some("a\nb\nc"))] #[case::scalars_ignored(None, Some(json!([1, true, null, "kept"])), Some("kept"))] +#[case::responses_function_call( + None, + Some(json!([ + {"role": "user", "content": "update the config"}, + {"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"), +)] fn prompt_from_context_matches_python( #[case] messages: Option, #[case] input: Option, diff --git a/litellm/caching/qdrant_semantic_cache.py b/litellm/caching/qdrant_semantic_cache.py index b99023c07fd..8fe03131760 100644 --- a/litellm/caching/qdrant_semantic_cache.py +++ b/litellm/caching/qdrant_semantic_cache.py @@ -24,7 +24,7 @@ from litellm.constants import ( ) from litellm.litellm_core_utils.asyncify import asyncify from litellm.litellm_core_utils.prompt_templates.common_utils import ( - get_str_from_messages, + get_str_from_messages_with_tools, ) from litellm.types.utils import EmbeddingResponse @@ -286,7 +286,7 @@ class QdrantSemanticCache(BaseCache): # get the prompt messages: Final = kwargs["messages"] - prompt: Final = get_str_from_messages(messages) + prompt: Final = get_str_from_messages_with_tools(messages) # create an embedding for prompt embedding_response: Final = cast( @@ -325,7 +325,7 @@ class QdrantSemanticCache(BaseCache): # get the messages messages: Final = kwargs["messages"] - prompt: Final = get_str_from_messages(messages) + prompt: Final = get_str_from_messages_with_tools(messages) # convert to embedding embedding_response: Final = cast( @@ -400,7 +400,7 @@ class QdrantSemanticCache(BaseCache): # get the prompt messages: Final = kwargs["messages"] - prompt: Final = get_str_from_messages(messages) + prompt: Final = get_str_from_messages_with_tools(messages) embedding_response: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) # get the embedding @@ -435,7 +435,7 @@ class QdrantSemanticCache(BaseCache): # get the messages messages: Final = kwargs["messages"] - prompt: Final = get_str_from_messages(messages) + prompt: Final = get_str_from_messages_with_tools(messages) embedding_response: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index d4c815e15b7..1cb1b457cd5 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -21,7 +21,8 @@ from litellm._logging import print_verbose, verbose_logger 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, + get_str_from_messages_with_tools, + tool_call_str, ) from litellm.types.utils import EmbeddingResponse @@ -263,7 +264,7 @@ class RedisSemanticCache(BaseCache): """ messages: Final = kwargs.get("messages") if messages: - return get_str_from_messages(messages) + return get_str_from_messages_with_tools(messages) if "input" not in kwargs: return None @@ -291,6 +292,10 @@ class RedisSemanticCache(BaseCache): return if isinstance(value, dict): + if value.get("type") == "function_call": + prompt_parts.append(tool_call_str(value.get("name"), value.get("arguments"))) + return + content = value.get("content") if content is not None: cls._collect_responses_input_text(content, prompt_parts) diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 3b6827375f3..90e54750fd1 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -13,6 +13,9 @@ from pathlib import Path from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, cast +from pydantic import BaseModel +from typing_extensions import TypeIs # noqa: TID251 # narrows untyped message payloads without a runtime conversion + import litellm from litellm import verbose_logger from litellm.router_utils.batch_utils import InMemoryFile @@ -192,6 +195,66 @@ def get_str_from_messages(messages: list[AllMessageValues]) -> str: return text +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 + """ + return "".join(_message_str_with_tools(message) for message in _str_mappings(messages)) + + +def tool_call_str(name: object, arguments: object) -> str: + return json.dumps({"name": name, "arguments": arguments}, separators=(",", ":"), default=str) + + +def _message_str_with_tools(message: Mapping[str, object]) -> str: + return ( + _content_str_with_tools(message.get("content")) + + "".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: + if isinstance(content, str): + return content + return "".join(_block_str_with_tools(block) 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 _openai_tool_call_str(tool_call: Mapping[str, object]) -> str: + function: Final = _as_str_mapping(tool_call.get("function")) or {} + return tool_call_str(function.get("name"), function.get("arguments")) + + +def _str_mappings(values: object) -> Iterator[Mapping[str, object]]: + items: Final = values if isinstance(values, (list, tuple)) else () + return (mapping for item in items if (mapping := _as_str_mapping(item)) is not None) + + +def _as_str_mapping(value: object) -> Mapping[str, object] | None: + if isinstance(value, BaseModel): + return value.model_dump() + if _is_str_mapping(value): + return value + return None + + +def _is_str_mapping(value: object) -> TypeIs[Mapping[str, object]]: # guard-ok: message and block keys are str + return isinstance(value, Mapping) + + def is_non_content_values_set(message: AllMessageValues) -> bool: ignore_keys: Final = ["content", "role", "name"] return any(message.get(key, None) is not None for key in message if key not in ignore_keys) diff --git a/tests/integration/_support/manifest.py b/tests/integration/_support/manifest.py index a4a86a21568..a37ae94c2ed 100644 --- a/tests/integration/_support/manifest.py +++ b/tests/integration/_support/manifest.py @@ -11,6 +11,7 @@ OWNED_DIRECTORIES: Final = frozenset( "providers", "streaming", "messages_endpoint", + "caching", "configuration", "mcp", "observability", diff --git a/tests/integration/caching/test_semantic_cache_tool_turns.py b/tests/integration/caching/test_semantic_cache_tool_turns.py new file mode 100644 index 00000000000..6c9a0b060fc --- /dev/null +++ b/tests/integration/caching/test_semantic_cache_tool_turns.py @@ -0,0 +1,277 @@ +import hashlib +import json +import math +import threading +import uuid +from collections.abc import Iterator, Mapping, Sequence +from dataclasses import dataclass, field +from pathlib import Path +from typing import Final +from urllib.parse import urlsplit + +import pytest +import yaml +from integration._support.client import Gateway, eventually, gateway_from_environment +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_VECTOR_SIZE: Final = 16 +_CHAT_MODEL: Final = "semantic-chat" +_CLAUDE_MODEL: Final = "semantic-claude" +_EMBEDDING_MODEL: Final = "semantic-embedder" +_COLLECTION: Final = "semantic-tool-turns" + + +def _vector(text: str) -> tuple[float, ...]: + digest: Final = hashlib.sha256(text.encode()).digest() + return tuple((byte - 127.5) / 127.5 for byte in digest[:_VECTOR_SIZE]) + + +def _cosine(left: Sequence[float], right: Sequence[float]) -> float: + dot: Final = sum(a * b for a, b in zip(left, right, strict=True)) + return dot / (math.sqrt(sum(a * a for a in left)) * math.sqrt(sum(b * b for b in right))) + + +def _answer(body: Mapping[str, JsonValue]) -> str: + return "answer-" + hashlib.sha256(json.dumps(body["messages"], sort_keys=True).encode()).hexdigest()[:16] + + +@dataclass(slots=True) +class _Peer: + """Embeddings, chat, Anthropic Messages and a Qdrant collection on one owned socket""" + + lock: threading.Lock = field(default_factory=threading.Lock) + points: list[Mapping[str, JsonValue]] = field(default_factory=list) # mutable-ok: the Qdrant collection + embedded: list[str] = field(default_factory=list) # mutable-ok: every prompt the proxy embedded + answered: list[str] = field(default_factory=list) # mutable-ok: every completion the provider served + + def stored(self) -> int: + with self.lock: + return len(self.points) + + def respond(self, request: Request) -> Reply: + path: Final = urlsplit(request.target).path + body: Final = _JSON_OBJECT.validate_json(request.body) if request.body else {} + collection: Final = f"/qdrant/collections/{_COLLECTION}" + if request.method == "GET" and path == f"{collection}/exists": + return self._json({"result": {"exists": False}, "status": "ok"}) + if request.method in {"GET", "PUT"} and path in {collection, f"{collection}/index"}: + return self._json({"result": True, "status": "ok"}) + if request.method == "PUT" and path == f"{collection}/points": + with self.lock: + self.points.extend(_JSON_OBJECT.validate_python(point) for point in body["points"]) + return self._json({"result": {"status": "completed"}, "status": "ok"}) + if request.method == "POST" and path == f"{collection}/points/search": + return self._json({"result": self._search(body), "status": "ok"}) + if request.method == "GET" and path == "/v1/models": + return self._json({"object": "list", "data": []}) + if request.method == "POST" and path == "/v1/embeddings": + text: Final = str(body["input"]) + with self.lock: + self.embedded.append(text) + return self._json( + { + "object": "list", + "model": "text-embedding-3-small", + "data": [{"object": "embedding", "index": 0, "embedding": list(_vector(text))}], + "usage": {"prompt_tokens": 4, "total_tokens": 4}, + } + ) + if request.method == "POST" and path == "/v1/chat/completions": + return self._json(self._chat_reply(_answer(body))) + if request.method == "POST" and path == "/v1/messages": + return self._json(self._messages_reply(_answer(body))) + raise AssertionError(f"unexpected peer request {request.method} {request.target}") + + def _search(self, body: Mapping[str, JsonValue]) -> list[JsonValue]: + query: Final = [float(str(value)) for value in _list(body["vector"])] + key: Final = _JSON_OBJECT.validate_python( + _JSON_OBJECT.validate_python(_list(_JSON_OBJECT.validate_python(body["filter"])["must"])[0])["match"] + )["value"] + with self.lock: + scoped: Final = [ + point + for point in self.points + if _JSON_OBJECT.validate_python(point["payload"])["litellm_cache_key"] == key + ] + ranked: Final = sorted( + ( + { + "id": point["id"], + "score": _cosine(query, [float(str(value)) for value in _list(point["vector"])]), + "payload": point["payload"], + } + for point in scoped + ), + key=lambda hit: -float(str(hit["score"])), + ) + return list(ranked[:1]) + + def _chat_reply(self, answer: str) -> Mapping[str, JsonValue]: + with self.lock: + self.answered.append(answer) + return { + "id": f"chatcmpl-{uuid.uuid4().hex}", + "object": "chat.completion", + "created": 1789788253, + "model": "gpt-5.4-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": answer}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12}, + } + + def _messages_reply(self, answer: str) -> Mapping[str, JsonValue]: + with self.lock: + self.answered.append(answer) + return { + "id": f"msg_{uuid.uuid4().hex}", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-5-5", + "content": [{"type": "text", "text": answer}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 10, "output_tokens": 2}, + } + + @staticmethod + def _json(value: Mapping[str, JsonValue]) -> Reply: + return Reply(body=json.dumps(value).encode()) + + +def _list(value: JsonValue) -> list[JsonValue]: + assert isinstance(value, list), value + return value + + +@dataclass(frozen=True, slots=True) +class _SemanticProxy: + gateway: Gateway + peer: _Peer + + +@pytest.fixture(scope="module") +def semantic_proxy(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_SemanticProxy]: + peer: Final = _Peer() + directory: Final = tmp_path_factory.mktemp("semantic-tool-turns") + with gateway_from_environment() as gateway, wire_server(peer.respond) as wire: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["model_list"] = [ + { + "model_name": _CHAT_MODEL, + "litellm_params": {"model": "openai/gpt-5.4-mini", "api_base": f"{wire.url}/v1", "api_key": "k"}, + }, + { + "model_name": _CLAUDE_MODEL, + "litellm_params": {"model": "anthropic/claude-sonnet-5-5", "api_base": wire.url, "api_key": "k"}, + }, + { + "model_name": _EMBEDDING_MODEL, + "litellm_params": { + "model": "openai/text-embedding-3-small", + "api_base": f"{wire.url}/v1", + "api_key": "k", + }, + }, + ] + config["litellm_settings"]["cache_params"] = { + "type": "qdrant-semantic", + "qdrant_api_base": f"{wire.url}/qdrant", + "qdrant_collection_name": _COLLECTION, + "qdrant_semantic_cache_embedding_model": _EMBEDDING_MODEL, + "qdrant_semantic_cache_vector_size": _VECTOR_SIZE, + "qdrant_quantization_config": "binary", + "similarity_threshold": 0.99, + } + path: Final = directory / "semantic_tool_turns.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, directory, {}, config=path) as candidate: + yield _SemanticProxy(candidate, peer) + + +def _send_turns(proxy: _SemanticProxy, route: str, model: str, turns: Sequence[list[JsonValue]]) -> list[str]: + def reply_text(turn: list[JsonValue]) -> str: + stored_before: Final = proxy.peer.stored() + answered_before: Final = len(proxy.peer.answered) + body: Final[dict[str, JsonValue]] = {"model": model, "max_tokens": 16, "messages": turn} + response: Final = proxy.gateway.request("POST", route, body) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + if len(proxy.peer.answered) > answered_before: + eventually(proxy.peer.stored, lambda count: count > stored_before) + if route == "/v1/messages": + return str(_JSON_OBJECT.validate_python(_list(payload["content"])[0])["text"]) + choice: Final = _JSON_OBJECT.validate_python(_list(payload["choices"])[0]) + return str(_JSON_OBJECT.validate_python(choice["message"])["content"]) + + return [reply_text(turn) for turn in turns] + + +def test_claude_code_tool_turns_on_messages_are_not_served_the_first_turn_answer( + semantic_proxy: _SemanticProxy, +) -> None: + task: Final[JsonValue] = {"role": "user", "content": f"fix the failing test {uuid.uuid4().hex}"} + list_files: Final[list[JsonValue]] = [ + task, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": "toolu_1", "name": "Bash", "input": {"command": "ls"}}], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "calc.py test_calc.py"}], + }, + ] + read_file: Final[list[JsonValue]] = [ + *list_files, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": "toolu_2", "name": "Read", "input": {"file_path": "calc.py"}}], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_2", + "content": [{"type": "text", "text": "def add(a, b): return a - b"}], + } + ], + }, + ] + + answers: Final = _send_turns(semantic_proxy, "/v1/messages", _CLAUDE_MODEL, ([task], list_files, read_file)) + + 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 + + +def test_openai_agent_loop_tool_calls_on_chat_completions_are_not_served_a_cached_answer( + semantic_proxy: _SemanticProxy, +) -> None: + task: Final[JsonValue] = {"role": "user", "content": f"update both config files {uuid.uuid4().hex}"} + + def wrote(path: str) -> list[JsonValue]: + return [ + task, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "write_file", "arguments": json.dumps({"path": path})}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_1", "content": "ok"}, + ] + + answers: Final = _send_turns( + semantic_proxy, "/v1/chat/completions", _CHAT_MODEL, (wrote("a.yaml"), wrote("b.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 diff --git a/tests/integration/run.py b/tests/integration/run.py index 19bce35f542..2e82fa4539d 100644 --- a/tests/integration/run.py +++ b/tests/integration/run.py @@ -14,7 +14,7 @@ GROUPS: Final = MappingProxyType( "management": ("management", "authorization", "configuration"), "accounting": ("pricing", "spend"), "database": ("database",), - "providers": ("providers", "routing", "streaming", "messages_endpoint"), + "providers": ("providers", "routing", "streaming", "messages_endpoint", "caching"), "extensions": ("observability", "compatibility"), "mcp": ("mcp",), "sdk": ("sdk",), diff --git a/tests/unit/caching/test_redis_semantic_cache.py b/tests/unit/caching/test_redis_semantic_cache.py index 461689165bb..1e4a6c2bdc6 100644 --- a/tests/unit/caching/test_redis_semantic_cache.py +++ b/tests/unit/caching/test_redis_semantic_cache.py @@ -579,6 +579,41 @@ def test_redis_semantic_cache_prompt_extraction_prefers_messages(): assert prompt == "message prompt" +def test_redis_semantic_cache_prompt_extraction_keeps_tool_turns_distinct(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + def turn(command: str) -> list[dict[str, object]]: + return [ + {"role": "user", "content": "fix the failing test"}, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": "t1", "name": "Bash", "input": {"cmd": command}}], + }, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "t1", "content": "ok"}]}, + ] + + assert RedisSemanticCache._get_prompt_from_kwargs(messages=turn("ls")) == ( + 'fix the failing test{"name":"Bash","arguments":{"cmd":"ls"}}ok' + ) + assert RedisSemanticCache._get_prompt_from_kwargs(messages=turn("pwd")) == ( + 'fix the failing test{"name":"Bash","arguments":{"cmd":"pwd"}}ok' + ) + + +def test_redis_semantic_cache_prompt_extraction_keeps_responses_function_calls(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + prompt = RedisSemanticCache._get_prompt_from_kwargs( + input=[ + {"role": "user", "content": "update the config"}, + {"type": "function_call", "call_id": "c1", "name": "write_file", "arguments": '{"path":"a.yaml"}'}, + {"type": "function_call_output", "call_id": "c1", "output": "ok"}, + ] + ) + + assert prompt == 'update the config\n{"name":"write_file","arguments":"{\\"path\\":\\"a.yaml\\"}"}\nok' + + 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 0375ff14852..4c1af95f4fe 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,6 +7,8 @@ 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, @@ -16,6 +18,8 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( encrypted_reasoning_signature, get_file_ids_from_messages, get_format_from_file_id, + get_str_from_messages, + get_str_from_messages_with_tools, handle_any_messages_to_chat_completion_str_messages_conversion, hoist_images_from_tool_messages, is_encrypted_reasoning_block, @@ -2004,3 +2008,115 @@ class TestMergeConsecutiveSystemMessages: ) assert merged == [{"role": "system"}, {"role": "user", "content": "Hi"}] + + +_TASK: Final = {"role": "user", "content": "fix the failing test"} + + +@pytest.mark.parametrize( + ("messages", "expected"), + [ + pytest.param( + [ + _TASK, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": "t1", "name": "Bash", "input": {"cmd": "ls"}}], + }, + ], + 'fix the failing test{"name":"Bash","arguments":{"cmd":"ls"}}', + id="anthropic-tool-use-name-and-input-without-id", + ), + pytest.param( + [_TASK, {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "t1", "content": "calc.py"}]}], + "fix the failing testcalc.py", + id="anthropic-string-tool-result", + ), + pytest.param( + [ + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "t1", + "content": [{"type": "text", "text": "a"}, {"type": "text", "text": "b"}], + }, + {"type": "text", "text": "next"}, + ], + } + ], + "abnext", + id="anthropic-nested-text-tool-result-then-text", + ), + pytest.param( + [ + _TASK, + { + "role": "assistant", + "content": "writing", + "tool_calls": [ + {"id": "c1", "type": "function", "function": {"name": "write", "arguments": '{"path": "a"}'}}, + {"id": "c2", "type": "function", "function": {"name": "write", "arguments": '{"path": "b"}'}}, + ], + }, + {"role": "tool", "tool_call_id": "c1", "content": "ok"}, + ], + 'fix the failing testwriting{"name":"write","arguments":"{\\"path\\": \\"a\\"}"}' + '{"name":"write","arguments":"{\\"path\\": \\"b\\"}"}ok', + id="openai-tool-calls-in-order-before-tool-result", + ), + pytest.param( + [ + Message( + content=None, + tool_calls=[ChatCompletionMessageToolCall(id="c1", function=Function(name="read", arguments="{}"))], + ) + ], + '{"name":"read","arguments":"{}"}', + id="openai-response-message-object", + ), + pytest.param( + [{"role": "assistant", "content": None, "tool_calls": [{"id": "c1", "type": "function"}, "junk"]}, "junk"], + '{"name":null,"arguments":null}', + id="malformed-tool-call-entries", + ), + ], +) +def test_get_str_from_messages_with_tools_keeps_tool_exchange(messages: list[object], expected: str) -> None: + assert get_str_from_messages_with_tools(messages) == expected + + +@pytest.mark.parametrize( + "messages", + [ + pytest.param([_TASK, {"role": "assistant", "content": "done"}], id="string-content"), + pytest.param( + [ + { + "role": "user", + "content": [ + {"type": "text", "text": "what is "}, + {"type": "image_url", "image_url": {"url": "https://example.com/a.png"}}, + {"type": "text", "text": "this"}, + ], + } + ], + id="text-and-image-parts", + ), + pytest.param( + [ + { + "role": "tool", + "tool_call_id": "c1", + "content": "small", + "search_results": [{"source": "s", "title": "t"}], + } + ], + id="search-results", + ), + pytest.param([{"role": "assistant", "content": None}, {"role": "user"}], id="missing-content"), + ], +) +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