Skip to content
Merged
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
27 changes: 20 additions & 7 deletions livekit-agents/livekit/agents/telemetry/gen_ai.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
import contextvars
import json
import os
from collections.abc import Iterable, Sequence
from collections.abc import Callable, Iterable, Sequence
from typing import TYPE_CHECKING, Any, TypeAlias

from opentelemetry import trace
Expand Down Expand Up @@ -51,26 +51,39 @@ def capture_content_enabled() -> bool:
# paths have no nested `llm_request` span to carry the convention's attributes, so the node
# span records them instead. LLMStream marks the context when it does create one, which is
# what tells the two cases apart.
_inference_recorded: contextvars.ContextVar[list[bool] | None] = contextvars.ContextVar(
"lk_inference_recorded", default=None
_on_inference_span_created: contextvars.ContextVar[Callable[[], None] | None] = (
contextvars.ContextVar("lk_inference_recorded", default=None)
)


def track_inference_span() -> list[bool]:
def track_inference_span(*, model: str | None = None, provider: str | None = None) -> list[bool]:
"""Start tracking, returning a marker that fills in if an ``llm_request`` span is created.

Record the node's request identity before a fallback can replace the provider with
the instance that served the request.

No reset: the caller runs as its own asyncio task, so the context copy — and this
variable with it — is discarded when that task finishes.
"""
span = trace.get_current_span()
recorded: list[bool] = []
_inference_recorded.set(recorded)

def on_created() -> None:
if not recorded and span.is_recording():
if model:
span.set_attribute(trace_types.ATTR_GEN_AI_REQUEST_MODEL, model)
if (normalized := trace_types.gen_ai_provider_name(provider)) is not None:
span.set_attribute(trace_types.ATTR_GEN_AI_PROVIDER_NAME, normalized)
recorded.append(True)

_on_inference_span_created.set(on_created)
return recorded


def mark_inference_span_recorded() -> None:
"""Called where an ``llm_request`` span is created, so the enclosing node stands down."""
if (recorded := _inference_recorded.get()) is not None:
recorded.append(True)
if (on_created := _on_inference_span_created.get()) is not None:
on_created()


def _text_part(content: str) -> dict[str, Any]:
Expand Down
19 changes: 4 additions & 15 deletions livekit-agents/livekit/agents/voice/generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -214,7 +214,7 @@ async def _llm_inference_task(
# provider call the convention describes — setting them here as well would make a
# backend summing gen_ai.usage.* report twice the calls and tokens. A custom node that
# never builds an LLMStream has no such span, and records them here instead.
inference_recorded = gen_ai_telemetry.track_inference_span()
inference_recorded = gen_ai_telemetry.track_inference_span(model=model, provider=provider)

llm_node = node(chat_ctx, tools, model_settings)
if asyncio.iscoroutine(llm_node):
Expand All @@ -238,8 +238,6 @@ async def _llm_inference_task(
chat_ctx,
tools,
data,
model,
provider,
streaming=False,
)
return True
Expand Down Expand Up @@ -319,7 +317,7 @@ async def _llm_inference_task(
except BaseException as exc:
# a node that raises still made a request; without this it leaves no inference span
_record_uninstrumented_inference(
current_span, inference_recorded, chat_ctx, tools, data, model, provider, error=exc
current_span, inference_recorded, chat_ctx, tools, data, error=exc
)
raise
finally:
Expand All @@ -343,7 +341,7 @@ async def _llm_inference_task(
if data.ttft is not None:
current_span.set_attribute(trace_types.ATTR_RESPONSE_TTFT, data.ttft)
_record_uninstrumented_inference(
current_span, inference_recorded, chat_ctx, tools, data, model, provider, usage=usage
current_span, inference_recorded, chat_ctx, tools, data, usage=usage
)
return True

Expand All @@ -354,8 +352,6 @@ def _record_uninstrumented_inference(
chat_ctx: ChatContext,
tools: list[llm.Tool],
data: _LLMGenerationData,
model: str | None,
provider: str | None,
*,
usage: CompletionUsage | None = None,
streaming: bool = True,
Expand All @@ -368,16 +364,9 @@ def _record_uninstrumented_inference(
nested ``llm_request`` span to carry the convention's attributes. When one was created,
this stands down so the counts are not reported twice.

The configured model and provider are only reported when that LLM served the request.
Reaching here means it did not, so a third-party engine is left unattributed rather
than credited to the model the agent happens to be configured with.
Custom nodes without an LLMStream are left unattributed to the configured model.
"""
if inference_recorded:
# the configured LLM served this, so its identity describes the call
if model:
span.set_attribute(trace_types.ATTR_GEN_AI_REQUEST_MODEL, model)
if (normalized := trace_types.gen_ai_provider_name(provider)) is not None:
span.set_attribute(trace_types.ATTR_GEN_AI_PROVIDER_NAME, normalized)
return

gen_ai_telemetry.set_request_attributes(
Expand Down
41 changes: 34 additions & 7 deletions tests/test_coverage_spans.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,8 @@
from livekit.agents.llm import ChatContext, FallbackAdapter, LLMStream, Tool
from livekit.agents.telemetry import set_tracer_provider, trace_types, tracer
from livekit.agents.types import DEFAULT_API_CONNECT_OPTIONS, APIConnectOptions
from livekit.agents.voice import generation
from livekit.agents.voice.io import ModelSettings
from livekit.agents.voice.transcription.synchronizer import _SyncedAudioOutput

from .fake_io import FakeAudioInput
Expand Down Expand Up @@ -203,8 +205,10 @@ def chat(
)


@pytest.mark.parametrize("through_node", [False, True], ids=["stream", "llm_node"])
async def test_llm_fallback_records_failed_and_serving_provider(
span_exporter: InMemorySpanExporter,
through_node: bool,
) -> None:
primary = _FailingLLM()
secondary = FakeLLM(
Expand All @@ -214,18 +218,36 @@ async def test_llm_fallback_records_failed_and_serving_provider(
chat_ctx = ChatContext()
chat_ctx.add_message(role="user", content="hi")
try:
# the span the request is made under (llm_node in the pipeline, open until the
# stream is consumed) is told who served
with tracer.start_as_current_span("caller") as caller:
stream = adapter.chat(chat_ctx=chat_ctx)
chunks = [chunk async for chunk in stream]
await stream.aclose()
if through_node:

def node(
chat_ctx: ChatContext, tools: list[Tool], model_settings: ModelSettings
) -> LLMStream:
return adapter.chat(chat_ctx=chat_ctx, tools=tools)

task, data = generation.perform_llm_inference(
node=node,
chat_ctx=chat_ctx,
tool_ctx=llm.ToolContext([]),
model_settings=ModelSettings(),
model=adapter.model,
provider=adapter.provider,
)
assert await task is True
response = data.generated_text
[caller] = _spans(span_exporter, "llm_node")
else:
with tracer.start_as_current_span("caller") as caller:
stream = adapter.chat(chat_ctx=chat_ctx)
chunks = [chunk async for chunk in stream]
await stream.aclose()
response = "".join(c.delta.content or "" for c in chunks if c.delta)
finally:
await adapter.aclose()
await primary.aclose()
await secondary.aclose()

assert "".join(c.delta.content or "" for c in chunks if c.delta) == "hello"
assert response == "hello"

# the adapter's request span nests the attempt span that ran the fallback loop
[request] = _spans(span_exporter, "llm_fallback_adapter")
Expand All @@ -250,6 +272,11 @@ async def test_llm_fallback_records_failed_and_serving_provider(
assert isinstance(caller, ReadableSpan)
caller_attrs = caller.attributes or {}
assert caller_attrs[trace_types.ATTR_GEN_AI_RESPONSE_MODEL] == secondary.model
assert caller_attrs[trace_types.ATTR_GEN_AI_PROVIDER_NAME] == trace_types.gen_ai_provider_name(
secondary.provider
)
if through_node:
assert caller_attrs[trace_types.ATTR_GEN_AI_REQUEST_MODEL] == primary.model
# and the adapter itself now reports who serves next
assert adapter.model == secondary.model and adapter.provider == secondary.provider

Expand Down
54 changes: 52 additions & 2 deletions tests/test_llm_telemetry.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
from __future__ import annotations

import json
from collections.abc import Iterator
from collections.abc import AsyncIterator, Iterator
from typing import Any

import pytest
Expand All @@ -10,7 +10,7 @@
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
from opentelemetry.sdk.trace.sampling import ALWAYS_OFF

from livekit.agents import llm
from livekit.agents import APIConnectionError, llm
from livekit.agents.telemetry import gen_ai, set_tracer_provider, trace_types, tracer
from livekit.agents.types import (
DEFAULT_API_CONNECT_OPTIONS,
Expand Down Expand Up @@ -249,6 +249,52 @@ async def test_llm_stream_skips_content_builders_for_nonrecording_span(
assert response.usage.prompt_tokens == 100


@pytest.mark.parametrize("fails", [False, True], ids=["success", "failure"])
async def test_llm_node_preserves_model_identity(
span_exporter: InMemorySpanExporter,
monkeypatch: pytest.MonkeyPatch,
fails: bool,
) -> None:
if fails:

async def fail(stream: _UsageLLMStream) -> None:
raise APIConnectionError("provider unavailable")

monkeypatch.setattr(_UsageLLMStream, "_run", fail)

async with _UsageLLM() as model:

async def node(
chat_ctx: llm.ChatContext, tools: list[llm.Tool], model_settings: ModelSettings
) -> AsyncIterator[llm.ChatChunk]:
with tracer.start_as_current_span("custom_llm_wrapper"):
async with model.chat(
chat_ctx=chat_ctx, tools=tools, conn_options=APIConnectOptions(max_retry=0)
) as stream:
async for chunk in stream:
yield chunk

task, _ = generation.perform_llm_inference(
node=node,
chat_ctx=llm.ChatContext.empty(),
tool_ctx=llm.ToolContext([]),
model_settings=ModelSettings(),
model=model.model,
provider=model.provider,
)
if fails:
with pytest.raises(APIConnectionError):
await task
else:
assert await task is True

[span] = [span for span in span_exporter.get_finished_spans() if span.name == "llm_node"]
assert span.attributes[trace_types.ATTR_GEN_AI_REQUEST_MODEL] == model.model
assert span.attributes[trace_types.ATTR_GEN_AI_PROVIDER_NAME] == model.provider
assert trace_types.ATTR_GEN_AI_OPERATION_NAME not in span.attributes
assert trace_types.ATTR_GEN_AI_USAGE_INPUT_TOKENS not in span.attributes


async def test_llm_node_skips_payloads_for_nonrecording_span(
nonrecording_tracer_provider: None,
monkeypatch: pytest.MonkeyPatch,
Expand Down Expand Up @@ -283,6 +329,8 @@ async def test_llm_node_preserves_noncontent_attributes_when_capture_is_disabled
chat_ctx=llm.ChatContext.empty(),
tool_ctx=llm.ToolContext([]),
model_settings=ModelSettings(),
model="unused-model",
provider="unused-provider",
)
assert await task is True
finally:
Expand All @@ -294,3 +342,5 @@ async def test_llm_node_preserves_noncontent_attributes_when_capture_is_disabled
assert spans[0].attributes[trace_types.ATTR_GEN_AI_OPERATION_NAME] == "chat"
assert trace_types.ATTR_GEN_AI_INPUT_MESSAGES not in spans[0].attributes
assert trace_types.ATTR_GEN_AI_OUTPUT_MESSAGES not in spans[0].attributes
assert trace_types.ATTR_GEN_AI_REQUEST_MODEL not in spans[0].attributes
assert trace_types.ATTR_GEN_AI_PROVIDER_NAME not in spans[0].attributes
Loading