From 394abff55a93c668dfe85510b871a83adda45836 Mon Sep 17 00:00:00 2001 From: kerry Date: Tue, 6 Oct 2026 23:22:19 +0000 Subject: [PATCH] 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> --- .../cache-qdrant-semantic/tests/qdrant.rs | 2 +- litellm-rust/crates/cache/src/semantic.rs | 184 ++---- litellm-rust/crates/cache/tests/semantic.rs | 171 ++---- litellm/caching/caching.py | 17 + litellm/caching/redis_semantic_cache.py | 61 +- .../prompt_templates/common_utils.py | 125 +--- .../code_coverage_tests/recursive_detector.py | 16 +- .../caching/test_cache_max_messages.py | 415 +++++++++++++ .../caching/test_semantic_cache_tool_turns.py | 544 ------------------ tests/unit/caching/test_caching.py | 65 ++- .../unit/caching/test_redis_semantic_cache.py | 170 ++---- ...ore_utils_prompt_templates_common_utils.py | 203 ++----- 12 files changed, 780 insertions(+), 1193 deletions(-) create mode 100644 tests/integration/caching/test_cache_max_messages.py delete mode 100644 tests/integration/caching/test_semantic_cache_tool_turns.py diff --git a/litellm-rust/crates/cache-qdrant-semantic/tests/qdrant.rs b/litellm-rust/crates/cache-qdrant-semantic/tests/qdrant.rs index 7a27519c269..fe25c503a36 100644 --- a/litellm-rust/crates/cache-qdrant-semantic/tests/qdrant.rs +++ b/litellm-rust/crates/cache-qdrant-semantic/tests/qdrant.rs @@ -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( diff --git a/litellm-rust/crates/cache/src/semantic.rs b/litellm-rust/crates/cache/src/semantic.rs index 26274fd4324..5aafb333b39 100644 --- a/litellm-rust/crates/cache/src/semantic.rs +++ b/litellm-rust/crates/cache/src/semantic.rs @@ -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) -> 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 { (!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 { 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 { 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, - call_ordinals: &HashMap<&str, usize>, -) { +fn collect_input_text(value: &Value, parts: &mut Vec) { 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; } } } diff --git a/litellm-rust/crates/cache/tests/semantic.rs b/litellm-rust/crates/cache/tests/semantic.rs index af87e6e0164..e0a39bb9d5a 100644 --- a/litellm-rust/crates/cache/tests/semantic.rs +++ b/litellm-rust/crates/cache/tests/semantic.rs @@ -30,6 +30,31 @@ fn context(messages: Option, input: Option) -> SemanticCacheContex ]}]), "What is this?", )] +#[case::tool_result_string( + json!([ + {"role": "user", "content": "list the files"}, + {"role": "assistant", "content": [ + {"type": "tool_use", "id": "toolu_1", "name": "Bash", "input": {"command": "ls"}}, + ]}, + {"role": "user", "content": [ + {"type": "tool_result", "tool_use_id": "toolu_1", "content": "calc.py test_calc.py"}, + ]}, + ]), + "list the filescalc.py test_calc.py", +)] +#[case::tool_result_blocks( + json!([{"role": "user", "content": [ + {"type": "tool_result", "tool_use_id": "toolu_1", "content": [ + {"type": "text", "text": "x = 1"}, + {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": ""}}, + ]}, + ]}]), + "x = 1", +)] +#[case::tool_result_without_content( + json!([{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_1"}]}]), + "", +)] #[case::missing_null_and_empty_content( json!([{"role": "assistant"}, {"role": "assistant", "content": null}, {"role": "user", "content": ""}]), "", @@ -38,143 +63,67 @@ fn context(messages: Option, input: Option) -> 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, #[case] input: Option, diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 4070f41fd32..e6ad0ebdebe 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -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 diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index 342dc77f88d..e39c2074739 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -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( diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 6f849966902..20dc2726de9 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -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: diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 21af1a65192..863e8befcf9 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -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) diff --git a/tests/integration/caching/test_cache_max_messages.py b/tests/integration/caching/test_cache_max_messages.py new file mode 100644 index 00000000000..33bd0c09b4c --- /dev/null +++ b/tests/integration/caching/test_cache_max_messages.py @@ -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" diff --git a/tests/integration/caching/test_semantic_cache_tool_turns.py b/tests/integration/caching/test_semantic_cache_tool_turns.py deleted file mode 100644 index ac80b806c40..00000000000 --- a/tests/integration/caching/test_semantic_cache_tool_turns.py +++ /dev/null @@ -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" diff --git a/tests/unit/caching/test_caching.py b/tests/unit/caching/test_caching.py index 0a7ac3ecad1..0adb6b9a6f4 100644 --- a/tests/unit/caching/test_caching.py +++ b/tests/unit/caching/test_caching.py @@ -1,6 +1,7 @@ import asyncio import logging import re +import uuid from typing import Final from unittest.mock import MagicMock @@ -9,7 +10,7 @@ import pytest import litellm import litellm.caching.redis_cache as redis_cache_module from litellm._internal_context import current_service_target -from litellm.caching.caching import Cache, response_cache_phase +from litellm.caching.caching import Cache, CacheMode, response_cache_phase from litellm.caching.caching_handler import _PENDING_CACHE_WRITES from litellm.caching.in_memory_cache import InMemoryCache from litellm.caching.redis_cache import RedisCache, _RedisTimeoutLogThrottle @@ -473,3 +474,65 @@ async def test_a_lookup_already_inside_the_phase_does_not_open_a_second_one(v2_s await cache.async_get_cache(dynamic_cache_object=backend, **_REQUEST) assert backend.seen == [("llm_response", "cache.get llm_response")] assert [s.name for s in v2_span_exporter.get_finished_spans()] == ["cache.get llm_response"] + + +_TOOL_TURN_ITEM: Final = {"role": "user", "content": "hi"} + + +@pytest.mark.parametrize( + ("kwargs", "expected"), + [ + pytest.param({"messages": [_TOOL_TURN_ITEM] * 4}, True, id="four-messages-are-cached"), + pytest.param({"messages": [_TOOL_TURN_ITEM] * 5}, False, id="five-messages-skip-the-cache"), + pytest.param({"input": [_TOOL_TURN_ITEM] * 4}, True, id="four-responses-items-are-cached"), + pytest.param({"input": [_TOOL_TURN_ITEM] * 5}, False, id="five-responses-items-skip-the-cache"), + pytest.param({"input": "one prompt"}, True, id="string-input-is-one-message"), + pytest.param({"input": ["a", "b", "c", "d", "e"]}, True, id="embedding-strings-are-not-messages"), + ], +) +def test_should_use_cache_stops_past_the_default_max_messages(kwargs: dict[str, object], expected: bool) -> None: + assert Cache(type=LiteLLMCacheType.LOCAL).should_use_cache(**kwargs) is expected + + +def test_responses_sdk_items_count_toward_max_messages() -> None: + from openai.types.responses import ResponseFunctionToolCall + + call: Final = ResponseFunctionToolCall(type="function_call", call_id="c1", name="ls", arguments="{}") + + assert Cache(type=LiteLLMCacheType.LOCAL).should_use_cache(input=[_TOOL_TURN_ITEM, call, call, call, call]) is False + + +def test_max_messages_is_configurable_and_none_disables_it() -> None: + three: Final = [_TOOL_TURN_ITEM] * 3 + + assert Cache(type=LiteLLMCacheType.LOCAL, max_messages=2).should_use_cache(messages=three) is False + assert Cache(type=LiteLLMCacheType.LOCAL, max_messages=3).should_use_cache(messages=three) is True + assert Cache(type=LiteLLMCacheType.LOCAL, max_messages=None).should_use_cache(messages=three * 50) is True + + +def test_max_messages_beats_an_explicit_use_cache_opt_in() -> None: + cache: Final = Cache(type=LiteLLMCacheType.LOCAL, mode=CacheMode.default_off) + + assert cache.should_use_cache(messages=[_TOOL_TURN_ITEM] * 4, cache={"use-cache": True}) is True + assert cache.should_use_cache(messages=[_TOOL_TURN_ITEM] * 5, cache={"use-cache": True}) is False + + +def test_completion_past_max_messages_is_neither_served_from_nor_written_to_the_cache( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(litellm, "cache", Cache(type=LiteLLMCacheType.LOCAL)) + tag: Final = uuid.uuid4().hex + four: Final = [{"role": "user", "content": f"{tag} turn {index}"} for index in range(4)] + five: Final = [*four, {"role": "user", "content": f"{tag} turn 4"}] + + def answer(messages: list[dict[str, str]], mock_response: str) -> str: + response: Final = litellm.completion(model="gpt-4o-mini", messages=messages, mock_response=mock_response) + assert isinstance(response, litellm.ModelResponse), response + choice: Final = response.choices[0] + assert isinstance(choice, litellm.Choices), choice + return str(choice.message.content) + + assert answer(four, "four first") == "four first" + assert answer(four, "four second") == "four first", "a 4-message repeat missed the cache" + assert answer(five, "five first") == "five first" + assert answer(five, "five second") == "five second", "a 5-message repeat was served from the cache" diff --git a/tests/unit/caching/test_redis_semantic_cache.py b/tests/unit/caching/test_redis_semantic_cache.py index 07ec324699b..61e5a727411 100644 --- a/tests/unit/caching/test_redis_semantic_cache.py +++ b/tests/unit/caching/test_redis_semantic_cache.py @@ -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" diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py index 61d7cdbd96c..1d19264fb9e 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py @@ -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)