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
23 changes: 17 additions & 6 deletions livekit-agents/livekit/agents/inference/interruption.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,8 +50,6 @@
)

SAMPLE_RATE = 16000
# local fallback when server/model side threshold is not available
THRESHOLD = 0.656
MIN_INTERRUPTION_DURATION = 0.025 * 2 # 25ms per frame, 2 consecutive frames
MAX_AUDIO_DURATION = 3 # 3 seconds
DETECTION_INTERVAL = 0.1 # 0.1 second
Expand Down Expand Up @@ -281,7 +279,7 @@ def __init__(
Initialize a AdaptiveInterruptionDetector instance.

Args:
threshold (float, optional): The threshold for the interruption detection. When not set, the server-recommended default (returned in session.created) is used, falling back to THRESHOLD if the server does not provide one.
threshold (float, optional): The threshold for the interruption detection. When not set, the server-recommended default (returned in session.created) is used.
min_interruption_duration (float, optional): The minimum duration, in seconds, of the interruption event, defaults to 50ms.
max_audio_duration (float, optional): The maximum audio duration, including the audio prefix, in seconds, for the interruption detection, defaults to 3s.
audio_prefix_duration (float, optional): The audio prefix duration, in seconds, for the interruption detection, defaults to 0.5s.
Expand Down Expand Up @@ -771,13 +769,16 @@ def update_options(
# opts are shared with the detector (self._opts is model._opts), no need to update them here
self._reconnect_event.set()

def _resolve_effective_threshold(self, default_threshold: float | None) -> float:
"""Return the effective threshold."""
def _resolve_effective_threshold(self, default_threshold: float | None) -> float | None:
"""Return the effective threshold for observability only.

Precedence: user override, then server default; None when neither is known.
"""
if is_given(self._opts.threshold):
return self._opts.threshold
if default_threshold is not None:
return default_threshold
return THRESHOLD
return None

async def _run(self) -> None:
closing_ws = False
Expand Down Expand Up @@ -846,6 +847,16 @@ async def recv_task(ws: aiohttp.ClientWebSocketResponse) -> None:

match msg:
case InterruptionWSSessionCreatedMessage():
if not is_given(self._opts.threshold) and msg.default_threshold is None:
raise APIStatusError(
message=(
"adaptive interruption session created without a threshold: "
"no user override and the server did not report a "
"default_threshold"
),
status_code=500,
retryable=False,
)
Comment thread
chenghao-mou marked this conversation as resolved.
# Observability only — the server makes the actual decision;
logger.debug(
"adaptive interruption session created",
Expand Down
52 changes: 51 additions & 1 deletion tests/test_interruption/test_interruption_failover.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from __future__ import annotations

import asyncio
import json
import time
from unittest.mock import AsyncMock, MagicMock

Expand All @@ -15,7 +16,7 @@
import pytest

from livekit import rtc
from livekit.agents._exceptions import APIError
from livekit.agents._exceptions import APIError, APIStatusError
from livekit.agents.inference.interruption import (
AdaptiveInterruptionDetector,
InterruptionDetectionError,
Expand Down Expand Up @@ -201,3 +202,52 @@ async def _receive_hang() -> aiohttp.WSMessage:
unrecoverable_errors = [e for e in errors if not e.recoverable]
assert len(recoverable_errors) == 0
assert len(unrecoverable_errors) == 1


class TestWsSessionCreatedMissingThreshold:
@pytest.mark.asyncio
async def test_immediate_unrecoverable_when_server_omits_threshold(self) -> None:
mock_session = AsyncMock(spec=aiohttp.ClientSession)

def _make_mock_ws() -> MagicMock:
mock_ws = MagicMock(spec=aiohttp.ClientWebSocketResponse)
mock_ws.send_str = AsyncMock()
mock_ws.send_bytes = AsyncMock()
mock_ws.closed = False
mock_ws.close_code = None

sent_created = False

async def _receive() -> aiohttp.WSMessage:
nonlocal sent_created
if not sent_created:
sent_created = True
return aiohttp.WSMessage(
type=aiohttp.WSMsgType.TEXT,
data=json.dumps({"type": "session.created"}),
extra=None,
)
await asyncio.sleep(3600)
return aiohttp.WSMessage(type=aiohttp.WSMsgType.CLOSED, data=None, extra=None)

mock_ws.receive = _receive
mock_ws.close = AsyncMock(return_value=True)
return mock_ws

mock_session.ws_connect = AsyncMock(side_effect=lambda *a, **kw: _make_mock_ws())

detector = _create_detector(mock_session)
errors = _collect_errors(detector)
stream = detector.stream(conn_options=CONN_OPTIONS)

exc = await _wait_for_stream_failure(stream)

assert isinstance(exc, APIStatusError)
assert exc.status_code == 500
assert exc.retryable is False

# retryable=False -> no retries, immediate unrecoverable
recoverable_errors = [e for e in errors if e.recoverable]
unrecoverable_errors = [e for e in errors if not e.recoverable]
assert len(recoverable_errors) == 0
assert len(unrecoverable_errors) == 1
7 changes: 3 additions & 4 deletions tests/test_interruption/test_interruption_session_create.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,6 @@
import pytest

from livekit.agents.inference.interruption import (
THRESHOLD,
AdaptiveInterruptionDetector,
InterruptionWebSocketStream,
InterruptionWSSessionCreatedMessage,
Expand Down Expand Up @@ -113,7 +112,7 @@ def test_default_threshold_optional(self) -> None:


class TestResolveEffectiveThreshold:
"""Observability-only resolution: user override > server default > THRESHOLD backup."""
"""Observability-only resolution: user override > server default > None."""

@pytest.mark.asyncio
async def test_user_override_wins(self) -> None:
Expand All @@ -132,9 +131,9 @@ async def test_falls_back_to_server_default(self) -> None:
await stream.aclose()

@pytest.mark.asyncio
async def test_falls_back_to_constant_when_server_silent(self) -> None:
async def test_returns_none_when_server_silent(self) -> None:
stream = await _make_idle_stream(_make_detector())
try:
assert stream._resolve_effective_threshold(None) == THRESHOLD
assert stream._resolve_effective_threshold(None) is None
finally:
await stream.aclose()
Loading