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>
This commit is contained in:
kerry 2026-10-06 18:18:12 +00:00
parent 4ac923b6e7
commit ffd3a85a18
2 changed files with 67 additions and 91 deletions

View file

@ -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:

View file

@ -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.