refactor(caching): split semantic cache prompt extraction by API format

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
kerry 2026-10-06 00:41:13 +00:00
parent b9a794291a
commit b27722ce58
3 changed files with 143 additions and 169 deletions

View file

@ -1,8 +1,9 @@
//! The embedding and prompt contract every semantic backend shares.
//!
//! Python's semantic caches all read their prompt through `get_semantic_cache_prompt_from_messages`, and
//! `RedisSemanticCache._get_prompt_from_kwargs` (inherited by Valkey) adds Responses API
//! `input`. Qdrant reads messages only. Each backend picks one of the two extractors here.
//! `RedisSemanticCache._get_prompt_from_kwargs` (inherited by Valkey) adds Responses API `input`
//! through `get_semantic_cache_prompt_from_responses_input`. Qdrant reads messages only. Each
//! backend picks one of the two extractors here.
use std::{collections::HashMap, future::Future, io};
@ -207,8 +208,9 @@ pub fn prompt_from_messages(context: &SemanticCacheContext) -> Option<String> {
(!messages.is_empty()).then(|| str_from_messages(messages))
}
/// `RedisSemanticCache._get_prompt_from_kwargs`: chat messages first, then the text parts of a
/// Responses API `input`. `None` when neither yields a prompt.
/// `RedisSemanticCache._get_prompt_from_kwargs`: chat messages first, then
/// `get_semantic_cache_prompt_from_responses_input` over a Responses API `input`. `None` when
/// neither yields a prompt.
pub fn prompt_from_context(context: &SemanticCacheContext) -> Option<String> {
if let Some(messages) = context.messages.as_ref().and_then(Value::as_array)
&& !messages.is_empty()

View file

@ -22,9 +22,7 @@ from litellm.constants import SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_semantic_cache_prompt_from_messages,
tool_call_ordinals,
tool_call_str,
tool_result_str,
get_semantic_cache_prompt_from_responses_input,
)
from litellm.types.utils import EmbeddingResponse
@ -271,101 +269,7 @@ class RedisSemanticCache(BaseCache):
if "input" not in kwargs:
return None
responses_input: Final = kwargs.get("input")
prompt: Final = cls._responses_input_prompt(responses_input, cls._responses_call_ordinals(responses_input))
return prompt or None
@classmethod
def _responses_input_prompt(cls, value: object, call_ordinals: Mapping[str, int]) -> str:
prompt_parts: Final[list[str]] = []
cls._collect_responses_input_text(value, prompt_parts, call_ordinals)
return "\n".join(prompt_parts).strip()
@classmethod
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
if isinstance(value, str):
stripped_value: Final = value.strip()
if stripped_value:
prompt_parts.append(stripped_value)
return
if isinstance(value, (list, tuple)):
for item in value:
cls._collect_responses_input_text(item, prompt_parts, call_ordinals)
return
if isinstance(value, dict):
content = value.get("content")
if content is not None:
cls._collect_responses_input_text(content, prompt_parts, call_ordinals)
return
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, call_ordinals)
return
for text_key in ("text", "output", "input_text", "output_text"):
text_value = getattr(value, text_key, None)
if isinstance(text_value, str):
stripped_text = text_value.strip()
if stripped_text:
prompt_parts.append(stripped_text)
return
@classmethod
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, 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"
)
@classmethod
def _function_call_as_prompt(cls, 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"))
if value.get("type") != "function_call_output":
return value
return tool_result_str(
value.get("call_id"), call_ordinals, cls._responses_input_prompt(value.get("output"), call_ordinals)
)
@staticmethod
def _coerce_response_input_value(value: object) -> object:
model_dump: Final = getattr(value, "model_dump", None)
if callable(model_dump):
return model_dump()
dict_method: Final = getattr(value, "dict", None)
if callable(dict_method):
return dict_method()
return value
return get_semantic_cache_prompt_from_responses_input(kwargs.get("input")) or None
def _embedding_input(self, prompt: str, router: "Router | None") -> str:
return truncate_embedding_input(

View file

@ -13,7 +13,6 @@ 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
@ -197,94 +196,163 @@ def get_str_from_messages(messages: list[AllMessageValues]) -> str:
def get_semantic_cache_prompt_from_messages(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. Each result is encoded with the position of the call it answers, since call ids are
random per session
The text the semantic cache embeds for ``messages``, shared by the Messages API and Chat Completions. Keeps
the text ``get_str_from_messages`` keeps, plus every tool call and tool result, so agent turns that differ
only in their tool exchange embed differently. Call ids are random per session, so each result names the
position of the call it answers instead
"""
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)
message_dicts: Final = _dumped_dicts(messages)
positions: Final = _tool_call_positions(_messages_tool_call_ids(message_dicts))
return "".join(_message_text(message, positions) for message in message_dicts)
def tool_call_str(name: object, arguments: object) -> str:
return f'{{"name":{_compact_json(name)},"arguments":{_compact_json(arguments)}}}'
def get_semantic_cache_prompt_from_responses_input(responses_input: object) -> str:
"""
The text the semantic cache embeds for a Responses API ``input``: each text part stripped and on its own
line, with ``function_call`` and ``function_call_output`` items encoded like tool calls and results above
"""
items: 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))
def tool_result_str(call_id: object, call_ordinals: Mapping[str, int], output: str) -> str:
ordinal: Final = call_ordinals.get(call_id) if isinstance(call_id, str) else None
return f'{{"result_of_call":{_compact_json(ordinal)},"output":{_compact_json(output)}}}'
def _message_text(message: Mapping[str, object], positions: Mapping[str, int]) -> str:
content_text: Final = _content_text(message.get("content"), positions)
return _chat_completions_message_text(message, content_text, positions) + extract_search_results_text(
message.get("search_results")
)
def tool_call_ordinals(call_ids: Iterable[object]) -> 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))
def _messages_api_block_text(block: Mapping[str, object], positions: Mapping[str, int]) -> str:
if block.get("type") == "tool_use":
return _tool_call_text(block.get("name"), block.get("input"))
if block.get("type") == "tool_result":
return _tool_result_text(block.get("tool_use_id"), _content_text(block.get("content"), positions), positions)
text: Final = block.get("text")
return text if isinstance(text, str) else ""
def _chat_completions_message_text(
message: Mapping[str, object], content_text: str, positions: Mapping[str, int]
) -> str:
result_or_content: Final = (
_tool_result_text(message.get("tool_call_id"), content_text, positions)
if message.get("role") == "tool"
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"))
)
def _chat_completions_tool_call_text(tool_call: Mapping[str, object]) -> str:
function: Final = _dumped(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"))
def _messages_tool_call_ids(messages: Iterable[Mapping[str, object]]) -> Iterator[object]:
for message in messages:
for block in _dumped_dicts(message.get("content")):
if block.get("type") == "tool_use":
yield block.get("id")
for tool_call in _dumped_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()
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":
yield _tool_call_text(item.get("name"), item.get("arguments"))
return
if isinstance(item, Mapping) and 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]:
for key in ("text", "output", "input_text", "output_text"):
text: Final = _field(item, key)
if isinstance(text, (list, tuple)):
yield from _responses_text_parts(text, positions)
return
if isinstance(text, str) and text.strip():
yield text.strip()
return
def _tool_call_text(name: object, arguments: object) -> str:
return _compact_json({"name": name, "arguments": arguments})
def _tool_result_text(call_id: object, output: str, positions: Mapping[str, int]) -> str:
position: Final = positions.get(call_id) if isinstance(call_id, str) else None
return _compact_json({"result_of_call": position, "output": output})
def _tool_call_positions(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)})
return MappingProxyType({call_id: position for position, 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_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 _dumped_dicts(values: object) -> tuple[Mapping[str, object], ...]:
return tuple(item for item in map(_dumped, _sequence(values)) if _is_str_mapping(item))
def _message_str_with_tools(message: Mapping[str, object], call_ordinals: Mapping[str, int]) -> str:
content: Final = _content_str_with_tools(message.get("content"), call_ordinals)
return (
(
tool_result_str(message.get("tool_call_id"), call_ordinals, content)
if message.get("role") == "tool"
else 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 _sequence(values: object) -> Sequence[object]:
return values if isinstance(values, (list, tuple)) else ()
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, call_ordinals) for block in _str_mappings(content))
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()
return value
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 _field(item: object, key: str) -> object:
if isinstance(item, Mapping):
return item.get(key)
return getattr(item, key, None)
def _openai_tool_call_str(tool_call: Mapping[str, object]) -> str:
function: Final = _as_str_mapping(tool_call.get("function"))
if function is None:
return tool_call_str(None, None)
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
def _is_str_mapping(value: object) -> TypeIs[Mapping[str, object]]:
return isinstance(value, Mapping)