From 77fc3315e55b5dc22a99b4f3c0a7d76eb7e8c367 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 7 Oct 2026 17:32:39 +0000 Subject: [PATCH] fix(caching): skip the cache past max_messages and keep tool_result text in semantic prompts (#43878) * 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> * refactor(caching): keep semantic tool prompt helpers within lint budgets Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): keep structured function_call_output text in semantic prompts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(caching): split Responses text-field collection to stay within complexity budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): tag each tool result with the position of the call it answers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): encode tool result position and output together so tool text cannot forge result tags Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(caching): expect encoded tool result record in qdrant semantic prompt parity case Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(caching): cover tool result arrangements, SDK clients, concurrency and qdrant outage for semantic cache Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * feat(caching): embed every semantic cache prompt field except volatile ones Replace the per-shape allowlist in the Python and Rust semantic cache prompt walkers with one include-by-default walker. Plain text keeps its old concatenation; any other block or message is embedded as compact JSON with call ids mapped to ordinals, cache_control dropped, and signatures, encrypted content and base64 data replaced with a short sha256 digest. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(code-quality): allow the bounded semantic cache prompt walkers in the recursion check Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(rust): expect structured JSON for unknown fields in redis and valkey semantic prompts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * Revert "test(rust): expect structured JSON for unknown fields in redis and valkey semantic prompts" This reverts commit 86c82b949b3da25461e72f317b0f54572ea5359f. * Revert "test(code-quality): allow the bounded semantic cache prompt walkers in the recursion check" This reverts commit 39efb5d9da8292844a9844e01d095702b11ecec4. * Revert "feat(caching): embed every semantic cache prompt field except volatile ones" This reverts commit 5aed3ab3de3597dec22d3afcb33db7196f0a2a24. * refactor(caching): rename get_str_from_messages_with_tools to get_semantic_cache_prompt_from_messages Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(caching): split semantic cache prompt extraction by API format Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): drop TypeIs guard and register Responses prompt walker with the recursion check Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): pick the Responses text field without a Final inside a loop Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(caching): walk semantic cache prompts as plain dicts, dumping pydantic items once up front Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(caching): write the semantic cache prompt builders as plain loops Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): skip the cache past max_messages and keep tool_result text in semantic prompts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(caching): drop formatting-only churn from the redis semantic cache tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci(integration): drop the caching group wiring that main already carries Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(caching): recurse into tool_result content in the semantic cache prompt helper Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): check the max_messages cap on the shared exact-cache proxy Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): read list-form function_call_output text in semantic cache prompts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(caching): extract nested Responses input lookup to keep walker under complexity limit Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * Revert "refactor(caching): extract nested Responses input lookup to keep walker under complexity limit" This reverts commit 0665296bf1d783583ffc7f494e111ef58fec2b18. * style(caching): suppress C901 on the Responses input walker instead of splitting it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/crates/cache/src/semantic.rs | 37 ++++- litellm-rust/crates/cache/tests/semantic.rs | 34 +++++ litellm/caching/caching.py | 17 +++ litellm/caching/qdrant_semantic_cache.py | 10 +- litellm/caching/redis_semantic_cache.py | 11 +- .../prompt_templates/common_utils.py | 27 ++++ .../caching/test_cache_max_messages.py | 144 ++++++++++++++++++ tests/unit/caching/test_caching.py | 65 +++++++- .../unit/caching/test_redis_semantic_cache.py | 34 +++++ ...ore_utils_prompt_templates_common_utils.py | 92 +++++++++++ 10 files changed, 456 insertions(+), 15 deletions(-) create mode 100644 tests/integration/caching/test_cache_max_messages.py diff --git a/litellm-rust/crates/cache/src/semantic.rs b/litellm-rust/crates/cache/src/semantic.rs index c88706a213b..6200645e555 100644 --- a/litellm-rust/crates/cache/src/semantic.rs +++ b/litellm-rust/crates/cache/src/semantic.rs @@ -83,17 +83,16 @@ impl Embedder for PreparedEmbedding { } } -/// `get_str_from_messages`: every message's text content followed by its search results. +/// `get_semantic_cache_prompt_from_messages`: every message's text content, including the text of +/// Messages API `tool_result` blocks, followed by 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); - } + Some(Value::Array(blocks)) => { + for block in blocks { + push_block_text(&mut text, block); } } _ => {} @@ -103,6 +102,28 @@ pub fn str_from_messages(messages: &[Value]) -> String { text } +fn push_block_text(text: &mut String, block: &Value) { + if block.get("type").and_then(Value::as_str) != Some("tool_result") { + push_text_field(text, block); + return; + } + match block.get("content") { + Some(Value::String(result)) => text.push_str(result), + Some(Value::Array(blocks)) => { + for inner in blocks { + push_text_field(text, inner); + } + } + _ => {} + } +} + +fn push_text_field(text: &mut String, block: &Value) { + if let Some(block_text) = block.get("text").and_then(Value::as_str) { + text.push_str(block_text); + } +} + /// 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()?; @@ -163,6 +184,10 @@ fn collect_input_text(value: &Value, parts: &mut Vec) { collect_input_text(content, parts); return; } + if let Some(output) = map.get("output").filter(|output| output.is_array()) { + collect_input_text(output, parts); + return; + } for key in ["text", "output", "input_text", "output_text"] { if let Some(Value::String(text)) = map.get(key) && push_trimmed(text, parts) diff --git a/litellm-rust/crates/cache/tests/semantic.rs b/litellm-rust/crates/cache/tests/semantic.rs index 97a552a8010..76af1863e1f 100644 --- a/litellm-rust/crates/cache/tests/semantic.rs +++ b/litellm-rust/crates/cache/tests/semantic.rs @@ -30,6 +30,31 @@ fn context(messages: Option, input: Option) -> SemanticCacheContex ]}]), "What is this?", )] +#[case::tool_result_string( + json!([ + {"role": "user", "content": "list the files"}, + {"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"}, + ]}, + ]), + "list the filescalc.py test_calc.py", +)] +#[case::tool_result_blocks( + json!([{"role": "user", "content": [ + {"type": "tool_result", "tool_use_id": "toolu_1", "content": [ + {"type": "text", "text": "x = 1"}, + {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": ""}}, + ]}, + ]}]), + "x = 1", +)] +#[case::tool_result_without_content( + json!([{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_1"}]}]), + "", +)] #[case::missing_null_and_empty_content( json!([{"role": "assistant"}, {"role": "assistant", "content": null}, {"role": "user", "content": ""}]), "", @@ -166,6 +191,15 @@ fn prompt_from_messages_reads_messages_only( ])), Some("model dump prompt\ndict prompt\ninline prompt"), )] +#[case::function_call_output_blocks( + None, + Some(json!([ + {"role": "user", "content": "update the config"}, + {"type": "function_call", "call_id": "c1", "name": "write_file", "arguments": "{\"path\": \"a\"}"}, + {"type": "function_call_output", "call_id": "c1", "output": [{"type": "input_text", "text": "wrote a"}]}, + ])), + Some("update the config\nwrote a"), +)] #[case::object_content( None, Some(json!({"content": [{"text": "object content prompt"}]})), diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 1dc4de04dc1..85ef5a93937 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -71,6 +71,17 @@ class CacheMode(str, Enum): #### LiteLLM.Completion / Embedding Cache #### +def _request_message_count(kwargs: Mapping[str, object]) -> int: + """Chat and Messages API `messages`, else Responses API `input` items; embedding `input` strings count as none""" + messages: Final = kwargs.get("messages") + if isinstance(messages, list): + return len(messages) + input_items: Final = kwargs.get("input") + if not isinstance(input_items, list): + return 0 + return sum(1 for item in input_items if isinstance(item, (Mapping, BaseModel))) + + class Cache: def __init__( self, @@ -119,6 +130,7 @@ class Cache: semantic_cache_embedding_max_input_tokens: int | None = None, semantic_cache_embedding_timeout: float | None = None, semantic_cache_scope: str = SemanticCacheScope.KEY.value, + max_messages: int | None = 4, # GCP IAM authentication parameters gcp_service_account: str | None = None, gcp_ssl_ca_certs: str | None = None, @@ -148,6 +160,7 @@ class Cache: semantic_cache_embedding_max_input_tokens (int, optional): Truncate prompts to this many tokens before embedding them for semantic caching. Defaults to the embedding deployment's configured max_input_tokens. semantic_cache_embedding_timeout (float, optional): Seconds a semantic-cache lookup may spend embedding the prompt before it gives up and lets the request continue to the LLM. Defaults to SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS. semantic_cache_scope (str, optional): "key" isolates semantic-cache buckets per key/team/org. "end_user" additionally isolates per end user (falls back to the key scope when the request carries no end-user id). Defaults to "key". + max_messages (int, optional): Requests with more `messages` (or Responses API `input` items) than this are neither looked up nor stored, so long agent conversations never serve or create a cache entry. None disables the limit. Defaults to 4. # Disk Cache Args disk_cache_dir (str, optional): The directory for the disk cache. Defaults to None. @@ -298,6 +311,7 @@ class Cache: self.ttl = ttl self.mode: CacheMode = mode or CacheMode.default_on self.semantic_cache_scope: str = SemanticCacheScope(semantic_cache_scope).value + self.max_messages: int | None = max_messages if self.type == LiteLLMCacheType.LOCAL and default_in_memory_ttl is not None: self.ttl = default_in_memory_ttl @@ -933,7 +947,10 @@ class Cache: If cache is default_on then this is True If cache is default_off then this is only true when user has opted in to use cache + Always False once the request carries more than `max_messages` messages """ + if self.max_messages is not None and _request_message_count(kwargs) > self.max_messages: + return False if self.mode == CacheMode.default_on: return True diff --git a/litellm/caching/qdrant_semantic_cache.py b/litellm/caching/qdrant_semantic_cache.py index 8dac3f2eef9..eb88df066f6 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_semantic_cache_prompt_from_messages, ) 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_semantic_cache_prompt_from_messages(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_semantic_cache_prompt_from_messages(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_semantic_cache_prompt_from_messages(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_semantic_cache_prompt_from_messages(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..8f99d76ba46 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -21,7 +21,7 @@ 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_semantic_cache_prompt_from_messages, ) from litellm.types.utils import EmbeddingResponse @@ -263,7 +263,7 @@ class RedisSemanticCache(BaseCache): """ messages: Final = kwargs.get("messages") if messages: - return get_str_from_messages(messages) + return get_semantic_cache_prompt_from_messages(messages) if "input" not in kwargs: return None @@ -274,7 +274,7 @@ class RedisSemanticCache(BaseCache): return prompt or None @classmethod - def _collect_responses_input_text(cls, value: object, prompt_parts: list[str]) -> None: + def _collect_responses_input_text(cls, value: object, prompt_parts: list[str]) -> None: # noqa: C901 # one branch per Responses input shape value = cls._coerce_response_input_value(value) if value is None: return @@ -296,6 +296,11 @@ class RedisSemanticCache(BaseCache): cls._collect_responses_input_text(content, prompt_parts) return + output = value.get("output") + if isinstance(output, list): + cls._collect_responses_input_text(output, prompt_parts) + return + for text_key in ("text", "output", "input_text", "output_text"): text_value = value.get(text_key) if isinstance(text_value, str): diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 49c4198c939..1f75125df09 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -192,6 +192,33 @@ def get_str_from_messages(messages: list[AllMessageValues]) -> str: return text +def get_semantic_cache_prompt_from_messages(messages: Sequence[Mapping[str, object]]) -> str: + """ + The text a semantic cache embeds for a request: `get_str_from_messages` plus the text inside + Messages API `tool_result` blocks, so a tool turn does not embed identically to the turn before it + """ + return "".join( + _semantic_cache_content_text(message.get("content")) + + extract_search_results_text(message.get("search_results")) + for message in messages + ) + + +def _semantic_cache_content_text(content: object) -> str: + if isinstance(content, str): + return content + if not isinstance(content, list): + return "" + return "".join(_semantic_cache_block_text(block) for block in content if isinstance(block, Mapping)) + + +def _semantic_cache_block_text(block: Mapping[str, object]) -> str: + if block.get("type") == "tool_result": + return _semantic_cache_content_text(block.get("content")) + text: Final = block.get("text") + return text if isinstance(text, str) else "" + + 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/caching/test_cache_max_messages.py b/tests/integration/caching/test_cache_max_messages.py new file mode 100644 index 00000000000..89401b7c4ec --- /dev/null +++ b/tests/integration/caching/test_cache_max_messages.py @@ -0,0 +1,144 @@ +import json +import os +import uuid +from collections.abc import Callable +from typing import Final + +import pytest +from integration._support.anthropic_sse import message_json +from integration._support.client import Gateway, eventually, object_value, string_value +from integration._support.openai_wire import chat_reply, responses_reply +from integration._support.provider import SharedProvider +from integration._support.wire import Reply +from pydantic import JsonValue +from redis import Redis + +_Turns = Callable[[str], tuple[list[JsonValue], list[JsonValue]]] + + +def _claude_code_turns(task: str) -> tuple[list[JsonValue], list[JsonValue]]: + """Turns 1 and 3 of a Claude Code session on /v1/messages: 1 and 5 messages""" + first: Final[list[JsonValue]] = [{"role": "user", "content": task}] + third: Final[list[JsonValue]] = [ + *first, + { + "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"}], + }, + { + "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": "def add(a, b): return a - b"}], + }, + ] + return first, third + + +def _agent_turns(task: str) -> tuple[list[JsonValue], list[JsonValue]]: + """Turns 1 and 3 of an OpenAI tool loop on /v1/chat/completions: 2 and 6 messages""" + + def call(call_id: str, path: str) -> list[JsonValue]: + function: Final[JsonValue] = {"name": "write_file", "arguments": json.dumps({"path": path})} + return [ + { + "role": "assistant", + "content": None, + "tool_calls": [{"id": call_id, "type": "function", "function": function}], + }, + {"role": "tool", "tool_call_id": call_id, "content": f"wrote {path}"}, + ] + + first: Final[list[JsonValue]] = [ + {"role": "system", "content": "You are a coding agent"}, + {"role": "user", "content": task}, + ] + return first, [*first, *call("call_1", "a.yaml"), *call("call_2", "b.yaml")] + + +def _responses_turns(task: str) -> tuple[list[JsonValue], list[JsonValue]]: + """Turns 1 and 3 of an agent on /v1/responses: 1 and 5 input items""" + + def call(call_id: str, path: str) -> list[JsonValue]: + return [ + { + "type": "function_call", + "call_id": call_id, + "name": "write_file", + "arguments": json.dumps({"path": path}), + }, + {"type": "function_call_output", "call_id": call_id, "output": f"wrote {path}"}, + ] + + first: Final[list[JsonValue]] = [{"role": "user", "content": task}] + return first, [*first, *call("call_1", "a.yaml"), *call("call_2", "b.yaml")] + + +def _body(path: str, model: str, conversation: list[JsonValue]) -> dict[str, JsonValue]: + if path == "/v1/responses": + return {"model": model, "input": conversation} + return {"model": model, "max_tokens": 16, "messages": conversation} + + +def _reply(path: str, text: str) -> Reply: + identity: Final = f"id_{uuid.uuid4().hex}" + if path == "/v1/messages": + return Reply(body=message_json(identity, "claude-sonnet-5-5", text)) + if path == "/v1/responses": + return responses_reply(identity, "gpt-5.6-sol", text, stream=False) + return chat_reply(identity, "gpt-5.4", text, stream=False) + + +def _answer(path: str, payload: dict[str, JsonValue]) -> str: + if path == "/v1/messages": + return string_value(_first(payload["content"])["text"]) + if path == "/v1/responses": + return string_value(_first(_first(payload["output"])["content"])["text"]) + return string_value(object_value(_first(payload["choices"])["message"])["content"]) + + +def _first(value: JsonValue) -> dict[str, JsonValue]: + assert isinstance(value, list), value + return object_value(value[0]) + + +def _cached_responses(redis: Redis) -> frozenset[bytes]: + digests: Final = tuple(key for key in redis.scan_iter() if len(key) == 64) + return frozenset(key for key in digests if b'"response"' in (redis.get(key) or b"")) + + +@pytest.mark.parametrize( + ("path", "model", "turns"), + [ + pytest.param("/v1/messages", "anthropic/claude-sonnet-5-5", _claude_code_turns, id="messages"), + pytest.param("/v1/chat/completions", "openai/gpt-5.4", _agent_turns, id="chat-completions"), + pytest.param("/v1/responses", "openai/responses/gpt-5.6-sol", _responses_turns, id="responses"), + ], +) +def test_cache_serves_a_turn_under_max_messages_and_skips_one_past_it( + gateway: Gateway, provider: SharedProvider, path: str, model: str, turns: _Turns +) -> None: + under_cap, past_cap = turns(f"update the config {uuid.uuid4().hex}") + provider.expect(_reply(path, "first answer"), _reply(path, "second answer"), _reply(path, "third answer")) + + with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as redis: + cached_before: Final = _cached_responses(redis) + gateway.post(path, _body(path, model, under_cap)) + eventually(lambda: _cached_responses(redis) - cached_before, lambda written: len(written) == 1) + repeated: Final = gateway.post(path, _body(path, model, under_cap)) + past_cap_twice: Final = ( + gateway.post(path, _body(path, model, past_cap)), + gateway.post(path, _body(path, model, past_cap)), + ) + + assert _answer(path, repeated) == "first answer", "a repeated turn under max_messages was not served from the cache" + assert tuple(_answer(path, answer) for answer in past_cap_twice) == ("second answer", "third answer"), ( + "a turn past max_messages was served from the cache" + ) + assert len(provider.received()) == 3 diff --git a/tests/unit/caching/test_caching.py b/tests/unit/caching/test_caching.py index 0a7ac3ecad1..0adb6b9a6f4 100644 --- a/tests/unit/caching/test_caching.py +++ b/tests/unit/caching/test_caching.py @@ -1,6 +1,7 @@ import asyncio import logging import re +import uuid from typing import Final from unittest.mock import MagicMock @@ -9,7 +10,7 @@ import pytest import litellm import litellm.caching.redis_cache as redis_cache_module from litellm._internal_context import current_service_target -from litellm.caching.caching import Cache, response_cache_phase +from litellm.caching.caching import Cache, CacheMode, response_cache_phase from litellm.caching.caching_handler import _PENDING_CACHE_WRITES from litellm.caching.in_memory_cache import InMemoryCache from litellm.caching.redis_cache import RedisCache, _RedisTimeoutLogThrottle @@ -473,3 +474,65 @@ async def test_a_lookup_already_inside_the_phase_does_not_open_a_second_one(v2_s await cache.async_get_cache(dynamic_cache_object=backend, **_REQUEST) assert backend.seen == [("llm_response", "cache.get llm_response")] assert [s.name for s in v2_span_exporter.get_finished_spans()] == ["cache.get llm_response"] + + +_TOOL_TURN_ITEM: Final = {"role": "user", "content": "hi"} + + +@pytest.mark.parametrize( + ("kwargs", "expected"), + [ + pytest.param({"messages": [_TOOL_TURN_ITEM] * 4}, True, id="four-messages-are-cached"), + pytest.param({"messages": [_TOOL_TURN_ITEM] * 5}, False, id="five-messages-skip-the-cache"), + pytest.param({"input": [_TOOL_TURN_ITEM] * 4}, True, id="four-responses-items-are-cached"), + pytest.param({"input": [_TOOL_TURN_ITEM] * 5}, False, id="five-responses-items-skip-the-cache"), + pytest.param({"input": "one prompt"}, True, id="string-input-is-one-message"), + pytest.param({"input": ["a", "b", "c", "d", "e"]}, True, id="embedding-strings-are-not-messages"), + ], +) +def test_should_use_cache_stops_past_the_default_max_messages(kwargs: dict[str, object], expected: bool) -> None: + assert Cache(type=LiteLLMCacheType.LOCAL).should_use_cache(**kwargs) is expected + + +def test_responses_sdk_items_count_toward_max_messages() -> None: + from openai.types.responses import ResponseFunctionToolCall + + call: Final = ResponseFunctionToolCall(type="function_call", call_id="c1", name="ls", arguments="{}") + + assert Cache(type=LiteLLMCacheType.LOCAL).should_use_cache(input=[_TOOL_TURN_ITEM, call, call, call, call]) is False + + +def test_max_messages_is_configurable_and_none_disables_it() -> None: + three: Final = [_TOOL_TURN_ITEM] * 3 + + assert Cache(type=LiteLLMCacheType.LOCAL, max_messages=2).should_use_cache(messages=three) is False + assert Cache(type=LiteLLMCacheType.LOCAL, max_messages=3).should_use_cache(messages=three) is True + assert Cache(type=LiteLLMCacheType.LOCAL, max_messages=None).should_use_cache(messages=three * 50) is True + + +def test_max_messages_beats_an_explicit_use_cache_opt_in() -> None: + cache: Final = Cache(type=LiteLLMCacheType.LOCAL, mode=CacheMode.default_off) + + assert cache.should_use_cache(messages=[_TOOL_TURN_ITEM] * 4, cache={"use-cache": True}) is True + assert cache.should_use_cache(messages=[_TOOL_TURN_ITEM] * 5, cache={"use-cache": True}) is False + + +def test_completion_past_max_messages_is_neither_served_from_nor_written_to_the_cache( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(litellm, "cache", Cache(type=LiteLLMCacheType.LOCAL)) + tag: Final = uuid.uuid4().hex + four: Final = [{"role": "user", "content": f"{tag} turn {index}"} for index in range(4)] + five: Final = [*four, {"role": "user", "content": f"{tag} turn 4"}] + + def answer(messages: list[dict[str, str]], mock_response: str) -> str: + response: Final = litellm.completion(model="gpt-4o-mini", messages=messages, mock_response=mock_response) + assert isinstance(response, litellm.ModelResponse), response + choice: Final = response.choices[0] + assert isinstance(choice, litellm.Choices), choice + return str(choice.message.content) + + assert answer(four, "four first") == "four first" + assert answer(four, "four second") == "four first", "a 4-message repeat missed the cache" + assert answer(five, "five first") == "five first" + assert answer(five, "five second") == "five second", "a 5-message repeat was served from the cache" diff --git a/tests/unit/caching/test_redis_semantic_cache.py b/tests/unit/caching/test_redis_semantic_cache.py index 99844c695cb..93db26fc9b2 100644 --- a/tests/unit/caching/test_redis_semantic_cache.py +++ b/tests/unit/caching/test_redis_semantic_cache.py @@ -568,6 +568,20 @@ def test_redis_semantic_cache_set_cache_flattens_structured_responses_input(): ) +def test_redis_semantic_cache_prompt_extraction_reads_function_call_output_blocks(): + 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"}'}, + {"type": "function_call_output", "call_id": "c1", "output": [{"type": "input_text", "text": "wrote a"}]}, + ] + ) + + assert prompt == "update the config\nwrote a" + + def test_redis_semantic_cache_prompt_extraction_prefers_messages(): from litellm.caching.redis_semantic_cache import RedisSemanticCache @@ -1416,3 +1430,23 @@ async def test_redis_async_embedding_truncates_off_the_event_loop(monkeypatch): assert embedding == [0.1, 0.2] assert _token_count("sem-embed", router.aembedding.call_args.kwargs["input"]) == 5 assert_loop_stayed_free(took, lags) + + +def test_redis_semantic_cache_prompt_extraction_keeps_tool_result_text(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + prompt = RedisSemanticCache._get_prompt_from_kwargs( + messages=[ + {"role": "user", "content": "list the files"}, + { + "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"}], + }, + ] + ) + + assert prompt == "list the filescalc.py test_calc.py" 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 7415c74226d..1d19264fb9e 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 @@ -16,6 +16,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_semantic_cache_prompt_from_messages, + get_str_from_messages, handle_any_messages_to_chat_completion_str_messages_conversion, hoist_images_from_tool_messages, is_encrypted_reasoning_block, @@ -2171,3 +2173,93 @@ class TestMergeConsecutiveSystemMessages: ) assert merged == [{"role": "system"}, {"role": "user", "content": "Hi"}] + + +_CLAUDE_CODE_TOOL_TURN: Final = [ + {"role": "user", "content": "list the files"}, + { + "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"}], + }, +] + + +@pytest.mark.parametrize( + ("messages", "expected"), + [ + pytest.param(_CLAUDE_CODE_TOOL_TURN, "list the filescalc.py test_calc.py", id="tool-result-string"), + pytest.param( + [ + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_1", + "content": [ + {"type": "text", "text": "x = 1"}, + {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": ""}}, + ], + } + ], + } + ], + "x = 1", + id="tool-result-blocks", + ), + pytest.param( + [{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_1"}]}], + "", + id="tool-result-without-content", + ), + ], +) +def test_get_semantic_cache_prompt_from_messages_keeps_tool_result_text( + messages: list[dict[str, object]], expected: str +) -> None: + assert get_semantic_cache_prompt_from_messages(messages) == expected + + +def test_get_semantic_cache_prompt_from_messages_differs_from_the_turn_before_it() -> None: + assert get_str_from_messages(_CLAUDE_CODE_TOOL_TURN) == get_str_from_messages(_CLAUDE_CODE_TOOL_TURN[:1]) + assert get_semantic_cache_prompt_from_messages(_CLAUDE_CODE_TOOL_TURN) != get_semantic_cache_prompt_from_messages( + _CLAUDE_CODE_TOOL_TURN[:1] + ) + + +@pytest.mark.parametrize( + "messages", + [ + pytest.param([{"role": "system", "content": "be brief. "}, {"role": "user", "content": "hello"}], id="strings"), + 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-parts", + ), + pytest.param( + [ + {"role": "assistant"}, + {"role": "assistant", "content": None}, + {"role": "user", "content": ""}, + {"role": "tool", "content": "small", "search_results": [{"source": "s", "title": "t", "content": []}]}, + ], + id="empty-content-and-search-results", + ), + ], +) +def test_get_semantic_cache_prompt_from_messages_matches_get_str_from_messages_without_tool_results( + messages: list[dict[str, object]], +) -> None: + assert get_semantic_cache_prompt_from_messages(messages) == get_str_from_messages(messages)