mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(caching): keep tool calls and tool results in semantic cache prompts
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
2ed9761921
commit
f462fbeb2e
10 changed files with 600 additions and 20 deletions
59
litellm-rust/crates/cache/src/semantic.rs
vendored
59
litellm-rust/crates/cache/src/semantic.rs
vendored
|
|
@ -1,6 +1,6 @@
|
|||
//! The embedding and prompt contract every semantic backend shares.
|
||||
//!
|
||||
//! Python's semantic caches all read their prompt through `get_str_from_messages`, and
|
||||
//! Python's semantic caches all read their prompt through `get_str_from_messages_with_tools`, 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.
|
||||
|
||||
|
|
@ -8,7 +8,7 @@ use std::{future::Future, io};
|
|||
|
||||
use serde::Serialize;
|
||||
use serde_json::{
|
||||
Value,
|
||||
Value, json,
|
||||
ser::{CharEscape, Formatter, Serializer},
|
||||
};
|
||||
|
||||
|
|
@ -83,26 +83,57 @@ impl Embedder for PreparedEmbedding {
|
|||
}
|
||||
}
|
||||
|
||||
/// `get_str_from_messages`: every message's text content followed by its search results.
|
||||
/// `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.
|
||||
pub fn str_from_messages(messages: &[Value]) -> String {
|
||||
let mut text = String::new();
|
||||
for message in messages.iter().filter_map(Value::as_object) {
|
||||
match message.get("content") {
|
||||
Some(Value::String(content)) => text.push_str(content),
|
||||
Some(Value::Array(parts)) => {
|
||||
for part in parts {
|
||||
if let Some(part_text) = part.get("text").and_then(Value::as_str) {
|
||||
text.push_str(part_text);
|
||||
}
|
||||
}
|
||||
push_content_text(&mut text, message.get("content"));
|
||||
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")),
|
||||
));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
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>) {
|
||||
match content {
|
||||
Some(Value::String(content)) => text.push_str(content),
|
||||
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") => push_content_text(text, block.get("content")),
|
||||
_ => {
|
||||
if let Some(block_text) = block.get("text").and_then(Value::as_str) {
|
||||
text.push_str(block_text);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
/// `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.
|
||||
pub fn prompt_from_messages(context: &SemanticCacheContext) -> Option<String> {
|
||||
let messages = context.messages.as_ref()?.as_array()?;
|
||||
|
|
@ -159,6 +190,10 @@ fn collect_input_text(value: &Value, parts: &mut Vec<String>) {
|
|||
}
|
||||
}
|
||||
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 let Some(content) = map.get("content").filter(|content| !content.is_null()) {
|
||||
collect_input_text(content, parts);
|
||||
return;
|
||||
|
|
|
|||
48
litellm-rust/crates/cache/tests/semantic.rs
vendored
48
litellm-rust/crates/cache/tests/semantic.rs
vendored
|
|
@ -100,6 +100,45 @@ fn context(messages: Option<Value>, input: Option<Value>) -> SemanticCacheContex
|
|||
json!([{"role": "tool", "search_results": [{"citations": false}, {"citations": 3}]}]),
|
||||
"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"},
|
||||
]}]),
|
||||
"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"},
|
||||
]}]),
|
||||
"abnext",
|
||||
)]
|
||||
#[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\"}"}ok"#,
|
||||
)]
|
||||
#[case::malformed_tool_call_entries(
|
||||
json!([{"role": "assistant", "content": null, "tool_calls": [{"id": "c1", "type": "function"}, "junk"]}]),
|
||||
r#"{"name":null,"arguments":null}"#,
|
||||
)]
|
||||
fn str_from_messages_matches_python(#[case] messages: Value, #[case] expected: &str) {
|
||||
assert_eq!(str_from_messages(messages.as_array().unwrap()), expected);
|
||||
}
|
||||
|
|
@ -192,6 +231,15 @@ 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\\\"}\"}\nok"),
|
||||
)]
|
||||
fn prompt_from_context_matches_python(
|
||||
#[case] messages: Option<Value>,
|
||||
#[case] input: Option<Value>,
|
||||
|
|
|
|||
|
|
@ -24,7 +24,7 @@ from litellm.constants import (
|
|||
)
|
||||
from litellm.litellm_core_utils.asyncify import asyncify
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_str_from_messages,
|
||||
get_str_from_messages_with_tools,
|
||||
)
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
|
|
@ -286,7 +286,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
|
||||
# get the prompt
|
||||
messages: Final = kwargs["messages"]
|
||||
prompt: Final = get_str_from_messages(messages)
|
||||
prompt: Final = get_str_from_messages_with_tools(messages)
|
||||
|
||||
# create an embedding for prompt
|
||||
embedding_response: Final = cast(
|
||||
|
|
@ -325,7 +325,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
|
||||
# get the messages
|
||||
messages: Final = kwargs["messages"]
|
||||
prompt: Final = get_str_from_messages(messages)
|
||||
prompt: Final = get_str_from_messages_with_tools(messages)
|
||||
|
||||
# convert to embedding
|
||||
embedding_response: Final = cast(
|
||||
|
|
@ -400,7 +400,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
|
||||
# get the prompt
|
||||
messages: Final = kwargs["messages"]
|
||||
prompt: Final = get_str_from_messages(messages)
|
||||
prompt: Final = get_str_from_messages_with_tools(messages)
|
||||
embedding_response: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
|
||||
# get the embedding
|
||||
|
|
@ -435,7 +435,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
|
||||
# get the messages
|
||||
messages: Final = kwargs["messages"]
|
||||
prompt: Final = get_str_from_messages(messages)
|
||||
prompt: Final = get_str_from_messages_with_tools(messages)
|
||||
|
||||
embedding_response: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
|
||||
|
|
|
|||
|
|
@ -21,7 +21,8 @@ from litellm._logging import print_verbose, verbose_logger
|
|||
from litellm.constants import SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS
|
||||
from litellm.litellm_core_utils.asyncify import asyncify
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_str_from_messages,
|
||||
get_str_from_messages_with_tools,
|
||||
tool_call_str,
|
||||
)
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
|
|
@ -263,7 +264,7 @@ class RedisSemanticCache(BaseCache):
|
|||
"""
|
||||
messages: Final = kwargs.get("messages")
|
||||
if messages:
|
||||
return get_str_from_messages(messages)
|
||||
return get_str_from_messages_with_tools(messages)
|
||||
|
||||
if "input" not in kwargs:
|
||||
return None
|
||||
|
|
@ -291,6 +292,10 @@ class RedisSemanticCache(BaseCache):
|
|||
return
|
||||
|
||||
if isinstance(value, dict):
|
||||
if value.get("type") == "function_call":
|
||||
prompt_parts.append(tool_call_str(value.get("name"), value.get("arguments")))
|
||||
return
|
||||
|
||||
content = value.get("content")
|
||||
if content is not None:
|
||||
cls._collect_responses_input_text(content, prompt_parts)
|
||||
|
|
|
|||
|
|
@ -13,6 +13,9 @@ from pathlib import Path
|
|||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, cast
|
||||
|
||||
from pydantic import BaseModel
|
||||
from typing_extensions import TypeIs # noqa: TID251 # narrows untyped message payloads without a runtime conversion
|
||||
|
||||
import litellm
|
||||
from litellm import verbose_logger
|
||||
from litellm.router_utils.batch_utils import InMemoryFile
|
||||
|
|
@ -192,6 +195,66 @@ def get_str_from_messages(messages: list[AllMessageValues]) -> str:
|
|||
return text
|
||||
|
||||
|
||||
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
|
||||
"""
|
||||
return "".join(_message_str_with_tools(message) for message in _str_mappings(messages))
|
||||
|
||||
|
||||
def tool_call_str(name: object, arguments: object) -> str:
|
||||
return json.dumps({"name": name, "arguments": arguments}, separators=(",", ":"), default=str)
|
||||
|
||||
|
||||
def _message_str_with_tools(message: Mapping[str, object]) -> str:
|
||||
return (
|
||||
_content_str_with_tools(message.get("content"))
|
||||
+ "".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:
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
return "".join(_block_str_with_tools(block) 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 _openai_tool_call_str(tool_call: Mapping[str, object]) -> str:
|
||||
function: Final = _as_str_mapping(tool_call.get("function")) or {}
|
||||
return tool_call_str(function.get("name"), function.get("arguments"))
|
||||
|
||||
|
||||
def _str_mappings(values: object) -> Iterator[Mapping[str, object]]:
|
||||
items: Final = values if isinstance(values, (list, tuple)) else ()
|
||||
return (mapping for item in items if (mapping := _as_str_mapping(item)) is not None)
|
||||
|
||||
|
||||
def _as_str_mapping(value: object) -> Mapping[str, object] | None:
|
||||
if isinstance(value, BaseModel):
|
||||
return value.model_dump()
|
||||
if _is_str_mapping(value):
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def _is_str_mapping(value: object) -> TypeIs[Mapping[str, object]]: # guard-ok: message and block keys are str
|
||||
return isinstance(value, Mapping)
|
||||
|
||||
|
||||
def is_non_content_values_set(message: AllMessageValues) -> bool:
|
||||
ignore_keys: Final = ["content", "role", "name"]
|
||||
return any(message.get(key, None) is not None for key in message if key not in ignore_keys)
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ OWNED_DIRECTORIES: Final = frozenset(
|
|||
"providers",
|
||||
"streaming",
|
||||
"messages_endpoint",
|
||||
"caching",
|
||||
"configuration",
|
||||
"mcp",
|
||||
"observability",
|
||||
|
|
|
|||
277
tests/integration/caching/test_semantic_cache_tool_turns.py
Normal file
277
tests/integration/caching/test_semantic_cache_tool_turns.py
Normal file
|
|
@ -0,0 +1,277 @@
|
|||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import 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 = "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
|
||||
|
||||
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 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))
|
||||
|
||||
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
|
||||
|
||||
|
||||
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"))
|
||||
)
|
||||
|
||||
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
|
||||
|
|
@ -14,7 +14,7 @@ GROUPS: Final = MappingProxyType(
|
|||
"management": ("management", "authorization", "configuration"),
|
||||
"accounting": ("pricing", "spend"),
|
||||
"database": ("database",),
|
||||
"providers": ("providers", "routing", "streaming", "messages_endpoint"),
|
||||
"providers": ("providers", "routing", "streaming", "messages_endpoint", "caching"),
|
||||
"extensions": ("observability", "compatibility"),
|
||||
"mcp": ("mcp",),
|
||||
"sdk": ("sdk",),
|
||||
|
|
|
|||
|
|
@ -579,6 +579,41 @@ def test_redis_semantic_cache_prompt_extraction_prefers_messages():
|
|||
assert prompt == "message prompt"
|
||||
|
||||
|
||||
def test_redis_semantic_cache_prompt_extraction_keeps_tool_turns_distinct():
|
||||
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"}]},
|
||||
]
|
||||
|
||||
assert RedisSemanticCache._get_prompt_from_kwargs(messages=turn("ls")) == (
|
||||
'fix the failing test{"name":"Bash","arguments":{"cmd":"ls"}}ok'
|
||||
)
|
||||
assert RedisSemanticCache._get_prompt_from_kwargs(messages=turn("pwd")) == (
|
||||
'fix the failing test{"name":"Bash","arguments":{"cmd":"pwd"}}ok'
|
||||
)
|
||||
|
||||
|
||||
def test_redis_semantic_cache_prompt_extraction_keeps_responses_function_calls():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
prompt = RedisSemanticCache._get_prompt_from_kwargs(
|
||||
input=[
|
||||
{"role": "user", "content": "update the config"},
|
||||
{"type": "function_call", "call_id": "c1", "name": "write_file", "arguments": '{"path":"a.yaml"}'},
|
||||
{"type": "function_call_output", "call_id": "c1", "output": "ok"},
|
||||
]
|
||||
)
|
||||
|
||||
assert prompt == 'update the config\n{"name":"write_file","arguments":"{\\"path\\":\\"a.yaml\\"}"}\nok'
|
||||
|
||||
|
||||
def test_redis_semantic_cache_prompt_extraction_handles_model_objects():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
|
|
|
|||
|
|
@ -7,6 +7,8 @@ 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,
|
||||
|
|
@ -16,6 +18,8 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
encrypted_reasoning_signature,
|
||||
get_file_ids_from_messages,
|
||||
get_format_from_file_id,
|
||||
get_str_from_messages,
|
||||
get_str_from_messages_with_tools,
|
||||
handle_any_messages_to_chat_completion_str_messages_conversion,
|
||||
hoist_images_from_tool_messages,
|
||||
is_encrypted_reasoning_block,
|
||||
|
|
@ -2004,3 +2008,115 @@ class TestMergeConsecutiveSystemMessages:
|
|||
)
|
||||
|
||||
assert merged == [{"role": "system"}, {"role": "user", "content": "Hi"}]
|
||||
|
||||
|
||||
_TASK: Final = {"role": "user", "content": "fix the failing test"}
|
||||
|
||||
|
||||
@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 testcalc.py",
|
||||
id="anthropic-string-tool-result",
|
||||
),
|
||||
pytest.param(
|
||||
[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "t1",
|
||||
"content": [{"type": "text", "text": "a"}, {"type": "text", "text": "b"}],
|
||||
},
|
||||
{"type": "text", "text": "next"},
|
||||
],
|
||||
}
|
||||
],
|
||||
"abnext",
|
||||
id="anthropic-nested-text-tool-result-then-text",
|
||||
),
|
||||
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\\"}"}ok',
|
||||
id="openai-tool-calls-in-order-before-tool-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",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_get_str_from_messages_with_tools_keeps_tool_exchange(messages: list[object], expected: str) -> None:
|
||||
assert get_str_from_messages_with_tools(messages) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"messages",
|
||||
[
|
||||
pytest.param([_TASK, {"role": "assistant", "content": "done"}], id="string-content"),
|
||||
pytest.param(
|
||||
[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "what is "},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.com/a.png"}},
|
||||
{"type": "text", "text": "this"},
|
||||
],
|
||||
}
|
||||
],
|
||||
id="text-and-image-parts",
|
||||
),
|
||||
pytest.param(
|
||||
[
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "c1",
|
||||
"content": "small",
|
||||
"search_results": [{"source": "s", "title": "t"}],
|
||||
}
|
||||
],
|
||||
id="search-results",
|
||||
),
|
||||
pytest.param([{"role": "assistant", "content": None}, {"role": "user"}], id="missing-content"),
|
||||
],
|
||||
)
|
||||
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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue