refactor(caching): walk semantic cache prompts as plain dicts, dumping pydantic items once up front

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
kerry 2026-10-06 18:06:38 +00:00
parent 1c4a69bfbe
commit 4ac923b6e7
3 changed files with 67 additions and 149 deletions

View file

@ -13,6 +13,8 @@ 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
@ -199,7 +201,7 @@ def get_semantic_cache_prompt_from_messages(messages: object) -> str:
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
"""
message_dicts: Final = _dumped_dicts(messages)
message_dicts: Final = _dicts(_plain(messages))
positions: Final = _tool_call_positions(_messages_tool_call_ids(message_dicts))
return "".join(_message_text(message, positions) for message in message_dicts)
@ -209,11 +211,9 @@ def get_semantic_cache_prompt_from_responses_input(responses_input: object) -> s
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: Final = tuple(_dumped(item) for item in _sequence(responses_input))
call_ids: Final = (
item.get("call_id") for item in items if isinstance(item, Mapping) and item.get("type") == "function_call"
)
return _responses_text(responses_input, _tool_call_positions(call_ids))
items: Final = _plain(responses_input)
call_ids: Final = (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 _message_text(message: Mapping[str, object], positions: Mapping[str, int]) -> str:
@ -226,7 +226,7 @@ def _message_text(message: Mapping[str, object], positions: Mapping[str, int]) -
def _content_text(content: object, positions: Mapping[str, int]) -> str:
if isinstance(content, str):
return content
return "".join(_messages_api_block_text(block, positions) for block in _dumped_dicts(content))
return "".join(_messages_api_block_text(block, positions) for block in _dicts(content))
def _messages_api_block_text(block: Mapping[str, object], positions: Mapping[str, int]) -> str:
@ -247,12 +247,12 @@ def _chat_completions_message_text(
else content_text
)
return result_or_content + "".join(
_chat_completions_tool_call_text(tool_call) for tool_call in _dumped_dicts(message.get("tool_calls"))
_chat_completions_tool_call_text(tool_call) for tool_call in _dicts(message.get("tool_calls"))
)
def _chat_completions_tool_call_text(tool_call: Mapping[str, object]) -> str:
function: Final = _dumped(tool_call.get("function"))
function: Final = tool_call.get("function")
if not isinstance(function, Mapping):
return _tool_call_text(None, None)
return _tool_call_text(function.get("name"), function.get("arguments"))
@ -260,53 +260,37 @@ def _chat_completions_tool_call_text(tool_call: Mapping[str, object]) -> str:
def _messages_tool_call_ids(messages: Iterable[Mapping[str, object]]) -> Iterator[object]:
for message in messages:
for block in _dumped_dicts(message.get("content")):
for block in _dicts(message.get("content")):
if block.get("type") == "tool_use":
yield block.get("id")
for tool_call in _dumped_dicts(message.get("tool_calls")):
for tool_call in _dicts(message.get("tool_calls")):
yield tool_call.get("id")
def _responses_text(value: object, positions: Mapping[str, int]) -> str:
return "\n".join(_responses_text_parts(value, positions)).strip()
return "\n".join(_responses_text_lines(value, positions)).strip()
def _responses_text_parts(value: object, positions: Mapping[str, int]) -> Iterator[str]:
item: Final = _dumped(value)
if item is None:
return
if isinstance(item, str):
if item.strip():
yield item.strip()
return
if isinstance(item, (list, tuple)):
for nested in item:
yield from _responses_text_parts(nested, positions)
return
if isinstance(item, Mapping) and item.get("type") == "function_call":
def _responses_text_lines(value: object, positions: Mapping[str, int]) -> Iterator[str]:
if isinstance(value, str):
if value.strip():
yield value.strip()
elif isinstance(value, list):
for item in value:
yield from _responses_text_lines(item, positions)
elif isinstance(value, Mapping):
yield from _responses_item_lines(value, positions)
def _responses_item_lines(item: Mapping[str, object], positions: Mapping[str, int]) -> Iterator[str]:
if item.get("type") == "function_call":
yield _tool_call_text(item.get("name"), item.get("arguments"))
return
if isinstance(item, Mapping) and item.get("type") == "function_call_output":
elif item.get("type") == "function_call_output":
yield _tool_result_text(item.get("call_id"), _responses_text(item.get("output"), positions), positions)
return
content: Final = _field(item, "content")
if content is not None:
yield from _responses_text_parts(content, positions)
return
yield from _responses_first_text_field(item, positions)
def _responses_first_text_field(item: object, positions: Mapping[str, int]) -> Iterator[str]:
fields: Final = (_field(item, key) for key in ("text", "output", "input_text", "output_text"))
text: Final = next((field for field in fields if _is_text_field(field)), None)
if isinstance(text, (list, tuple)):
yield from _responses_text_parts(text, positions)
elif isinstance(text, str):
yield text.strip()
def _is_text_field(field: object) -> bool:
return isinstance(field, (list, tuple)) or (isinstance(field, str) and bool(field.strip()))
elif item.get("content") is not None:
yield from _responses_text_lines(item.get("content"), positions)
else:
yield from _responses_text_lines(item.get("text"), positions)
def _tool_call_text(name: object, arguments: object) -> str:
@ -327,29 +311,21 @@ def _compact_json(value: object) -> str:
return json.dumps(value, separators=(",", ":"), default=str)
def _dumped_dicts(values: object) -> tuple[Mapping[str, object], ...]:
return tuple(item for item in map(_dumped, _sequence(values)) if isinstance(item, Mapping))
def _sequence(values: object) -> Sequence[object]:
return values if isinstance(values, (list, tuple)) else ()
def _dumped(value: object) -> object:
"""pydantic models, and anything else that dumps itself, as the dict the request carried"""
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()
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 _field(item: object, key: str) -> object:
if isinstance(item, Mapping):
return item.get(key)
return getattr(item, key, None)
def _dicts(values: object) -> tuple[Mapping[str, object], ...]:
if not isinstance(values, list):
return ()
return tuple(value for value in values if isinstance(value, Mapping))
def is_non_content_values_set(message: AllMessageValues) -> bool:

View file

@ -23,7 +23,8 @@ IGNORE_FUNCTIONS = [
"_can_object_call_model", # max depth set.
"encode_unserializable_types", # max depth set.
"filter_value_from_dict", # max depth set.
"_responses_text_parts", # walks only the nesting a Responses `input` carries; same walk the RedisSemanticCache classmethod did before it moved here.
"_responses_text_lines", # walks only the nesting a Responses `input` 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.,

View file

@ -203,10 +203,7 @@ def test_redis_semantic_cache_uses_isolated_index_for_old_schema(monkeypatch):
assert redis_semantic_cache.llmcache is fallback_cache_mock
assert semantic_cache_mock.call_args_list[0].kwargs["name"] == "existing_index"
assert (
semantic_cache_mock.call_args_list[1].kwargs["name"]
== "existing_index_isolated"
)
assert semantic_cache_mock.call_args_list[1].kwargs["name"] == "existing_index_isolated"
assert semantic_cache_mock.call_args_list[1].kwargs["filterable_fields"] == [
RedisSemanticCache._cache_key_filterable_field()
]
@ -236,10 +233,7 @@ def test_redis_semantic_cache_overwrites_stale_isolated_index(monkeypatch):
)
assert redis_semantic_cache.llmcache is fallback_cache_mock
assert (
semantic_cache_mock.call_args_list[2].kwargs["name"]
== "existing_index_isolated"
)
assert semantic_cache_mock.call_args_list[2].kwargs["name"] == "existing_index_isolated"
assert semantic_cache_mock.call_args_list[2].kwargs["overwrite"] is True
assert semantic_cache_mock.call_args_list[2].kwargs["filterable_fields"] == [
RedisSemanticCache._cache_key_filterable_field()
@ -363,9 +357,7 @@ async def test_redis_semantic_cache_async_get_cache(monkeypatch):
]
redis_semantic_cache.llmcache.acheck = AsyncMock(return_value=mock_result)
redis_semantic_cache._get_async_embedding = AsyncMock(
return_value=[0.1, 0.2, 0.3]
)
redis_semantic_cache._get_async_embedding = AsyncMock(return_value=[0.1, 0.2, 0.3])
with patch.object(
redis_semantic_cache,
@ -413,9 +405,7 @@ async def test_redis_semantic_cache_async_get_cache_rejects_unscoped_hit(monkeyp
}
]
)
redis_semantic_cache._get_async_embedding = AsyncMock(
return_value=[0.1, 0.2, 0.3]
)
redis_semantic_cache._get_async_embedding = AsyncMock(return_value=[0.1, 0.2, 0.3])
with patch.object(
redis_semantic_cache,
@ -447,9 +437,7 @@ async def test_redis_semantic_cache_async_set_cache_stores_cache_key_filter(
redis_semantic_cache = RedisSemanticCache(similarity_threshold=0.8)
redis_semantic_cache.llmcache.astore = AsyncMock()
redis_semantic_cache._get_async_embedding = AsyncMock(
return_value=[0.1, 0.2, 0.3]
)
redis_semantic_cache._get_async_embedding = AsyncMock(return_value=[0.1, 0.2, 0.3])
await redis_semantic_cache.async_set_cache(
key="test_key",
@ -676,27 +664,25 @@ def test_redis_semantic_cache_prompt_extraction_tells_apart_parallel_outputs_ans
assert prompt_for("c1", "c2") != prompt_for("c2", "c1")
def test_redis_semantic_cache_prompt_extraction_handles_model_objects():
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
class ModelDumpInput:
def model_dump(self):
return {"content": [{"text": "model dump prompt"}]}
class DictInput:
def dict(self):
return {"content": [{"output_text": "dict prompt"}]}
prompt = RedisSemanticCache._get_prompt_from_kwargs(
input=[
ModelDumpInput(),
DictInput(),
{"content": [{"input_text": "inline prompt"}]},
{"content": [{"type": "input_image", "image_url": "https://example.com"}]},
{"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 == "model dump prompt\ndict prompt\ninline prompt"
assert (
prompt
== 'write hello\n{"name":"write_file","arguments":"{\\"path\\":\\"a.txt\\"}"}\n{"result_of_call":1,"output":"ok"}'
)
def test_redis_semantic_cache_prompt_extraction_returns_none_without_text():
@ -706,46 +692,11 @@ def test_redis_semantic_cache_prompt_extraction_returns_none_without_text():
assert RedisSemanticCache._get_prompt_from_kwargs(input=None) is None
assert RedisSemanticCache._get_prompt_from_kwargs(input=" ") is None
assert (
RedisSemanticCache._get_prompt_from_kwargs(
input=[{"type": "input_image", "image_url": "https://example.com"}]
)
RedisSemanticCache._get_prompt_from_kwargs(input=[{"type": "input_image", "image_url": "https://example.com"}])
is None
)
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
@ -964,9 +915,7 @@ def test_redis_get_embedding_falls_back_to_direct(monkeypatch):
fake_proxy.llm_model_list = None
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy)
with patch(
"litellm.embedding", return_value={"data": [{"embedding": [0.1, 0.2]}]}
) as direct_embed:
with patch("litellm.embedding", return_value={"data": [{"embedding": [0.1, 0.2]}]}) as direct_embed:
vec = cache._get_embedding("hello")
assert vec == [0.1, 0.2]
@ -1102,9 +1051,7 @@ def test_redis_sync_set_cache_passes_precomputed_vector():
cache = RedisSemanticCache.__new__(RedisSemanticCache)
cache.llmcache = MagicMock()
cache._get_cache_filters = MagicMock(
return_value={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"}
)
cache._get_cache_filters = MagicMock(return_value={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"})
cache._get_ttl = MagicMock(return_value=None)
cache._get_embedding = MagicMock(return_value=[0.1, 0.2, 0.3])
@ -1141,9 +1088,7 @@ def test_redis_sync_get_cache_passes_precomputed_vector():
)
cache._get_embedding = MagicMock(return_value=[0.1, 0.2, 0.3])
with patch.object(
cache, "_get_cache_key_filter_expression", return_value="cache-key-filter"
):
with patch.object(cache, "_get_cache_key_filter_expression", return_value="cache-key-filter"):
result = cache.get_cache(
key="test_key",
messages=[{"content": "What is the capital of France?"}],
@ -1257,9 +1202,7 @@ def test_redis_get_embedding_truncates_direct_path_with_explicit_limit(monkeypat
fake_proxy.llm_model_list = None
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy)
with patch(
"litellm.embedding", return_value={"data": [{"embedding": [0.1, 0.2]}]}
) as direct_embed:
with patch("litellm.embedding", return_value={"data": [{"embedding": [0.1, 0.2]}]}) as direct_embed:
cache._get_embedding(LONG_PROMPT)
sent_input = direct_embed.call_args.kwargs["input"]
@ -1306,9 +1249,7 @@ def test_redis_init_defers_redisvl_construction(monkeypatch):
def test_redis_failed_llmcache_build_is_not_memoized(monkeypatch):
built_cache = MagicMock()
semantic_cache_mock = MagicMock(
side_effect=[ConnectionError("redis down"), built_cache]
)
semantic_cache_mock = MagicMock(side_effect=[ConnectionError("redis down"), built_cache])
custom_vectorizer_mock = MagicMock()
with _fake_redisvl_modules(semantic_cache_mock, custom_vectorizer_mock):