diff --git a/litellm-rust/crates/cache/src/semantic.rs b/litellm-rust/crates/cache/src/semantic.rs index 5af6a54d894..26274fd4324 100644 --- a/litellm-rust/crates/cache/src/semantic.rs +++ b/litellm-rust/crates/cache/src/semantic.rs @@ -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 { (!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 { if let Some(messages) = context.messages.as_ref().and_then(Value::as_array) && !messages.is_empty() diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index dfbdcce8ae8..342dc77f88d 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -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( diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index f191286144f..59c582e181e 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -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)