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
11 changes: 11 additions & 0 deletions livekit-agents/livekit/agents/llm/realtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -264,6 +264,17 @@ def truncate(
audio_transcript: NotGivenOr[str] = NOT_GIVEN,
) -> None: ...

async def drain_pending_metrics(self) -> None:
"""Wait (bounded) for metrics the provider has not delivered yet.

Called by the framework right before the session's event listeners are
detached at teardown, so a usage/metrics event that is still in flight can
be emitted and collected. Providers that report usage asynchronously
relative to playback (e.g. Gemini Live's end-of-turn usageMetadata)
override this; the default is a no-op.
"""
return None

@abstractmethod
async def aclose(self) -> None: ...

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -429,6 +429,9 @@ def truncate(
audio_transcript=audio_transcript,
)

async def drain_pending_metrics(self) -> None:
await self._active.drain_pending_metrics()

async def aclose(self) -> None:
# cancel an in-flight swap first, else its fresh child would leak past aclose
if self._swap_task is not None:
Expand Down
10 changes: 10 additions & 0 deletions livekit-agents/livekit/agents/voice/agent_activity.py
Original file line number Diff line number Diff line change
Expand Up @@ -1093,6 +1093,16 @@ async def _close_session(self) -> None:
self.llm.off("error", self._on_error)

if isinstance(self.llm, llm.RealtimeModel) and self._rt_session is not None:
# give the session a bounded chance to flush an in-flight usage/metrics
# event (e.g. Gemini's end-of-turn usageMetadata) BEFORE detaching the
# listeners below — otherwise the last generation's usage is emitted into
# the void and never reaches session.usage.
try:
await self._rt_session.drain_pending_metrics()
except Exception:
logger.warning(
"error draining pending realtime metrics before close", exc_info=True
)
self._rt_session.off("generation_created", self._on_generation_created)
self._rt_session.off("input_speech_started", self._on_input_speech_started)
self._rt_session.off("input_speech_stopped", self._on_input_speech_stopped)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,13 @@

lk_google_debug = int(os.getenv("LK_GOOGLE_DEBUG", 0))

# Bounded grace (seconds) to wait for the current generation's usageMetadata when the
# session is being torn down. Gemini Live reports token usage only at the end of a
# generation, so a caller hanging up during (or right at the end of) an utterance would
# otherwise close the websocket before the usage event arrives — silently dropping the
# last generation's token usage (on a one-turn call: all of it). <= 0 disables the wait.
lk_google_usage_drain_timeout = float(os.getenv("LK_GOOGLE_USAGE_DRAIN_TIMEOUT", "2.0"))

# stop rejecting tool calls after this many in a row to avoid a loop (tool_choice="none")
MAX_TOOL_CALL_REJECTIONS = 3

Expand Down Expand Up @@ -178,6 +185,8 @@ class _ResponseGeneration:
"""The timestamp when the generation is completed"""
_done: bool = False
"""Whether the generation is done (set when the turn is complete)"""
_usage_received: bool = False
"""Whether this generation's usageMetadata has been received"""

def push_text(self, text: str) -> None:
if self.output_text:
Expand Down Expand Up @@ -496,6 +505,10 @@ def __init__(self, realtime_model: RealtimeModel) -> None:
# means we're draining that turn's trailing events (which have no generation to attach
# to). reset when the next generation starts.
self._rejected_tool_calls = 0
# set when the current generation's usageMetadata lands; cleared on each new
# generation. drain_pending_metrics() waits on it at teardown so the final
# generation's token usage isn't lost when the call ends mid-turn.
self._usage_received_ev = asyncio.Event()

self._session_resumption_handle: str | None = (
self._opts.session_resumption.handle
Expand Down Expand Up @@ -828,6 +841,43 @@ def truncate(
logger.warning("truncate is not supported by the Google Realtime API.")
pass

async def drain_pending_metrics(self) -> None:
"""Wait (bounded) for the current generation's usageMetadata before teardown.

Gemini Live only reports token usage at the end of a generation, so when the
call ends mid-turn (e.g. the caller hangs up while — or right after — the agent
speaks) the usage event is still in flight. The framework calls this before
detaching its ``metrics_collected`` listener and closing the session, which is
what lets that final usage event still be emitted and collected. No-op when
nothing is pending or the underlying session is already gone.
"""
if lk_google_usage_drain_timeout <= 0:
return
gen = self._current_generation
if gen is None or gen._usage_received:
return
if (
self._main_atask is None
or self._main_atask.done()
or self._msg_ch.closed
or self._session_should_close.is_set()
or self._active_session is None
):
# not connected (or already tearing down) — the usage event can no longer
# arrive, so waiting would only delay the close.
return
try:
await asyncio.wait_for(self._usage_received_ev.wait(), lk_google_usage_drain_timeout)
except asyncio.TimeoutError:
logger.warning(
"closing Gemini realtime session without usage metadata for the last "
"generation; its token usage will be missing",
extra={
"response_id": gen.response_id,
"timeout": lk_google_usage_drain_timeout,
},
)

async def aclose(self) -> None:
self._msg_ch.close()
self._session_should_close.set()
Expand Down Expand Up @@ -1205,6 +1255,7 @@ def _start_new_generation(self) -> None:
audio_ch=utils.aio.Chan[rtc.AudioFrame](),
_created_timestamp=time.time(),
)
self._usage_received_ev.clear()
if not self._realtime_model.capabilities.audio_output:
self._current_generation.audio_ch.close()

Expand Down Expand Up @@ -1498,6 +1549,8 @@ def _token_details_map(
model_name=self._realtime_model.model, model_provider=self._realtime_model.provider
),
)
current_gen._usage_received = True
self._usage_received_ev.set()
self.emit("metrics_collected", metrics)

def _handle_go_away(self, go_away: types.LiveServerGoAway) -> None:
Expand Down
228 changes: 228 additions & 0 deletions tests/test_realtime_usage_drain.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,228 @@
"""Tests for the realtime usage-metadata drain at session teardown.

Gemini Live reports token usage only at end-of-turn; when a call ends mid-turn
(e.g. the caller hangs up while the agent is speaking its first line) the usage
event is still in flight. The fix under test has two halves:

1. the Gemini plugin's ``RealtimeSession.drain_pending_metrics()`` waits (bounded)
for the current generation's usageMetadata;
2. ``AgentActivity._close_session`` awaits ``drain_pending_metrics()`` BEFORE it
detaches the ``metrics_collected`` listener, so the drained event still lands
in ``session.usage``.
"""

from __future__ import annotations

import asyncio
import time
from unittest.mock import MagicMock

import pytest
from google.genai import types as genai_types

from livekit.agents import Agent, AgentSession, llm, utils
from livekit.agents.metrics import LLMModelUsage, RealtimeModelMetrics
from livekit.agents.metrics.base import Metadata
from livekit.plugins.google.realtime import realtime_api

from .fake_realtime import FakeRealtimeModel

pytestmark = pytest.mark.unit


def _make_gemini_session() -> realtime_api.RealtimeSession:
"""Build a hermetic plugin RealtimeSession: no client, no websocket, no tasks.

Only the state ``drain_pending_metrics`` / ``_handle_usage_metadata`` touch is
initialized; ``_main_atask`` is a placeholder task standing in for a live
connection loop.
"""
sess = realtime_api.RealtimeSession.__new__(realtime_api.RealtimeSession)
fake_model = MagicMock()
fake_model.label = "google.realtime.RealtimeModel"
fake_model.model = "gemini-test"
fake_model.provider = "google"
llm.RealtimeSession.__init__(sess, fake_model)
sess._msg_ch = utils.aio.Chan()
sess._session_should_close = asyncio.Event()
sess._usage_received_ev = asyncio.Event()
sess._current_generation = None
sess._rejected_tool_calls = 0
sess._active_session = MagicMock()
sess._main_atask = asyncio.create_task(asyncio.sleep(30))
return sess


def _make_generation() -> realtime_api._ResponseGeneration:
return realtime_api._ResponseGeneration(
message_ch=utils.aio.Chan(),
function_ch=utils.aio.Chan(),
input_id="GI_test",
response_id="GR_test",
text_ch=utils.aio.Chan(),
audio_ch=utils.aio.Chan(),
_created_timestamp=time.time(),
)


async def _cleanup(sess: realtime_api.RealtimeSession) -> None:
sess._main_atask.cancel()
try:
await sess._main_atask
except asyncio.CancelledError:
pass


async def test_drain_noop_without_generation() -> None:
sess = _make_gemini_session()
try:
started = time.monotonic()
await sess.drain_pending_metrics()
assert time.monotonic() - started < 0.5
finally:
await _cleanup(sess)


async def test_drain_noop_when_usage_already_received() -> None:
sess = _make_gemini_session()
try:
gen = _make_generation()
gen._usage_received = True
sess._current_generation = gen
started = time.monotonic()
await sess.drain_pending_metrics()
assert time.monotonic() - started < 0.5
finally:
await _cleanup(sess)


async def test_drain_noop_when_not_connected() -> None:
sess = _make_gemini_session()
try:
sess._current_generation = _make_generation()
sess._active_session = None # never connected / already torn down
started = time.monotonic()
await sess.drain_pending_metrics()
assert time.monotonic() - started < 0.5
finally:
await _cleanup(sess)


async def test_drain_waits_for_usage_metadata() -> None:
"""A late usageMetadata delivered during the drain window is still emitted."""
sess = _make_gemini_session()
try:
gen = _make_generation()
sess._current_generation = gen

collected: list[RealtimeModelMetrics] = []
sess.on("metrics_collected", collected.append)

usage = genai_types.UsageMetadata(
prompt_token_count=100, response_token_count=25, total_token_count=125
)

async def _deliver_late() -> None:
await asyncio.sleep(0.05)
sess._handle_usage_metadata(usage)

deliver_task = asyncio.create_task(_deliver_late())
await sess.drain_pending_metrics()
await deliver_task

assert gen._usage_received
assert len(collected) == 1
assert collected[0].input_tokens == 100
assert collected[0].output_tokens == 25
finally:
await _cleanup(sess)


async def test_drain_times_out_when_usage_never_arrives(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(realtime_api, "lk_google_usage_drain_timeout", 0.1)
sess = _make_gemini_session()
try:
sess._current_generation = _make_generation()
started = time.monotonic()
await sess.drain_pending_metrics()
elapsed = time.monotonic() - started
assert 0.05 < elapsed < 1.0
finally:
await _cleanup(sess)


async def test_drain_disabled_via_timeout() -> None:
sess = _make_gemini_session()
try:
sess._current_generation = _make_generation()
orig = realtime_api.lk_google_usage_drain_timeout
realtime_api.lk_google_usage_drain_timeout = 0.0
try:
started = time.monotonic()
await sess.drain_pending_metrics()
assert time.monotonic() - started < 0.5
finally:
realtime_api.lk_google_usage_drain_timeout = orig
finally:
await _cleanup(sess)


async def test_usage_metadata_marks_generation_and_event() -> None:
sess = _make_gemini_session()
try:
gen = _make_generation()
sess._current_generation = gen
assert not sess._usage_received_ev.is_set()

sess._handle_usage_metadata(
genai_types.UsageMetadata(
prompt_token_count=10, response_token_count=5, total_token_count=15
)
)
assert gen._usage_received
assert sess._usage_received_ev.is_set()
finally:
await _cleanup(sess)


async def test_close_session_drains_before_detaching_metrics_listener() -> None:
"""AgentSession.aclose() must collect a usage event flushed by the drain hook.

This is the ordering the whole fix depends on: _close_session awaits
drain_pending_metrics() BEFORE off("metrics_collected"), so a usage event
emitted from the drain still reaches session.usage.
"""
model = FakeRealtimeModel()
session: AgentSession = AgentSession(llm=model)
await session.start(Agent(instructions="test"))

rt = model.active_session
metric = RealtimeModelMetrics(
label=model.label,
request_id="GR_last_turn",
timestamp=time.time(),
input_tokens=111,
output_tokens=22,
total_tokens=133,
input_token_details=RealtimeModelMetrics.InputTokenDetails(audio_tokens=90, text_tokens=10),
output_token_details=RealtimeModelMetrics.OutputTokenDetails(audio_tokens=22),
metadata=Metadata(model_name="fake-realtime", model_provider="fake"),
)

drained = asyncio.Event()

async def _drain_pending_metrics() -> None:
# simulate the provider flushing the in-flight usage event during the grace
rt.emit("metrics_collected", metric)
drained.set()

rt.drain_pending_metrics = _drain_pending_metrics # type: ignore[method-assign]

await session.aclose()

assert drained.is_set()
llm_usage = [u for u in session.usage.model_usage if isinstance(u, LLMModelUsage)]
assert sum(u.input_tokens for u in llm_usage) == 111
assert sum(u.output_tokens for u in llm_usage) == 22
Loading