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 @@ -1659,7 +1659,10 @@ async def _process_responses(self) -> None:

finally:
logger.info("main output response stream processing task exiting")
self._is_sess_active.clear()
# A retry can replace this task before the old task reaches its
# cleanup block. Keep the session active for the replacement task.
if asyncio.current_task() is self._response_task:
self._is_sess_active.clear()

async def _restart_session(self, ex: Exception) -> None:
# Get restart attempts from current generation, or 0 if no generation
Expand Down
Original file line number Diff line number Diff line change
@@ -1,9 +1,15 @@
import asyncio
from types import SimpleNamespace

import pytest

from livekit.plugins.aws.experimental.realtime import realtime_model
from livekit.plugins.aws.experimental.realtime.realtime_model import (
_is_recoverable_validation_error,
)

pytestmark = pytest.mark.unit


def test_system_instability_validation_error_is_recoverable() -> None:
exc = SimpleNamespace(message="System instability detected. Please retry your request.")
Expand All @@ -15,3 +21,73 @@ def test_unrecognized_validation_error_is_not_recoverable() -> None:
exc = SimpleNamespace(message="The provided request is invalid.")

assert _is_recoverable_validation_error(exc) is False


async def test_old_response_task_does_not_deactivate_restarted_session(
monkeypatch: pytest.MonkeyPatch,
) -> None:
class RetryableStreamError(Exception):
message = "stream dropped"

class OtherStreamError(Exception):
pass

for exception_name in (
"ValidationException",
"ThrottlingException",
"ModelNotReadyException",
"ModelErrorException",
"InvalidEventBytes",
"ModelTimeoutException",
):
monkeypatch.setattr(realtime_model, exception_name, OtherStreamError)
monkeypatch.setattr(realtime_model, "ModelStreamErrorException", RetryableStreamError)

class OutputStream:
async def receive(self):
raise RetryableStreamError("stream dropped")

class StreamResponse:
async def await_output(self):
return None, OutputStream()

session = object.__new__(realtime_model.RealtimeSession)
session._is_sess_active = asyncio.Event()
session._is_sess_active.set()
session._stream_ready = asyncio.Event()
session._stream_response = StreamResponse()
session._realtime_model = SimpleNamespace(_label="test")
session._events = {}
current_task = asyncio.current_task()
replacement_task = object()
session._response_task = current_task

async def restart_session(_exception):
session._response_task = replacement_task

session._restart_session = restart_session

await session._process_responses()

assert session._is_sess_active.is_set()


async def test_current_response_task_deactivates_session() -> None:
class OutputStream:
async def receive(self):
return None

class StreamResponse:
async def await_output(self):
return None, OutputStream()

session = object.__new__(realtime_model.RealtimeSession)
session._is_sess_active = asyncio.Event()
session._is_sess_active.set()
session._stream_ready = asyncio.Event()
session._stream_response = StreamResponse()
session._response_task = asyncio.current_task()

await session._process_responses()

assert not session._is_sess_active.is_set()