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
Original file line number Diff line number Diff line change
Expand Up @@ -281,6 +281,24 @@ class _DiscardedGeneration:
pass


# error codes that can never succeed on retry: the socket reconnects fine (the key is
# valid), but every generation fails and the server closes the connection again, so
# retrying loops forever. When one of these is received, the reconnect loop must stop.
_FATAL_ERROR_CODES = frozenset(
{
"insufficient_quota",
"invalid_api_key",
"account_deactivated",
"billing_hard_limit_reached",
}
)


def _is_fatal_error(error: object | None) -> bool:
code = getattr(error, "code", None) or getattr(error, "type", None)
return isinstance(code, str) and code in _FATAL_ERROR_CODES


class RealtimeModel(llm.RealtimeModel):
@overload
def __init__(
Expand Down Expand Up @@ -933,40 +951,51 @@ async def _reconnect() -> None:
self.emit("session_reconnected", llm.RealtimeSessionReconnectedEvent())

reconnecting = False
while not self._msg_ch.closed:
try:
ws_conn = await self._create_ws_conn()
if reconnecting:
await _reconnect()
num_retries = 0 # reset the retry counter
await self._run_ws(ws_conn)

except APIError as e:
if max_retries == 0 or not e.retryable:
try:
while not self._msg_ch.closed:
try:
ws_conn = await self._create_ws_conn()
if reconnecting:
await _reconnect()
num_retries = 0 # reset the retry counter
await self._run_ws(ws_conn)

except APIError as e:
if max_retries == 0 or not e.retryable:
self._emit_error(e, recoverable=False)
raise
elif num_retries == max_retries:
self._emit_error(e, recoverable=False)
raise APIConnectionError(
f"{self._realtime_model._provider_label} connection failed after {num_retries} attempts",
) from e
else:
self._emit_error(e, recoverable=True)

retry_interval = self._opts.conn_options._interval_for_retry(num_retries)
logger.warning(
f"{self._realtime_model._provider_label} connection failed, retrying in {retry_interval}s",
exc_info=e,
extra={"attempt": num_retries, "max_retries": max_retries},
)
await asyncio.sleep(retry_interval)
num_retries += 1

except Exception as e:
self._emit_error(e, recoverable=False)
raise
elif num_retries == max_retries:
self._emit_error(e, recoverable=False)
raise APIConnectionError(
f"{self._realtime_model._provider_label} connection failed after {num_retries} attempts",
) from e
else:
self._emit_error(e, recoverable=True)

retry_interval = self._opts.conn_options._interval_for_retry(num_retries)
logger.warning(
f"{self._realtime_model._provider_label} connection failed, retrying in {retry_interval}s",
exc_info=e,
extra={"attempt": num_retries, "max_retries": max_retries},
)
await asyncio.sleep(retry_interval)
num_retries += 1

except Exception as e:
self._emit_error(e, recoverable=False)
raise

reconnecting = True
reconnecting = True
finally:
# the session loop has exited (fatal server error, retries exhausted, or
# close); close any in-progress generation and fail any pending
# generate_reply futures so consumers don't hang and callers don't wait
# out their timeout
self._close_current_generation("session closed")
for fut in self._response_created_futures.values():
if not fut.done():
fut.set_exception(llm.RealtimeError("realtime session closed"))
self._response_created_futures.clear()

async def _create_ws_conn(self) -> aiohttp.ClientWebSocketResponse:
headers = {"User-Agent": "LiveKit Agents"}
Expand Down Expand Up @@ -1145,7 +1174,12 @@ async def _recv_task() -> None:
self._handle_error(RealtimeErrorEvent.construct(**event))
elif lk_oai_debug:
logger.debug(f"unhandled event: {event['type']}", extra={"event": event})
except Exception:
except Exception as e:
# terminal server errors (e.g. insufficient_quota) must break the recv
# loop so _main_task stops reconnecting; every other handler failure is
# logged and skipped
if isinstance(e, APIError) and not e.retryable:
raise
if event["type"] == "response.output_audio.delta":
event["delta"] = event["delta"][:10] + "..."
logger.exception("failed to handle event", extra={"event": event})
Expand Down Expand Up @@ -2138,16 +2172,18 @@ def _handle_response_done_but_not_complete(self, event: ResponseDoneEvent) -> No
else:
error_body = None
message = f"{provider_label} response failed with unknown error"
self._emit_error(
APIError(
message=message,
body=error_body,
retryable=True,
),
# all possible faulures undocumented by openai,
# so we assume optimistically all retryable/recoverable
recoverable=True,
# failures are largely undocumented by openai, so we assume optimistically
# recoverable unless the code is a known-fatal one (quota / auth / billing),
# which is raised so the recv loop breaks and _main_task stops reconnecting
recoverable = not _is_fatal_error(error_body)
error = APIError(
message=message,
body=error_body,
retryable=recoverable,
)
if not recoverable:
raise error
self._emit_error(error, recoverable=True)
elif event.response.status in {"cancelled", "incomplete"}:
status_details = event.response.status_details
if isinstance(status_details, str):
Expand Down Expand Up @@ -2181,14 +2217,18 @@ def _handle_error(self, event: RealtimeErrorEvent) -> None:
f"{provider_label} returned an error: {event.error}",
extra={"error": event.error},
)
self._emit_error(
APIError(
message=f"{provider_label} returned an error",
body=event.error,
retryable=True,
),
recoverable=True,
recoverable = not _is_fatal_error(event.error)
error = APIError(
message=f"{provider_label} returned an error",
body=event.error,
retryable=recoverable,
)
if not recoverable:
# terminal (e.g. insufficient_quota): raise instead of emitting; the recv loop
# re-raises it so _main_task emits it with recoverable=False and stops
# reconnecting
raise error
self._emit_error(error, recoverable=True)

# response errors are handled by _handle_response_done via _done_fut.
# error events here are for non-response errors (e.g. invalid request).
Expand Down
90 changes: 89 additions & 1 deletion tests/test_realtime/test_openai_realtime_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,9 @@
import pytest

from livekit.agents import llm
from livekit.agents._exceptions import APIError
from livekit.agents.llm.remote_chat_context import RemoteChatContext
from livekit.plugins.openai.realtime.realtime_model import RealtimeSession
from livekit.plugins.openai.realtime.realtime_model import RealtimeSession, _is_fatal_error

pytestmark = pytest.mark.unit

Expand All @@ -31,3 +32,90 @@ def test_update_chat_ctx_deletes_empty_remote_items() -> None:
if getattr(event, "type", None) == "conversation.item.delete"
]
assert delete_ids == ["audio_item"]


# --------------------------------------------------------------------------- #
# fatal error classification: a fatal error must break the recv loop so that
# _main_task stops reconnecting (raised as APIError(retryable=False))
# --------------------------------------------------------------------------- #


def test_is_fatal_error_matches_known_codes() -> None:
assert _is_fatal_error(SimpleNamespace(code="insufficient_quota"))
assert _is_fatal_error(SimpleNamespace(code=None, type="invalid_api_key"))
assert not _is_fatal_error(SimpleNamespace(code="server_error"))
assert not _is_fatal_error(SimpleNamespace())
assert not _is_fatal_error(None)


def _handle_error_session(capture: dict[str, object]) -> RealtimeSession:
return cast(
RealtimeSession,
SimpleNamespace(
_realtime_model=SimpleNamespace(_provider_label="openai"),
_emit_error=lambda error, recoverable: capture.update(recoverable=recoverable),
),
)


def test_handle_error_raises_on_fatal() -> None:
# a fatal code is raised (not emitted here): the recv loop re-raises it so
# _main_task emits it once with recoverable=False and stops reconnecting
captured: dict[str, object] = {}
session = _handle_error_session(captured)
event = SimpleNamespace(
error=SimpleNamespace(message="quota exceeded", code="insufficient_quota")
)
with pytest.raises(APIError) as exc_info:
RealtimeSession._handle_error(session, event)
assert exc_info.value.retryable is False
assert captured == {} # not emitted by the handler; _main_task owns the emit


def test_handle_error_emits_transient_as_recoverable() -> None:
captured: dict[str, object] = {}
session = _handle_error_session(captured)
event = SimpleNamespace(error=SimpleNamespace(message="server hiccup", code="server_error"))
RealtimeSession._handle_error(session, event)
assert captured["recoverable"] is True


def test_handle_error_ignores_cancellation_failed() -> None:
captured: dict[str, object] = {}
event = SimpleNamespace(error=SimpleNamespace(message="Cancellation failed: no response"))
RealtimeSession._handle_error(_handle_error_session(captured), event)
assert captured == {} # early return, nothing emitted


def test_response_done_failed_fatal_raises() -> None:
captured: dict[str, object] = {}
session = _handle_error_session(captured)
event = SimpleNamespace(
response=SimpleNamespace(
id="resp_1",
status="failed",
status_details=SimpleNamespace(
error=SimpleNamespace(type="insufficient_quota", code="insufficient_quota")
),
)
)
with pytest.raises(APIError) as exc_info:
RealtimeSession._handle_response_done_but_not_complete(session, event)
assert exc_info.value.retryable is False
assert captured == {}


def test_response_done_failed_transient_stays_recoverable() -> None:
captured: dict[str, object] = {}
session = _handle_error_session(captured)
event = SimpleNamespace(
response=SimpleNamespace(
id="resp_1",
status="failed",
status_details=SimpleNamespace(
error=SimpleNamespace(type="invalid_request_error", code="rate_limit_exceeded")
),
)
)
RealtimeSession._handle_response_done_but_not_complete(session, event)
assert captured["recoverable"] is True
Loading