mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(caching): skip the cache past max_messages and keep tool_result text in semantic prompts
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
ffd3a85a18
commit
394abff55a
12 changed files with 780 additions and 1193 deletions
|
|
@ -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(
|
||||
|
|
|
|||
184
litellm-rust/crates/cache/src/semantic.rs
vendored
184
litellm-rust/crates/cache/src/semantic.rs
vendored
|
|
@ -1,15 +1,14 @@
|
|||
//! The embedding and prompt contract every semantic backend shares.
|
||||
//!
|
||||
//! Python's semantic caches all read their prompt through `get_semantic_cache_prompt_from_messages`, and
|
||||
//! `RedisSemanticCache._get_prompt_from_kwargs` (inherited by Valkey) adds Responses API `input`
|
||||
//! through `get_semantic_cache_prompt_from_responses_input`. Qdrant reads messages only. Each
|
||||
//! backend picks one of the two extractors here.
|
||||
//! Python's semantic caches all read their prompt through `get_str_from_messages`, and
|
||||
//! `RedisSemanticCache._get_prompt_from_kwargs` (inherited by Valkey) adds Responses API
|
||||
//! `input`. Qdrant reads messages only. Each backend picks one of the two extractors here.
|
||||
|
||||
use std::{collections::HashMap, future::Future, io};
|
||||
use std::{future::Future, io};
|
||||
|
||||
use serde::Serialize;
|
||||
use serde_json::{
|
||||
Value, json,
|
||||
Value,
|
||||
ser::{CharEscape, Formatter, Serializer},
|
||||
};
|
||||
|
||||
|
|
@ -84,122 +83,45 @@ impl Embedder for PreparedEmbedding {
|
|||
}
|
||||
}
|
||||
|
||||
/// `get_semantic_cache_prompt_from_messages`: every message's content text, tool calls and tool results,
|
||||
/// then its OpenAI `tool_calls`, then its search results. Each tool result is encoded with the
|
||||
/// position of the call it answers.
|
||||
/// `get_semantic_cache_prompt_from_messages`: every message's text content, including the text of
|
||||
/// Messages API `tool_result` blocks, followed by its search results.
|
||||
pub fn str_from_messages(messages: &[Value]) -> String {
|
||||
let messages: Vec<_> = messages.iter().filter_map(Value::as_object).collect();
|
||||
let call_ordinals = tool_call_ordinals(messages.iter().flat_map(|message| {
|
||||
let block_ids = message
|
||||
.get("content")
|
||||
.and_then(Value::as_array)
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.filter(|block| block.get("type").and_then(Value::as_str) == Some("tool_use"))
|
||||
.filter_map(|block| block.get("id"));
|
||||
let tool_call_ids = message
|
||||
.get("tool_calls")
|
||||
.and_then(Value::as_array)
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.filter_map(|tool_call| tool_call.get("id"));
|
||||
block_ids.chain(tool_call_ids)
|
||||
}));
|
||||
let mut text = String::new();
|
||||
for message in messages {
|
||||
if message.get("role").and_then(Value::as_str) == Some("tool") {
|
||||
let mut output = String::new();
|
||||
push_content_text(&mut output, message.get("content"), &call_ordinals);
|
||||
text.push_str(&tool_result_json(
|
||||
message.get("tool_call_id"),
|
||||
&call_ordinals,
|
||||
&output,
|
||||
));
|
||||
} else {
|
||||
push_content_text(&mut text, message.get("content"), &call_ordinals);
|
||||
}
|
||||
if let Some(Value::Array(tool_calls)) = message.get("tool_calls") {
|
||||
for tool_call in tool_calls.iter().filter_map(Value::as_object) {
|
||||
let function = tool_call.get("function");
|
||||
text.push_str(&tool_call_json(
|
||||
function.and_then(|function| function.get("name")),
|
||||
function.and_then(|function| function.get("arguments")),
|
||||
));
|
||||
for message in messages.iter().filter_map(Value::as_object) {
|
||||
match message.get("content") {
|
||||
Some(Value::String(content)) => text.push_str(content),
|
||||
Some(Value::Array(blocks)) => {
|
||||
for block in blocks {
|
||||
push_block_text(&mut text, block);
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
push_search_results_text(&mut text, message.get("search_results"));
|
||||
}
|
||||
text
|
||||
}
|
||||
|
||||
/// `_content_str_with_tools`: text parts, Anthropic `tool_use` blocks and `tool_result` content.
|
||||
fn push_content_text(
|
||||
text: &mut String,
|
||||
content: Option<&Value>,
|
||||
call_ordinals: &HashMap<&str, usize>,
|
||||
) {
|
||||
match content {
|
||||
Some(Value::String(content)) => text.push_str(content),
|
||||
fn push_block_text(text: &mut String, block: &Value) {
|
||||
if block.get("type").and_then(Value::as_str) != Some("tool_result") {
|
||||
push_text_field(text, block);
|
||||
return;
|
||||
}
|
||||
match block.get("content") {
|
||||
Some(Value::String(result)) => text.push_str(result),
|
||||
Some(Value::Array(blocks)) => {
|
||||
for block in blocks.iter().filter_map(Value::as_object) {
|
||||
match block.get("type").and_then(Value::as_str) {
|
||||
Some("tool_use") => {
|
||||
text.push_str(&tool_call_json(block.get("name"), block.get("input")));
|
||||
}
|
||||
Some("tool_result") => {
|
||||
let mut output = String::new();
|
||||
push_content_text(&mut output, block.get("content"), call_ordinals);
|
||||
text.push_str(&tool_result_json(
|
||||
block.get("tool_use_id"),
|
||||
call_ordinals,
|
||||
&output,
|
||||
));
|
||||
}
|
||||
_ => {
|
||||
if let Some(block_text) = block.get("text").and_then(Value::as_str) {
|
||||
text.push_str(block_text);
|
||||
}
|
||||
}
|
||||
}
|
||||
for inner in blocks {
|
||||
push_text_field(text, inner);
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
/// `tool_call_ordinals`: the 1-based position of each distinct string call id, first seen first.
|
||||
fn tool_call_ordinals<'a>(call_ids: impl Iterator<Item = &'a Value>) -> HashMap<&'a str, usize> {
|
||||
let mut ordinals = HashMap::new();
|
||||
for call_id in call_ids.filter_map(Value::as_str) {
|
||||
let next = ordinals.len() + 1;
|
||||
ordinals.entry(call_id).or_insert(next);
|
||||
fn push_text_field(text: &mut String, block: &Value) {
|
||||
if let Some(block_text) = block.get("text").and_then(Value::as_str) {
|
||||
text.push_str(block_text);
|
||||
}
|
||||
ordinals
|
||||
}
|
||||
|
||||
/// `tool_result_str`: `{"result_of_call":N,"output":...}`, with a `null` position when the result
|
||||
/// answers no known call.
|
||||
fn tool_result_json(
|
||||
call_id: Option<&Value>,
|
||||
call_ordinals: &HashMap<&str, usize>,
|
||||
output: &str,
|
||||
) -> String {
|
||||
let ordinal = call_id
|
||||
.and_then(Value::as_str)
|
||||
.and_then(|call_id| call_ordinals.get(call_id));
|
||||
format!(
|
||||
"{{\"result_of_call\":{},\"output\":{}}}",
|
||||
compact_json(&json!(ordinal)),
|
||||
compact_json(&Value::String(output.to_owned())),
|
||||
)
|
||||
}
|
||||
|
||||
/// `tool_call_str`: the compact `{"name":...,"arguments":...}` a tool call contributes.
|
||||
fn tool_call_json(name: Option<&Value>, arguments: Option<&Value>) -> String {
|
||||
compact_json(&json!({
|
||||
"name": name.unwrap_or(&Value::Null),
|
||||
"arguments": arguments.unwrap_or(&Value::Null),
|
||||
}))
|
||||
}
|
||||
|
||||
/// The messages prompt Qdrant embeds: `None` when the request carries no messages.
|
||||
|
|
@ -208,9 +130,8 @@ pub fn prompt_from_messages(context: &SemanticCacheContext) -> Option<String> {
|
|||
(!messages.is_empty()).then(|| str_from_messages(messages))
|
||||
}
|
||||
|
||||
/// `RedisSemanticCache._get_prompt_from_kwargs`: chat messages first, then
|
||||
/// `get_semantic_cache_prompt_from_responses_input` over a Responses API `input`. `None` when
|
||||
/// neither yields a prompt.
|
||||
/// `RedisSemanticCache._get_prompt_from_kwargs`: chat messages first, then the text parts of a
|
||||
/// Responses API `input`. `None` when neither yields a prompt.
|
||||
pub fn prompt_from_context(context: &SemanticCacheContext) -> Option<String> {
|
||||
if let Some(messages) = context.messages.as_ref().and_then(Value::as_array)
|
||||
&& !messages.is_empty()
|
||||
|
|
@ -218,16 +139,8 @@ pub fn prompt_from_context(context: &SemanticCacheContext) -> Option<String> {
|
|||
return Some(str_from_messages(messages));
|
||||
}
|
||||
let input = context.input.as_ref()?;
|
||||
let call_ordinals = tool_call_ordinals(
|
||||
input
|
||||
.as_array()
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.filter(|item| item.get("type").and_then(Value::as_str) == Some("function_call"))
|
||||
.filter_map(|item| item.get("call_id")),
|
||||
);
|
||||
let mut parts = Vec::new();
|
||||
collect_input_text(input, &mut parts, &call_ordinals);
|
||||
collect_input_text(input, &mut parts);
|
||||
let prompt = python_strip(&parts.join("\n")).to_owned();
|
||||
(!prompt.is_empty()).then_some(prompt)
|
||||
}
|
||||
|
|
@ -256,49 +169,26 @@ fn push_search_results_text(text: &mut String, search_results: Option<&Value>) {
|
|||
}
|
||||
}
|
||||
|
||||
fn collect_input_text(
|
||||
value: &Value,
|
||||
parts: &mut Vec<String>,
|
||||
call_ordinals: &HashMap<&str, usize>,
|
||||
) {
|
||||
fn collect_input_text(value: &Value, parts: &mut Vec<String>) {
|
||||
match value {
|
||||
Value::String(text) => {
|
||||
push_trimmed(text, parts);
|
||||
}
|
||||
Value::Array(items) => {
|
||||
for item in items {
|
||||
collect_input_text(item, parts, call_ordinals);
|
||||
collect_input_text(item, parts);
|
||||
}
|
||||
}
|
||||
Value::Object(map) => {
|
||||
if map.get("type").and_then(Value::as_str) == Some("function_call") {
|
||||
parts.push(tool_call_json(map.get("name"), map.get("arguments")));
|
||||
return;
|
||||
}
|
||||
if map.get("type").and_then(Value::as_str) == Some("function_call_output") {
|
||||
let mut output_parts = Vec::new();
|
||||
if let Some(output) = map.get("output") {
|
||||
collect_input_text(output, &mut output_parts, call_ordinals);
|
||||
}
|
||||
parts.push(tool_result_json(
|
||||
map.get("call_id"),
|
||||
call_ordinals,
|
||||
python_strip(&output_parts.join("\n")),
|
||||
));
|
||||
return;
|
||||
}
|
||||
if let Some(content) = map.get("content").filter(|content| !content.is_null()) {
|
||||
collect_input_text(content, parts, call_ordinals);
|
||||
collect_input_text(content, parts);
|
||||
return;
|
||||
}
|
||||
for key in ["text", "output", "input_text", "output_text"] {
|
||||
match map.get(key) {
|
||||
Some(nested @ Value::Array(_)) => {
|
||||
collect_input_text(nested, parts, call_ordinals);
|
||||
return;
|
||||
}
|
||||
Some(Value::String(text)) if push_trimmed(text, parts) => return,
|
||||
_ => {}
|
||||
if let Some(Value::String(text)) = map.get(key)
|
||||
&& push_trimmed(text, parts)
|
||||
{
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
171
litellm-rust/crates/cache/tests/semantic.rs
vendored
171
litellm-rust/crates/cache/tests/semantic.rs
vendored
|
|
@ -30,6 +30,31 @@ fn context(messages: Option<Value>, input: Option<Value>) -> SemanticCacheContex
|
|||
]}]),
|
||||
"What is this?",
|
||||
)]
|
||||
#[case::tool_result_string(
|
||||
json!([
|
||||
{"role": "user", "content": "list the files"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "tool_use", "id": "toolu_1", "name": "Bash", "input": {"command": "ls"}},
|
||||
]},
|
||||
{"role": "user", "content": [
|
||||
{"type": "tool_result", "tool_use_id": "toolu_1", "content": "calc.py test_calc.py"},
|
||||
]},
|
||||
]),
|
||||
"list the filescalc.py test_calc.py",
|
||||
)]
|
||||
#[case::tool_result_blocks(
|
||||
json!([{"role": "user", "content": [
|
||||
{"type": "tool_result", "tool_use_id": "toolu_1", "content": [
|
||||
{"type": "text", "text": "x = 1"},
|
||||
{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": ""}},
|
||||
]},
|
||||
]}]),
|
||||
"x = 1",
|
||||
)]
|
||||
#[case::tool_result_without_content(
|
||||
json!([{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_1"}]}]),
|
||||
"",
|
||||
)]
|
||||
#[case::missing_null_and_empty_content(
|
||||
json!([{"role": "assistant"}, {"role": "assistant", "content": null}, {"role": "user", "content": ""}]),
|
||||
"",
|
||||
|
|
@ -38,143 +63,67 @@ fn context(messages: Option<Value>, input: Option<Value>) -> SemanticCacheContex
|
|||
json!([{"role": "tool", "content": "small", "search_results": [
|
||||
{"source": "s", "title": "t", "content": [{"text": "hidden payload"}]},
|
||||
]}]),
|
||||
r#"{"result_of_call":null,"output":"small"}sthidden payload"#,
|
||||
"smallsthidden payload",
|
||||
)]
|
||||
#[case::title_only_search_result(
|
||||
json!([{"role": "tool", "content": "small", "search_results": [
|
||||
{"source": "s", "title": "long title", "content": []},
|
||||
]}]),
|
||||
r#"{"result_of_call":null,"output":"small"}slong title"#,
|
||||
"smallslong title",
|
||||
)]
|
||||
#[case::search_results_without_content(
|
||||
json!([{"role": "tool", "search_results": [{"source": "s", "title": "t"}]}]),
|
||||
r#"{"result_of_call":null,"output":""}st"#,
|
||||
"st",
|
||||
)]
|
||||
#[case::search_result_fields_in_python_order(
|
||||
json!([{"role": "tool", "content": "c", "search_results": [
|
||||
{"citations": {"enabled": true}, "content": [{"text": "body"}], "title": "t", "source": "s"},
|
||||
{"source": "s2"},
|
||||
]}]),
|
||||
r#"{"result_of_call":null,"output":"c"}stbody{"enabled":true}s2"#,
|
||||
r#"cstbody{"enabled":true}s2"#,
|
||||
)]
|
||||
#[case::null_citations_skipped(
|
||||
json!([{"role": "tool", "content": "c", "search_results": [
|
||||
{"source": "s", "citations": null},
|
||||
]}]),
|
||||
r#"{"result_of_call":null,"output":"c"}s"#,
|
||||
"cs",
|
||||
)]
|
||||
#[case::non_string_and_non_object_entries_skipped(
|
||||
json!([{"role": "tool", "content": "c", "search_results": [
|
||||
"junk",
|
||||
{"source": 1, "title": null, "content": ["junk", {"text": 3}, {"text": "kept"}]},
|
||||
]}]),
|
||||
r#"{"result_of_call":null,"output":"c"}kept"#,
|
||||
"ckept",
|
||||
)]
|
||||
#[case::non_list_search_results_skipped(
|
||||
json!([{"role": "tool", "content": "c", "search_results": {"source": "s"}}]),
|
||||
r#"{"result_of_call":null,"output":"c"}"#,
|
||||
"c",
|
||||
)]
|
||||
#[case::citations_compact_in_insertion_order(
|
||||
json!([{"role": "tool", "search_results": [
|
||||
{"citations": {"z": 1, "a": [1.5, true, null], "m": {"k": "v"}}},
|
||||
]}]),
|
||||
r#"{"result_of_call":null,"output":""}{"z":1,"a":[1.5,true,null],"m":{"k":"v"}}"#,
|
||||
r#"{"z":1,"a":[1.5,true,null],"m":{"k":"v"}}"#,
|
||||
)]
|
||||
#[case::citations_ensure_ascii(
|
||||
json!([{"role": "tool", "search_results": [{"citations": ["caf\u{e9}", "\u{4e2d}"]}]}]),
|
||||
r#"{"result_of_call":null,"output":""}["caf\u00e9","\u4e2d"]"#,
|
||||
r#"["caf\u00e9","\u4e2d"]"#,
|
||||
)]
|
||||
#[case::citations_astral_chars_as_surrogate_pairs(
|
||||
json!([{"role": "tool", "search_results": [{"citations": "\u{1f600}"}]}]),
|
||||
r#"{"result_of_call":null,"output":""}"\ud83d\ude00""#,
|
||||
r#""\ud83d\ude00""#,
|
||||
)]
|
||||
#[case::citations_escapes(
|
||||
json!([{"role": "tool", "search_results": [{"citations": "q\"\\\n\t\u{1}/"}]}]),
|
||||
r#"{"result_of_call":null,"output":""}"q\"\\\n\t\u0001/""#,
|
||||
r#""q\"\\\n\t\u0001/""#,
|
||||
)]
|
||||
#[case::citations_large_float_exponent(
|
||||
json!([{"role": "tool", "search_results": [{"citations": [1e20, 1.0]}]}]),
|
||||
r#"{"result_of_call":null,"output":""}[1e+20,1.0]"#,
|
||||
"[1e+20,1.0]",
|
||||
)]
|
||||
#[case::citations_scalars(
|
||||
json!([{"role": "tool", "search_results": [{"citations": false}, {"citations": 3}]}]),
|
||||
r#"{"result_of_call":null,"output":""}false3"#,
|
||||
)]
|
||||
#[case::anthropic_tool_use_name_and_input_without_id(
|
||||
json!([
|
||||
{"role": "user", "content": "fix the failing test"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "tool_use", "id": "t1", "name": "Bash", "input": {"cmd": "ls"}},
|
||||
]},
|
||||
]),
|
||||
r#"fix the failing test{"name":"Bash","arguments":{"cmd":"ls"}}"#,
|
||||
)]
|
||||
#[case::anthropic_string_tool_result(
|
||||
json!([{"role": "user", "content": [
|
||||
{"type": "tool_result", "tool_use_id": "t1", "content": "calc.py"},
|
||||
]}]),
|
||||
r#"{"result_of_call":null,"output":"calc.py"}"#,
|
||||
)]
|
||||
#[case::anthropic_nested_text_tool_result_then_text(
|
||||
json!([{"role": "user", "content": [
|
||||
{"type": "tool_result", "tool_use_id": "t1", "content": [
|
||||
{"type": "text", "text": "a"},
|
||||
{"type": "text", "text": "b"},
|
||||
]},
|
||||
{"type": "text", "text": "next"},
|
||||
]}]),
|
||||
r#"{"result_of_call":null,"output":"ab"}next"#,
|
||||
)]
|
||||
#[case::openai_tool_calls_in_order_before_tool_result(
|
||||
json!([
|
||||
{"role": "assistant", "content": "writing", "tool_calls": [
|
||||
{"id": "c1", "type": "function", "function": {"name": "write", "arguments": "{\"path\": \"a\"}"}},
|
||||
{"id": "c2", "type": "function", "function": {"name": "write", "arguments": "{\"path\": \"b\"}"}},
|
||||
]},
|
||||
{"role": "tool", "tool_call_id": "c1", "content": "ok"},
|
||||
]),
|
||||
r#"writing{"name":"write","arguments":"{\"path\": \"a\"}"}{"name":"write","arguments":"{\"path\": \"b\"}"}{"result_of_call":1,"output":"ok"}"#,
|
||||
)]
|
||||
#[case::anthropic_parallel_tool_results_tagged_with_their_call(
|
||||
json!([
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "tool_use", "id": "t1", "name": "Read", "input": {"path": "a"}},
|
||||
{"type": "tool_use", "id": "t2", "name": "Read", "input": {"path": "b"}},
|
||||
]},
|
||||
{"role": "user", "content": [
|
||||
{"type": "tool_result", "tool_use_id": "t2", "content": "B"},
|
||||
{"type": "tool_result", "tool_use_id": "t1", "content": "A"},
|
||||
]},
|
||||
]),
|
||||
r#"{"name":"Read","arguments":{"path":"a"}}{"name":"Read","arguments":{"path":"b"}}{"result_of_call":2,"output":"B"}{"result_of_call":1,"output":"A"}"#,
|
||||
)]
|
||||
#[case::reused_call_ids_keep_first_position(
|
||||
json!([
|
||||
{"role": "assistant", "content": null, "tool_calls": [
|
||||
{"id": "call_0", "type": "function", "function": {"name": "a", "arguments": "{}"}},
|
||||
]},
|
||||
{"role": "assistant", "content": null, "tool_calls": [
|
||||
{"id": "call_0", "type": "function", "function": {"name": "b", "arguments": "{}"}},
|
||||
{"id": "call_1", "type": "function", "function": {"name": "c", "arguments": "{}"}},
|
||||
]},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": "C"},
|
||||
]),
|
||||
r#"{"name":"a","arguments":"{}"}{"name":"b","arguments":"{}"}{"name":"c","arguments":"{}"}{"result_of_call":2,"output":"C"}"#,
|
||||
)]
|
||||
#[case::unknown_call_ids_encode_a_null_position(
|
||||
json!([
|
||||
{"role": "tool", "tool_call_id": "c9", "content": "ok"},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "t9", "content": "done"}]},
|
||||
]),
|
||||
r#"{"result_of_call":null,"output":"ok"}{"result_of_call":null,"output":"done"}"#,
|
||||
)]
|
||||
#[case::tool_output_cannot_forge_an_encoded_result(
|
||||
json!([{"role": "tool", "tool_call_id": "c1", "content": "\"},{\"result_of_call\":2,\"output\":\""}]),
|
||||
r#"{"result_of_call":null,"output":"\"},{\"result_of_call\":2,\"output\":\""}"#,
|
||||
)]
|
||||
#[case::malformed_tool_call_entries(
|
||||
json!([{"role": "assistant", "content": null, "tool_calls": [{"id": "c1", "type": "function"}, "junk"]}]),
|
||||
r#"{"name":null,"arguments":null}"#,
|
||||
"false3",
|
||||
)]
|
||||
fn str_from_messages_matches_python(#[case] messages: Value, #[case] expected: &str) {
|
||||
assert_eq!(str_from_messages(messages.as_array().unwrap()), expected);
|
||||
|
|
@ -268,50 +217,6 @@ fn prompt_from_messages_reads_messages_only(
|
|||
)]
|
||||
#[case::nested_lists(None, Some(json!([["a", [" b "]], "", "c"])), Some("a\nb\nc"))]
|
||||
#[case::scalars_ignored(None, Some(json!([1, true, null, "kept"])), Some("kept"))]
|
||||
#[case::responses_function_call(
|
||||
None,
|
||||
Some(json!([
|
||||
{"role": "user", "content": "update the config"},
|
||||
{"type": "function_call", "call_id": "c1", "name": "write_file", "arguments": "{\"path\":\"a.yaml\"}"},
|
||||
{"type": "function_call_output", "call_id": "c1", "output": "ok"},
|
||||
])),
|
||||
Some("update the config\n{\"name\":\"write_file\",\"arguments\":\"{\\\"path\\\":\\\"a.yaml\\\"}\"}\n{\"result_of_call\":1,\"output\":\"ok\"}"),
|
||||
)]
|
||||
#[case::responses_structured_function_call_output(
|
||||
None,
|
||||
Some(json!([
|
||||
{"type": "function_call", "call_id": "c1", "name": "write_file", "arguments": "{}"},
|
||||
{"type": "function_call_output", "call_id": "c1", "output": [{"type": "input_text", "text": "denied"}]},
|
||||
])),
|
||||
Some("{\"name\":\"write_file\",\"arguments\":\"{}\"}\n{\"result_of_call\":1,\"output\":\"denied\"}"),
|
||||
)]
|
||||
#[case::responses_multi_part_output_joined_by_lines(
|
||||
None,
|
||||
Some(json!([
|
||||
{"type": "function_call", "call_id": "c1", "name": "run", "arguments": "{}"},
|
||||
{"type": "function_call_output", "call_id": "c1", "output": [
|
||||
{"type": "input_text", "text": " line one "},
|
||||
{"type": "input_text", "text": "line two"},
|
||||
]},
|
||||
])),
|
||||
Some(r#"{"name":"run","arguments":"{}"}
|
||||
{"result_of_call":1,"output":"line one\nline two"}"#),
|
||||
)]
|
||||
#[case::responses_parallel_outputs_tagged_with_their_call(
|
||||
None,
|
||||
Some(json!([
|
||||
{"type": "function_call", "call_id": "c1", "name": "read", "arguments": "a"},
|
||||
{"type": "function_call", "call_id": "c2", "name": "read", "arguments": "b"},
|
||||
{"type": "function_call_output", "call_id": "c2", "output": "B"},
|
||||
{"type": "function_call_output", "call_id": "c1", "output": "A"},
|
||||
])),
|
||||
Some("{\"name\":\"read\",\"arguments\":\"a\"}\n{\"name\":\"read\",\"arguments\":\"b\"}\n{\"result_of_call\":2,\"output\":\"B\"}\n{\"result_of_call\":1,\"output\":\"A\"}"),
|
||||
)]
|
||||
#[case::responses_unknown_call_id_encodes_a_null_position(
|
||||
None,
|
||||
Some(json!([{"type": "function_call_output", "call_id": "c9", "output": "orphan"}])),
|
||||
Some(r#"{"result_of_call":null,"output":"orphan"}"#),
|
||||
)]
|
||||
fn prompt_from_context_matches_python(
|
||||
#[case] messages: Option<Value>,
|
||||
#[case] input: Option<Value>,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
415
tests/integration/caching/test_cache_max_messages.py
Normal file
415
tests/integration/caching/test_cache_max_messages.py
Normal file
|
|
@ -0,0 +1,415 @@
|
|||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, eventually, gateway_from_environment
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_VECTOR_SIZE: Final = 16
|
||||
_CHAT_MODEL: Final = "capped-chat"
|
||||
_CLAUDE_MODEL: Final = "capped-claude"
|
||||
_EMBEDDING_MODEL: Final = "capped-embedder"
|
||||
_COLLECTION: Final = "cache-max-messages"
|
||||
|
||||
|
||||
def _vector(text: str) -> tuple[float, ...]:
|
||||
digest: Final = hashlib.sha256(text.encode()).digest()
|
||||
return tuple((byte - 127.5) / 127.5 for byte in digest[:_VECTOR_SIZE])
|
||||
|
||||
|
||||
def _cosine(left: Sequence[float], right: Sequence[float]) -> float:
|
||||
dot: Final = sum(a * b for a, b in zip(left, right, strict=True))
|
||||
return dot / (math.sqrt(sum(a * a for a in left)) * math.sqrt(sum(b * b for b in right)))
|
||||
|
||||
|
||||
def _answer() -> str:
|
||||
return f"answer-{uuid.uuid4().hex[:16]}"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _Peer:
|
||||
"""Embeddings, chat, Responses, Anthropic Messages and a Qdrant collection on one owned socket"""
|
||||
|
||||
lock: threading.Lock = field(default_factory=threading.Lock)
|
||||
points: list[Mapping[str, JsonValue]] = field(default_factory=list) # mutable-ok: the Qdrant collection
|
||||
embedded: list[str] = field(default_factory=list) # mutable-ok: every prompt the proxy embedded
|
||||
answered: list[str] = field(default_factory=list) # mutable-ok: every completion the provider served
|
||||
|
||||
def stored(self) -> int:
|
||||
with self.lock:
|
||||
return len(self.points)
|
||||
|
||||
def respond(self, request: Request) -> Reply:
|
||||
path: Final = urlsplit(request.target).path
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body) if request.body else {}
|
||||
collection: Final = f"/qdrant/collections/{_COLLECTION}"
|
||||
if request.method == "GET" and path == f"{collection}/exists":
|
||||
return self._json({"result": {"exists": False}, "status": "ok"})
|
||||
if request.method in {"GET", "PUT"} and path in {collection, f"{collection}/index"}:
|
||||
return self._json({"result": True, "status": "ok"})
|
||||
if request.method == "PUT" and path == f"{collection}/points":
|
||||
with self.lock:
|
||||
self.points.extend(_JSON_OBJECT.validate_python(point) for point in _list(body["points"]))
|
||||
return self._json({"result": {"status": "completed"}, "status": "ok"})
|
||||
if request.method == "POST" and path == f"{collection}/points/search":
|
||||
return self._json({"result": self._search(body), "status": "ok"})
|
||||
if request.method == "GET" and path == "/v1/models":
|
||||
return self._json({"object": "list", "data": []})
|
||||
if request.method == "POST" and path == "/v1/embeddings":
|
||||
text: Final = str(body["input"])
|
||||
with self.lock:
|
||||
self.embedded.append(text)
|
||||
return self._json(
|
||||
{
|
||||
"object": "list",
|
||||
"model": "text-embedding-3-small",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": list(_vector(text))}],
|
||||
"usage": {"prompt_tokens": 4, "total_tokens": 4},
|
||||
}
|
||||
)
|
||||
if request.method == "POST" and path == "/v1/chat/completions":
|
||||
return self._json(self._chat_reply(_answer()))
|
||||
if request.method == "POST" and path == "/v1/responses":
|
||||
return self._json(self._responses_reply(_answer()))
|
||||
if request.method == "POST" and path == "/v1/messages":
|
||||
return self._json(self._messages_reply(_answer()))
|
||||
raise AssertionError(f"unexpected peer request {request.method} {request.target}")
|
||||
|
||||
def _search(self, body: Mapping[str, JsonValue]) -> list[JsonValue]:
|
||||
query: Final = [float(str(value)) for value in _list(body["vector"])]
|
||||
key: Final = _JSON_OBJECT.validate_python(
|
||||
_JSON_OBJECT.validate_python(_list(_JSON_OBJECT.validate_python(body["filter"])["must"])[0])["match"]
|
||||
)["value"]
|
||||
with self.lock:
|
||||
scoped: Final = [
|
||||
point
|
||||
for point in self.points
|
||||
if _JSON_OBJECT.validate_python(point["payload"])["litellm_cache_key"] == key
|
||||
]
|
||||
ranked: Final = sorted(
|
||||
(
|
||||
{
|
||||
"id": point["id"],
|
||||
"score": _cosine(query, [float(str(value)) for value in _list(point["vector"])]),
|
||||
"payload": point["payload"],
|
||||
}
|
||||
for point in scoped
|
||||
),
|
||||
key=lambda hit: -float(str(hit["score"])),
|
||||
)
|
||||
return list(ranked[:1])
|
||||
|
||||
def _chat_reply(self, answer: str) -> Mapping[str, JsonValue]:
|
||||
with self.lock:
|
||||
self.answered.append(answer)
|
||||
return {
|
||||
"id": f"chatcmpl-{uuid.uuid4().hex}",
|
||||
"object": "chat.completion",
|
||||
"created": 1789788253,
|
||||
"model": "gpt-5.4-mini",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": answer}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12},
|
||||
}
|
||||
|
||||
def _responses_reply(self, answer: str) -> Mapping[str, JsonValue]:
|
||||
with self.lock:
|
||||
self.answered.append(answer)
|
||||
identity: Final = uuid.uuid4().hex
|
||||
return {
|
||||
"id": f"resp_{identity}",
|
||||
"object": "response",
|
||||
"created_at": 1789788253,
|
||||
"status": "completed",
|
||||
"model": "gpt-5.4-mini",
|
||||
"output": [
|
||||
{
|
||||
"id": f"msg_{identity}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": answer, "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 10, "output_tokens": 2, "total_tokens": 12},
|
||||
}
|
||||
|
||||
def _messages_reply(self, answer: str) -> Mapping[str, JsonValue]:
|
||||
with self.lock:
|
||||
self.answered.append(answer)
|
||||
return {
|
||||
"id": f"msg_{uuid.uuid4().hex}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-5-5",
|
||||
"content": [{"type": "text", "text": answer}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 10, "output_tokens": 2},
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _json(value: Mapping[str, JsonValue]) -> Reply:
|
||||
return Reply(body=json.dumps(value).encode())
|
||||
|
||||
|
||||
def _list(value: JsonValue) -> list[JsonValue]:
|
||||
assert isinstance(value, list), value
|
||||
return value
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Proxy:
|
||||
gateway: Gateway
|
||||
peer: _Peer
|
||||
|
||||
def reply(self, route: str, conversation: list[JsonValue]) -> str:
|
||||
"""The answer text for one request, as the client sees it"""
|
||||
model: Final = _CLAUDE_MODEL if route == "/v1/messages" else _CHAT_MODEL
|
||||
body: Final[dict[str, JsonValue]] = (
|
||||
{"model": model, "input": conversation}
|
||||
if route == "/v1/responses"
|
||||
else {"model": model, "max_tokens": 16, "messages": conversation}
|
||||
)
|
||||
response: Final = self.gateway.request("POST", route, body)
|
||||
assert response.status_code == 200, response.text
|
||||
payload: Final = _JSON_OBJECT.validate_json(response.content)
|
||||
if route == "/v1/messages":
|
||||
return str(_JSON_OBJECT.validate_python(_list(payload["content"])[0])["text"])
|
||||
if route == "/v1/responses":
|
||||
message: Final = _JSON_OBJECT.validate_python(_list(payload["output"])[0])
|
||||
return str(_JSON_OBJECT.validate_python(_list(message["content"])[0])["text"])
|
||||
choice: Final = _JSON_OBJECT.validate_python(_list(payload["choices"])[0])
|
||||
return str(_JSON_OBJECT.validate_python(choice["message"])["content"])
|
||||
|
||||
def miss(self, route: str, conversation: list[JsonValue]) -> str:
|
||||
answered_before: Final = len(self.peer.answered)
|
||||
answer: Final = self.reply(route, conversation)
|
||||
assert len(self.peer.answered) == answered_before + 1, f"{answer} was served from the cache"
|
||||
return answer
|
||||
|
||||
def hit(self, route: str, conversation: list[JsonValue]) -> str:
|
||||
"""Repeats the request until the cache serves it, since the store after a miss is asynchronous"""
|
||||
|
||||
def attempt() -> tuple[str, bool]:
|
||||
answered_before: Final = len(self.peer.answered)
|
||||
answer: Final = self.reply(route, conversation)
|
||||
return answer, len(self.peer.answered) == answered_before
|
||||
|
||||
return eventually(attempt, lambda outcome: outcome[1])[0]
|
||||
|
||||
|
||||
def _proxy(tmp_path_factory: pytest.TempPathFactory, cache_params: Mapping[str, JsonValue]) -> Iterator[_Proxy]:
|
||||
peer: Final = _Peer()
|
||||
directory: Final = tmp_path_factory.mktemp("cache-max-messages")
|
||||
with gateway_from_environment() as gateway, wire_server(peer.respond) as wire:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["model_list"] = [
|
||||
{
|
||||
"model_name": _CHAT_MODEL,
|
||||
"litellm_params": {"model": "openai/gpt-5.4-mini", "api_base": f"{wire.url}/v1", "api_key": "k"},
|
||||
},
|
||||
{
|
||||
"model_name": _CLAUDE_MODEL,
|
||||
"litellm_params": {"model": "anthropic/claude-sonnet-5-5", "api_base": wire.url, "api_key": "k"},
|
||||
},
|
||||
{
|
||||
"model_name": _EMBEDDING_MODEL,
|
||||
"litellm_params": {
|
||||
"model": "openai/text-embedding-3-small",
|
||||
"api_base": f"{wire.url}/v1",
|
||||
"api_key": "k",
|
||||
},
|
||||
},
|
||||
]
|
||||
config["litellm_settings"]["cache_params"] = {
|
||||
**cache_params,
|
||||
**({"qdrant_api_base": f"{wire.url}/qdrant"} if cache_params["type"] == "qdrant-semantic" else {}),
|
||||
}
|
||||
path: Final = directory / "cache_max_messages.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, directory, {}, config=path) as candidate:
|
||||
yield _Proxy(candidate, peer)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def qdrant_proxy(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Proxy]:
|
||||
"""Qdrant semantic cache with the default max_messages of 4"""
|
||||
yield from _proxy(
|
||||
tmp_path_factory,
|
||||
{
|
||||
"type": "qdrant-semantic",
|
||||
"qdrant_collection_name": _COLLECTION,
|
||||
"qdrant_semantic_cache_embedding_model": _EMBEDDING_MODEL,
|
||||
"qdrant_semantic_cache_vector_size": _VECTOR_SIZE,
|
||||
"qdrant_quantization_config": "binary",
|
||||
"similarity_threshold": 0.99,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def exact_proxy(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Proxy]:
|
||||
"""Exact cache with max_messages lowered to 2 in cache_params"""
|
||||
yield from _proxy(tmp_path_factory, {"type": "local", "max_messages": 2})
|
||||
|
||||
|
||||
def _claude_code_turns(task: str) -> tuple[list[JsonValue], list[JsonValue], list[JsonValue]]:
|
||||
"""Turns 1 to 3 of a Claude Code session on /v1/messages: 1, 3 and 5 messages"""
|
||||
first: Final[list[JsonValue]] = [{"role": "user", "content": task}]
|
||||
second: Final[list[JsonValue]] = [
|
||||
*first,
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "tool_use", "id": "toolu_1", "name": "Bash", "input": {"command": "ls"}}],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "calc.py test_calc.py"}],
|
||||
},
|
||||
]
|
||||
third: Final[list[JsonValue]] = [
|
||||
*second,
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "tool_use", "id": "toolu_2", "name": "Read", "input": {"file_path": "calc.py"}}],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "toolu_2",
|
||||
"content": [{"type": "text", "text": "def add(a, b): return a - b"}],
|
||||
}
|
||||
],
|
||||
},
|
||||
]
|
||||
return first, second, third
|
||||
|
||||
|
||||
def _agent_turns(task: str) -> tuple[list[JsonValue], list[JsonValue], list[JsonValue]]:
|
||||
"""Turns 1 to 3 of an OpenAI tool loop on /v1/chat/completions: 2, 4 and 6 messages"""
|
||||
|
||||
def call(call_id: str, path: str) -> list[JsonValue]:
|
||||
return [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": call_id,
|
||||
"type": "function",
|
||||
"function": {"name": "write_file", "arguments": json.dumps({"path": path})},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": call_id, "content": f"wrote {path}"},
|
||||
]
|
||||
|
||||
first: Final[list[JsonValue]] = [
|
||||
{"role": "system", "content": "You are a coding agent"},
|
||||
{"role": "user", "content": task},
|
||||
]
|
||||
second: Final[list[JsonValue]] = [*first, *call("call_1", "a.yaml")]
|
||||
third: Final[list[JsonValue]] = [*second, *call("call_2", "b.yaml")]
|
||||
return first, second, third
|
||||
|
||||
|
||||
def _responses_turns(task: str) -> tuple[list[JsonValue], list[JsonValue], list[JsonValue]]:
|
||||
"""Turns 1 to 3 of an agent on /v1/responses: 1, 3 and 5 input items"""
|
||||
|
||||
def call(call_id: str, path: str) -> list[JsonValue]:
|
||||
return [
|
||||
{
|
||||
"type": "function_call",
|
||||
"call_id": call_id,
|
||||
"name": "write_file",
|
||||
"arguments": json.dumps({"path": path}),
|
||||
},
|
||||
{"type": "function_call_output", "call_id": call_id, "output": f"wrote {path}"},
|
||||
]
|
||||
|
||||
first: Final[list[JsonValue]] = [{"role": "user", "content": task}]
|
||||
second: Final[list[JsonValue]] = [*first, *call("call_1", "a.yaml")]
|
||||
third: Final[list[JsonValue]] = [*second, *call("call_2", "b.yaml")]
|
||||
return first, second, third
|
||||
|
||||
|
||||
def _assert_turns_are_cached_up_to_four_messages(
|
||||
proxy: _Proxy, route: str, turns: Callable[[str], tuple[list[JsonValue], list[JsonValue], list[JsonValue]]]
|
||||
) -> None:
|
||||
first, second, third = turns(f"update the config {uuid.uuid4().hex}")
|
||||
stored_before: Final = proxy.peer.stored()
|
||||
|
||||
short_answers: Final = (proxy.miss(route, first), proxy.miss(route, second))
|
||||
repeated: Final = (proxy.hit(route, first), proxy.hit(route, second))
|
||||
long_answers: Final = (proxy.miss(route, third), proxy.miss(route, third))
|
||||
proxy.hit(route, first)
|
||||
|
||||
assert repeated == short_answers, f"a repeated short turn got another turn's answer: {short_answers} {repeated}"
|
||||
assert long_answers[0] != long_answers[1], "a turn past max_messages was served from the cache"
|
||||
assert all("b.yaml" not in text and "def add" not in text for text in proxy.peer.embedded), (
|
||||
"a turn past max_messages was embedded"
|
||||
)
|
||||
assert proxy.peer.stored() == stored_before + 2, "a turn past max_messages was written to the cache"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("route", "turns"),
|
||||
[
|
||||
pytest.param("/v1/messages", _claude_code_turns, id="messages"),
|
||||
pytest.param("/v1/chat/completions", _agent_turns, id="chat-completions"),
|
||||
],
|
||||
)
|
||||
def test_agent_turns_are_cached_up_to_four_messages_and_bypass_the_cache_past_it(
|
||||
qdrant_proxy: _Proxy,
|
||||
route: str,
|
||||
turns: Callable[[str], tuple[list[JsonValue], list[JsonValue], list[JsonValue]]],
|
||||
) -> None:
|
||||
_assert_turns_are_cached_up_to_four_messages(qdrant_proxy, route, turns)
|
||||
|
||||
|
||||
def test_tool_result_text_tells_a_tool_turn_apart_from_the_turn_before_it(qdrant_proxy: _Proxy) -> None:
|
||||
task: Final = f"list the files {uuid.uuid4().hex}"
|
||||
first, second, _ = _claude_code_turns(task)
|
||||
|
||||
qdrant_proxy.miss("/v1/messages", first)
|
||||
qdrant_proxy.miss("/v1/messages", second)
|
||||
|
||||
assert task in qdrant_proxy.peer.embedded[-1], qdrant_proxy.peer.embedded[-1]
|
||||
assert "calc.py test_calc.py" in qdrant_proxy.peer.embedded[-1], qdrant_proxy.peer.embedded[-1]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("route", "turns"),
|
||||
[
|
||||
pytest.param("/v1/messages", _claude_code_turns, id="messages"),
|
||||
pytest.param("/v1/chat/completions", _agent_turns, id="chat-completions"),
|
||||
pytest.param("/v1/responses", _responses_turns, id="responses"),
|
||||
],
|
||||
)
|
||||
def test_exact_cache_honours_max_messages_from_cache_params(
|
||||
exact_proxy: _Proxy,
|
||||
route: str,
|
||||
turns: Callable[[str], tuple[list[JsonValue], list[JsonValue], list[JsonValue]]],
|
||||
) -> None:
|
||||
first, second, _ = turns(f"what is {uuid.uuid4().hex}")
|
||||
|
||||
answer: Final = exact_proxy.miss(route, first)
|
||||
repeated: Final = exact_proxy.hit(route, first)
|
||||
long_answers: Final = (exact_proxy.miss(route, second), exact_proxy.miss(route, second))
|
||||
|
||||
assert repeated == answer
|
||||
assert long_answers[0] != long_answers[1], "a turn past max_messages was served from a cache capped at 2"
|
||||
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue