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:
kerry 2026-09-30 16:43:47 +00:00
parent 2ed9761921
commit f462fbeb2e
10 changed files with 600 additions and 20 deletions

View file

@ -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;

View file

@ -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>,

View file

@ -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"))

View file

@ -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)

View file

@ -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)

View file

@ -11,6 +11,7 @@ OWNED_DIRECTORIES: Final = frozenset(
"providers",
"streaming",
"messages_endpoint",
"caching",
"configuration",
"mcp",
"observability",

View 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

View file

@ -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",),

View file

@ -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

View file

@ -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