Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
Tuple,
TypeVar,
Union,
cast,
)

import wrapt
Expand Down Expand Up @@ -1076,10 +1077,32 @@ async def aclose(self) -> None:
await wrapped_aclose()


def _set_streaming_output_attributes(
span: trace_api.Span,
output_messages: Dict[int, Dict[str, Any]],
usage_stats: Any,
) -> str:
if reasoning_items := output_messages.get(0, {}).get("reasoning_items"):
_remove_redundant_reasoning_entries(output_messages, reasoning_items)
aggregated_output = cast(str, output_messages.get(0, {}).get("content", ""))
_set_span_attribute(span, SpanAttributes.OUTPUT_VALUE, aggregated_output)
if finish_reason := output_messages.get(0, {}).get("finish_reason"):
_set_span_attribute(span, SpanAttributes.LLM_FINISH_REASON, finish_reason)
for idx, msg in output_messages.items():
message = _build_message_from_accumulated(msg)
for key, value in _get_attributes_from_message_param(message):
_set_span_attribute(span, f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.{idx}.{key}", value)

if usage_stats:
_set_token_counts_from_usage(span, SimpleNamespace(usage=usage_stats))

return aggregated_output


def _finalize_sync_streaming_span(span: trace_api.Span, stream: Any) -> Any:
output_messages: Dict[int, Dict[str, Any]] = {}
usage_stats = None
aggregated_output = None
aggregated_output: Optional[str] = None
try:
for token in stream:
if token.choices:
Expand Down Expand Up @@ -1117,33 +1140,25 @@ def _finalize_sync_streaming_span(span: trace_api.Span, stream: Any) -> Any:
if usage_attrs:
usage_stats = usage_attrs
yield token
if reasoning_items := output_messages.get(0, {}).get("reasoning_items"):
_remove_redundant_reasoning_entries(output_messages, reasoning_items)
aggregated_output = output_messages.get(0, {}).get("content", "")
_set_span_attribute(span, SpanAttributes.OUTPUT_VALUE, aggregated_output)
if finish_reason := output_messages.get(0, {}).get("finish_reason"):
_set_span_attribute(span, SpanAttributes.LLM_FINISH_REASON, finish_reason)
for idx, msg in output_messages.items():
message = _build_message_from_accumulated(msg)
for key, value in _get_attributes_from_message_param(message):
_set_span_attribute(
span, f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.{idx}.{key}", value
)

if usage_stats:
_set_token_counts_from_usage(span, SimpleNamespace(usage=usage_stats))
aggregated_output = _set_streaming_output_attributes(span, output_messages, usage_stats)
except Exception as e:
span.record_exception(e)
raise
else:
_set_span_status(span, aggregated_output)
finally:
if aggregated_output is None:
try:
_set_streaming_output_attributes(span, output_messages, usage_stats)
except Exception:
logger.exception("Failed to record partial streaming span output")
span.end()


async def _finalize_streaming_span(span: trace_api.Span, stream: Any) -> Any:
output_messages: Dict[int, Dict[str, Any]] = {}
usage_stats = None
aggregated_output: Optional[str] = None
try:
async for token in stream:
if token.choices:
Expand Down Expand Up @@ -1181,26 +1196,18 @@ async def _finalize_streaming_span(span: trace_api.Span, stream: Any) -> Any:
if usage_attrs:
usage_stats = usage_attrs
yield token
if reasoning_items := output_messages.get(0, {}).get("reasoning_items"):
_remove_redundant_reasoning_entries(output_messages, reasoning_items)
aggregated_output = output_messages.get(0, {}).get("content", "")
_set_span_attribute(span, SpanAttributes.OUTPUT_VALUE, aggregated_output)
if finish_reason := output_messages.get(0, {}).get("finish_reason"):
_set_span_attribute(span, SpanAttributes.LLM_FINISH_REASON, finish_reason)
for idx, msg in output_messages.items():
message = _build_message_from_accumulated(msg)
for key, value in _get_attributes_from_message_param(message):
_set_span_attribute(
span, f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.{idx}.{key}", value
)
if usage_stats:
_set_token_counts_from_usage(span, SimpleNamespace(usage=usage_stats))
aggregated_output = _set_streaming_output_attributes(span, output_messages, usage_stats)
except Exception as e:
span.record_exception(e)
raise
else:
_set_span_status(span, aggregated_output)
finally:
if aggregated_output is None:
try:
_set_streaming_output_attributes(span, output_messages, usage_stats)
except Exception:
logger.exception("Failed to record partial streaming span output")
span.end()


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
_instrument_func_type_completion,
_instrument_func_type_responses,
)
from openinference.semconv.trace import MessageAttributes, SpanAttributes


@pytest.fixture()
Expand Down Expand Up @@ -77,6 +78,69 @@ async def run() -> str:
assert spans[0].name == "acompletion"


def test_sync_streaming_early_close_records_partial_output(
in_memory_span_exporter: InMemorySpanExporter,
setup_litellm_instrumentation: Any,
) -> None:
in_memory_span_exporter.clear()

response = litellm.completion(
model="gpt-3.5-turbo",
messages=[{"content": "What's the capital of China?", "role": "user"}],
mock_response="The capital of China is Beijing",
stream=True,
)

partial_output = ""
for chunk in response:
if content := chunk.choices[0].delta.content:
partial_output += content
break
response.close()

assert partial_output
spans = in_memory_span_exporter.get_finished_spans()
assert len(spans) == 1
attributes = dict(spans[0].attributes or {})
assert attributes.get(SpanAttributes.OUTPUT_VALUE) == partial_output
assert (
attributes.get(
f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_CONTENT}"
)
== partial_output
)


async def test_async_streaming_early_close_records_partial_output(
in_memory_span_exporter: InMemorySpanExporter,
setup_litellm_instrumentation: Any,
) -> None:
in_memory_span_exporter.clear()

response = await litellm.acompletion(
model="gpt-3.5-turbo",
messages=[{"content": "What's the capital of China?", "role": "user"}],
mock_response="The capital of China is Beijing",
stream=True,
)

chunk = await response.__anext__()
partial_output = chunk.choices[0].delta.content
assert partial_output
await response.aclose()

spans = in_memory_span_exporter.get_finished_spans()
assert len(spans) == 1
attributes = dict(spans[0].attributes or {})
assert attributes.get(SpanAttributes.OUTPUT_VALUE) == partial_output
assert (
attributes.get(
f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_CONTENT}"
)
== partial_output
)


def test_responses_foreign_stream_type_passes_through(
in_memory_span_exporter: InMemorySpanExporter,
setup_litellm_instrumentation: Any,
Expand Down
Loading