mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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 commit86c82b949b. * Revert "test(code-quality): allow the bounded semantic cache prompt walkers in the recursion check" This reverts commit39efb5d9da. * Revert "feat(caching): embed every semantic cache prompt field except volatile ones" This reverts commit5aed3ab3de. * 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 commit0665296bf1. * 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:
parent
79209b92a1
commit
77fc3315e5
10 changed files with 456 additions and 15 deletions
37
litellm-rust/crates/cache/src/semantic.rs
vendored
37
litellm-rust/crates/cache/src/semantic.rs
vendored
|
|
@ -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)
|
||||
|
|
|
|||
34
litellm-rust/crates/cache/tests/semantic.rs
vendored
34
litellm-rust/crates/cache/tests/semantic.rs
vendored
|
|
@ -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"}]})),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
144
tests/integration/caching/test_cache_max_messages.py
Normal file
144
tests/integration/caching/test_cache_max_messages.py
Normal 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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue