项目文件夹

文件
2026-07-13 13:39:38 +08:00

723 行
27 KiB
Python

from __future__ import annotations
import asyncio
import json
import logging
import threading
import time
from collections.abc import Iterator
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any
import aiofiles
import aiohttp
import requests
from google.protobuf.json_format import MessageToDict
from opentelemetry import context as otel_context, metrics as metrics_api, trace as trace_api
from opentelemetry._logs import LogRecord as OTelLogRecord, get_logger_provider, set_logger_provider
from opentelemetry._logs.severity import SeverityNumber
from opentelemetry.exporter.otlp.proto.http import Compression
from opentelemetry.exporter.otlp.proto.http._log_exporter import OTLPLogExporter
from opentelemetry.exporter.otlp.proto.http.metric_exporter import OTLPMetricExporter
from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter
from opentelemetry.sdk import trace as trace_sdk
from opentelemetry.sdk._logs import (
LoggerProvider,
LoggingHandler,
LogRecordProcessor,
ReadWriteLogRecord,
)
from opentelemetry.sdk._logs.export import BatchLogRecordProcessor
from opentelemetry.sdk.metrics import (
Counter as SdkCounter,
Histogram as SdkHistogram,
MeterProvider as SdkMeterProvider,
ObservableCounter as SdkObservableCounter,
ObservableGauge as SdkObservableGauge,
ObservableUpDownCounter as SdkObservableUpDownCounter,
UpDownCounter as SdkUpDownCounter,
)
from opentelemetry.sdk.metrics.export import AggregationTemporality, PeriodicExportingMetricReader
from opentelemetry.sdk.resources import SERVICE_NAME, Resource
from opentelemetry.sdk.trace import SpanProcessor
from opentelemetry.sdk.trace.export import BatchSpanProcessor
from opentelemetry.trace import Span, Tracer
from opentelemetry.util._decorator import _agnosticcontextmanager
from opentelemetry.util.types import Attributes, AttributeValue
from livekit import api
from livekit.protocol import agent_pb, metrics as proto_metrics
from ..log import TRACE_LEVEL, logger
from . import trace_types
if TYPE_CHECKING:
from ..llm import ChatContext, ChatItem
from ..observability import Tagger
from ..voice.report import SessionReport
class _DynamicTracer(Tracer):
def __init__(self, instrumenting_module_name: str) -> None:
self._instrumenting_module_name = instrumenting_module_name
self._tracer_provider: trace_api.TracerProvider = trace_api.get_tracer_provider()
self._tracer = trace_api.get_tracer(instrumenting_module_name)
def set_provider(self, tracer_provider: trace_api.TracerProvider) -> None:
self._tracer_provider = tracer_provider
self._tracer = trace_api.get_tracer(
self._instrumenting_module_name,
tracer_provider=self._tracer_provider,
)
def start_span(self, *args: Any, **kwargs: Any) -> Span:
return self._tracer.start_span(*args, **kwargs)
@_agnosticcontextmanager
def start_as_current_span(self, *args: Any, **kwargs: Any) -> Iterator[Span]:
with self._tracer.start_as_current_span(*args, **kwargs) as span:
yield span
tracer: _DynamicTracer = _DynamicTracer("livekit-agents")
class _MetadataSpanProcessor(SpanProcessor):
def __init__(self, metadata: dict[str, AttributeValue]) -> None:
self._metadata = metadata
def on_start(self, span: Span, parent_context: otel_context.Context | None = None) -> None:
span.set_attributes(self._metadata)
class _MetadataLogProcessor(LogRecordProcessor):
def __init__(self, metadata: dict[str, AttributeValue]) -> None:
self._metadata = metadata
def on_emit(self, log_data: ReadWriteLogRecord) -> None:
if log_data.log_record.attributes:
log_data.log_record.attributes.update(self._metadata) # type: ignore
else:
log_data.log_record.attributes = self._metadata
if log_data.instrumentation_scope:
log_data.log_record.attributes.update( # type: ignore
{"logger.name": log_data.instrumentation_scope.name}
)
def shutdown(self) -> None:
pass
def force_flush(self, timeout_millis: int = 30000) -> bool:
return True
class _BufferingHandler(logging.Handler):
"""Buffers log records in memory for later replay through OTLP."""
def __init__(self) -> None:
super().__init__()
self.buffer: list[logging.LogRecord] = []
def emit(self, record: logging.LogRecord) -> None:
self.buffer.append(record)
class _TraceLevelLoggingHandler(LoggingHandler):
"""Custom LoggingHandler that properly maps TRACE_LEVEL to OTel TRACE severity.
The default OTel LoggingHandler maps any log level < 10 to UNSPECIFIED,
but we want TRACE_LEVEL (5) to map to TRACE for proper severity in exports.
"""
def _translate(self, record: logging.LogRecord) -> OTelLogRecord:
log_record = super()._translate(record)
# OTel's std_to_otel returns UNSPECIFIED for levels < 10
# Map our TRACE_LEVEL to OTel's TRACE
if record.levelno == TRACE_LEVEL:
log_record.severity_number = SeverityNumber.TRACE
return log_record
def set_tracer_provider(
tracer_provider: trace_api.TracerProvider, *, metadata: dict[str, AttributeValue] | None = None
) -> None:
"""Set the tracer provider for the livekit-agents.
Args:
tracer_provider (TracerProvider): The tracer provider to set.
metadata (dict[str, AttributeValue] | None, optional): Metadata to set on all spans. Defaults to None.
"""
if metadata and isinstance(tracer_provider, trace_sdk.TracerProvider):
tracer_provider.add_span_processor(_MetadataSpanProcessor(metadata))
tracer.set_provider(tracer_provider)
def _setup_cloud_tracer(
*,
room_id: str,
job_id: str,
observability_url: str,
enable_traces: bool = True,
enable_logs: bool = True,
) -> None:
token_ttl = timedelta(hours=6)
refresh_margin = timedelta(minutes=5)
class _AuthRefreshingSession(requests.Session):
def __init__(self, header_provider: _AuthHeaderProvider) -> None:
super().__init__()
self._header_provider = header_provider
def request(self, *args: Any, **kwargs: Any) -> requests.Response:
self.headers.update(self._header_provider())
return super().request(*args, **kwargs)
class _AuthHeaderProvider:
def __init__(self) -> None:
self._lock = threading.Lock()
self._auth_header = ""
self._expires_at = datetime.min.replace(tzinfo=timezone.utc)
self._refresh()
def _refresh(self) -> None:
access_token = (
api.AccessToken()
.with_observability_grants(api.ObservabilityGrants(write=True))
.with_ttl(token_ttl)
)
self._auth_header = f"Bearer {access_token.to_jwt()}"
self._expires_at = datetime.now(timezone.utc) + token_ttl
def __call__(self) -> dict[str, str]:
now = datetime.now(timezone.utc)
if now >= self._expires_at - refresh_margin:
with self._lock:
if now >= self._expires_at - refresh_margin:
self._refresh()
return {"Authorization": self._auth_header}
header_provider = _AuthHeaderProvider()
session = _AuthRefreshingSession(header_provider)
otlp_compression = Compression.Gzip
metadata: dict[str, AttributeValue] = {"room_id": room_id, "job_id": job_id}
resource = Resource.create(
{
SERVICE_NAME: "livekit-agents",
"room_id": room_id,
"job_id": job_id,
}
)
if enable_traces:
# Check if a tracer provider is not set and set one up
# below shows how the ProxyTracerProvider is returned when none have been setup
# https://github.com/open-telemetry/opentelemetry-python/blob/0018c0030bac9bdce4487fe5fcb3ec6a542ec904/opentelemetry-api/src/opentelemetry/trace/__init__.py#L555
tracer_provider: trace_api.TracerProvider
if isinstance(
tracer._tracer_provider,
(trace_api.ProxyTracerProvider, trace_api.NoOpTracerProvider),
):
tracer_provider = trace_sdk.TracerProvider(resource=resource)
set_tracer_provider(tracer_provider)
else:
# attach the processor to the existing tracer provider
tracer_provider = tracer._tracer_provider
if isinstance(tracer_provider, trace_sdk.TracerProvider):
tracer_provider.resource.merge(resource)
span_exporter = OTLPSpanExporter(
endpoint=f"{observability_url}/observability/traces/otlp/v0",
compression=otlp_compression,
session=session,
)
if isinstance(tracer_provider, trace_sdk.TracerProvider):
tracer_provider.add_span_processor(_MetadataSpanProcessor(metadata))
tracer_provider.add_span_processor(BatchSpanProcessor(span_exporter))
# Always set up the logger provider — it's needed for session reports,
# evaluations, and chat history, not just Python log export.
logger_provider = get_logger_provider()
if not isinstance(logger_provider, LoggerProvider):
logger_provider = LoggerProvider()
set_logger_provider(logger_provider)
if enable_logs:
log_exporter = OTLPLogExporter(
endpoint=f"{observability_url}/observability/logs/otlp/v0",
compression=otlp_compression,
session=session,
)
logger_provider.add_log_record_processor(_MetadataLogProcessor(metadata))
logger_provider.add_log_record_processor(BatchLogRecordProcessor(log_exporter))
handler = _TraceLevelLoggingHandler(level=logging.NOTSET, logger_provider=logger_provider)
root = logging.getLogger()
root.addHandler(handler)
# Set up the MeterProvider for OTEL metrics export
current_meter_provider = metrics_api.get_meter_provider()
if not isinstance(current_meter_provider, SdkMeterProvider):
metric_exporter = OTLPMetricExporter(
endpoint=f"{observability_url}/observability/metrics/otlp/v0",
compression=otlp_compression,
session=session,
preferred_temporality={
SdkCounter: AggregationTemporality.DELTA,
SdkUpDownCounter: AggregationTemporality.DELTA,
SdkHistogram: AggregationTemporality.DELTA,
SdkObservableCounter: AggregationTemporality.DELTA,
SdkObservableUpDownCounter: AggregationTemporality.DELTA,
SdkObservableGauge: AggregationTemporality.DELTA,
},
)
reader = PeriodicExportingMetricReader(metric_exporter, export_interval_millis=30000)
meter_provider = SdkMeterProvider(resource=resource, metric_readers=[reader])
metrics_api.set_meter_provider(meter_provider)
def _chat_ctx_to_otel_events(chat_ctx: ChatContext) -> list[tuple[str, Attributes]]:
role_to_event = {
"system": trace_types.EVENT_GEN_AI_SYSTEM_MESSAGE,
# OpenAI's `developer` role is the successor to `system` on the
# Chat Completions API and carries equivalent instructional content,
# so surface it as the system-message span event rather than dropping
# it on the floor.
"developer": trace_types.EVENT_GEN_AI_SYSTEM_MESSAGE,
"user": trace_types.EVENT_GEN_AI_USER_MESSAGE,
"assistant": trace_types.EVENT_GEN_AI_ASSISTANT_MESSAGE,
}
events: list[tuple[str, Attributes]] = []
for item in chat_ctx.items:
if item.type == "message" and (event_name := role_to_event.get(item.role)):
# only support text content for now
events.append((event_name, {"content": item.raw_text_content or ""}))
elif item.type == "function_call":
events.append(
(
trace_types.EVENT_GEN_AI_ASSISTANT_MESSAGE,
{
"role": "assistant",
"tool_calls": [
json.dumps(
{
"function": {"name": item.name, "arguments": item.arguments},
"id": item.call_id,
"type": "function",
}
)
],
},
)
)
elif item.type == "function_call_output":
events.append(
(
trace_types.EVENT_GEN_AI_TOOL_MESSAGE,
{"content": item.output, "name": item.name, "id": item.call_id},
)
)
return events
def _build_proto_chat_item(
item: ChatItem,
) -> agent_pb.agent_session.ChatContext.ChatItem:
item_pb = agent_pb.agent_session.ChatContext.ChatItem()
if item.type == "message":
msg = item_pb.message
msg.id = item.id
role_map = {
"developer": agent_pb.agent_session.DEVELOPER,
"system": agent_pb.agent_session.SYSTEM,
"user": agent_pb.agent_session.USER,
"assistant": agent_pb.agent_session.ASSISTANT,
}
msg.role = role_map[item.role]
from ..llm.chat_context import Instructions
for content in item.content:
if isinstance(content, (str, Instructions)):
content_pb = msg.content.add()
content_pb.text = str(content)
msg.interrupted = item.interrupted
if item.transcript_confidence is not None:
msg.transcript_confidence = item.transcript_confidence
for key, value in item.extra.items():
msg.extra[key] = str(value)
metrics = item.metrics
if "started_speaking_at" in metrics:
msg.metrics.started_speaking_at.FromMilliseconds(
int(metrics["started_speaking_at"] * 1000)
)
if "stopped_speaking_at" in metrics:
msg.metrics.stopped_speaking_at.FromMilliseconds(
int(metrics["stopped_speaking_at"] * 1000)
)
if "transcription_delay" in metrics:
msg.metrics.transcription_delay = metrics["transcription_delay"]
if "end_of_turn_delay" in metrics:
msg.metrics.end_of_turn_delay = metrics["end_of_turn_delay"]
if "on_user_turn_completed_delay" in metrics:
msg.metrics.on_user_turn_completed_delay = metrics["on_user_turn_completed_delay"]
if "llm_node_ttft" in metrics:
msg.metrics.llm_node_ttft = metrics["llm_node_ttft"]
if "tts_node_ttfb" in metrics:
msg.metrics.tts_node_ttfb = metrics["tts_node_ttfb"]
if "e2e_latency" in metrics:
msg.metrics.e2e_latency = metrics["e2e_latency"]
msg.created_at.FromMilliseconds(int(item.created_at * 1000))
elif item.type == "function_call":
fc = item_pb.function_call
fc.id = item.id
fc.call_id = item.call_id
fc.arguments = item.arguments
fc.name = item.name
fc.created_at.FromMilliseconds(int(item.created_at * 1000))
elif item.type == "function_call_output":
fco = item_pb.function_call_output
fco.id = item.id
fco.name = item.name
fco.call_id = item.call_id
fco.output = item.output
fco.is_error = item.is_error
fco.created_at.FromMilliseconds(int(item.created_at * 1000))
elif item.type == "agent_handoff":
ah = item_pb.agent_handoff
ah.id = item.id
if item.old_agent_id is not None:
ah.old_agent_id = item.old_agent_id
ah.new_agent_id = item.new_agent_id
ah.created_at.FromMilliseconds(int(item.created_at * 1000))
elif item.type == "agent_config_update":
acu = item_pb.agent_config_update
acu.id = item.id
if item.instructions is not None:
acu.instructions = item.instructions
if item.tools_added:
acu.tools_added.extend(item.tools_added)
if item.tools_removed:
acu.tools_removed.extend(item.tools_removed)
acu.created_at.FromMilliseconds(int(item.created_at * 1000))
return item_pb
def _to_proto_chat_item(item: ChatItem) -> dict:
return MessageToDict(_build_proto_chat_item(item), preserving_proto_field_name=True)
async def _parse_retry_delay(resp: aiohttp.ClientResponse) -> float | None:
"""Parse a protobuf Status error response for RetryInfo and return the retry delay in seconds,
or None if the error is not retryable."""
from google.rpc import error_details_pb2, status_pb2 # type: ignore[import-untyped]
try:
body = await resp.read()
status = status_pb2.Status()
status.ParseFromString(body)
for detail in status.details:
retry_info = error_details_pb2.RetryInfo()
if detail.Unpack(retry_info):
delay = retry_info.retry_delay
return float(delay.seconds + delay.nanos / 1e9)
except Exception:
pass
return None
async def _upload_session_report(
*,
agent_name: str,
observability_url: str,
report: SessionReport,
tagger: Tagger,
http_session: aiohttp.ClientSession,
) -> None:
def _get_logger(name: str) -> Any:
return get_logger_provider().get_logger(
name=name,
attributes={
"room_id": report.room_id,
"job_id": report.job_id,
"room": report.room,
},
)
def _log(
otel_logger: Any,
body: str,
timestamp: int,
attributes: dict,
severity: SeverityNumber = SeverityNumber.UNSPECIFIED,
severity_text: str = "unspecified",
) -> None:
otel_logger.emit(
body=body,
timestamp=timestamp,
attributes=attributes,
severity_number=severity,
severity_text=severity_text,
)
chat_logger = _get_logger("chat_history")
recording_options = report.recording_options
if any(recording_options.values()):
_log(
chat_logger,
body="session report",
timestamp=int((report.started_at or report.timestamp or 0) * 1e9),
attributes={
"session.options": vars(report.options),
"session.report_timestamp": report.timestamp,
"session.tags": sorted(tagger.tags) if tagger.tags else None,
"agent_name": agent_name,
"sdk_version": report.sdk_version,
"usage": [
{k: v for k, v in u.model_dump().items() if v != 0 and v != 0.0}
for u in report.model_usage
]
if report.model_usage
else None,
},
)
if recording_options["transcript"]:
for item in report.chat_history.items:
item_log = _to_proto_chat_item(item)
severity: SeverityNumber = SeverityNumber.UNSPECIFIED
severity_text: str = "unspecified"
if item.type == "function_call_output" and item.is_error:
severity = SeverityNumber.ERROR
severity_text = "error"
_log(
chat_logger,
body="chat item",
timestamp=int(item.created_at * 1e9),
attributes={"chat.item": item_log},
severity=severity,
severity_text=severity_text,
)
eval_logger = _get_logger("evaluations")
if tagger.evaluations:
for evaluation in tagger.evaluations:
severity = SeverityNumber.UNSPECIFIED
severity_text = "unspecified"
if evaluation.get("verdict") == "fail":
severity = SeverityNumber.ERROR
severity_text = "error"
_log(
eval_logger,
body="evaluation",
timestamp=int(report.timestamp * 1e9),
attributes={"evaluation": evaluation},
severity=severity,
severity_text=severity_text,
)
for tag, entry in tagger._tags.items():
if entry.metadata:
_log(
eval_logger,
body="tag",
timestamp=int(entry.timestamp * 1e9),
attributes={"tag": {"name": tag, "metadata": entry.metadata}},
)
if tagger.outcome:
is_fail = tagger.outcome == "fail"
outcome_data: dict[str, Any] = {"outcome": tagger.outcome}
if tagger.outcome_reason:
outcome_data["reason"] = tagger.outcome_reason
_log(
eval_logger,
body="outcome",
timestamp=int(report.timestamp * 1e9),
attributes={"outcome": outcome_data},
severity=SeverityNumber.ERROR if is_fail else SeverityNumber.UNSPECIFIED,
severity_text="error" if is_fail else "unspecified",
)
has_audio = (
recording_options["audio"]
and report.audio_recording_path
and report.audio_recording_started_at
)
if not recording_options["transcript"] and not has_audio:
return
# emit recording
access_token = (
api.AccessToken()
.with_observability_grants(api.ObservabilityGrants(write=True))
.with_ttl(timedelta(hours=6))
)
jwt = access_token.to_jwt()
header_msg = proto_metrics.MetricsRecordingHeader(
room_id=report.room_id,
)
header_msg.start_time.FromMilliseconds(int((report.audio_recording_started_at or 0) * 1000))
header_bytes = header_msg.SerializeToString()
chat_history_json = ""
if recording_options["transcript"]:
chat_history_json = json.dumps(report.chat_history.to_dict(exclude_timestamp=False))
audio_bytes = b""
if has_audio and report.audio_recording_path:
try:
async with aiofiles.open(report.audio_recording_path, "rb") as f:
audio_bytes = await f.read()
except Exception:
audio_bytes = b""
url = f"{observability_url}/observability/recordings/v0"
def _build_multipart() -> aiohttp.MultipartWriter:
mp = aiohttp.MultipartWriter("form-data")
part = mp.append(header_bytes)
part.set_content_disposition("form-data", name="header", filename="header.binpb")
part.headers["Content-Type"] = "application/protobuf"
part.headers["Content-Length"] = str(len(header_bytes))
if recording_options["transcript"]:
part = mp.append(chat_history_json)
part.set_content_disposition(
"form-data", name="chat_history", filename="chat_history.json"
)
part.headers["Content-Type"] = "application/json"
part.headers["Content-Length"] = str(len(chat_history_json))
if audio_bytes:
part = mp.append(audio_bytes)
part.set_content_disposition("form-data", name="audio", filename="recording.ogg")
part.headers["Content-Type"] = "audio/ogg"
part.headers["Content-Length"] = str(len(audio_bytes))
return mp
max_retries = 3
for attempt in range(max_retries + 1):
mp = _build_multipart()
headers = {
"Authorization": f"Bearer {jwt}",
"Content-Type": mp.content_type,
}
logger.debug("uploading session report to LiveKit Cloud")
async with http_session.post(url, data=mp, headers=headers) as resp:
if resp.status < 400:
break
retry_delay = await _parse_retry_delay(resp)
if retry_delay is None or attempt == max_retries:
resp.raise_for_status()
raise RuntimeError(f"recording upload failed: status {resp.status}")
logger.warning(
"recording upload failed (attempt %d/%d), retrying in %.1fs",
attempt + 1,
max_retries + 1,
retry_delay,
)
await asyncio.sleep(retry_delay)
logger.debug("finished uploading")
_TELEMETRY_SHUTDOWN_TIMEOUT = 10.0
def _shutdown_telemetry(timeout: float = _TELEMETRY_SHUTDOWN_TIMEOUT) -> None:
"""Shut down OTel providers with a hard wall-clock bound.
``provider.shutdown()`` internally joins its exporter worker with a 30s
default timeout per provider (and ``force_flush`` ignores its timeout arg
in the current SDK — see #4623). Across tracer/logger/meter that's up to
~90s, enough to stall the caller's event loop past the supervisor's 60s
ping/pong deadline when the OTLP endpoint is rate-limiting or unreachable.
Each provider is shut down in its *own* daemon thread, run in parallel.
That matters for two reasons:
1) Main-thread wait is bounded by ``max`` of the three, not the ``sum``.
2) ``BatchProcessor.shutdown()`` sets ``_shutdown = True`` as its first
action; running in parallel guarantees that flag gets set on every
provider within milliseconds, even if one hangs in ``worker_thread.join``.
Any later atexit re-entry (OTel registers one, and Python's
``logging.shutdown()`` may spawn a *non-daemon* thread via
``LoggingHandler.flush`` → ``force_flush`` — see opentelemetry-python
PR #4636) then short-circuits instead of hanging process exit.
Any unfinished work stays on existing daemon threads and is discarded at
process exit.
Upstream context:
- https://github.com/open-telemetry/opentelemetry-python/issues/4623
(TracerProvider.shutdown() has no configurable timeout — still open)
"""
# Detach the OTLP LoggingHandler from the root logger — belt to the
# suspenders of the parallel shutdown below.
root = logging.getLogger()
for h in list(root.handlers):
if isinstance(h, LoggingHandler):
root.removeHandler(h)
providers: list[Any] = []
if isinstance(lp := get_logger_provider(), LoggerProvider):
providers.append(lp)
if isinstance(tp := tracer._tracer_provider, trace_sdk.TracerProvider):
providers.append(tp)
if isinstance(mp := metrics_api.get_meter_provider(), SdkMeterProvider):
providers.append(mp)
def _shutdown_one(provider: Any) -> None:
try:
provider.shutdown()
except Exception:
logger.exception("failed to shut down telemetry provider")
threads = [
threading.Thread(
target=_shutdown_one,
args=(p,),
name=f"livekit-telemetry-shutdown-{type(p).__name__}",
daemon=True,
)
for p in providers
]
for t in threads:
t.start()
deadline = time.monotonic() + timeout
for t in threads:
t.join(max(0.0, deadline - time.monotonic()))
if any(t.is_alive() for t in threads):
logger.warning("telemetry shutdown exceeded %.1fs; continuing", timeout)