From ffd3a85a1811b6e7f8420e06b306bc4b2aa436ba Mon Sep 17 00:00:00 2001 From: kerry Date: Tue, 6 Oct 2026 18:18:12 +0000 Subject: [PATCH] refactor(caching): write the semantic cache prompt builders as plain loops Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../prompt_templates/common_utils.py | 155 ++++++++---------- .../code_coverage_tests/recursive_detector.py | 3 +- 2 files changed, 67 insertions(+), 91 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index ae611c44664..6f849966902 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -201,64 +201,24 @@ 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 = _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) + message_dicts = _dicts(_plain(messages)) + positions = _tool_call_positions(_messages_tool_call_ids(message_dicts)) + parts = [] + for message in message_dicts: + content = message.get("content") + text = content if isinstance(content, str) else _messages_api_blocks_text(content, positions) + if message.get("role") == "tool": + text = _tool_result_text(message.get("tool_call_id"), text, positions) + for tool_call in _dicts(message.get("tool_calls")): + function = tool_call.get("function") + if not isinstance(function, Mapping): + function = {} + text += _tool_call_text(function.get("name"), function.get("arguments")) + parts.append(text + extract_search_results_text(message.get("search_results"))) + return "".join(parts) -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 = _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: - 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 _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 _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 _dicts(message.get("tool_calls")) - ) - - -def _chat_completions_tool_call_text(tool_call: Mapping[str, object]) -> str: - 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")) - - -def _messages_tool_call_ids(messages: Iterable[Mapping[str, object]]) -> Iterator[object]: +def _messages_tool_call_ids(messages: Sequence[Mapping[str, object]]) -> Iterator[object]: for message in messages: for block in _dicts(message.get("content")): if block.get("type") == "tool_use": @@ -267,48 +227,63 @@ def _messages_tool_call_ids(messages: Iterable[Mapping[str, object]]) -> Iterato yield tool_call.get("id") +def _messages_api_blocks_text(blocks: object, positions: Mapping[str, int]) -> str: + text = "" + for block in _dicts(blocks): + if block.get("type") == "tool_use": + text += _tool_call_text(block.get("name"), block.get("input")) + elif block.get("type") == "tool_result": + content = block.get("content") + output = content if isinstance(content, str) else _messages_api_blocks_text(content, positions) + text += _tool_result_text(block.get("tool_use_id"), output, positions) + elif isinstance(block.get("text"), str): + text += str(block["text"]) + return text + + +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 = _plain(responses_input) + call_ids = [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 _responses_text(value: object, positions: Mapping[str, int]) -> str: - return "\n".join(_responses_text_lines(value, positions)).strip() - - -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")) - elif item.get("type") == "function_call_output": - yield _tool_result_text(item.get("call_id"), _responses_text(item.get("output"), positions), positions) - 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) + return value.strip() + if isinstance(value, list): + lines = [_responses_text(item, positions) for item in value] + return "\n".join(line for line in lines if line) + if not isinstance(value, Mapping): + return "" + if value.get("type") == "function_call": + return _tool_call_text(value.get("name"), value.get("arguments")) + if value.get("type") == "function_call_output": + output = _responses_text(value.get("output"), positions) + return _tool_result_text(value.get("call_id"), output, positions) + if value.get("content") is not None: + return _responses_text(value.get("content"), positions) + return _responses_text(value.get("text"), positions) def _tool_call_text(name: object, arguments: object) -> str: - return _compact_json({"name": name, "arguments": arguments}) + return json.dumps({"name": name, "arguments": arguments}, separators=(",", ":"), default=str) 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}) + position = positions.get(call_id) if isinstance(call_id, str) else None + return json.dumps({"result_of_call": position, "output": output}, separators=(",", ":"), default=str) 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: 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) + positions = {} + for call_id in call_ids: + if isinstance(call_id, str) and call_id not in positions: + positions[call_id] = len(positions) + 1 + return positions def _plain(value: object) -> object: @@ -322,10 +297,10 @@ def _plain(value: object) -> object: return value -def _dicts(values: object) -> tuple[Mapping[str, object], ...]: +def _dicts(values: object) -> list[Mapping[str, object]]: if not isinstance(values, list): - return () - return tuple(value for value in values if isinstance(value, Mapping)) + return [] + return [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 457f9fb6234..21af1a65192 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_lines", # walks only the nesting a Responses `input` carries. + "_responses_text", # walks only the nesting a Responses `input` carries. + "_messages_api_blocks_text", # walks only the nesting a `tool_result` block 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.