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
10 changes: 7 additions & 3 deletions livekit-agents/livekit/agents/_exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -135,14 +135,18 @@ def __str__(self) -> str:
# naming our own type by .message keeps a cycle that ends on one out of __str__, which
# would otherwise re-enter here through the same chain.
if isinstance(root, APIConnectionError):
detail = f"{type(root).__name__}: {root.message}"
root_message = root.message
else:
# a third-party __str__ may raise (e.g. aiohttp.ClientConnectorError on a partially
# initialized connection key); the type name alone still beats losing the message.
try:
detail = f"{type(root).__name__}: {root}"
root_message = str(root)
except Exception:
detail = type(root).__name__
root_message = ""

detail = type(root).__name__
if root_message:
detail += f": {root_message}"

return f"{self.message} (caused by {detail})"

Expand Down
31 changes: 29 additions & 2 deletions livekit-agents/livekit/agents/llm/realtime_fallback_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
import weakref
from collections.abc import AsyncIterable, Callable
from dataclasses import dataclass
from typing import TYPE_CHECKING, Literal
from typing import TYPE_CHECKING, Any, Literal

from livekit import rtc

Expand Down Expand Up @@ -156,7 +156,7 @@ async def aclose(self) -> None:
await model.aclose()


class _FallbackRealtimeSession(RealtimeSession[Literal["realtime_availability_changed"]]):
class _FallbackRealtimeSession(RealtimeSession[str]):
"""Bound once by AgentActivity; swaps the inner child session internally."""

def __init__(
Expand All @@ -183,6 +183,7 @@ def _forward(ev: object) -> None:
self._forwarders: dict[EventTypes, Callable[[object], None]] = {
event: _make_forwarder(event) for event in _FORWARDED_EVENTS
}
self._extra_forwarders: dict[str, Callable[..., None]] = {}

# per-model availability, with a cooldown after a failure
self._available = [True] * len(adapter._models)
Expand All @@ -197,6 +198,7 @@ def _forward(ev: object) -> None:
self._swapping = False

self._active_index = 0
self._active_bound = False
self._active = adapter._models[0].session(
turn_detection_disabled=self._turn_detection_disabled
)
Expand All @@ -207,13 +209,38 @@ def _forward(ev: object) -> None:
def _bind(self, child: RealtimeSession) -> None:
for event, forwarder in self._forwarders.items():
child.on(event, forwarder)
for extra_event, forwarder in self._extra_forwarders.items():
child.on(extra_event, forwarder)
child.on("error", self._on_child_error)
self._active_bound = True

def _unbind(self, child: RealtimeSession) -> None:
self._active_bound = False
for event, forwarder in self._forwarders.items():
child.off(event, forwarder)
for extra_event, forwarder in self._extra_forwarders.items():
child.off(extra_event, forwarder)
child.off("error", self._on_child_error)

def on(
self,
event: EventTypes | str,
callback: Callable[..., Any] | None = None,
) -> Callable[..., Any]:
if event not in _FORWARDED_EVENTS and event != "error":
forwarder = self._extra_forwarders.get(event)
if forwarder is None:

def _forward(*args: object) -> None:
self.emit(event, *args)

forwarder = _forward
self._extra_forwarders[event] = forwarder
if self._active_bound:
self._active.on(event, forwarder)

return super().on(event, callback)

def _set_available(self, index: int, available: bool) -> None:
if self._available[index] == available:
return
Expand Down
9 changes: 7 additions & 2 deletions tests/test_exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,9 +76,14 @@ def test_str_stops_after_ten_distinct_causes() -> None:
assert str(_wrap(root)) == "Connection error. (caused by APIConnectionError: wrapped 10)"


def test_timeout_subclass_keeps_its_own_message() -> None:
def test_str_omits_separator_for_empty_cause_message() -> None:
err = _wrap(TimeoutError(), APITimeoutError())
assert str(err).startswith("Request timed out. (caused by TimeoutError")
assert str(err) == "Request timed out. (caused by TimeoutError)"


def test_str_omits_separator_for_empty_connection_error_message() -> None:
err = _wrap(APIConnectionError(""))
assert str(err) == "Connection error. (caused by APIConnectionError)"


def test_raise_from_populates_the_message() -> None:
Expand Down
36 changes: 36 additions & 0 deletions tests/test_realtime_fallback.py
Original file line number Diff line number Diff line change
Expand Up @@ -262,6 +262,42 @@ async def test_restart_preserves_wrapper_subscribers() -> None:
assert received == ["after-restart"]


async def test_restart_preserves_provider_event_subscribers() -> None:
primary = FakeRealtimeModel()
adapter = RealtimeModelFallbackAdapter([primary])
session = adapter.session()
received: list[object] = []
session.on("provider_event", lambda ev: received.append(ev))

primary.active_session.emit("provider_event", "before-restart")
await adapter.restart_session()
primary.active_session.emit("provider_event", "after-restart")

assert received == ["before-restart", "after-restart"]


async def test_provider_event_subscribed_during_restart_skips_old_child() -> None:
primary = FakeRealtimeModel()
adapter = RealtimeModelFallbackAdapter([primary])
session = adapter.session()
old_child = primary.active_session
close_gate = asyncio.Event()
old_child.block_aclose = close_gate

restart_task = asyncio.create_task(adapter.restart_session())
await old_child.aclose_entered.wait()

received: list[object] = []
session.on("provider_event", lambda ev: received.append(ev))
old_child.emit("provider_event", "old-child")

close_gate.set()
await restart_task
primary.active_session.emit("provider_event", "new-child")

assert received == ["new-child"]


async def test_restart_emits_no_error() -> None:
primary = FakeRealtimeModel()
adapter = RealtimeModelFallbackAdapter([primary])
Expand Down
Loading