mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
1c4a69bfbe
commit
4ac923b6e7
3 changed files with 67 additions and 149 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue