mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(caching): tag each tool result with the position of the call it answers
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
516586008c
commit
19bc489cee
7 changed files with 322 additions and 54 deletions
95
litellm-rust/crates/cache/src/semantic.rs
vendored
95
litellm-rust/crates/cache/src/semantic.rs
vendored
|
|
@ -4,7 +4,7 @@
|
|||
//! `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::{future::Future, io};
|
||||
use std::{collections::HashMap, future::Future, io};
|
||||
|
||||
use serde::Serialize;
|
||||
use serde_json::{
|
||||
|
|
@ -84,11 +84,35 @@ impl Embedder for PreparedEmbedding {
|
|||
}
|
||||
|
||||
/// `get_str_from_messages_with_tools`: every message's content text, tool calls and tool results,
|
||||
/// then its OpenAI `tool_calls`, then its search results.
|
||||
/// then its OpenAI `tool_calls`, then its search results. Each tool result is tagged with the
|
||||
/// position of the call it answers.
|
||||
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.iter().filter_map(Value::as_object) {
|
||||
push_content_text(&mut text, message.get("content"));
|
||||
for message in messages {
|
||||
if message.get("role").and_then(Value::as_str) == Some("tool") {
|
||||
text.push_str(&tool_result_tag(
|
||||
message.get("tool_call_id"),
|
||||
&call_ordinals,
|
||||
));
|
||||
}
|
||||
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");
|
||||
|
|
@ -104,7 +128,11 @@ pub fn str_from_messages(messages: &[Value]) -> String {
|
|||
}
|
||||
|
||||
/// `_content_str_with_tools`: text parts, Anthropic `tool_use` blocks and `tool_result` content.
|
||||
fn push_content_text(text: &mut String, content: Option<&Value>) {
|
||||
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),
|
||||
Some(Value::Array(blocks)) => {
|
||||
|
|
@ -113,7 +141,10 @@ fn push_content_text(text: &mut String, content: Option<&Value>) {
|
|||
Some("tool_use") => {
|
||||
text.push_str(&tool_call_json(block.get("name"), block.get("input")));
|
||||
}
|
||||
Some("tool_result") => push_content_text(text, block.get("content")),
|
||||
Some("tool_result") => {
|
||||
text.push_str(&tool_result_tag(block.get("tool_use_id"), call_ordinals));
|
||||
push_content_text(text, block.get("content"), call_ordinals);
|
||||
}
|
||||
_ => {
|
||||
if let Some(block_text) = block.get("text").and_then(Value::as_str) {
|
||||
text.push_str(block_text);
|
||||
|
|
@ -126,6 +157,26 @@ fn push_content_text(text: &mut String, content: Option<&Value>) {
|
|||
}
|
||||
}
|
||||
|
||||
/// `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);
|
||||
}
|
||||
ordinals
|
||||
}
|
||||
|
||||
/// `tool_result_str`: `{"result_of_call":N}` for a result answering a known call, else empty.
|
||||
fn tool_result_tag(call_id: Option<&Value>, call_ordinals: &HashMap<&str, usize>) -> String {
|
||||
call_id
|
||||
.and_then(Value::as_str)
|
||||
.and_then(|call_id| call_ordinals.get(call_id))
|
||||
.map_or_else(String::new, |ordinal| {
|
||||
format!("{{\"result_of_call\":{ordinal}}}")
|
||||
})
|
||||
}
|
||||
|
||||
/// `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!({
|
||||
|
|
@ -149,8 +200,16 @@ 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);
|
||||
collect_input_text(input, &mut parts, &call_ordinals);
|
||||
let prompt = python_strip(&parts.join("\n")).to_owned();
|
||||
(!prompt.is_empty()).then_some(prompt)
|
||||
}
|
||||
|
|
@ -179,14 +238,18 @@ fn push_search_results_text(text: &mut String, search_results: Option<&Value>) {
|
|||
}
|
||||
}
|
||||
|
||||
fn collect_input_text(value: &Value, parts: &mut Vec<String>) {
|
||||
fn collect_input_text(
|
||||
value: &Value,
|
||||
parts: &mut Vec<String>,
|
||||
call_ordinals: &HashMap<&str, usize>,
|
||||
) {
|
||||
match value {
|
||||
Value::String(text) => {
|
||||
push_trimmed(text, parts);
|
||||
}
|
||||
Value::Array(items) => {
|
||||
for item in items {
|
||||
collect_input_text(item, parts);
|
||||
collect_input_text(item, parts, call_ordinals);
|
||||
}
|
||||
}
|
||||
Value::Object(map) => {
|
||||
|
|
@ -194,14 +257,24 @@ fn collect_input_text(value: &Value, parts: &mut Vec<String>) {
|
|||
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 tag = tool_result_tag(map.get("call_id"), call_ordinals);
|
||||
if !tag.is_empty() {
|
||||
parts.push(tag);
|
||||
if let Some(output) = map.get("output") {
|
||||
collect_input_text(output, parts, call_ordinals);
|
||||
}
|
||||
return;
|
||||
}
|
||||
}
|
||||
if let Some(content) = map.get("content").filter(|content| !content.is_null()) {
|
||||
collect_input_text(content, parts);
|
||||
collect_input_text(content, parts, call_ordinals);
|
||||
return;
|
||||
}
|
||||
for key in ["text", "output", "input_text", "output_text"] {
|
||||
match map.get(key) {
|
||||
Some(nested @ Value::Array(_)) => {
|
||||
collect_input_text(nested, parts);
|
||||
collect_input_text(nested, parts, call_ordinals);
|
||||
return;
|
||||
}
|
||||
Some(Value::String(text)) if push_trimmed(text, parts) => return,
|
||||
|
|
|
|||
54
litellm-rust/crates/cache/tests/semantic.rs
vendored
54
litellm-rust/crates/cache/tests/semantic.rs
vendored
|
|
@ -133,7 +133,40 @@ fn context(messages: Option<Value>, input: Option<Value>) -> SemanticCacheContex
|
|||
]},
|
||||
{"role": "tool", "tool_call_id": "c1", "content": "ok"},
|
||||
]),
|
||||
r#"writing{"name":"write","arguments":"{\"path\": \"a\"}"}{"name":"write","arguments":"{\"path\": \"b\"}"}ok"#,
|
||||
r#"writing{"name":"write","arguments":"{\"path\": \"a\"}"}{"name":"write","arguments":"{\"path\": \"b\"}"}{"result_of_call":1}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}B{"result_of_call":1}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}C"#,
|
||||
)]
|
||||
#[case::unknown_call_ids_left_untagged(
|
||||
json!([
|
||||
{"role": "tool", "tool_call_id": "c9", "content": "ok"},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "t9", "content": "done"}]},
|
||||
]),
|
||||
"okdone",
|
||||
)]
|
||||
#[case::malformed_tool_call_entries(
|
||||
json!([{"role": "assistant", "content": null, "tool_calls": [{"id": "c1", "type": "function"}, "junk"]}]),
|
||||
|
|
@ -238,7 +271,7 @@ fn prompt_from_messages_reads_messages_only(
|
|||
{"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\\\"}\"}\nok"),
|
||||
Some("update the config\n{\"name\":\"write_file\",\"arguments\":\"{\\\"path\\\":\\\"a.yaml\\\"}\"}\n{\"result_of_call\":1}\nok"),
|
||||
)]
|
||||
#[case::responses_structured_function_call_output(
|
||||
None,
|
||||
|
|
@ -246,7 +279,22 @@ fn prompt_from_messages_reads_messages_only(
|
|||
{"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\":\"{}\"}\ndenied"),
|
||||
Some("{\"name\":\"write_file\",\"arguments\":\"{}\"}\n{\"result_of_call\":1}\ndenied"),
|
||||
)]
|
||||
#[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}\nB\n{\"result_of_call\":1}\nA"),
|
||||
)]
|
||||
#[case::responses_unknown_call_id_output_untagged(
|
||||
None,
|
||||
Some(json!([{"type": "function_call_output", "call_id": "c9", "output": "orphan"}])),
|
||||
Some("orphan"),
|
||||
)]
|
||||
fn prompt_from_context_matches_python(
|
||||
#[case] messages: Option<Value>,
|
||||
|
|
|
|||
|
|
@ -22,7 +22,9 @@ from litellm.constants import SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS
|
|||
from litellm.litellm_core_utils.asyncify import asyncify
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_str_from_messages_with_tools,
|
||||
tool_call_ordinals,
|
||||
tool_call_str,
|
||||
tool_result_str,
|
||||
)
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
|
|
@ -269,14 +271,17 @@ class RedisSemanticCache(BaseCache):
|
|||
if "input" not in kwargs:
|
||||
return None
|
||||
|
||||
responses_input: Final = kwargs.get("input")
|
||||
prompt_parts: Final[list[str]] = []
|
||||
cls._collect_responses_input_text(kwargs.get("input"), prompt_parts)
|
||||
cls._collect_responses_input_text(responses_input, prompt_parts, cls._responses_call_ordinals(responses_input))
|
||||
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._function_call_as_prompt(cls._coerce_response_input_value(value))
|
||||
def _collect_responses_input_text(
|
||||
cls, value: object, prompt_parts: list[str], call_ordinals: Mapping[str, int]
|
||||
) -> None:
|
||||
value = cls._function_call_as_prompt(cls._coerce_response_input_value(value), call_ordinals)
|
||||
if value is None:
|
||||
return
|
||||
|
||||
|
|
@ -288,21 +293,21 @@ class RedisSemanticCache(BaseCache):
|
|||
|
||||
if isinstance(value, (list, tuple)):
|
||||
for item in value:
|
||||
cls._collect_responses_input_text(item, prompt_parts)
|
||||
cls._collect_responses_input_text(item, prompt_parts, call_ordinals)
|
||||
return
|
||||
|
||||
if isinstance(value, dict):
|
||||
content = value.get("content")
|
||||
if content is not None:
|
||||
cls._collect_responses_input_text(content, prompt_parts)
|
||||
cls._collect_responses_input_text(content, prompt_parts, call_ordinals)
|
||||
return
|
||||
|
||||
cls._collect_responses_text_fields(value, prompt_parts)
|
||||
cls._collect_responses_text_fields(value, prompt_parts, call_ordinals)
|
||||
return
|
||||
|
||||
content = getattr(value, "content", None)
|
||||
if content is not None:
|
||||
cls._collect_responses_input_text(content, prompt_parts)
|
||||
cls._collect_responses_input_text(content, prompt_parts, call_ordinals)
|
||||
return
|
||||
|
||||
for text_key in ("text", "output", "input_text", "output_text"):
|
||||
|
|
@ -314,21 +319,38 @@ class RedisSemanticCache(BaseCache):
|
|||
return
|
||||
|
||||
@classmethod
|
||||
def _collect_responses_text_fields(cls, value: dict, prompt_parts: list[str]) -> None:
|
||||
def _collect_responses_text_fields(
|
||||
cls, value: dict, prompt_parts: list[str], call_ordinals: Mapping[str, int]
|
||||
) -> None:
|
||||
for text_key in ("text", "output", "input_text", "output_text"):
|
||||
text_value = value.get(text_key)
|
||||
if isinstance(text_value, (list, tuple)):
|
||||
cls._collect_responses_input_text(text_value, prompt_parts)
|
||||
cls._collect_responses_input_text(text_value, prompt_parts, call_ordinals)
|
||||
return
|
||||
if isinstance(text_value, str) and (stripped_text := text_value.strip()):
|
||||
prompt_parts.append(stripped_text)
|
||||
return
|
||||
|
||||
@classmethod
|
||||
def _responses_call_ordinals(cls, responses_input: object) -> Mapping[str, int]:
|
||||
items: Final = responses_input if isinstance(responses_input, (list, tuple)) else ()
|
||||
dumped_items: Final = (cls._coerce_response_input_value(item) for item in items)
|
||||
return tool_call_ordinals(
|
||||
item.get("call_id")
|
||||
for item in dumped_items
|
||||
if isinstance(item, dict) and item.get("type") == "function_call"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _function_call_as_prompt(value: object) -> object:
|
||||
if isinstance(value, dict) and value.get("type") == "function_call":
|
||||
def _function_call_as_prompt(value: object, call_ordinals: Mapping[str, int]) -> object:
|
||||
if not isinstance(value, dict):
|
||||
return value
|
||||
if value.get("type") == "function_call":
|
||||
return tool_call_str(value.get("name"), value.get("arguments"))
|
||||
return value
|
||||
result_tag: Final = (
|
||||
tool_result_str(value.get("call_id"), call_ordinals) if value.get("type") == "function_call_output" else ""
|
||||
)
|
||||
return (result_tag, value.get("output")) if result_tag else value
|
||||
|
||||
@staticmethod
|
||||
def _coerce_response_input_value(value: object) -> object:
|
||||
|
|
|
|||
|
|
@ -199,42 +199,68 @@ def get_str_from_messages_with_tools(messages: object) -> str:
|
|||
"""
|
||||
``get_str_from_messages`` that also keeps each conversation's tool calls and tool results, so agent turns
|
||||
that differ only in their tool exchange (Anthropic ``tool_use`` / ``tool_result``, OpenAI ``tool_calls``)
|
||||
produce different text
|
||||
produce different text. Each result is tagged with the position of the call it answers, since call ids are
|
||||
random per session
|
||||
"""
|
||||
return "".join(_message_str_with_tools(message) for message in _str_mappings(messages))
|
||||
message_mappings: Final = tuple(_str_mappings(messages))
|
||||
call_ordinals: Final = tool_call_ordinals(_message_tool_call_ids(message_mappings))
|
||||
return "".join(_message_str_with_tools(message, call_ordinals) for message in message_mappings)
|
||||
|
||||
|
||||
def tool_call_str(name: object, arguments: object) -> str:
|
||||
return f'{{"name":{_compact_json(name)},"arguments":{_compact_json(arguments)}}}'
|
||||
|
||||
|
||||
def tool_result_str(call_id: object, call_ordinals: Mapping[str, int]) -> str:
|
||||
ordinal: Final = call_ordinals.get(call_id) if isinstance(call_id, str) else None
|
||||
return "" if ordinal is None else f'{{"result_of_call":{ordinal}}}'
|
||||
|
||||
|
||||
def tool_call_ordinals(call_ids: Iterable[object]) -> Mapping[str, int]:
|
||||
string_ids: Final = (call_id for call_id in call_ids if isinstance(call_id, str))
|
||||
return MappingProxyType({call_id: ordinal for ordinal, call_id in enumerate(dict.fromkeys(string_ids), start=1)})
|
||||
|
||||
|
||||
def _compact_json(value: object) -> str:
|
||||
return json.dumps(value, separators=(",", ":"), default=str)
|
||||
|
||||
|
||||
def _message_str_with_tools(message: Mapping[str, object]) -> str:
|
||||
def _message_tool_call_ids(messages: Iterable[Mapping[str, object]]) -> Iterator[object]:
|
||||
for message in messages:
|
||||
yield from (
|
||||
block.get("id") for block in _str_mappings(message.get("content")) if block.get("type") == "tool_use"
|
||||
)
|
||||
yield from (tool_call.get("id") for tool_call in _str_mappings(message.get("tool_calls")))
|
||||
|
||||
|
||||
def _message_str_with_tools(message: Mapping[str, object], call_ordinals: Mapping[str, int]) -> str:
|
||||
result_tag: Final = (
|
||||
tool_result_str(message.get("tool_call_id"), call_ordinals) if message.get("role") == "tool" else ""
|
||||
)
|
||||
return (
|
||||
_content_str_with_tools(message.get("content"))
|
||||
result_tag
|
||||
+ _content_str_with_tools(message.get("content"), call_ordinals)
|
||||
+ "".join(_openai_tool_call_str(tool_call) for tool_call in _str_mappings(message.get("tool_calls")))
|
||||
+ extract_search_results_text(message.get("search_results"))
|
||||
)
|
||||
|
||||
|
||||
def _content_str_with_tools(content: object) -> str:
|
||||
def _content_str_with_tools(content: object, call_ordinals: Mapping[str, int]) -> str:
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
return "".join(_block_str_with_tools(block) for block in _str_mappings(content))
|
||||
return "".join(_block_str_with_tools(block, call_ordinals) for block in _str_mappings(content))
|
||||
|
||||
|
||||
def _block_str_with_tools(block: Mapping[str, object]) -> str:
|
||||
match block.get("type"):
|
||||
case "tool_use":
|
||||
return tool_call_str(block.get("name"), block.get("input"))
|
||||
case "tool_result":
|
||||
return _content_str_with_tools(block.get("content"))
|
||||
case _:
|
||||
text: Final = block.get("text")
|
||||
return text if isinstance(text, str) else ""
|
||||
def _block_str_with_tools(block: Mapping[str, object], call_ordinals: Mapping[str, int]) -> str:
|
||||
block_type: Final = block.get("type")
|
||||
if block_type == "tool_use":
|
||||
return tool_call_str(block.get("name"), block.get("input"))
|
||||
if block_type == "tool_result":
|
||||
return tool_result_str(block.get("tool_use_id"), call_ordinals) + _content_str_with_tools(
|
||||
block.get("content"), call_ordinals
|
||||
)
|
||||
text: Final = block.get("text")
|
||||
return text if isinstance(text, str) else ""
|
||||
|
||||
|
||||
def _openai_tool_call_str(tool_call: Mapping[str, object]) -> str:
|
||||
|
|
|
|||
|
|
@ -241,10 +241,13 @@ def test_claude_code_tool_turns_on_messages_are_not_served_the_first_turn_answer
|
|||
},
|
||||
]
|
||||
|
||||
answers: Final = _send_turns(semantic_proxy, "/v1/messages", _CLAUDE_MODEL, ([task], list_files, read_file))
|
||||
answers: Final = _send_turns(
|
||||
semantic_proxy, "/v1/messages", _CLAUDE_MODEL, ([task], list_files, read_file, list_files)
|
||||
)
|
||||
|
||||
assert len(set(answers)) == 3, f"a later tool turn replayed an earlier cached answer: {answers}"
|
||||
assert answers == semantic_proxy.peer.answered[-3:], semantic_proxy.peer.answered
|
||||
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(
|
||||
|
|
@ -270,8 +273,9 @@ def test_openai_agent_loop_tool_calls_on_chat_completions_are_not_served_a_cache
|
|||
]
|
||||
|
||||
answers: Final = _send_turns(
|
||||
semantic_proxy, "/v1/chat/completions", _CHAT_MODEL, (wrote("a.yaml"), wrote("b.yaml"))
|
||||
semantic_proxy, "/v1/chat/completions", _CHAT_MODEL, (wrote("a.yaml"), wrote("b.yaml"), wrote("a.yaml"))
|
||||
)
|
||||
|
||||
assert len(set(answers)) == 2, f"a different tool call replayed an earlier cached answer: {answers}"
|
||||
assert answers == semantic_proxy.peer.answered[-2:], semantic_proxy.peer.answered
|
||||
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}"
|
||||
|
|
|
|||
|
|
@ -593,10 +593,10 @@ def test_redis_semantic_cache_prompt_extraction_keeps_tool_turns_distinct():
|
|||
]
|
||||
|
||||
assert RedisSemanticCache._get_prompt_from_kwargs(messages=turn("ls")) == (
|
||||
'fix the failing test{"name":"Bash","arguments":{"cmd":"ls"}}ok'
|
||||
'fix the failing test{"name":"Bash","arguments":{"cmd":"ls"}}{"result_of_call":1}ok'
|
||||
)
|
||||
assert RedisSemanticCache._get_prompt_from_kwargs(messages=turn("pwd")) == (
|
||||
'fix the failing test{"name":"Bash","arguments":{"cmd":"pwd"}}ok'
|
||||
'fix the failing test{"name":"Bash","arguments":{"cmd":"pwd"}}{"result_of_call":1}ok'
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -611,7 +611,9 @@ def test_redis_semantic_cache_prompt_extraction_keeps_responses_function_calls()
|
|||
]
|
||||
)
|
||||
|
||||
assert prompt == 'update the config\n{"name":"write_file","arguments":"{\\"path\\":\\"a.yaml\\"}"}\nok'
|
||||
assert prompt == (
|
||||
'update the config\n{"name":"write_file","arguments":"{\\"path\\":\\"a.yaml\\"}"}\n{"result_of_call":1}\nok'
|
||||
)
|
||||
|
||||
|
||||
def test_redis_semantic_cache_prompt_extraction_keeps_structured_function_call_outputs():
|
||||
|
|
@ -631,10 +633,30 @@ def test_redis_semantic_cache_prompt_extraction_keeps_structured_function_call_o
|
|||
)
|
||||
|
||||
expected_call = '{"name":"write_file","arguments":"{\\"path\\":\\"a.txt\\"}"}'
|
||||
assert prompt_for("wrote 5 bytes") == f"write hello\n{expected_call}\nwrote 5 bytes"
|
||||
assert prompt_for("wrote 5 bytes") == f'write hello\n{expected_call}\n{{"result_of_call":1}}\nwrote 5 bytes'
|
||||
assert prompt_for("wrote 5 bytes") != prompt_for("PermissionError")
|
||||
|
||||
|
||||
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}\nempty\n{"result_of_call":1}\nsecret'
|
||||
)
|
||||
assert prompt_for("c1", "c2") != prompt_for("c2", "c1")
|
||||
|
||||
|
||||
def test_redis_semantic_cache_prompt_extraction_handles_model_objects():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
|
|
|
|||
|
|
@ -7,8 +7,6 @@ from typing import Final
|
|||
|
||||
import pytest
|
||||
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
ENCRYPTED_REASONING_SIGNATURE_PREFIX,
|
||||
TOOL_RESULT_IMAGE_BOUNDARY,
|
||||
|
|
@ -32,6 +30,7 @@ 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}$'
|
||||
|
||||
|
|
@ -2063,9 +2062,51 @@ _TASK: Final = {"role": "user", "content": "fix the failing test"}
|
|||
{"role": "tool", "tool_call_id": "c1", "content": "ok"},
|
||||
],
|
||||
'fix the failing testwriting{"name":"write","arguments":"{\\"path\\": \\"a\\"}"}'
|
||||
'{"name":"write","arguments":"{\\"path\\": \\"b\\"}"}ok',
|
||||
'{"name":"write","arguments":"{\\"path\\": \\"b\\"}"}{"result_of_call":1}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}B{"result_of_call":1}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}C',
|
||||
id="reused-call-ids-keep-first-position",
|
||||
),
|
||||
pytest.param(
|
||||
[
|
||||
Message(
|
||||
|
|
@ -2120,3 +2161,35 @@ def test_get_str_from_messages_with_tools_keeps_tool_exchange(messages: list[obj
|
|||
)
|
||||
def test_get_str_from_messages_with_tools_matches_get_str_from_messages_without_tools(messages: list[object]) -> None:
|
||||
assert get_str_from_messages_with_tools(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_str_from_messages_with_tools_tells_apart_parallel_results_answering_different_calls() -> None:
|
||||
assert get_str_from_messages_with_tools(_parallel_reads("empty", "secret")) != get_str_from_messages_with_tools(
|
||||
_results_in_swapped_order("secret", "empty")
|
||||
)
|
||||
|
||||
|
||||
def test_get_str_from_messages_with_tools_ignores_call_ids_that_differ_between_sessions() -> None:
|
||||
assert get_str_from_messages_with_tools(
|
||||
_parallel_reads("A", "B", call_id_prefix="toolu_")
|
||||
) == get_str_from_messages_with_tools(_parallel_reads("A", "B"))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue