mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
4ac923b6e7
commit
ffd3a85a18
2 changed files with 67 additions and 91 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue