Skip to content
Draft
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
32 changes: 27 additions & 5 deletions livekit-agents/livekit/agents/voice/generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -457,17 +457,32 @@ async def _tts_inference_task(
provider: str | None = None,
) -> bool:
current_span = trace.get_current_span()
if model:
current_span.set_attribute(trace_types.ATTR_GEN_AI_REQUEST_MODEL, model)
if provider:
current_span.set_attribute(trace_types.ATTR_GEN_AI_PROVIDER_NAME, provider)
gen_ai_telemetry.set_request_attributes(
current_span,
operation=None,
model=model,
provider=provider,
output_type=trace_types.GenAIOutputType.SPEECH,
)

audio_ch, timed_texts_fut = data.audio_ch, data.timed_texts_fut
if text_transforms:
input = _apply_text_transforms(input, text_transforms)

start_time: float | None = None
input_tee = itertools.tee(input, 2)
record_content = current_span.is_recording() and gen_ai_telemetry.capture_content_enabled()
input_text: list[str] = []

async def _capture_input() -> AsyncIterable[str]:
nonlocal record_content
async for chunk in input_tee[1]:
record_content = record_content and gen_ai_telemetry.capture_content_enabled()
if record_content:
input_text.append(chunk)
yield chunk

observed_input = _capture_input()

async def _get_start_time() -> None:
nonlocal start_time
Expand All @@ -477,7 +492,7 @@ async def _get_start_time() -> None:

_start_time_task = asyncio.create_task(_get_start_time())
try:
tts_node = node(input_tee[1], model_settings)
tts_node = node(observed_input, model_settings)
if asyncio.iscoroutine(tts_node):
tts_node = await tts_node

Expand Down Expand Up @@ -513,6 +528,13 @@ async def _get_start_time() -> None:
return audio_duration > 0
finally:
await aio.gracefully_cancel(_start_time_task)
if record_content and gen_ai_telemetry.capture_content_enabled():
gen_ai_telemetry.set_content_attributes(
current_span,
input_messages=gen_ai_telemetry.to_speech_messages(
"".join(input_text), role="assistant"
),
)
await input_tee.aclose()


Expand Down
95 changes: 95 additions & 0 deletions tests/test_tts_telemetry.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,95 @@
from __future__ import annotations

import asyncio
import json
from collections.abc import AsyncIterable, Iterator

import pytest
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter

from livekit import rtc
from livekit.agents.telemetry import gen_ai, set_tracer_provider, tracer
from livekit.agents.voice.generation import perform_tts_inference
from livekit.agents.voice.io import ModelSettings

pytestmark = [pytest.mark.unit, pytest.mark.no_concurrent]


@pytest.fixture
def exporter() -> Iterator[InMemorySpanExporter]:
original = tracer._tracer_provider
provider = TracerProvider()
exporter = InMemorySpanExporter()
provider.add_span_processor(SimpleSpanProcessor(exporter))
set_tracer_provider(provider)
try:
yield exporter
finally:
set_tracer_provider(original)
provider.shutdown()


@pytest.mark.parametrize("capture", [True, False])
@pytest.mark.parametrize("finish", ["complete", "early", "cancel"])
async def test_tts_capture_observes_consumption_without_draining(
exporter: InMemorySpanExporter,
monkeypatch: pytest.MonkeyPatch,
capture: bool,
finish: str,
) -> None:
monkeypatch.setattr(gen_ai, "_capture_content", capture)
produced: list[str] = []
consumed: list[str] = []
first_consumed = asyncio.Event()

async def source() -> AsyncIterable[str]:
for chunk in ["one ", "two ", "three"]:
produced.append(chunk)
yield chunk

async def uppercase(source: AsyncIterable[str]) -> AsyncIterable[str]:
async for chunk in source:
yield chunk.upper()

async def node(
text: AsyncIterable[str], settings: ModelSettings
) -> AsyncIterable[rtc.AudioFrame]:
async for chunk in text:
consumed.append(chunk)
first_consumed.set()
if finish == "cancel":
await asyncio.Event().wait()
# Give a background reader a chance to consume ahead of this node.
await asyncio.sleep(0.01)
yield rtc.AudioFrame.create(sample_rate=24000, num_channels=1, samples_per_channel=240)
if finish == "early":
break

task, _ = perform_tts_inference(
node=node,
input=source(),
model_settings=ModelSettings(),
text_transforms=[uppercase],
provider="google",
model="test-tts",
)
if finish == "cancel":
await asyncio.wait_for(first_consumed.wait(), 5)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
else:
assert await task
expected = ["one ", "two ", "three"] if finish == "complete" else ["one "]
assert produced == expected
assert consumed == [chunk.upper() for chunk in expected]
[span] = [span for span in exporter.get_finished_spans() if span.name == "tts_node"]
assert span.attributes["gen_ai.provider.name"] == "gcp.gen_ai"
if capture:
assert json.loads(span.attributes["gen_ai.input.messages"]) == [
{"role": "assistant", "parts": [{"type": "text", "content": "".join(consumed)}]}
]
else:
assert "gen_ai.input.messages" not in span.attributes
Loading