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 @@ -5,6 +5,7 @@
import contextlib
import json
import os
import sys
import time
from collections.abc import AsyncIterable
from dataclasses import dataclass, field, replace
Expand Down Expand Up @@ -76,6 +77,15 @@
lk_oai_debug = int(os.getenv("LK_OPENAI_DEBUG", 0))


def _valid_usage_seconds(value: object) -> bool:
"""Accept durations representable as finite, nonnegative metric values."""
return (
isinstance(value, (int, float))
and not isinstance(value, bool)
and 0 <= value <= sys.float_info.max
)


class ResponsesDelegationOptions(TypedDict, total=False):
"""The backend Responses model delegated work runs on, under ``delegation="responses"``.

Expand Down Expand Up @@ -579,6 +589,12 @@ def _handle_event(self, event: dict[str, Any]) -> None:
if lk_oai_debug and etype != "session.output_audio.delta":
logger.debug("gpt-live server event", extra={"lk.pii.event": event})

if etype in ("session.usage.updated", "session.closed"):
usage = event.get("usage")
if not isinstance(usage, dict) or not _valid_usage_seconds(usage.get("seconds", 0)):
# OpenAI's construct() still coerces integers to floats, which can overflow.
event = {**event, "usage": {}}

if etype == "session.started":
self._handle_session_started(types.SessionStartedEvent.construct(**event))
elif etype == "session.output_audio.delta":
Expand Down Expand Up @@ -837,6 +853,9 @@ def _handle_session_closed(self, event: types.SessionClosedEvent) -> None:

def _handle_usage(self, usage: types.Usage) -> None:
# reported cumulatively for the whole session, so only the delta goes to the collectors
seconds = usage.seconds
if not _valid_usage_seconds(seconds) or seconds <= self._usage_total.seconds:
return
previous, self._usage_total = self._usage_total, usage
self.emit(
"metrics_collected",
Expand Down
54 changes: 54 additions & 0 deletions tests/test_gpt_live_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -1039,6 +1039,60 @@ async def test_a_delegated_model_is_billed_under_its_own_name(
await model.aclose()


@pytest.mark.parametrize(
"seconds", [10, 8, 0, -1, True, float("nan"), float("inf"), float("-inf"), "invalid", None]
)
async def test_invalid_usage_does_not_lower_or_poison_the_cumulative_watermark(
monkeypatch: pytest.MonkeyPatch, seconds: Any
) -> None:
"""Only increasing, finite duration reports can advance the cumulative watermark."""
_connect_hook(monkeypatch)
model = GPTLiveModel(api_key="sk-test")
session = model.session()
collected: list[RealtimeModelMetrics] = []
session.on("metrics_collected", collected.append)
try:
await session._update_session()
await session._session_started_fut
for value in (10, seconds, 12, 12):
session._handle_event({"type": "session.usage.updated", "usage": {"seconds": value}})
assert [metric.session_duration for metric in collected] == [10, 2]
session._handle_event({"type": "session.closed", "usage": {"seconds": 12.5}})
assert [metric.session_duration for metric in collected] == [10, 2, 0.5]
finally:
await session.aclose()
await model.aclose()


@pytest.mark.parametrize(
"seconds", [10**400, -(10**400), True, False, "12", float("nan"), float("inf")]
)
async def test_invalid_wire_usage_preserves_metrics_and_acknowledges_close(
monkeypatch: pytest.MonkeyPatch, seconds: Any
) -> None:
"""Malformed usage must be rejected before the event parser coerces numeric fields."""
_connect_hook(monkeypatch)
model = GPTLiveModel(api_key="sk-test")
session = model.session()
collected: list[RealtimeModelMetrics] = []
session.on("metrics_collected", collected.append)
try:
await session._update_session()
await session._session_started_fut
for value in (seconds, 10, seconds, 12):
event = {"type": "session.usage.updated", "usage": {"seconds": value}}
original = event["usage"].copy()
session._handle_event(event)
assert event["usage"] == original
assert [metric.session_duration for metric in collected] == [10, 2]
session._handle_event({"type": "session.closed", "usage": {"seconds": seconds}})
assert session._session_closed_fut.done()
assert [metric.session_duration for metric in collected] == [10, 2]
finally:
await session.aclose()
await model.aclose()


def _user_events(session: GPTLiveSession) -> list[tuple[str, Any]]:
events: list[tuple[str, Any]] = []
for name in (
Expand Down