From cecca16305f32471f2d4eca317a9853dcf7ccfbc Mon Sep 17 00:00:00 2001 From: Henry Su Date: Thu, 20 Aug 2026 18:44:33 -0500 Subject: [PATCH] fix: validate realtime playback durations --- src/agents/realtime/model.py | 11 +++++++++-- tests/realtime/test_playback_tracker.py | 23 +++++++++++++++++++++++ 2 files changed, 32 insertions(+), 2 deletions(-) diff --git a/src/agents/realtime/model.py b/src/agents/realtime/model.py index 1a5efa7f93..7f86982ddd 100644 --- a/src/agents/realtime/model.py +++ b/src/agents/realtime/model.py @@ -1,6 +1,7 @@ from __future__ import annotations import abc +import math from collections.abc import Callable from typing_extensions import NotRequired, TypedDict @@ -57,12 +58,18 @@ def on_play_ms(self, item_id: str, item_content_index: int, ms: float) -> None: item_content_index: The index of the audio content in `item.content` ms: The number of milliseconds of audio that have been played. """ + if not math.isfinite(ms) or ms < 0: + raise ValueError("Playback duration must be a finite, non-negative number.") + if self._current_item != (item_id, item_content_index): + new_elapsed_ms = ms self._current_item = (item_id, item_content_index) - self._elapsed_ms = ms else: assert self._elapsed_ms is not None - self._elapsed_ms += ms + new_elapsed_ms = self._elapsed_ms + ms + if not math.isfinite(new_elapsed_ms): + raise ValueError("Playback duration total must be finite.") + self._elapsed_ms = new_elapsed_ms def on_interrupted(self) -> None: """Called by the model when the audio playback has been interrupted.""" diff --git a/tests/realtime/test_playback_tracker.py b/tests/realtime/test_playback_tracker.py index 1dd70e22c2..322c539f39 100644 --- a/tests/realtime/test_playback_tracker.py +++ b/tests/realtime/test_playback_tracker.py @@ -1,3 +1,4 @@ +import math from unittest.mock import AsyncMock, patch import pytest @@ -17,6 +18,28 @@ def model(self): """Create a fresh model instance for each test.""" return OpenAIRealtimeWebSocketModel() + @pytest.mark.parametrize("duration", [-1.0, math.nan, math.inf, -math.inf]) + def test_tracker_rejects_invalid_playback_duration(self, duration: float) -> None: + tracker = RealtimePlaybackTracker() + + with pytest.raises(ValueError, match="finite, non-negative"): + tracker.on_play_ms("item", 0, duration) + + assert tracker.get_state() == { + "current_item_id": None, + "current_item_content_index": None, + "elapsed_ms": None, + } + + def test_tracker_rejects_overflowing_playback_total(self) -> None: + tracker = RealtimePlaybackTracker() + tracker.on_play_ms("item", 0, 1e308) + + with pytest.raises(ValueError, match="total must be finite"): + tracker.on_play_ms("item", 0, 1e308) + + assert tracker.get_state()["elapsed_ms"] == 1e308 + @pytest.mark.asyncio async def test_interrupt_timing_with_custom_playback_tracker(self, model): """Test interrupt uses custom playback tracker elapsed time instead of default timing."""