mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
b9a794291a
commit
b27722ce58
3 changed files with 143 additions and 169 deletions
10
litellm-rust/crates/cache/src/semantic.rs
vendored
10
litellm-rust/crates/cache/src/semantic.rs
vendored
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue