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 @@ -179,9 +179,11 @@ async def _post(self, endpoint: str, payload: dict[str, Any]) -> dict[str, Any]:
Response data as a dictionary

Raises:
APIStatusError: If Tavus returns a non-retryable error, or a retryable one
persists after all retries
APIConnectionError: If the request fails after all retries
"""
for i in range(self._conn_options.max_retry):
for attempt in range(self._conn_options.max_retry + 1):
try:
async with self._session.post(
f"{self._api_url}/{endpoint}",
Expand All @@ -197,14 +199,30 @@ async def _post(self, endpoint: str, payload: dict[str, Any]) -> dict[str, Any]:
raise APIStatusError(
"Server returned an error", status_code=response.status, body=text
)
return await response.json() # type: ignore
except Exception as e:
if isinstance(e, APIConnectionError):
logger.warning("failed to call tavus api", extra={"error": str(e)})
else:
logger.exception("failed to call tavus api")

if i < self._conn_options.max_retry - 1:
await asyncio.sleep(self._conn_options.retry_interval)
try:
return await response.json() # type: ignore
except ValueError as e:
# Tavus already accepted the POST; retrying could create a duplicate.
raise APIConnectionError(
"Tavus returned an invalid response", retryable=False
) from e
except APIStatusError as e:
# A 4xx such as a bad API key will fail the same way every time.
if not e.retryable:
raise
logger.warning(
"failed to call tavus api",
extra={"attempt": attempt + 1, "status_code": e.status_code},
)
if attempt >= self._conn_options.max_retry:
raise
await asyncio.sleep(self._conn_options.retry_interval)
except (aiohttp.ClientError, asyncio.TimeoutError) as e:
logger.warning(
"failed to call tavus api", extra={"attempt": attempt + 1, "error": str(e)}
)
if attempt >= self._conn_options.max_retry:
raise APIConnectionError("Failed to call Tavus API after all retries") from e
await asyncio.sleep(self._conn_options.retry_interval)

raise APIConnectionError("Failed to call Tavus API after all retries")
130 changes: 130 additions & 0 deletions tests/test_plugin_tavus.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,130 @@
"""Tests for the Tavus plugin: retry classification when calling the Tavus API.

Same behaviour as the Anam fix in #7314: a 4xx fails on the first attempt with the
provider's status, a 5xx is retried `max_retry` times, and a 5xx that persists is
raised as the provider's error rather than a generic connection error.
"""

from __future__ import annotations

from typing import Any

import aiohttp
import pytest

from livekit.agents import APIConnectionError, APIConnectOptions, APIStatusError

pytestmark = pytest.mark.plugin("tavus")

_OPTS = APIConnectOptions(max_retry=2, retry_interval=0.0, timeout=1.0)


class _Response:
def __init__(self, status: int) -> None:
self.status = status
self.ok = status < 400

async def text(self) -> str:
return f"status {self.status}"

async def json(self) -> dict[str, Any]:
return {"conversation_id": "c1"}

async def __aenter__(self) -> _Response:
return self

async def __aexit__(self, *exc: object) -> None:
return None


class _ScriptedSession:
"""Answers each POST with the next status in `script`; an exception instance in
the script is raised instead."""

def __init__(self, script: list[int | BaseException]) -> None:
self._script = list(script)
self.posts = 0

def post(self, url: str, **kwargs: Any) -> _Response:
self.posts += 1
step = self._script.pop(0) if len(self._script) > 1 else self._script[0]
if isinstance(step, BaseException):
raise step
return _Response(step)


def _api(session: _ScriptedSession):
from livekit.plugins.tavus.api import TavusAPI

return TavusAPI(
api_key="test-key",
api_url="http://tavus.test",
conn_options=_OPTS,
session=session, # type: ignore[arg-type]
)


async def test_client_error_is_not_retried_and_keeps_its_status():
"""A bad API key fails the same way every time; retrying it only delays the
error and replaced the 401 with a generic 'after all retries' message."""
session = _ScriptedSession([401])

with pytest.raises(APIStatusError) as exc_info:
await _api(session)._post("conversations", {})

assert session.posts == 1
assert exc_info.value.status_code == 401


async def test_server_error_is_retried_until_success():
session = _ScriptedSession([503, 503, 200])

data = await _api(session)._post("conversations", {})

assert data == {"conversation_id": "c1"}
assert session.posts == 3


async def test_persistent_server_error_is_raised_after_max_retry_retries():
"""One initial attempt plus `max_retry` retries, like Anam, and the final error
is the provider's status rather than a generic connection error."""
session = _ScriptedSession([503])

with pytest.raises(APIStatusError) as exc_info:
await _api(session)._post("conversations", {})

assert session.posts == _OPTS.max_retry + 1
assert exc_info.value.status_code == 503


async def test_malformed_success_body_is_not_retried():
"""Tavus already accepted the POST, so retrying could create a duplicate
conversation; fail once with a non-retryable error chained to the decode error."""

class _BadJSONResponse(_Response):
async def json(self) -> dict[str, Any]:
raise ValueError("Expecting value: line 1 column 1 (char 0)")

class _BadJSONSession(_ScriptedSession):
def post(self, url: str, **kwargs: Any) -> _Response:
self.posts += 1
return _BadJSONResponse(200)

session = _BadJSONSession([200])

with pytest.raises(APIConnectionError) as exc_info:
await _api(session)._post("conversations", {})

assert session.posts == 1
assert exc_info.value.retryable is False
assert isinstance(exc_info.value.__cause__, ValueError)


async def test_network_error_is_retried_and_chained():
session = _ScriptedSession([aiohttp.ClientConnectionError("connection refused")])

with pytest.raises(APIConnectionError) as exc_info:
await _api(session)._post("conversations", {})

assert session.posts == _OPTS.max_retry + 1
assert isinstance(exc_info.value.__cause__, aiohttp.ClientConnectionError)
Loading