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>
This commit is contained in:
kerry 2026-10-06 23:22:19 +00:00
parent ffd3a85a18
commit 394abff55a
12 changed files with 780 additions and 1193 deletions

View file

@ -250,7 +250,7 @@ async fn async_and_sync_set_get_store_exact_payload(entry: JsonValue) {
"citations": {"page": 1, "section": "intro"},
}],
}]),
r#"{"result_of_call":null,"output":""}sourcetitlebody{"page":1,"section":"intro"}"#
r#"sourcetitlebody{"page":1,"section":"intro"}"#
)]
#[tokio::test(flavor = "multi_thread")]
async fn prompt_matches_python_message_rules(

View file

@ -1,15 +1,14 @@
//! The embedding and prompt contract every semantic backend shares.
//!
//! Python's semantic caches all read their prompt through `get_semantic_cache_prompt_from_messages`, and
//! `RedisSemanticCache._get_prompt_from_kwargs` (inherited by Valkey) adds Responses API `input`
//! through `get_semantic_cache_prompt_from_responses_input`. Qdrant reads messages only. Each
//! backend picks one of the two extractors here.
//! Python's semantic caches all read their prompt through `get_str_from_messages`, and
//! `RedisSemanticCache._get_prompt_from_kwargs` (inherited by Valkey) adds Responses API
//! `input`. Qdrant reads messages only. Each backend picks one of the two extractors here.
use std::{collections::HashMap, future::Future, io};
use std::{future::Future, io};
use serde::Serialize;
use serde_json::{
Value, json,
Value,
ser::{CharEscape, Formatter, Serializer},
};
@ -84,122 +83,45 @@ impl Embedder for PreparedEmbedding {
}
}
/// `get_semantic_cache_prompt_from_messages`: every message's content text, tool calls and tool results,
/// then its OpenAI `tool_calls`, then its search results. Each tool result is encoded with the
/// position of the call it answers.
/// `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 messages: Vec<_> = messages.iter().filter_map(Value::as_object).collect();
let call_ordinals = tool_call_ordinals(messages.iter().flat_map(|message| {
let block_ids = message
.get("content")
.and_then(Value::as_array)
.into_iter()
.flatten()
.filter(|block| block.get("type").and_then(Value::as_str) == Some("tool_use"))
.filter_map(|block| block.get("id"));
let tool_call_ids = message
.get("tool_calls")
.and_then(Value::as_array)
.into_iter()
.flatten()
.filter_map(|tool_call| tool_call.get("id"));
block_ids.chain(tool_call_ids)
}));
let mut text = String::new();
for message in messages {
if message.get("role").and_then(Value::as_str) == Some("tool") {
let mut output = String::new();
push_content_text(&mut output, message.get("content"), &call_ordinals);
text.push_str(&tool_result_json(
message.get("tool_call_id"),
&call_ordinals,
&output,
));
} else {
push_content_text(&mut text, message.get("content"), &call_ordinals);
}
if let Some(Value::Array(tool_calls)) = message.get("tool_calls") {
for tool_call in tool_calls.iter().filter_map(Value::as_object) {
let function = tool_call.get("function");
text.push_str(&tool_call_json(
function.and_then(|function| function.get("name")),
function.and_then(|function| function.get("arguments")),
));
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(blocks)) => {
for block in blocks {
push_block_text(&mut text, block);
}
}
_ => {}
}
push_search_results_text(&mut text, message.get("search_results"));
}
text
}
/// `_content_str_with_tools`: text parts, Anthropic `tool_use` blocks and `tool_result` content.
fn push_content_text(
text: &mut String,
content: Option<&Value>,
call_ordinals: &HashMap<&str, usize>,
) {
match content {
Some(Value::String(content)) => text.push_str(content),
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 block in blocks.iter().filter_map(Value::as_object) {
match block.get("type").and_then(Value::as_str) {
Some("tool_use") => {
text.push_str(&tool_call_json(block.get("name"), block.get("input")));
}
Some("tool_result") => {
let mut output = String::new();
push_content_text(&mut output, block.get("content"), call_ordinals);
text.push_str(&tool_result_json(
block.get("tool_use_id"),
call_ordinals,
&output,
));
}
_ => {
if let Some(block_text) = block.get("text").and_then(Value::as_str) {
text.push_str(block_text);
}
}
}
for inner in blocks {
push_text_field(text, inner);
}
}
_ => {}
}
}
/// `tool_call_ordinals`: the 1-based position of each distinct string call id, first seen first.
fn tool_call_ordinals<'a>(call_ids: impl Iterator<Item = &'a Value>) -> HashMap<&'a str, usize> {
let mut ordinals = HashMap::new();
for call_id in call_ids.filter_map(Value::as_str) {
let next = ordinals.len() + 1;
ordinals.entry(call_id).or_insert(next);
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);
}
ordinals
}
/// `tool_result_str`: `{"result_of_call":N,"output":...}`, with a `null` position when the result
/// answers no known call.
fn tool_result_json(
call_id: Option<&Value>,
call_ordinals: &HashMap<&str, usize>,
output: &str,
) -> String {
let ordinal = call_id
.and_then(Value::as_str)
.and_then(|call_id| call_ordinals.get(call_id));
format!(
"{{\"result_of_call\":{},\"output\":{}}}",
compact_json(&json!(ordinal)),
compact_json(&Value::String(output.to_owned())),
)
}
/// `tool_call_str`: the compact `{"name":...,"arguments":...}` a tool call contributes.
fn tool_call_json(name: Option<&Value>, arguments: Option<&Value>) -> String {
compact_json(&json!({
"name": name.unwrap_or(&Value::Null),
"arguments": arguments.unwrap_or(&Value::Null),
}))
}
/// The messages prompt Qdrant embeds: `None` when the request carries no messages.
@ -208,9 +130,8 @@ pub fn prompt_from_messages(context: &SemanticCacheContext) -> Option<String> {
(!messages.is_empty()).then(|| str_from_messages(messages))
}
/// `RedisSemanticCache._get_prompt_from_kwargs`: chat messages first, then
/// `get_semantic_cache_prompt_from_responses_input` over a Responses API `input`. `None` when
/// neither yields a prompt.
/// `RedisSemanticCache._get_prompt_from_kwargs`: chat messages first, then the text parts of a
/// Responses API `input`. `None` when neither yields a prompt.
pub fn prompt_from_context(context: &SemanticCacheContext) -> Option<String> {
if let Some(messages) = context.messages.as_ref().and_then(Value::as_array)
&& !messages.is_empty()
@ -218,16 +139,8 @@ pub fn prompt_from_context(context: &SemanticCacheContext) -> Option<String> {
return Some(str_from_messages(messages));
}
let input = context.input.as_ref()?;
let call_ordinals = tool_call_ordinals(
input
.as_array()
.into_iter()
.flatten()
.filter(|item| item.get("type").and_then(Value::as_str) == Some("function_call"))
.filter_map(|item| item.get("call_id")),
);
let mut parts = Vec::new();
collect_input_text(input, &mut parts, &call_ordinals);
collect_input_text(input, &mut parts);
let prompt = python_strip(&parts.join("\n")).to_owned();
(!prompt.is_empty()).then_some(prompt)
}
@ -256,49 +169,26 @@ fn push_search_results_text(text: &mut String, search_results: Option<&Value>) {
}
}
fn collect_input_text(
value: &Value,
parts: &mut Vec<String>,
call_ordinals: &HashMap<&str, usize>,
) {
fn collect_input_text(value: &Value, parts: &mut Vec<String>) {
match value {
Value::String(text) => {
push_trimmed(text, parts);
}
Value::Array(items) => {
for item in items {
collect_input_text(item, parts, call_ordinals);
collect_input_text(item, parts);
}
}
Value::Object(map) => {
if map.get("type").and_then(Value::as_str) == Some("function_call") {
parts.push(tool_call_json(map.get("name"), map.get("arguments")));
return;
}
if map.get("type").and_then(Value::as_str) == Some("function_call_output") {
let mut output_parts = Vec::new();
if let Some(output) = map.get("output") {
collect_input_text(output, &mut output_parts, call_ordinals);
}
parts.push(tool_result_json(
map.get("call_id"),
call_ordinals,
python_strip(&output_parts.join("\n")),
));
return;
}
if let Some(content) = map.get("content").filter(|content| !content.is_null()) {
collect_input_text(content, parts, call_ordinals);
collect_input_text(content, parts);
return;
}
for key in ["text", "output", "input_text", "output_text"] {
match map.get(key) {
Some(nested @ Value::Array(_)) => {
collect_input_text(nested, parts, call_ordinals);
return;
}
Some(Value::String(text)) if push_trimmed(text, parts) => return,
_ => {}
if let Some(Value::String(text)) = map.get(key)
&& push_trimmed(text, parts)
{
return;
}
}
}

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": ""}]),
"",
@ -38,143 +63,67 @@ fn context(messages: Option<Value>, input: Option<Value>) -> SemanticCacheContex
json!([{"role": "tool", "content": "small", "search_results": [
{"source": "s", "title": "t", "content": [{"text": "hidden payload"}]},
]}]),
r#"{"result_of_call":null,"output":"small"}sthidden payload"#,
"smallsthidden payload",
)]
#[case::title_only_search_result(
json!([{"role": "tool", "content": "small", "search_results": [
{"source": "s", "title": "long title", "content": []},
]}]),
r#"{"result_of_call":null,"output":"small"}slong title"#,
"smallslong title",
)]
#[case::search_results_without_content(
json!([{"role": "tool", "search_results": [{"source": "s", "title": "t"}]}]),
r#"{"result_of_call":null,"output":""}st"#,
"st",
)]
#[case::search_result_fields_in_python_order(
json!([{"role": "tool", "content": "c", "search_results": [
{"citations": {"enabled": true}, "content": [{"text": "body"}], "title": "t", "source": "s"},
{"source": "s2"},
]}]),
r#"{"result_of_call":null,"output":"c"}stbody{"enabled":true}s2"#,
r#"cstbody{"enabled":true}s2"#,
)]
#[case::null_citations_skipped(
json!([{"role": "tool", "content": "c", "search_results": [
{"source": "s", "citations": null},
]}]),
r#"{"result_of_call":null,"output":"c"}s"#,
"cs",
)]
#[case::non_string_and_non_object_entries_skipped(
json!([{"role": "tool", "content": "c", "search_results": [
"junk",
{"source": 1, "title": null, "content": ["junk", {"text": 3}, {"text": "kept"}]},
]}]),
r#"{"result_of_call":null,"output":"c"}kept"#,
"ckept",
)]
#[case::non_list_search_results_skipped(
json!([{"role": "tool", "content": "c", "search_results": {"source": "s"}}]),
r#"{"result_of_call":null,"output":"c"}"#,
"c",
)]
#[case::citations_compact_in_insertion_order(
json!([{"role": "tool", "search_results": [
{"citations": {"z": 1, "a": [1.5, true, null], "m": {"k": "v"}}},
]}]),
r#"{"result_of_call":null,"output":""}{"z":1,"a":[1.5,true,null],"m":{"k":"v"}}"#,
r#"{"z":1,"a":[1.5,true,null],"m":{"k":"v"}}"#,
)]
#[case::citations_ensure_ascii(
json!([{"role": "tool", "search_results": [{"citations": ["caf\u{e9}", "\u{4e2d}"]}]}]),
r#"{"result_of_call":null,"output":""}["caf\u00e9","\u4e2d"]"#,
r#"["caf\u00e9","\u4e2d"]"#,
)]
#[case::citations_astral_chars_as_surrogate_pairs(
json!([{"role": "tool", "search_results": [{"citations": "\u{1f600}"}]}]),
r#"{"result_of_call":null,"output":""}"\ud83d\ude00""#,
r#""\ud83d\ude00""#,
)]
#[case::citations_escapes(
json!([{"role": "tool", "search_results": [{"citations": "q\"\\\n\t\u{1}/"}]}]),
r#"{"result_of_call":null,"output":""}"q\"\\\n\t\u0001/""#,
r#""q\"\\\n\t\u0001/""#,
)]
#[case::citations_large_float_exponent(
json!([{"role": "tool", "search_results": [{"citations": [1e20, 1.0]}]}]),
r#"{"result_of_call":null,"output":""}[1e+20,1.0]"#,
"[1e+20,1.0]",
)]
#[case::citations_scalars(
json!([{"role": "tool", "search_results": [{"citations": false}, {"citations": 3}]}]),
r#"{"result_of_call":null,"output":""}false3"#,
)]
#[case::anthropic_tool_use_name_and_input_without_id(
json!([
{"role": "user", "content": "fix the failing test"},
{"role": "assistant", "content": [
{"type": "tool_use", "id": "t1", "name": "Bash", "input": {"cmd": "ls"}},
]},
]),
r#"fix the failing test{"name":"Bash","arguments":{"cmd":"ls"}}"#,
)]
#[case::anthropic_string_tool_result(
json!([{"role": "user", "content": [
{"type": "tool_result", "tool_use_id": "t1", "content": "calc.py"},
]}]),
r#"{"result_of_call":null,"output":"calc.py"}"#,
)]
#[case::anthropic_nested_text_tool_result_then_text(
json!([{"role": "user", "content": [
{"type": "tool_result", "tool_use_id": "t1", "content": [
{"type": "text", "text": "a"},
{"type": "text", "text": "b"},
]},
{"type": "text", "text": "next"},
]}]),
r#"{"result_of_call":null,"output":"ab"}next"#,
)]
#[case::openai_tool_calls_in_order_before_tool_result(
json!([
{"role": "assistant", "content": "writing", "tool_calls": [
{"id": "c1", "type": "function", "function": {"name": "write", "arguments": "{\"path\": \"a\"}"}},
{"id": "c2", "type": "function", "function": {"name": "write", "arguments": "{\"path\": \"b\"}"}},
]},
{"role": "tool", "tool_call_id": "c1", "content": "ok"},
]),
r#"writing{"name":"write","arguments":"{\"path\": \"a\"}"}{"name":"write","arguments":"{\"path\": \"b\"}"}{"result_of_call":1,"output":"ok"}"#,
)]
#[case::anthropic_parallel_tool_results_tagged_with_their_call(
json!([
{"role": "assistant", "content": [
{"type": "tool_use", "id": "t1", "name": "Read", "input": {"path": "a"}},
{"type": "tool_use", "id": "t2", "name": "Read", "input": {"path": "b"}},
]},
{"role": "user", "content": [
{"type": "tool_result", "tool_use_id": "t2", "content": "B"},
{"type": "tool_result", "tool_use_id": "t1", "content": "A"},
]},
]),
r#"{"name":"Read","arguments":{"path":"a"}}{"name":"Read","arguments":{"path":"b"}}{"result_of_call":2,"output":"B"}{"result_of_call":1,"output":"A"}"#,
)]
#[case::reused_call_ids_keep_first_position(
json!([
{"role": "assistant", "content": null, "tool_calls": [
{"id": "call_0", "type": "function", "function": {"name": "a", "arguments": "{}"}},
]},
{"role": "assistant", "content": null, "tool_calls": [
{"id": "call_0", "type": "function", "function": {"name": "b", "arguments": "{}"}},
{"id": "call_1", "type": "function", "function": {"name": "c", "arguments": "{}"}},
]},
{"role": "tool", "tool_call_id": "call_1", "content": "C"},
]),
r#"{"name":"a","arguments":"{}"}{"name":"b","arguments":"{}"}{"name":"c","arguments":"{}"}{"result_of_call":2,"output":"C"}"#,
)]
#[case::unknown_call_ids_encode_a_null_position(
json!([
{"role": "tool", "tool_call_id": "c9", "content": "ok"},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "t9", "content": "done"}]},
]),
r#"{"result_of_call":null,"output":"ok"}{"result_of_call":null,"output":"done"}"#,
)]
#[case::tool_output_cannot_forge_an_encoded_result(
json!([{"role": "tool", "tool_call_id": "c1", "content": "\"},{\"result_of_call\":2,\"output\":\""}]),
r#"{"result_of_call":null,"output":"\"},{\"result_of_call\":2,\"output\":\""}"#,
)]
#[case::malformed_tool_call_entries(
json!([{"role": "assistant", "content": null, "tool_calls": [{"id": "c1", "type": "function"}, "junk"]}]),
r#"{"name":null,"arguments":null}"#,
"false3",
)]
fn str_from_messages_matches_python(#[case] messages: Value, #[case] expected: &str) {
assert_eq!(str_from_messages(messages.as_array().unwrap()), expected);
@ -268,50 +217,6 @@ fn prompt_from_messages_reads_messages_only(
)]
#[case::nested_lists(None, Some(json!([["a", [" b "]], "", "c"])), Some("a\nb\nc"))]
#[case::scalars_ignored(None, Some(json!([1, true, null, "kept"])), Some("kept"))]
#[case::responses_function_call(
None,
Some(json!([
{"role": "user", "content": "update the config"},
{"type": "function_call", "call_id": "c1", "name": "write_file", "arguments": "{\"path\":\"a.yaml\"}"},
{"type": "function_call_output", "call_id": "c1", "output": "ok"},
])),
Some("update the config\n{\"name\":\"write_file\",\"arguments\":\"{\\\"path\\\":\\\"a.yaml\\\"}\"}\n{\"result_of_call\":1,\"output\":\"ok\"}"),
)]
#[case::responses_structured_function_call_output(
None,
Some(json!([
{"type": "function_call", "call_id": "c1", "name": "write_file", "arguments": "{}"},
{"type": "function_call_output", "call_id": "c1", "output": [{"type": "input_text", "text": "denied"}]},
])),
Some("{\"name\":\"write_file\",\"arguments\":\"{}\"}\n{\"result_of_call\":1,\"output\":\"denied\"}"),
)]
#[case::responses_multi_part_output_joined_by_lines(
None,
Some(json!([
{"type": "function_call", "call_id": "c1", "name": "run", "arguments": "{}"},
{"type": "function_call_output", "call_id": "c1", "output": [
{"type": "input_text", "text": " line one "},
{"type": "input_text", "text": "line two"},
]},
])),
Some(r#"{"name":"run","arguments":"{}"}
{"result_of_call":1,"output":"line one\nline two"}"#),
)]
#[case::responses_parallel_outputs_tagged_with_their_call(
None,
Some(json!([
{"type": "function_call", "call_id": "c1", "name": "read", "arguments": "a"},
{"type": "function_call", "call_id": "c2", "name": "read", "arguments": "b"},
{"type": "function_call_output", "call_id": "c2", "output": "B"},
{"type": "function_call_output", "call_id": "c1", "output": "A"},
])),
Some("{\"name\":\"read\",\"arguments\":\"a\"}\n{\"name\":\"read\",\"arguments\":\"b\"}\n{\"result_of_call\":2,\"output\":\"B\"}\n{\"result_of_call\":1,\"output\":\"A\"}"),
)]
#[case::responses_unknown_call_id_encodes_a_null_position(
None,
Some(json!([{"type": "function_call_output", "call_id": "c9", "output": "orphan"}])),
Some(r#"{"result_of_call":null,"output":"orphan"}"#),
)]
fn prompt_from_context_matches_python(
#[case] messages: Option<Value>,
#[case] input: Option<Value>,

View file

@ -92,6 +92,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:
_native_cache: "ResponseCacheRuntime | None" = None
@ -142,6 +153,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,
@ -171,6 +183,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.
@ -321,6 +334,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
@ -1005,7 +1019,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

@ -22,7 +22,6 @@ 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_semantic_cache_prompt_from_messages,
get_semantic_cache_prompt_from_responses_input,
)
from litellm.types.utils import EmbeddingResponse
@ -269,7 +268,65 @@ class RedisSemanticCache(BaseCache):
if "input" not in kwargs:
return None
return get_semantic_cache_prompt_from_responses_input(kwargs.get("input")) or None
prompt_parts: Final[list[str]] = []
cls._collect_responses_input_text(kwargs.get("input"), prompt_parts)
prompt: Final = "\n".join(prompt_parts).strip()
return prompt or None
@classmethod
def _collect_responses_input_text(cls, value: object, prompt_parts: list[str]) -> None:
value = cls._coerce_response_input_value(value)
if value is None:
return
if isinstance(value, str):
stripped_value: Final = value.strip()
if stripped_value:
prompt_parts.append(stripped_value)
return
if isinstance(value, (list, tuple)):
for item in value:
cls._collect_responses_input_text(item, prompt_parts)
return
if isinstance(value, dict):
content = value.get("content")
if content is not None:
cls._collect_responses_input_text(content, prompt_parts)
return
for text_key in ("text", "output", "input_text", "output_text"):
text_value = value.get(text_key)
if isinstance(text_value, str):
stripped_text = text_value.strip()
if stripped_text:
prompt_parts.append(stripped_text)
return
return
content = getattr(value, "content", None)
if content is not None:
cls._collect_responses_input_text(content, prompt_parts)
return
for text_key in ("text", "output", "input_text", "output_text"):
text_value = getattr(value, text_key, None)
if isinstance(text_value, str):
stripped_text = text_value.strip()
if stripped_text:
prompt_parts.append(stripped_text)
return
@staticmethod
def _coerce_response_input_value(value: object) -> object:
model_dump: Final = getattr(value, "model_dump", None)
if callable(model_dump):
return model_dump()
dict_method: Final = getattr(value, "dict", None)
if callable(dict_method):
return dict_method()
return value
def _embedding_input(self, prompt: str, router: "Router | None") -> str:
return truncate_embedding_input(

View file

@ -13,8 +13,6 @@ from pathlib import Path
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, cast
from pydantic import BaseModel
import litellm
from litellm import verbose_logger
from litellm.router_utils.batch_utils import InMemoryFile
@ -194,113 +192,44 @@ def get_str_from_messages(messages: list[AllMessageValues]) -> str:
return text
def get_semantic_cache_prompt_from_messages(messages: object) -> str:
def get_semantic_cache_prompt_from_messages(messages: Sequence[Mapping[str, object]]) -> str:
"""
The text the semantic cache embeds for ``messages``, shared by the Messages API and Chat Completions. Keeps
the text ``get_str_from_messages`` keeps, plus every tool call and tool result, so agent turns that differ
only in their tool exchange embed differently. Call ids are random per session, so each result names the
position of the call it answers instead
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
"""
message_dicts = _dicts(_plain(messages))
positions = _tool_call_positions(_messages_tool_call_ids(message_dicts))
parts = []
for message in message_dicts:
content = message.get("content")
text = content if isinstance(content, str) else _messages_api_blocks_text(content, positions)
if message.get("role") == "tool":
text = _tool_result_text(message.get("tool_call_id"), text, positions)
for tool_call in _dicts(message.get("tool_calls")):
function = tool_call.get("function")
if not isinstance(function, Mapping):
function = {}
text += _tool_call_text(function.get("name"), function.get("arguments"))
parts.append(text + extract_search_results_text(message.get("search_results")))
return "".join(parts)
return "".join(_semantic_cache_message_text(message) for message in messages)
def _messages_tool_call_ids(messages: Sequence[Mapping[str, object]]) -> Iterator[object]:
for message in messages:
for block in _dicts(message.get("content")):
if block.get("type") == "tool_use":
yield block.get("id")
for tool_call in _dicts(message.get("tool_calls")):
yield tool_call.get("id")
def _semantic_cache_message_text(message: Mapping[str, object]) -> str:
return _semantic_cache_content_text(message.get("content")) + extract_search_results_text(
message.get("search_results")
)
def _messages_api_blocks_text(blocks: object, positions: Mapping[str, int]) -> str:
text = ""
for block in _dicts(blocks):
if block.get("type") == "tool_use":
text += _tool_call_text(block.get("name"), block.get("input"))
elif block.get("type") == "tool_result":
content = block.get("content")
output = content if isinstance(content, str) else _messages_api_blocks_text(content, positions)
text += _tool_result_text(block.get("tool_use_id"), output, positions)
elif isinstance(block.get("text"), str):
text += str(block["text"])
return text
def get_semantic_cache_prompt_from_responses_input(responses_input: object) -> str:
"""
The text the semantic cache embeds for a Responses API ``input``: each text part stripped and on its own
line, with ``function_call`` and ``function_call_output`` items encoded like tool calls and results above
"""
items = _plain(responses_input)
call_ids = [item.get("call_id") for item in _dicts(items) if item.get("type") == "function_call"]
return _responses_text(items, _tool_call_positions(call_ids))
def _responses_text(value: object, positions: Mapping[str, int]) -> str:
if isinstance(value, str):
return value.strip()
if isinstance(value, list):
lines = [_responses_text(item, positions) for item in value]
return "\n".join(line for line in lines if line)
if not isinstance(value, Mapping):
def _semantic_cache_content_text(content: object) -> str:
if isinstance(content, str):
return content
if not isinstance(content, list):
return ""
if value.get("type") == "function_call":
return _tool_call_text(value.get("name"), value.get("arguments"))
if value.get("type") == "function_call_output":
output = _responses_text(value.get("output"), positions)
return _tool_result_text(value.get("call_id"), output, positions)
if value.get("content") is not None:
return _responses_text(value.get("content"), positions)
return _responses_text(value.get("text"), positions)
return "".join(_semantic_cache_block_text(block) for block in content)
def _tool_call_text(name: object, arguments: object) -> str:
return json.dumps({"name": name, "arguments": arguments}, separators=(",", ":"), default=str)
def _semantic_cache_block_text(block: object) -> str:
if not isinstance(block, Mapping):
return ""
if block.get("type") != "tool_result":
return _text_field(block)
result: Final = block.get("content")
if isinstance(result, str):
return result
if isinstance(result, list):
return "".join(_text_field(inner) for inner in result)
return ""
def _tool_result_text(call_id: object, output: str, positions: Mapping[str, int]) -> str:
position = positions.get(call_id) if isinstance(call_id, str) else None
return json.dumps({"result_of_call": position, "output": output}, separators=(",", ":"), default=str)
def _tool_call_positions(call_ids: Iterable[object]) -> Mapping[str, int]:
positions = {}
for call_id in call_ids:
if isinstance(call_id, str) and call_id not in positions:
positions[call_id] = len(positions) + 1
return positions
def _plain(value: object) -> object:
"""SDK callers pass pydantic items (``input += response.output``); dump them to the dicts the request carried"""
if isinstance(value, BaseModel):
return value.model_dump()
if isinstance(value, (list, tuple)):
return [_plain(item) for item in value]
if isinstance(value, Mapping):
return {key: _plain(item) for key, item in value.items()}
return value
def _dicts(values: object) -> list[Mapping[str, object]]:
if not isinstance(values, list):
return []
return [value for value in values if isinstance(value, Mapping)]
def _text_field(block: object) -> str:
text: Final = block.get("text") if isinstance(block, Mapping) else None
return text if isinstance(text, str) else ""
def is_non_content_values_set(message: AllMessageValues) -> bool:

View file

@ -23,9 +23,6 @@ IGNORE_FUNCTIONS = [
"_can_object_call_model", # max depth set.
"encode_unserializable_types", # max depth set.
"filter_value_from_dict", # max depth set.
"_responses_text", # walks only the nesting a Responses `input` carries.
"_messages_api_blocks_text", # walks only the nesting a `tool_result` block carries.
"_plain", # walks only the nesting a request `messages` or `input` carries.
"normalize_json_schema_types", # max depth set.
"_extract_fields_recursive", # max depth set.
"_remove_json_schema_refs", # max depth set.,
@ -109,8 +106,13 @@ class RecursiveFunctionFinder(ast.NodeVisitor):
return True
# Case 2: Method call with self (e.g., self.my_func())
if isinstance(call_node.func, ast.Attribute) and isinstance(call_node.func.value, ast.Name):
return call_node.func.value.id == "self" and call_node.func.attr == func_node.name
if isinstance(call_node.func, ast.Attribute) and isinstance(
call_node.func.value, ast.Name
):
return (
call_node.func.value.id == "self"
and call_node.func.attr == func_node.name
)
return False
@ -145,7 +147,9 @@ if __name__ == "__main__":
# this is used in the CI/CD pipeline to prevent recursive functions from being merged
directory_path = "./litellm"
recursive_functions, ignored_recursive_functions = find_recursive_functions_in_directory(directory_path)
recursive_functions, ignored_recursive_functions = (
find_recursive_functions_in_directory(directory_path)
)
print("UNIGNORED RECURSIVE FUNCTIONS: ", recursive_functions)
print("IGNORED RECURSIVE FUNCTIONS: ", ignored_recursive_functions)

View file

@ -0,0 +1,415 @@
import hashlib
import json
import math
import threading
import uuid
from collections.abc import Callable, Iterator, Mapping, Sequence
from dataclasses import dataclass, field
from pathlib import Path
from typing import Final
from urllib.parse import urlsplit
import pytest
import yaml
from integration._support.client import Gateway, eventually, gateway_from_environment
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_VECTOR_SIZE: Final = 16
_CHAT_MODEL: Final = "capped-chat"
_CLAUDE_MODEL: Final = "capped-claude"
_EMBEDDING_MODEL: Final = "capped-embedder"
_COLLECTION: Final = "cache-max-messages"
def _vector(text: str) -> tuple[float, ...]:
digest: Final = hashlib.sha256(text.encode()).digest()
return tuple((byte - 127.5) / 127.5 for byte in digest[:_VECTOR_SIZE])
def _cosine(left: Sequence[float], right: Sequence[float]) -> float:
dot: Final = sum(a * b for a, b in zip(left, right, strict=True))
return dot / (math.sqrt(sum(a * a for a in left)) * math.sqrt(sum(b * b for b in right)))
def _answer() -> str:
return f"answer-{uuid.uuid4().hex[:16]}"
@dataclass(slots=True)
class _Peer:
"""Embeddings, chat, Responses, Anthropic Messages and a Qdrant collection on one owned socket"""
lock: threading.Lock = field(default_factory=threading.Lock)
points: list[Mapping[str, JsonValue]] = field(default_factory=list) # mutable-ok: the Qdrant collection
embedded: list[str] = field(default_factory=list) # mutable-ok: every prompt the proxy embedded
answered: list[str] = field(default_factory=list) # mutable-ok: every completion the provider served
def stored(self) -> int:
with self.lock:
return len(self.points)
def respond(self, request: Request) -> Reply:
path: Final = urlsplit(request.target).path
body: Final = _JSON_OBJECT.validate_json(request.body) if request.body else {}
collection: Final = f"/qdrant/collections/{_COLLECTION}"
if request.method == "GET" and path == f"{collection}/exists":
return self._json({"result": {"exists": False}, "status": "ok"})
if request.method in {"GET", "PUT"} and path in {collection, f"{collection}/index"}:
return self._json({"result": True, "status": "ok"})
if request.method == "PUT" and path == f"{collection}/points":
with self.lock:
self.points.extend(_JSON_OBJECT.validate_python(point) for point in _list(body["points"]))
return self._json({"result": {"status": "completed"}, "status": "ok"})
if request.method == "POST" and path == f"{collection}/points/search":
return self._json({"result": self._search(body), "status": "ok"})
if request.method == "GET" and path == "/v1/models":
return self._json({"object": "list", "data": []})
if request.method == "POST" and path == "/v1/embeddings":
text: Final = str(body["input"])
with self.lock:
self.embedded.append(text)
return self._json(
{
"object": "list",
"model": "text-embedding-3-small",
"data": [{"object": "embedding", "index": 0, "embedding": list(_vector(text))}],
"usage": {"prompt_tokens": 4, "total_tokens": 4},
}
)
if request.method == "POST" and path == "/v1/chat/completions":
return self._json(self._chat_reply(_answer()))
if request.method == "POST" and path == "/v1/responses":
return self._json(self._responses_reply(_answer()))
if request.method == "POST" and path == "/v1/messages":
return self._json(self._messages_reply(_answer()))
raise AssertionError(f"unexpected peer request {request.method} {request.target}")
def _search(self, body: Mapping[str, JsonValue]) -> list[JsonValue]:
query: Final = [float(str(value)) for value in _list(body["vector"])]
key: Final = _JSON_OBJECT.validate_python(
_JSON_OBJECT.validate_python(_list(_JSON_OBJECT.validate_python(body["filter"])["must"])[0])["match"]
)["value"]
with self.lock:
scoped: Final = [
point
for point in self.points
if _JSON_OBJECT.validate_python(point["payload"])["litellm_cache_key"] == key
]
ranked: Final = sorted(
(
{
"id": point["id"],
"score": _cosine(query, [float(str(value)) for value in _list(point["vector"])]),
"payload": point["payload"],
}
for point in scoped
),
key=lambda hit: -float(str(hit["score"])),
)
return list(ranked[:1])
def _chat_reply(self, answer: str) -> Mapping[str, JsonValue]:
with self.lock:
self.answered.append(answer)
return {
"id": f"chatcmpl-{uuid.uuid4().hex}",
"object": "chat.completion",
"created": 1789788253,
"model": "gpt-5.4-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": answer}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12},
}
def _responses_reply(self, answer: str) -> Mapping[str, JsonValue]:
with self.lock:
self.answered.append(answer)
identity: Final = uuid.uuid4().hex
return {
"id": f"resp_{identity}",
"object": "response",
"created_at": 1789788253,
"status": "completed",
"model": "gpt-5.4-mini",
"output": [
{
"id": f"msg_{identity}",
"type": "message",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": answer, "annotations": []}],
}
],
"usage": {"input_tokens": 10, "output_tokens": 2, "total_tokens": 12},
}
def _messages_reply(self, answer: str) -> Mapping[str, JsonValue]:
with self.lock:
self.answered.append(answer)
return {
"id": f"msg_{uuid.uuid4().hex}",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-5-5",
"content": [{"type": "text", "text": answer}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 10, "output_tokens": 2},
}
@staticmethod
def _json(value: Mapping[str, JsonValue]) -> Reply:
return Reply(body=json.dumps(value).encode())
def _list(value: JsonValue) -> list[JsonValue]:
assert isinstance(value, list), value
return value
@dataclass(frozen=True, slots=True)
class _Proxy:
gateway: Gateway
peer: _Peer
def reply(self, route: str, conversation: list[JsonValue]) -> str:
"""The answer text for one request, as the client sees it"""
model: Final = _CLAUDE_MODEL if route == "/v1/messages" else _CHAT_MODEL
body: Final[dict[str, JsonValue]] = (
{"model": model, "input": conversation}
if route == "/v1/responses"
else {"model": model, "max_tokens": 16, "messages": conversation}
)
response: Final = self.gateway.request("POST", route, body)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
if route == "/v1/messages":
return str(_JSON_OBJECT.validate_python(_list(payload["content"])[0])["text"])
if route == "/v1/responses":
message: Final = _JSON_OBJECT.validate_python(_list(payload["output"])[0])
return str(_JSON_OBJECT.validate_python(_list(message["content"])[0])["text"])
choice: Final = _JSON_OBJECT.validate_python(_list(payload["choices"])[0])
return str(_JSON_OBJECT.validate_python(choice["message"])["content"])
def miss(self, route: str, conversation: list[JsonValue]) -> str:
answered_before: Final = len(self.peer.answered)
answer: Final = self.reply(route, conversation)
assert len(self.peer.answered) == answered_before + 1, f"{answer} was served from the cache"
return answer
def hit(self, route: str, conversation: list[JsonValue]) -> str:
"""Repeats the request until the cache serves it, since the store after a miss is asynchronous"""
def attempt() -> tuple[str, bool]:
answered_before: Final = len(self.peer.answered)
answer: Final = self.reply(route, conversation)
return answer, len(self.peer.answered) == answered_before
return eventually(attempt, lambda outcome: outcome[1])[0]
def _proxy(tmp_path_factory: pytest.TempPathFactory, cache_params: Mapping[str, JsonValue]) -> Iterator[_Proxy]:
peer: Final = _Peer()
directory: Final = tmp_path_factory.mktemp("cache-max-messages")
with gateway_from_environment() as gateway, wire_server(peer.respond) as wire:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["model_list"] = [
{
"model_name": _CHAT_MODEL,
"litellm_params": {"model": "openai/gpt-5.4-mini", "api_base": f"{wire.url}/v1", "api_key": "k"},
},
{
"model_name": _CLAUDE_MODEL,
"litellm_params": {"model": "anthropic/claude-sonnet-5-5", "api_base": wire.url, "api_key": "k"},
},
{
"model_name": _EMBEDDING_MODEL,
"litellm_params": {
"model": "openai/text-embedding-3-small",
"api_base": f"{wire.url}/v1",
"api_key": "k",
},
},
]
config["litellm_settings"]["cache_params"] = {
**cache_params,
**({"qdrant_api_base": f"{wire.url}/qdrant"} if cache_params["type"] == "qdrant-semantic" else {}),
}
path: Final = directory / "cache_max_messages.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, directory, {}, config=path) as candidate:
yield _Proxy(candidate, peer)
@pytest.fixture(scope="module")
def qdrant_proxy(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Proxy]:
"""Qdrant semantic cache with the default max_messages of 4"""
yield from _proxy(
tmp_path_factory,
{
"type": "qdrant-semantic",
"qdrant_collection_name": _COLLECTION,
"qdrant_semantic_cache_embedding_model": _EMBEDDING_MODEL,
"qdrant_semantic_cache_vector_size": _VECTOR_SIZE,
"qdrant_quantization_config": "binary",
"similarity_threshold": 0.99,
},
)
@pytest.fixture(scope="module")
def exact_proxy(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Proxy]:
"""Exact cache with max_messages lowered to 2 in cache_params"""
yield from _proxy(tmp_path_factory, {"type": "local", "max_messages": 2})
def _claude_code_turns(task: str) -> tuple[list[JsonValue], list[JsonValue], list[JsonValue]]:
"""Turns 1 to 3 of a Claude Code session on /v1/messages: 1, 3 and 5 messages"""
first: Final[list[JsonValue]] = [{"role": "user", "content": task}]
second: 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"}],
},
]
third: Final[list[JsonValue]] = [
*second,
{
"role": "assistant",
"content": [{"type": "tool_use", "id": "toolu_2", "name": "Read", "input": {"file_path": "calc.py"}}],
},
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "toolu_2",
"content": [{"type": "text", "text": "def add(a, b): return a - b"}],
}
],
},
]
return first, second, third
def _agent_turns(task: str) -> tuple[list[JsonValue], list[JsonValue], list[JsonValue]]:
"""Turns 1 to 3 of an OpenAI tool loop on /v1/chat/completions: 2, 4 and 6 messages"""
def call(call_id: str, path: str) -> list[JsonValue]:
return [
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": call_id,
"type": "function",
"function": {"name": "write_file", "arguments": json.dumps({"path": path})},
}
],
},
{"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},
]
second: Final[list[JsonValue]] = [*first, *call("call_1", "a.yaml")]
third: Final[list[JsonValue]] = [*second, *call("call_2", "b.yaml")]
return first, second, third
def _responses_turns(task: str) -> tuple[list[JsonValue], list[JsonValue], list[JsonValue]]:
"""Turns 1 to 3 of an agent on /v1/responses: 1, 3 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}]
second: Final[list[JsonValue]] = [*first, *call("call_1", "a.yaml")]
third: Final[list[JsonValue]] = [*second, *call("call_2", "b.yaml")]
return first, second, third
def _assert_turns_are_cached_up_to_four_messages(
proxy: _Proxy, route: str, turns: Callable[[str], tuple[list[JsonValue], list[JsonValue], list[JsonValue]]]
) -> None:
first, second, third = turns(f"update the config {uuid.uuid4().hex}")
stored_before: Final = proxy.peer.stored()
short_answers: Final = (proxy.miss(route, first), proxy.miss(route, second))
repeated: Final = (proxy.hit(route, first), proxy.hit(route, second))
long_answers: Final = (proxy.miss(route, third), proxy.miss(route, third))
proxy.hit(route, first)
assert repeated == short_answers, f"a repeated short turn got another turn's answer: {short_answers} {repeated}"
assert long_answers[0] != long_answers[1], "a turn past max_messages was served from the cache"
assert all("b.yaml" not in text and "def add" not in text for text in proxy.peer.embedded), (
"a turn past max_messages was embedded"
)
assert proxy.peer.stored() == stored_before + 2, "a turn past max_messages was written to the cache"
@pytest.mark.parametrize(
("route", "turns"),
[
pytest.param("/v1/messages", _claude_code_turns, id="messages"),
pytest.param("/v1/chat/completions", _agent_turns, id="chat-completions"),
],
)
def test_agent_turns_are_cached_up_to_four_messages_and_bypass_the_cache_past_it(
qdrant_proxy: _Proxy,
route: str,
turns: Callable[[str], tuple[list[JsonValue], list[JsonValue], list[JsonValue]]],
) -> None:
_assert_turns_are_cached_up_to_four_messages(qdrant_proxy, route, turns)
def test_tool_result_text_tells_a_tool_turn_apart_from_the_turn_before_it(qdrant_proxy: _Proxy) -> None:
task: Final = f"list the files {uuid.uuid4().hex}"
first, second, _ = _claude_code_turns(task)
qdrant_proxy.miss("/v1/messages", first)
qdrant_proxy.miss("/v1/messages", second)
assert task in qdrant_proxy.peer.embedded[-1], qdrant_proxy.peer.embedded[-1]
assert "calc.py test_calc.py" in qdrant_proxy.peer.embedded[-1], qdrant_proxy.peer.embedded[-1]
@pytest.mark.parametrize(
("route", "turns"),
[
pytest.param("/v1/messages", _claude_code_turns, id="messages"),
pytest.param("/v1/chat/completions", _agent_turns, id="chat-completions"),
pytest.param("/v1/responses", _responses_turns, id="responses"),
],
)
def test_exact_cache_honours_max_messages_from_cache_params(
exact_proxy: _Proxy,
route: str,
turns: Callable[[str], tuple[list[JsonValue], list[JsonValue], list[JsonValue]]],
) -> None:
first, second, _ = turns(f"what is {uuid.uuid4().hex}")
answer: Final = exact_proxy.miss(route, first)
repeated: Final = exact_proxy.hit(route, first)
long_answers: Final = (exact_proxy.miss(route, second), exact_proxy.miss(route, second))
assert repeated == answer
assert long_answers[0] != long_answers[1], "a turn past max_messages was served from a cache capped at 2"

View file

@ -1,544 +0,0 @@
import asyncio
import hashlib
import json
import math
import threading
import uuid
from collections.abc import Callable, Iterator, Mapping, Sequence
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, field
from pathlib import Path
from typing import Final
from urllib.parse import urlsplit
import anthropic
import httpx
import openai
import pytest
import yaml
from integration._support.client import Gateway, eventually, gateway_from_environment
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_VECTOR_SIZE: Final = 16
_CHAT_MODEL: Final = "semantic-chat"
_CLAUDE_MODEL: Final = "semantic-claude"
_EMBEDDING_MODEL: Final = "semantic-embedder"
_COLLECTION: Final = "semantic-tool-turns"
def _vector(text: str) -> tuple[float, ...]:
digest: Final = hashlib.sha256(text.encode()).digest()
return tuple((byte - 127.5) / 127.5 for byte in digest[:_VECTOR_SIZE])
def _cosine(left: Sequence[float], right: Sequence[float]) -> float:
dot: Final = sum(a * b for a, b in zip(left, right, strict=True))
return dot / (math.sqrt(sum(a * a for a in left)) * math.sqrt(sum(b * b for b in right)))
def _answer(body: Mapping[str, JsonValue]) -> str:
return "answer-" + hashlib.sha256(json.dumps(body["messages"], sort_keys=True).encode()).hexdigest()[:16]
@dataclass(slots=True)
class _Peer:
"""Embeddings, chat, Anthropic Messages and a Qdrant collection on one owned socket"""
lock: threading.Lock = field(default_factory=threading.Lock)
points: list[Mapping[str, JsonValue]] = field(default_factory=list) # mutable-ok: the Qdrant collection
embedded: list[str] = field(default_factory=list) # mutable-ok: every prompt the proxy embedded
answered: list[str] = field(default_factory=list) # mutable-ok: every completion the provider served
qdrant_down: threading.Event = field(default_factory=threading.Event)
refused_writes: list[str] = field(default_factory=list) # mutable-ok: upserts refused while Qdrant is down
def stored(self) -> int:
with self.lock:
return len(self.points)
def respond(self, request: Request) -> Reply:
path: Final = urlsplit(request.target).path
body: Final = _JSON_OBJECT.validate_json(request.body) if request.body else {}
collection: Final = f"/qdrant/collections/{_COLLECTION}"
if path.startswith("/qdrant/") and self.qdrant_down.is_set():
if path == f"{collection}/points":
with self.lock:
self.refused_writes.append(path)
return Reply(status=503, body=b'{"status":{"error":"qdrant unavailable"}}')
if request.method == "GET" and path == f"{collection}/exists":
return self._json({"result": {"exists": False}, "status": "ok"})
if request.method in {"GET", "PUT"} and path in {collection, f"{collection}/index"}:
return self._json({"result": True, "status": "ok"})
if request.method == "PUT" and path == f"{collection}/points":
with self.lock:
self.points.extend(_JSON_OBJECT.validate_python(point) for point in body["points"])
return self._json({"result": {"status": "completed"}, "status": "ok"})
if request.method == "POST" and path == f"{collection}/points/search":
return self._json({"result": self._search(body), "status": "ok"})
if request.method == "GET" and path == "/v1/models":
return self._json({"object": "list", "data": []})
if request.method == "POST" and path == "/v1/embeddings":
text: Final = str(body["input"])
with self.lock:
self.embedded.append(text)
return self._json(
{
"object": "list",
"model": "text-embedding-3-small",
"data": [{"object": "embedding", "index": 0, "embedding": list(_vector(text))}],
"usage": {"prompt_tokens": 4, "total_tokens": 4},
}
)
if request.method == "POST" and path == "/v1/chat/completions":
return self._json(self._chat_reply(_answer(body)))
if request.method == "POST" and path == "/v1/messages":
return self._json(self._messages_reply(_answer(body)))
raise AssertionError(f"unexpected peer request {request.method} {request.target}")
def _search(self, body: Mapping[str, JsonValue]) -> list[JsonValue]:
query: Final = [float(str(value)) for value in _list(body["vector"])]
key: Final = _JSON_OBJECT.validate_python(
_JSON_OBJECT.validate_python(_list(_JSON_OBJECT.validate_python(body["filter"])["must"])[0])["match"]
)["value"]
with self.lock:
scoped: Final = [
point
for point in self.points
if _JSON_OBJECT.validate_python(point["payload"])["litellm_cache_key"] == key
]
ranked: Final = sorted(
(
{
"id": point["id"],
"score": _cosine(query, [float(str(value)) for value in _list(point["vector"])]),
"payload": point["payload"],
}
for point in scoped
),
key=lambda hit: -float(str(hit["score"])),
)
return list(ranked[:1])
def _chat_reply(self, answer: str) -> Mapping[str, JsonValue]:
with self.lock:
self.answered.append(answer)
return {
"id": f"chatcmpl-{uuid.uuid4().hex}",
"object": "chat.completion",
"created": 1789788253,
"model": "gpt-5.4-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": answer}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12},
}
def _messages_reply(self, answer: str) -> Mapping[str, JsonValue]:
with self.lock:
self.answered.append(answer)
return {
"id": f"msg_{uuid.uuid4().hex}",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-5-5",
"content": [{"type": "text", "text": answer}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 10, "output_tokens": 2},
}
@staticmethod
def _json(value: Mapping[str, JsonValue]) -> Reply:
return Reply(body=json.dumps(value).encode())
def _list(value: JsonValue) -> list[JsonValue]:
assert isinstance(value, list), value
return value
@dataclass(frozen=True, slots=True)
class _SemanticProxy:
gateway: Gateway
peer: _Peer
@pytest.fixture(scope="module")
def semantic_proxy(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_SemanticProxy]:
peer: Final = _Peer()
directory: Final = tmp_path_factory.mktemp("semantic-tool-turns")
with gateway_from_environment() as gateway, wire_server(peer.respond) as wire:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["model_list"] = [
{
"model_name": _CHAT_MODEL,
"litellm_params": {"model": "openai/gpt-5.4-mini", "api_base": f"{wire.url}/v1", "api_key": "k"},
},
{
"model_name": _CLAUDE_MODEL,
"litellm_params": {"model": "anthropic/claude-sonnet-5-5", "api_base": wire.url, "api_key": "k"},
},
{
"model_name": _EMBEDDING_MODEL,
"litellm_params": {
"model": "openai/text-embedding-3-small",
"api_base": f"{wire.url}/v1",
"api_key": "k",
},
},
]
config["litellm_settings"]["cache_params"] = {
"type": "qdrant-semantic",
"qdrant_api_base": f"{wire.url}/qdrant",
"qdrant_collection_name": _COLLECTION,
"qdrant_semantic_cache_embedding_model": _EMBEDDING_MODEL,
"qdrant_semantic_cache_vector_size": _VECTOR_SIZE,
"qdrant_quantization_config": "binary",
"similarity_threshold": 0.99,
}
path: Final = directory / "semantic_tool_turns.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, directory, {}, config=path) as candidate:
yield _SemanticProxy(candidate, peer)
def _send_turns(proxy: _SemanticProxy, route: str, model: str, turns: Sequence[list[JsonValue]]) -> list[str]:
def reply_text(turn: list[JsonValue]) -> str:
stored_before: Final = proxy.peer.stored()
answered_before: Final = len(proxy.peer.answered)
body: Final[dict[str, JsonValue]] = {"model": model, "max_tokens": 16, "messages": turn}
response: Final = proxy.gateway.request("POST", route, body)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
if len(proxy.peer.answered) > answered_before:
eventually(proxy.peer.stored, lambda count: count > stored_before)
if route == "/v1/messages":
return str(_JSON_OBJECT.validate_python(_list(payload["content"])[0])["text"])
choice: Final = _JSON_OBJECT.validate_python(_list(payload["choices"])[0])
return str(_JSON_OBJECT.validate_python(choice["message"])["content"])
return [reply_text(turn) for turn in turns]
def test_claude_code_tool_turns_on_messages_are_not_served_the_first_turn_answer(
semantic_proxy: _SemanticProxy,
) -> None:
task: Final[JsonValue] = {"role": "user", "content": f"fix the failing test {uuid.uuid4().hex}"}
list_files: Final[list[JsonValue]] = [
task,
{
"role": "assistant",
"content": [{"type": "tool_use", "id": "toolu_1", "name": "Bash", "input": {"command": "ls"}}],
},
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "calc.py test_calc.py"}],
},
]
read_file: Final[list[JsonValue]] = [
*list_files,
{
"role": "assistant",
"content": [{"type": "tool_use", "id": "toolu_2", "name": "Read", "input": {"file_path": "calc.py"}}],
},
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "toolu_2",
"content": [{"type": "text", "text": "def add(a, b): return a - b"}],
}
],
},
]
answers: Final = _send_turns(
semantic_proxy, "/v1/messages", _CLAUDE_MODEL, ([task], list_files, read_file, list_files)
)
assert len(set(answers[:3])) == 3, f"a later tool turn replayed an earlier cached answer: {answers}"
assert answers[:3] == semantic_proxy.peer.answered[-3:], semantic_proxy.peer.answered
assert answers[3] == answers[1], f"a repeated tool turn missed the cache: {answers}"
def test_openai_agent_loop_tool_calls_on_chat_completions_are_not_served_a_cached_answer(
semantic_proxy: _SemanticProxy,
) -> None:
task: Final[JsonValue] = {"role": "user", "content": f"update both config files {uuid.uuid4().hex}"}
def wrote(path: str) -> list[JsonValue]:
return [
task,
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "write_file", "arguments": json.dumps({"path": path})},
}
],
},
{"role": "tool", "tool_call_id": "call_1", "content": "ok"},
]
answers: Final = _send_turns(
semantic_proxy, "/v1/chat/completions", _CHAT_MODEL, (wrote("a.yaml"), wrote("b.yaml"), wrote("a.yaml"))
)
assert len(set(answers[:2])) == 2, f"a different tool call replayed an earlier cached answer: {answers}"
assert answers[:2] == semantic_proxy.peer.answered[-2:], semantic_proxy.peer.answered
assert answers[2] == answers[0], f"a repeated tool call missed the cache: {answers}"
def _tool_turn(task: str, calls: Sequence[tuple[str, str]], results: Sequence[tuple[str, str]]) -> list[JsonValue]:
return [
{"role": "user", "content": task},
{
"role": "assistant",
"content": None,
"tool_calls": [
{"id": call_id, "type": "function", "function": {"name": "read_file", "arguments": arguments}}
for call_id, arguments in calls
],
},
*({"role": "tool", "tool_call_id": call_id, "content": output} for call_id, output in results),
]
def _claude_turn(task: str, calls: Sequence[tuple[str, str]], results: Sequence[tuple[str, str]]) -> list[JsonValue]:
return [
{"role": "user", "content": task},
{
"role": "assistant",
"content": [
{"type": "tool_use", "id": call_id, "name": "Read", "input": {"file_path": file_path}}
for call_id, file_path in calls
],
},
{
"role": "user",
"content": [
{"type": "tool_result", "tool_use_id": call_id, "content": output} for call_id, output in results
],
},
]
_TWO_READS: Final = (("c1", '{"path":"a.py"}'), ("c2", '{"path":"b.py"}'))
_TWO_CLAUDE_READS: Final = (("toolu_1", "a.py"), ("toolu_2", "b.py"))
@pytest.mark.parametrize(
("route", "model", "build", "calls", "first", "second"),
[
pytest.param(
"/v1/chat/completions",
_CHAT_MODEL,
_tool_turn,
_TWO_READS,
(("c1", "x = 1"), ("c2", "y = 2")),
(("c2", "x = 1"), ("c1", "y = 2")),
id="chat-parallel-results-answer-the-other-call",
),
pytest.param(
"/v1/messages",
_CLAUDE_MODEL,
_claude_turn,
_TWO_CLAUDE_READS,
(("toolu_1", "x = 1"), ("toolu_2", "y = 2")),
(("toolu_2", "x = 1"), ("toolu_1", "y = 2")),
id="messages-parallel-results-answer-the-other-call",
),
pytest.param(
"/v1/chat/completions",
_CHAT_MODEL,
_tool_turn,
_TWO_READS[:1],
(("c1", "x = 1"),),
(("c9", "x = 1"),),
id="chat-result-for-an-unknown-call",
),
pytest.param(
"/v1/chat/completions",
_CHAT_MODEL,
_tool_turn,
_TWO_READS,
(("c1", '"},{"result_of_call":2,"output":"y = 2'),),
(("c1", ""), ("c2", "y = 2")),
id="chat-output-shaped-like-a-second-result",
),
pytest.param(
"/v1/chat/completions",
_CHAT_MODEL,
_tool_turn,
_TWO_READS[:1],
(("c1", ""),),
(("c1", "x" * 5000),),
id="chat-empty-and-5kb-output",
),
],
)
def test_tool_turns_that_differ_only_in_their_results_do_not_share_a_cached_answer(
semantic_proxy: _SemanticProxy,
route: str,
model: str,
build: Callable[[str, Sequence[tuple[str, str]], Sequence[tuple[str, str]]], list[JsonValue]],
calls: Sequence[tuple[str, str]],
first: Sequence[tuple[str, str]],
second: Sequence[tuple[str, str]],
) -> None:
task: Final = f"read both files {uuid.uuid4().hex}"
answers: Final = _send_turns(
semantic_proxy, route, model, (build(task, calls, first), build(task, calls, second), build(task, calls, first))
)
assert answers[0] != answers[1], f"a different tool result replayed an earlier cached answer: {answers}"
assert answers[:2] == semantic_proxy.peer.answered[-2:], semantic_proxy.peer.answered
assert answers[2] == answers[0], f"a repeated tool turn missed the cache: {answers}"
def _openai_sdk(proxy: _SemanticProxy, turn: list[JsonValue]) -> str:
sdk: Final = openai.OpenAI(
base_url=str(proxy.gateway.client.base_url) + "/v1",
api_key=proxy.gateway.key,
http_client=httpx.Client(trust_env=False, timeout=15),
)
completion: Final = sdk.chat.completions.create(model=_CHAT_MODEL, messages=turn) # pyright: ignore[reportArgumentType, reportCallIssue] # JSON turns are validated by the proxy, not the SDK types
return str(completion.choices[0].message.content)
def _async_openai_sdk(proxy: _SemanticProxy, turn: list[JsonValue]) -> str:
async def create() -> str:
async with httpx.AsyncClient(trust_env=False, timeout=15) as http_client:
sdk: Final = openai.AsyncOpenAI(
base_url=str(proxy.gateway.client.base_url) + "/v1", api_key=proxy.gateway.key, http_client=http_client
)
completion: Final = await sdk.chat.completions.create(model=_CHAT_MODEL, messages=turn) # pyright: ignore[reportArgumentType, reportCallIssue] # JSON turns are validated by the proxy, not the SDK types
return str(completion.choices[0].message.content)
return asyncio.run(create())
def _anthropic_text(message: anthropic.types.Message) -> str:
block: Final = message.content[0]
assert isinstance(block, anthropic.types.TextBlock), message
return block.text
def _anthropic_sdk(proxy: _SemanticProxy, turn: list[JsonValue]) -> str:
sdk: Final = anthropic.Anthropic(
base_url=str(proxy.gateway.client.base_url),
api_key=proxy.gateway.key,
http_client=httpx.Client(trust_env=False, timeout=15),
)
return _anthropic_text(sdk.messages.create(model=_CLAUDE_MODEL, max_tokens=16, messages=turn)) # pyright: ignore[reportArgumentType] # JSON turns are validated by the proxy, not the SDK types
def _async_anthropic_sdk(proxy: _SemanticProxy, turn: list[JsonValue]) -> str:
async def create() -> str:
async with httpx.AsyncClient(trust_env=False, timeout=15) as http_client:
sdk: Final = anthropic.AsyncAnthropic(
base_url=str(proxy.gateway.client.base_url), api_key=proxy.gateway.key, http_client=http_client
)
message: Final = await sdk.messages.create(model=_CLAUDE_MODEL, max_tokens=16, messages=turn) # pyright: ignore[reportArgumentType] # JSON turns are validated by the proxy, not the SDK types
return _anthropic_text(message)
return asyncio.run(create())
@pytest.mark.parametrize(
("send", "build"),
[
pytest.param(_openai_sdk, _tool_turn, id="openai-sdk"),
pytest.param(_async_openai_sdk, _tool_turn, id="async-openai-sdk"),
pytest.param(_anthropic_sdk, _claude_turn, id="anthropic-sdk"),
pytest.param(_async_anthropic_sdk, _claude_turn, id="async-anthropic-sdk"),
],
)
def test_sdk_clients_get_a_fresh_answer_for_each_tool_turn_and_a_hit_for_a_repeat(
semantic_proxy: _SemanticProxy,
send: Callable[[_SemanticProxy, list[JsonValue]], str],
build: Callable[[str, Sequence[tuple[str, str]], Sequence[tuple[str, str]]], list[JsonValue]],
) -> None:
task: Final = f"read the file {uuid.uuid4().hex}"
calls: Final = _TWO_READS[:1] if build is _tool_turn else _TWO_CLAUDE_READS[:1]
call_id: Final = calls[0][0]
alpha: Final = build(task, calls, ((call_id, "x = 1"),))
beta: Final = build(task, calls, ((call_id, "PermissionError"),))
def answer(turn: list[JsonValue]) -> str:
stored_before: Final = semantic_proxy.peer.stored()
answered_before: Final = len(semantic_proxy.peer.answered)
text: Final = send(semantic_proxy, turn)
if len(semantic_proxy.peer.answered) > answered_before:
eventually(semantic_proxy.peer.stored, lambda count: count > stored_before)
return text
answers: Final = [answer(turn) for turn in (alpha, beta, alpha)]
assert answers[0] != answers[1], f"a different tool result replayed an earlier cached answer: {answers}"
assert answers[:2] == semantic_proxy.peer.answered[-2:], semantic_proxy.peer.answered
assert answers[2] == answers[0], f"a repeated tool turn missed the cache: {answers}"
def test_concurrent_agent_loops_each_get_their_own_answer_and_their_own_hit(semantic_proxy: _SemanticProxy) -> None:
task: Final = f"read the file {uuid.uuid4().hex}"
turns: Final = tuple(_tool_turn(task, _TWO_READS[:1], (("c1", f"line {index}"),)) for index in range(8))
def send(turn: list[JsonValue]) -> str:
response: Final = semantic_proxy.gateway.request(
"POST", "/v1/chat/completions", {"model": _CHAT_MODEL, "messages": turn}
)
assert response.status_code == 200, response.text
choice: Final = _JSON_OBJECT.validate_python(_list(_JSON_OBJECT.validate_json(response.content)["choices"])[0])
return str(_JSON_OBJECT.validate_python(choice["message"])["content"])
stored_before: Final = semantic_proxy.peer.stored()
with ThreadPoolExecutor(max_workers=len(turns)) as pool:
first: Final = tuple(pool.map(send, turns))
eventually(semantic_proxy.peer.stored, lambda count: count >= stored_before + len(turns))
answered_before: Final = len(semantic_proxy.peer.answered)
with ThreadPoolExecutor(max_workers=len(turns)) as pool:
second: Final = tuple(pool.map(send, turns))
assert len(set(first)) == len(turns), f"concurrent tool turns shared a cached answer: {first}"
assert second == first, f"a repeated tool turn got another turn's answer: {first} {second}"
assert len(semantic_proxy.peer.answered) == answered_before, "a repeated tool turn missed the cache"
def test_tool_turns_are_answered_while_qdrant_is_down_and_cached_once_it_recovers(
semantic_proxy: _SemanticProxy,
) -> None:
task: Final = f"read the file {uuid.uuid4().hex}"
turn: Final = _tool_turn(task, _TWO_READS[:1], (("c1", "x = 1"),))
body: Final[dict[str, JsonValue]] = {"model": _CHAT_MODEL, "messages": turn}
def send() -> httpx.Response:
return semantic_proxy.gateway.request("POST", "/v1/chat/completions", body)
refused_before: Final = len(semantic_proxy.peer.refused_writes)
answered_before: Final = len(semantic_proxy.peer.answered)
semantic_proxy.peer.qdrant_down.set()
try:
outage: Final = (send(), send())
eventually(lambda: len(semantic_proxy.peer.refused_writes), lambda count: count >= refused_before + 2)
finally:
semantic_proxy.peer.qdrant_down.clear()
answered_during_outage: Final = len(semantic_proxy.peer.answered) - answered_before
stored_before: Final = semantic_proxy.peer.stored()
recovered: Final = send()
eventually(semantic_proxy.peer.stored, lambda count: count > stored_before)
answered_after_recovery: Final = len(semantic_proxy.peer.answered)
repeat: Final = send()
assert [response.status_code for response in outage] == [200, 200], [response.text for response in outage]
assert answered_during_outage == 2, "an outage request was not sent to the provider"
assert recovered.status_code == 200, recovered.text
assert answered_after_recovery - answered_before == 3, "the first request after recovery hit a stale entry"
assert repeat.status_code == 200, repeat.text
assert repeat.json()["choices"] == recovered.json()["choices"]
assert len(semantic_proxy.peer.answered) == answered_after_recovery, "the repeat after recovery missed the cache"

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

@ -567,122 +567,27 @@ def test_redis_semantic_cache_prompt_extraction_prefers_messages():
assert prompt == "message prompt"
def test_redis_semantic_cache_prompt_extraction_keeps_tool_turns_distinct():
def test_redis_semantic_cache_prompt_extraction_handles_model_objects():
from litellm.caching.redis_semantic_cache import RedisSemanticCache
def turn(command: str) -> list[dict[str, object]]:
return [
{"role": "user", "content": "fix the failing test"},
{
"role": "assistant",
"content": [{"type": "tool_use", "id": "t1", "name": "Bash", "input": {"cmd": command}}],
},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "t1", "content": "ok"}]},
]
class ModelDumpInput:
def model_dump(self):
return {"content": [{"text": "model dump prompt"}]}
assert RedisSemanticCache._get_prompt_from_kwargs(messages=turn("ls")) == (
'fix the failing test{"name":"Bash","arguments":{"cmd":"ls"}}{"result_of_call":1,"output":"ok"}'
)
assert RedisSemanticCache._get_prompt_from_kwargs(messages=turn("pwd")) == (
'fix the failing test{"name":"Bash","arguments":{"cmd":"pwd"}}{"result_of_call":1,"output":"ok"}'
)
def test_redis_semantic_cache_prompt_extraction_keeps_responses_function_calls():
from litellm.caching.redis_semantic_cache import RedisSemanticCache
class DictInput:
def dict(self):
return {"content": [{"output_text": "dict prompt"}]}
prompt = RedisSemanticCache._get_prompt_from_kwargs(
input=[
{"role": "user", "content": "update the config"},
{"type": "function_call", "call_id": "c1", "name": "write_file", "arguments": '{"path":"a.yaml"}'},
{"type": "function_call_output", "call_id": "c1", "output": "ok"},
ModelDumpInput(),
DictInput(),
{"content": [{"input_text": "inline prompt"}]},
{"content": [{"type": "input_image", "image_url": "https://example.com"}]},
]
)
assert prompt == (
'update the config\n{"name":"write_file","arguments":"{\\"path\\":\\"a.yaml\\"}"}\n{"result_of_call":1,"output":"ok"}'
)
def test_redis_semantic_cache_prompt_extraction_keeps_structured_function_call_outputs():
from litellm.caching.redis_semantic_cache import RedisSemanticCache
def prompt_for(output_text: str) -> str | None:
return RedisSemanticCache._get_prompt_from_kwargs(
input=[
{"type": "message", "role": "user", "content": [{"type": "input_text", "text": "write hello"}]},
{"type": "function_call", "call_id": "c1", "name": "write_file", "arguments": '{"path":"a.txt"}'},
{
"type": "function_call_output",
"call_id": "c1",
"output": [{"type": "input_text", "text": output_text}],
},
]
)
expected_call = '{"name":"write_file","arguments":"{\\"path\\":\\"a.txt\\"}"}'
assert (
prompt_for("wrote 5 bytes") == f'write hello\n{expected_call}\n{{"result_of_call":1,"output":"wrote 5 bytes"}}'
)
assert prompt_for("wrote 5 bytes") != prompt_for("PermissionError")
def test_redis_semantic_cache_prompt_extraction_joins_multi_part_function_call_output_lines():
from litellm.caching.redis_semantic_cache import RedisSemanticCache
prompt = RedisSemanticCache._get_prompt_from_kwargs(
input=[
{"type": "function_call", "call_id": "c1", "name": "run", "arguments": "{}"},
{
"type": "function_call_output",
"call_id": "c1",
"output": [{"type": "input_text", "text": " line one "}, {"type": "input_text", "text": "line two"}],
},
]
)
assert prompt == '{"name":"run","arguments":"{}"}\n{"result_of_call":1,"output":"line one\\nline two"}'
def test_redis_semantic_cache_prompt_extraction_tells_apart_parallel_outputs_answering_different_calls():
from litellm.caching.redis_semantic_cache import RedisSemanticCache
def prompt_for(first_output_call_id: str, second_output_call_id: str) -> str | None:
return RedisSemanticCache._get_prompt_from_kwargs(
input=[
{"type": "function_call", "call_id": "c1", "name": "read", "arguments": "a"},
{"type": "function_call", "call_id": "c2", "name": "read", "arguments": "b"},
{"type": "function_call_output", "call_id": first_output_call_id, "output": "empty"},
{"type": "function_call_output", "call_id": second_output_call_id, "output": "secret"},
]
)
assert prompt_for("c2", "c1") == (
'{"name":"read","arguments":"a"}\n{"name":"read","arguments":"b"}\n'
'{"result_of_call":2,"output":"empty"}\n{"result_of_call":1,"output":"secret"}'
)
assert prompt_for("c1", "c2") != prompt_for("c2", "c1")
def test_redis_semantic_cache_prompt_extraction_dumps_sdk_response_items_appended_to_input():
from openai.types.responses import ResponseFunctionToolCall
from litellm.caching.redis_semantic_cache import RedisSemanticCache
prompt = RedisSemanticCache._get_prompt_from_kwargs(
input=[
{"role": "user", "content": "write hello"},
ResponseFunctionToolCall(
type="function_call", call_id="c1", name="write_file", arguments='{"path":"a.txt"}'
),
{"type": "function_call_output", "call_id": "c1", "output": "ok"},
]
)
assert (
prompt
== 'write hello\n{"name":"write_file","arguments":"{\\"path\\":\\"a.txt\\"}"}\n{"result_of_call":1,"output":"ok"}'
)
assert prompt == "model dump prompt\ndict prompt\ninline prompt"
def test_redis_semantic_cache_prompt_extraction_returns_none_without_text():
@ -697,6 +602,37 @@ def test_redis_semantic_cache_prompt_extraction_returns_none_without_text():
)
def test_redis_semantic_cache_prompt_extraction_skips_blank_dict_text_keys():
from litellm.caching.redis_semantic_cache import RedisSemanticCache
prompt = RedisSemanticCache._get_prompt_from_kwargs(input={"text": " ", "input_text": "fallback prompt"})
assert prompt == "fallback prompt"
def test_redis_semantic_cache_prompt_extraction_skips_blank_object_text_keys():
from litellm.caching.redis_semantic_cache import RedisSemanticCache
class ResponseInput:
text = " "
input_text = "fallback prompt"
prompt = RedisSemanticCache._get_prompt_from_kwargs(input=ResponseInput())
assert prompt == "fallback prompt"
def test_redis_semantic_cache_prompt_extraction_handles_object_content():
from litellm.caching.redis_semantic_cache import RedisSemanticCache
class ResponseInput:
content = [{"text": "object content prompt"}]
prompt = RedisSemanticCache._get_prompt_from_kwargs(input=ResponseInput())
assert prompt == "object content prompt"
def test_redis_semantic_cache_set_cache_skips_blank_responses_input():
from litellm.caching.redis_semantic_cache import RedisSemanticCache
@ -1454,3 +1390,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

@ -30,7 +30,6 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
system_messages_first,
update_messages_with_model_file_ids,
)
from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message
_ARTIFACT_FIELD_PATTERN: Final = r'^(?!__.*__$)[^\p{Cc}\p{Cf}\p{Zl}\p{Zp}"\\./[\]]{1,200}$'
_ARTIFACT_DATA_ID_PATTERN: Final = r"^(?!\.\.?(?:\/|$))[A-Za-z0-9_\-.~:@+]{1,200}$"
@ -2176,28 +2175,23 @@ class TestMergeConsecutiveSystemMessages:
assert merged == [{"role": "system"}, {"role": "user", "content": "Hi"}]
_TASK: Final = {"role": "user", "content": "fix the failing test"}
_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(
[
_TASK,
{
"role": "assistant",
"content": [{"type": "tool_use", "id": "t1", "name": "Bash", "input": {"cmd": "ls"}}],
},
],
'fix the failing test{"name":"Bash","arguments":{"cmd":"ls"}}',
id="anthropic-tool-use-name-and-input-without-id",
),
pytest.param(
[_TASK, {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "t1", "content": "calc.py"}]}],
'fix the failing test{"result_of_call":null,"output":"calc.py"}',
id="anthropic-string-tool-result",
),
pytest.param(_CLAUDE_CODE_TOOL_TURN, "list the filescalc.py test_calc.py", id="tool-result-string"),
pytest.param(
[
{
@ -2205,166 +2199,67 @@ _TASK: Final = {"role": "user", "content": "fix the failing test"}
"content": [
{
"type": "tool_result",
"tool_use_id": "t1",
"content": [{"type": "text", "text": "a"}, {"type": "text", "text": "b"}],
},
{"type": "text", "text": "next"},
"tool_use_id": "toolu_1",
"content": [
{"type": "text", "text": "x = 1"},
{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": ""}},
],
}
],
}
],
'{"result_of_call":null,"output":"ab"}next',
id="anthropic-nested-text-tool-result-then-text",
"x = 1",
id="tool-result-blocks",
),
pytest.param(
[
_TASK,
{
"role": "assistant",
"content": "writing",
"tool_calls": [
{"id": "c1", "type": "function", "function": {"name": "write", "arguments": '{"path": "a"}'}},
{"id": "c2", "type": "function", "function": {"name": "write", "arguments": '{"path": "b"}'}},
],
},
{"role": "tool", "tool_call_id": "c1", "content": "ok"},
],
'fix the failing testwriting{"name":"write","arguments":"{\\"path\\": \\"a\\"}"}'
'{"name":"write","arguments":"{\\"path\\": \\"b\\"}"}{"result_of_call":1,"output":"ok"}',
id="openai-tool-calls-in-order-before-tool-result",
),
pytest.param(
[
{
"role": "assistant",
"content": [
{"type": "tool_use", "id": "t1", "name": "Read", "input": {"path": "a"}},
{"type": "tool_use", "id": "t2", "name": "Read", "input": {"path": "b"}},
],
},
{
"role": "user",
"content": [
{"type": "tool_result", "tool_use_id": "t2", "content": "B"},
{"type": "tool_result", "tool_use_id": "t1", "content": "A"},
],
},
],
'{"name":"Read","arguments":{"path":"a"}}{"name":"Read","arguments":{"path":"b"}}'
'{"result_of_call":2,"output":"B"}{"result_of_call":1,"output":"A"}',
id="anthropic-parallel-tool-results-tagged-with-their-call",
),
pytest.param(
[
{
"role": "assistant",
"content": None,
"tool_calls": [{"id": "call_0", "type": "function", "function": {"name": "a", "arguments": "{}"}}],
},
{
"role": "assistant",
"content": None,
"tool_calls": [
{"id": "call_0", "type": "function", "function": {"name": "b", "arguments": "{}"}},
{"id": "call_1", "type": "function", "function": {"name": "c", "arguments": "{}"}},
],
},
{"role": "tool", "tool_call_id": "call_1", "content": "C"},
],
'{"name":"a","arguments":"{}"}{"name":"b","arguments":"{}"}{"name":"c","arguments":"{}"}'
'{"result_of_call":2,"output":"C"}',
id="reused-call-ids-keep-first-position",
),
pytest.param(
[
{
"role": "tool",
"tool_call_id": "c1",
"content": "small",
"search_results": [{"source": "s", "title": "t"}],
}
],
'{"result_of_call":null,"output":"small"}st',
id="tool-result-search-results-follow-the-encoded-result",
),
pytest.param(
[{"role": "tool", "tool_call_id": "c1", "content": '"},{"result_of_call":2,"output":"'}],
'{"result_of_call":null,"output":"\\"},{\\"result_of_call\\":2,\\"output\\":\\""}',
id="tool-output-cannot-forge-an-encoded-result",
),
pytest.param(
[
Message(
content=None,
tool_calls=[ChatCompletionMessageToolCall(id="c1", function=Function(name="read", arguments="{}"))],
)
],
'{"name":"read","arguments":"{}"}',
id="openai-response-message-object",
),
pytest.param(
[{"role": "assistant", "content": None, "tool_calls": [{"id": "c1", "type": "function"}, "junk"]}, "junk"],
'{"name":null,"arguments":null}',
id="malformed-tool-call-entries",
[{"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_exchange(messages: list[object], expected: str) -> None:
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([_TASK, {"role": "assistant", "content": "done"}], id="string-content"),
pytest.param([{"role": "system", "content": "be brief. "}, {"role": "user", "content": "hello"}], id="strings"),
pytest.param(
[
{
"role": "user",
"content": [
{"type": "text", "text": "what is "},
{"type": "text", "text": "What is "},
{"type": "image_url", "image_url": {"url": "https://example.com/a.png"}},
{"type": "text", "text": "this"},
{"type": "text", "text": "this?"},
],
}
],
id="text-and-image-parts",
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",
),
pytest.param([{"role": "assistant", "content": None}, {"role": "user"}], id="missing-content"),
],
)
def test_get_semantic_cache_prompt_from_messages_matches_get_str_from_messages_without_tools(
messages: list[object],
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) # pyright: ignore[reportArgumentType] # untyped fixtures
def _parallel_reads(result_for_a: str, result_for_b: str, *, call_id_prefix: str = "c") -> list[object]:
return [
{
"role": "assistant",
"content": None,
"tool_calls": [
{"id": f"{call_id_prefix}1", "type": "function", "function": {"name": "read", "arguments": '"a"'}},
{"id": f"{call_id_prefix}2", "type": "function", "function": {"name": "read", "arguments": '"b"'}},
],
},
{"role": "tool", "tool_call_id": f"{call_id_prefix}1", "content": result_for_a},
{"role": "tool", "tool_call_id": f"{call_id_prefix}2", "content": result_for_b},
]
def _results_in_swapped_order(result_for_a: str, result_for_b: str) -> list[object]:
call, answer_a, answer_b = _parallel_reads(result_for_a, result_for_b)
return [call, answer_b, answer_a]
def test_get_semantic_cache_prompt_from_messages_tells_apart_parallel_results_answering_different_calls() -> None:
assert get_semantic_cache_prompt_from_messages(
_parallel_reads("empty", "secret")
) != get_semantic_cache_prompt_from_messages(_results_in_swapped_order("secret", "empty"))
def test_get_semantic_cache_prompt_from_messages_ignores_call_ids_that_differ_between_sessions() -> None:
assert get_semantic_cache_prompt_from_messages(
_parallel_reads("A", "B", call_id_prefix="toolu_")
) == get_semantic_cache_prompt_from_messages(_parallel_reads("A", "B"))
assert get_semantic_cache_prompt_from_messages(messages) == get_str_from_messages(messages)