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 86c82b949b.

* Revert "test(code-quality): allow the bounded semantic cache prompt walkers in the recursion check"

This reverts commit 39efb5d9da.

* Revert "feat(caching): embed every semantic cache prompt field except volatile ones"

This reverts commit 5aed3ab3de.

* 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 0665296bf1.

* 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 <kerry@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-07 17:32:39 +00:00 • committed by GitHub
parent 79209b92a1
commit 77fc3315e5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 456 additions and 15 deletions

View file

@ -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<String> {
let messages = context.messages.as_ref()?.as_array()?;
@ -163,6 +184,10 @@ fn collect_input_text(value: &Value, parts: &mut Vec<String>) {
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)

View file

@ -30,6 +30,31 @@ fn context(messages: Option<Value>, input: Option<Value>) -> 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"}]})),

View file

@ -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

View file

@ -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"))

View file

@ -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):

View file

@ -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)

View file

@ -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

View file

@ -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"

View file

@ -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"

View file

@ -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)