From 99fdb4856dc2a78268c8468050b68d90ca93803c Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 29 Sep 2026 10:35:48 +0000 Subject: [PATCH] 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> --- .../otel/mappers/openinference.py | 57 ++++++++++++++----- .../otel/test_otel_v2_vendor_mappers.py | 33 +++++++++++ 2 files changed, 75 insertions(+), 15 deletions(-) diff --git a/litellm/integrations/otel/mappers/openinference.py b/litellm/integrations/otel/mappers/openinference.py index bcb2933eafd..971c9a33664 100644 --- a/litellm/integrations/otel/mappers/openinference.py +++ b/litellm/integrations/otel/mappers/openinference.py @@ -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. diff --git a/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py b/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py index 26902fbb841..5169e6b9d5f 100644 --- a/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py +++ b/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py @@ -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)