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
19 changes: 19 additions & 0 deletions livekit-agents/livekit/agents/llm/_provider_format/mistralai.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,9 +16,15 @@ class MistralFormatData:

def to_conversations_ctx(
chat_ctx: llm.ChatContext,
*,
inject_dummy_user_message: bool = True,
) -> tuple[list[dict], MistralFormatData]:
"""Convert ChatContext to Mistral Conversations API entry format.

Mistral requires the last serving turn to be User or Tool (HTTP 400 /
code 3230 otherwise). When ``inject_dummy_user_message`` is True, append
a dummy user turn if the context ends on an assistant message.

Returns:
A tuple of (entries, instructions) where instructions is the extracted
system/developer message content (or None if absent).
Expand Down Expand Up @@ -65,9 +71,22 @@ def to_conversations_ctx(
}
)

if inject_dummy_user_message and not _last_serving_role_ok(entries):
entries.append({"type": "message.input", "role": "user", "content": "."})

return entries, MistralFormatData(instructions=instructions)


def _last_serving_role_ok(entries: list[dict[str, Any]]) -> bool:
"""True when the last entry is a user message or a tool result."""
if not entries:
return False
last = entries[-1]
if last.get("type") == "function.result":
return True
return last.get("type") == "message.input" or last.get("role") == "user"


def _to_entry(item: llm.ChatItem) -> dict[str, Any] | None:
if not isinstance(item, llm.ChatMessage):
return None
Expand Down
4 changes: 2 additions & 2 deletions livekit-agents/livekit/agents/llm/chat_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -711,7 +711,7 @@ def to_provider_format(

@overload
def to_provider_format(
self, format: Literal["mistralai"]
self, format: Literal["mistralai"], *, inject_dummy_user_message: bool = True
) -> tuple[list[dict], _provider_format.mistralai.MistralFormatData]: ...

@overload
Expand Down Expand Up @@ -746,7 +746,7 @@ def to_provider_format(
elif format == "anthropic":
return _provider_format.anthropic.to_chat_ctx(self, **kwargs)
elif format == "mistralai":
return _provider_format.mistralai.to_conversations_ctx(self)
return _provider_format.mistralai.to_conversations_ctx(self, **kwargs)
else:
raise ValueError(f"Unsupported provider format: {format}")

Expand Down
50 changes: 50 additions & 0 deletions tests/test_chat_ctx.py
Original file line number Diff line number Diff line change
Expand Up @@ -875,3 +875,53 @@ def test_to_provider_format_non_object_tool_arguments(fmt: str, arguments: str):

messages, _ = ctx.to_provider_format(format=fmt)
assert _tool_call_input(fmt, messages) == {}


def test_mistralai_injects_dummy_user_when_last_message_is_assistant():
"""Mistral serving requires last role User or Tool; a trailing assistant 400s."""
ctx = ChatContext.empty()
ctx.add_message(role="user", content="Hello")
ctx.add_message(role="assistant", content="One moment...")

entries, _ = ctx.to_provider_format(format="mistralai")

assert entries[-1] == {"type": "message.input", "role": "user", "content": "."}
assert entries[-2] == {
"type": "message.output",
"role": "assistant",
"content": "One moment...",
}


def test_mistralai_does_not_inject_when_last_is_user():
ctx = ChatContext.empty()
ctx.add_message(role="user", content="Hello")

entries, _ = ctx.to_provider_format(format="mistralai")

assert entries == [{"type": "message.input", "role": "user", "content": "Hello"}]


def test_mistralai_does_not_inject_when_last_is_tool_result():
ctx = ChatContext.empty()
ctx.add_message(role="user", content="lookup")
ctx.insert(FunctionCall(call_id="c1", name="lookup", arguments="{}"))
ctx.insert(FunctionCallOutput(call_id="c1", name="lookup", output="ok", is_error=False))

entries, _ = ctx.to_provider_format(format="mistralai")

assert entries[-1] == {
"type": "function.result",
"tool_call_id": "c1",
"result": "ok",
}


def test_mistralai_can_disable_dummy_user_injection():
ctx = ChatContext.empty()
ctx.add_message(role="user", content="Hello")
ctx.add_message(role="assistant", content="Hi")

entries, _ = ctx.to_provider_format(format="mistralai", inject_dummy_user_message=False)

assert entries[-1] == {"type": "message.output", "role": "assistant", "content": "Hi"}