From 4ac923b6e7bd7516a03d21187080b221482bee16 Mon Sep 17 00:00:00 2001 From: kerry Date: Tue, 6 Oct 2026 18:06:38 +0000 Subject: [PATCH] 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> --- .../prompt_templates/common_utils.py | 108 +++++++----------- .../code_coverage_tests/recursive_detector.py | 3 +- .../unit/caching/test_redis_semantic_cache.py | 105 ++++------------- 3 files changed, 67 insertions(+), 149 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 683e02e9e8f..ae611c44664 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -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: diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 4a024f6f7cb..457f9fb6234 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -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., diff --git a/tests/unit/caching/test_redis_semantic_cache.py b/tests/unit/caching/test_redis_semantic_cache.py index 3359796ec84..07ec324699b 100644 --- a/tests/unit/caching/test_redis_semantic_cache.py +++ b/tests/unit/caching/test_redis_semantic_cache.py @@ -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):