mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(otel): shed OpenInference output tool calls individually under the span attribute budget
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
f13bf491c8
commit
99fdb4856d
2 changed files with 75 additions and 15 deletions
|
|
@ -33,6 +33,7 @@ from litellm.integrations.otel.model.payloads import (
|
|||
_INPUT_MESSAGES: Final = "llm.input_messages"
|
||||
_OUTPUT_MESSAGES: Final = "llm.output_messages"
|
||||
_MESSAGE_FAMILIES: Final = (_INPUT_MESSAGES, _OUTPUT_MESSAGES)
|
||||
_MESSAGE_BASE: Final = -1
|
||||
|
||||
_ParsedMessage = tuple[object, str | None, tuple[MessageToolCall, ...]]
|
||||
|
||||
|
|
@ -77,35 +78,61 @@ def _message_value(messages: Sequence[_ParsedMessage]) -> str:
|
|||
)
|
||||
|
||||
|
||||
def _message_key_groups(attrs: Mapping[str, AttrValue]) -> Mapping[tuple[str, int], tuple[str, ...]]:
|
||||
"""Per-index message keys in ``attrs`` grouped by ``(family, index)``."""
|
||||
tagged: Final = sorted(
|
||||
(family, int(key.split(".")[2]), key)
|
||||
for key in attrs
|
||||
for family in _MESSAGE_FAMILIES
|
||||
if key.startswith(f"{family}.")
|
||||
def _message_key_group(key: str) -> tuple[str, int, int, str] | None:
|
||||
family: Final = next(
|
||||
(family for family in _MESSAGE_FAMILIES if key.startswith(f"{family}.")),
|
||||
None,
|
||||
)
|
||||
if family is None:
|
||||
return None
|
||||
parts: Final = key.split(".")
|
||||
message_idx: Final = int(parts[2])
|
||||
tool_idx: Final = int(parts[5]) if parts[4] == "tool_calls" else _MESSAGE_BASE
|
||||
return family, message_idx, tool_idx, key
|
||||
|
||||
|
||||
def _message_key_groups(attrs: Mapping[str, AttrValue]) -> Mapping[tuple[str, int, int], tuple[str, ...]]:
|
||||
"""Message and tool-call keys in ``attrs`` grouped by family, message index, and tool index."""
|
||||
tagged: Final = tuple(tag for key in attrs if (tag := _message_key_group(key)) is not None)
|
||||
return MappingProxyType(
|
||||
{group: tuple(key for _, _, key in keys) for group, keys in groupby(tagged, key=lambda tag: tag[:2])}
|
||||
{group: tuple(key for _, _, _, key in keys) for group, keys in groupby(sorted(tagged), key=lambda tag: tag[:3])}
|
||||
)
|
||||
|
||||
|
||||
def _shed_order(groups: Mapping[tuple[str, int], tuple[str, ...]]) -> tuple[tuple[str, int], ...]:
|
||||
"""Message groups least valuable first: middle prompt turns, extra choices, then the opener, the newest turn
|
||||
and the first choice."""
|
||||
inputs: Final = sorted(idx for family, idx in groups if family == _INPUT_MESSAGES)
|
||||
outputs: Final = sorted(idx for family, idx in groups if family == _OUTPUT_MESSAGES)
|
||||
def _message_shed_groups(
|
||||
groups: Mapping[tuple[str, int, int], tuple[str, ...]], family: str, message_idx: int
|
||||
) -> Iterator[tuple[str, int, int]]:
|
||||
tool_call_groups: Final = tuple(
|
||||
sorted(
|
||||
(group for group in groups if group[:2] == (family, message_idx) and group[2] != _MESSAGE_BASE),
|
||||
key=lambda group: group[2],
|
||||
reverse=True,
|
||||
)
|
||||
)
|
||||
yield from tool_call_groups
|
||||
base_group: Final = (family, message_idx, _MESSAGE_BASE)
|
||||
if base_group in groups:
|
||||
yield base_group
|
||||
|
||||
|
||||
def _shed_order(groups: Mapping[tuple[str, int, int], tuple[str, ...]]) -> tuple[tuple[str, int, int], ...]:
|
||||
"""Middle inputs, extra choices, pinned inputs, then the first choice, with tool calls before message keys."""
|
||||
inputs: Final = sorted(frozenset(idx for family, idx, _ in groups if family == _INPUT_MESSAGES))
|
||||
outputs: Final = sorted(frozenset(idx for family, idx, _ in groups if family == _OUTPUT_MESSAGES))
|
||||
pinned_inputs: Final = tuple(dict.fromkeys((*inputs[:1], *inputs[-1:])))
|
||||
return (
|
||||
message_order: Final = (
|
||||
*((_INPUT_MESSAGES, idx) for idx in inputs[1:-1]),
|
||||
*((_OUTPUT_MESSAGES, idx) for idx in reversed(outputs[1:])),
|
||||
*((_INPUT_MESSAGES, idx) for idx in pinned_inputs),
|
||||
*((_OUTPUT_MESSAGES, idx) for idx in outputs[:1]),
|
||||
)
|
||||
return tuple(
|
||||
chain.from_iterable(_message_shed_groups(groups, family, message_idx) for family, message_idx in message_order)
|
||||
)
|
||||
|
||||
|
||||
def fit_indexed_messages(attrs: Mapping[str, AttrValue], budget: int | None) -> Mapping[str, AttrValue]:
|
||||
"""``attrs`` with whole per-index messages shed, least valuable first, until at most ``budget`` keys remain.
|
||||
"""``attrs`` with indexed message attributes shed, least valuable first, until at most ``budget`` keys remain.
|
||||
|
||||
``None`` means the span has no attribute count limit. Every message still rides the ``input.value`` and
|
||||
``output.value`` blobs, so shedding a per-index pair loses no content.
|
||||
|
|
|
|||
|
|
@ -300,6 +300,39 @@ def test_openinference_output_tool_calls_do_not_shed_input_roles_under_budget():
|
|||
}
|
||||
|
||||
|
||||
def test_openinference_budget_sheds_trailing_output_tool_calls_before_the_message():
|
||||
tool_calls: Final = tuple(
|
||||
{
|
||||
"id": f"call_{index}",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "lookup_weather",
|
||||
"arguments": f'{{"city": "C{index}"}}',
|
||||
},
|
||||
}
|
||||
for index in range(60)
|
||||
)
|
||||
data: Final = _llm_call(
|
||||
choices_out=(
|
||||
{
|
||||
"finish_reason": "tool_calls",
|
||||
"message": {"role": "assistant", "content": None, "tool_calls": tool_calls},
|
||||
},
|
||||
)
|
||||
)
|
||||
mapped: Final = OpenInferenceMapper().map(data)
|
||||
attrs: Final = fit_indexed_messages(mapped, len(mapped) - 34)
|
||||
retained_tool_call_keys: Final = tuple(key for key in attrs if ".tool_calls." in key)
|
||||
retained_tool_call_indices: Final = frozenset(int(key.split(".")[5]) for key in retained_tool_call_keys)
|
||||
|
||||
assert retained_tool_call_indices == frozenset(range(50))
|
||||
assert len(retained_tool_call_keys) == 150
|
||||
assert attrs["llm.output_messages.0.message.role"] == "assistant"
|
||||
assert not any(key.startswith("llm.input_messages.0.") for key in attrs)
|
||||
assert len(attrs) == len(mapped) - 34
|
||||
assert json.loads(attrs["output.value"]) == [{"role": "assistant", "content": None, "tool_calls": list(tool_calls)}]
|
||||
|
||||
|
||||
def test_openinference_plain_output_messages_keep_the_existing_value_shape():
|
||||
attrs: Final = OpenInferenceMapper().map(_llm_call())
|
||||
assert all(".tool_calls." not in key for key in attrs)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue