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
@@ -1,9 +1,12 @@
from __future__ import annotations

import asyncio
import os
from dataclasses import dataclass
from typing import Any

import httpx

from livekit.agents import APIConnectionError, APIStatusError, APITimeoutError, llm
from livekit.agents.llm import (
ChatChunk,
Expand All @@ -18,6 +21,7 @@
)
from livekit.agents.utils import is_given, shortuuid
from mistralai.client import Mistral
from mistralai.client.errors import HTTPValidationError, SDKError
from mistralai.client.models import (
CompletionArgs,
ConversationEvents,
Expand Down Expand Up @@ -260,6 +264,8 @@ async def _run(self) -> None:

except APITimeoutError:
raise APITimeoutError(retryable=retryable) from None
except (asyncio.TimeoutError, httpx.TimeoutException) as e:
raise APITimeoutError(retryable=retryable) from e
except APIStatusError as e:
raise APIStatusError(
e.message,
Expand All @@ -268,6 +274,14 @@ async def _run(self) -> None:
body=e.body,
retryable=retryable,
) from None
except (SDKError, HTTPValidationError) as e:
raise APIStatusError(
e.message,
status_code=e.status_code,
request_id=e.headers.get("x-request-id"),
body=e.body,
retryable=retryable,
) from e
except Exception as e:
raise APIConnectionError(retryable=retryable) from e

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,8 @@
import weakref
from dataclasses import dataclass

import httpx

from livekit import rtc
from livekit.agents import (
APIConnectionError,
Expand All @@ -24,7 +26,7 @@
from livekit.agents.utils import AudioBuffer, is_given
from livekit.agents.voice.io import TimedString
from mistralai.client import Mistral
from mistralai.client.errors import SDKError
from mistralai.client.errors import HTTPValidationError, SDKError
from mistralai.client.models import (
RealtimeTranscriptionError,
RealtimeTranscriptionSessionCreated,
Expand Down Expand Up @@ -232,10 +234,15 @@ async def _recognize_impl(
],
)

except SDKError as e:
if e.status_code in (408, 504):
raise APITimeoutError() from e
raise APIStatusError(e.message, status_code=e.status_code, body=e.body) from e
except (asyncio.TimeoutError, httpx.TimeoutException) as e:
raise APITimeoutError() from e
except (SDKError, HTTPValidationError) as e:
raise APIStatusError(
e.message,
status_code=e.status_code,
request_id=e.headers.get("x-request-id"),
body=e.body,
) from e
Comment thread
jeanprbt marked this conversation as resolved.
except Exception as e:
raise APIConnectionError() from e

Expand Down Expand Up @@ -449,9 +456,14 @@ async def _run(self) -> None:

except (APIStatusError, APITimeoutError, APIConnectionError):
raise
except SDKError as e:
if e.status_code in (408, 504):
raise APITimeoutError() from e
raise APIStatusError(e.message, status_code=e.status_code, body=e.body) from e
except (asyncio.TimeoutError, httpx.TimeoutException) as e:
raise APITimeoutError() from e
except (SDKError, HTTPValidationError) as e:
raise APIStatusError(
e.message,
status_code=e.status_code,
request_id=e.headers.get("x-request-id"),
body=e.body,
) from e
except Exception as e:
raise APIConnectionError() from e
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

import asyncio
import base64
import os
import struct
Expand All @@ -19,7 +20,7 @@
from livekit.agents.types import DEFAULT_API_CONNECT_OPTIONS, NOT_GIVEN, NotGivenOr
from livekit.agents.utils import is_given
from mistralai.client import Mistral
from mistralai.client.errors import SDKError
from mistralai.client.errors import HTTPValidationError, SDKError

from .models import TTSModels, TTSVoices

Expand Down Expand Up @@ -184,9 +185,14 @@ async def _run(self, output_emitter: tts.AudioEmitter) -> None:

output_emitter.flush()

except httpx.TimeoutException as e:
except (asyncio.TimeoutError, httpx.TimeoutException) as e:
raise APITimeoutError() from e
except SDKError as e:
raise APIStatusError(e.message, status_code=e.status_code, body=e.body) from e
except (SDKError, HTTPValidationError) as e:
raise APIStatusError(
e.message,
status_code=e.status_code,
request_id=e.headers.get("x-request-id"),
body=e.body,
) from e
except Exception as e:
raise APIConnectionError() from e
144 changes: 144 additions & 0 deletions tests/test_plugin_mistralai_llm.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,144 @@
from __future__ import annotations

import types
from collections.abc import Callable

import httpx
import pytest
from mistralai.client.errors import HTTPValidationError, HTTPValidationErrorData, SDKError
from mistralai.client.models import MessageOutputEvent

from livekit.agents import APIStatusError, llm
from livekit.agents.types import APIConnectOptions
from livekit.plugins.mistralai.llm import LLM

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


def _event(data: object) -> object:
return types.SimpleNamespace(data=data)


class _FakeConversations:
def __init__(self, handler: Callable[[int], object]) -> None:
self._handler = handler
self.calls = 0

async def start_stream_async(self, **kwargs: object) -> object:
self.calls += 1
result = self._handler(self.calls)
if isinstance(result, BaseException):
raise result
return result


def _model(conversations: _FakeConversations) -> LLM:
client = types.SimpleNamespace(
beta=types.SimpleNamespace(conversations=conversations),
)
return LLM(client=client, model="test-model") # type: ignore[arg-type]


def _sdk_error(status_code: int) -> SDKError:
body = '{"message":"provider error"}'
response = httpx.Response(
status_code,
headers={"content-type": "application/json", "x-request-id": "req_test"},
content=body,
request=httpx.Request("POST", "https://api.mistral.ai/v1/conversations"),
)
return SDKError("provider error", response, body)


def _validation_error() -> HTTPValidationError:
body = '{"detail":[{"msg":"invalid request"}]}'
response = httpx.Response(
422,
headers={"content-type": "application/json", "x-request-id": "req_validation"},
content=body,
request=httpx.Request("POST", "https://api.mistral.ai/v1/conversations"),
)
return HTTPValidationError(HTTPValidationErrorData(), response, body)


class TestLLMErrorHandling:
@pytest.mark.asyncio
async def test_client_error_is_not_retried(self) -> None:
conversations = _FakeConversations(lambda _: _sdk_error(400))
model = _model(conversations)
errors: list[llm.LLMError] = []
model.on("error", errors.append)

try:
with pytest.raises(APIStatusError) as excinfo:
async with model.chat(
chat_ctx=llm.ChatContext.empty(),
conn_options=APIConnectOptions(max_retry=3, retry_interval=0.0, timeout=1.0),
) as stream:
async for _ in stream:
pass

error = excinfo.value
assert conversations.calls == 1
assert error.status_code == 400
assert error.request_id == "req_test"
assert error.body == '{"message":"provider error"}'
assert error.retryable is False
assert len(errors) == 1
assert errors[0].error is error
assert errors[0].recoverable is False
finally:
await model.aclose()

@pytest.mark.asyncio
async def test_validation_error_is_not_retried(self) -> None:
conversations = _FakeConversations(lambda _: _validation_error())
model = _model(conversations)

try:
with pytest.raises(APIStatusError) as excinfo:
async with model.chat(
chat_ctx=llm.ChatContext.empty(),
conn_options=APIConnectOptions(max_retry=3, retry_interval=0.0, timeout=1.0),
) as stream:
async for _ in stream:
pass

error = excinfo.value
assert conversations.calls == 1
assert error.status_code == 422
assert error.request_id == "req_validation"
assert error.body == '{"detail":[{"msg":"invalid request"}]}'
assert error.retryable is False
finally:
await model.aclose()

@pytest.mark.asyncio
async def test_status_error_after_output_is_not_retried(self) -> None:
async def _events():
yield _event(MessageOutputEvent(id="msg_1", content="partial"))
raise _sdk_error(500)

conversations = _FakeConversations(lambda _: _events())
model = _model(conversations)
errors: list[llm.LLMError] = []
model.on("error", errors.append)
chunks: list[llm.ChatChunk] = []

try:
with pytest.raises(APIStatusError) as excinfo:
async with model.chat(
chat_ctx=llm.ChatContext.empty(),
conn_options=APIConnectOptions(max_retry=3, retry_interval=0.0, timeout=1.0),
) as stream:
async for chunk in stream:
chunks.append(chunk)

assert [chunk.delta.content for chunk in chunks if chunk.delta] == ["partial"]
assert conversations.calls == 1
assert excinfo.value.status_code == 500
assert excinfo.value.retryable is False
assert len(errors) == 1
assert errors[0].recoverable is False
finally:
await model.aclose()