From 2199e503522c03d0bd10bda102323545710b0f8d Mon Sep 17 00:00:00 2001 From: NekoPunch Date: Wed, 29 Jul 2026 04:56:31 -0700 Subject: [PATCH 1/8] feat(python): add Mistral chat client Implements native Mistral support (#7366) with streaming, tool calling, and structured output. Talks to the REST API directly over httpx: the mistralai SDK's pinned OpenTelemetry deps conflict with the workspace. --- .../workflows/integration-tests-manual.yml | 1 + .../workflows/python-integration-tests.yml | 8 +- .github/workflows/python-merge-tests.yml | 7 +- python/.env.example | 2 + .../core/agent_framework/mistral/__init__.py | 4 + .../core/agent_framework/mistral/__init__.pyi | 14 +- .../core/tests/core/test_mistral_namespace.py | 4 + python/packages/mistral/AGENTS.md | 33 +- python/packages/mistral/README.md | 56 +- .../agent_framework_mistral/__init__.py | 5 + .../agent_framework_mistral/_chat_client.py | 987 ++++++++++++++++++ .../_embedding_client.py | 146 ++- python/packages/mistral/pyproject.toml | 5 +- .../tests/mistral/test_mistral_chat_client.py | 978 +++++++++++++++++ .../mistral/test_mistral_embedding_client.py | 384 +++---- .../02-agents/providers/mistral/README.md | 6 +- .../providers/mistral/mistral_agent_basic.py | 99 ++ .../providers/mistral/mistral_embeddings.py | 39 +- python/samples/README.md | 4 + python/uv.lock | 43 +- 20 files changed, 2500 insertions(+), 325 deletions(-) create mode 100644 python/packages/mistral/agent_framework_mistral/_chat_client.py create mode 100644 python/packages/mistral/tests/mistral/test_mistral_chat_client.py create mode 100644 python/samples/02-agents/providers/mistral/mistral_agent_basic.py diff --git a/.github/workflows/integration-tests-manual.yml b/.github/workflows/integration-tests-manual.yml index f43e7bbbba6..209853d7749 100644 --- a/.github/workflows/integration-tests-manual.yml +++ b/.github/workflows/integration-tests-manual.yml @@ -132,4 +132,5 @@ jobs: AZURE_TENANT_ID: ${{ secrets.AZURE_TENANT_ID }} AZURE_SUBSCRIPTION_ID: ${{ secrets.AZURE_SUBSCRIPTION_ID }} FOUNDRY_MODELS_API_KEY: ${{ secrets.FOUNDRY_MODELS_API_KEY }} + MISTRAL_API_KEY: ${{ secrets.MISTRAL_API_KEY }} OPENAI__APIKEY: ${{ secrets.OPENAI__APIKEY }} diff --git a/.github/workflows/python-integration-tests.yml b/.github/workflows/python-integration-tests.yml index 2007583cdb0..62052fd9a7b 100644 --- a/.github/workflows/python-integration-tests.yml +++ b/.github/workflows/python-integration-tests.yml @@ -27,6 +27,8 @@ on: required: true FOUNDRY_MODELS_API_KEY: required: false + MISTRAL_API_KEY: + required: false OPENAI__APIKEY: required: true @@ -172,6 +174,9 @@ jobs: ANTHROPIC_API_KEY: ${{ secrets.ANTHROPIC_API_KEY }} ANTHROPIC_CHAT_MODEL: ${{ vars.ANTHROPIC_CHAT_MODEL_ID }} LOCAL_MCP_URL: ${{ vars.LOCAL_MCP__URL }} + MISTRAL_API_KEY: ${{ secrets.MISTRAL_API_KEY }} + MISTRAL_CHAT_MODEL: ${{ vars.MISTRAL_CHAT_MODEL_ID }} + MISTRAL_EMBEDDING_MODEL: ${{ vars.MISTRAL_EMBEDDING_MODEL_ID }} OLLAMA_MODEL: qwen2.5:1.5b OLLAMA_EMBEDDING_MODEL: nomic-embed-text defaults: @@ -232,11 +237,12 @@ jobs: fallback_url: ${{ env.LOCAL_MCP_URL }} - name: Prefer local MCP URL when available run: echo "LOCAL_MCP_URL=${{ steps.local-mcp.outputs.effective_url }}" >> "$GITHUB_ENV" - - name: Test with pytest (Anthropic, Hyperlight, Ollama, MCP integration) + - name: Test with pytest (Anthropic, Hyperlight, Mistral, Ollama, MCP integration) run: > uv run pytest --import-mode=importlib packages/anthropic/tests packages/hyperlight/tests + packages/mistral/tests packages/ollama/tests packages/core/tests/core/test_mcp.py packages/hosting-mcp/tests diff --git a/.github/workflows/python-merge-tests.yml b/.github/workflows/python-merge-tests.yml index 88ca8f642d8..3113ca1664a 100644 --- a/.github/workflows/python-merge-tests.yml +++ b/.github/workflows/python-merge-tests.yml @@ -67,6 +67,7 @@ jobs: misc: - 'python/packages/anthropic/**' - 'python/packages/hyperlight/**' + - 'python/packages/mistral/**' - 'python/packages/ollama/**' - 'python/packages/core/agent_framework/_mcp.py' - 'python/packages/core/tests/core/test_mcp.py' @@ -278,6 +279,9 @@ jobs: ANTHROPIC_API_KEY: ${{ secrets.ANTHROPIC_API_KEY }} ANTHROPIC_CHAT_MODEL: ${{ vars.ANTHROPIC_CHAT_MODEL_ID }} LOCAL_MCP_URL: ${{ vars.LOCAL_MCP__URL }} + MISTRAL_API_KEY: ${{ secrets.MISTRAL_API_KEY }} + MISTRAL_CHAT_MODEL: ${{ vars.MISTRAL_CHAT_MODEL_ID }} + MISTRAL_EMBEDDING_MODEL: ${{ vars.MISTRAL_EMBEDDING_MODEL_ID }} OLLAMA_MODEL: qwen2.5:1.5b OLLAMA_EMBEDDING_MODEL: nomic-embed-text defaults: @@ -335,11 +339,12 @@ jobs: fallback_url: ${{ env.LOCAL_MCP_URL }} - name: Prefer local MCP URL when available run: echo "LOCAL_MCP_URL=${{ steps.local-mcp.outputs.effective_url }}" >> "$GITHUB_ENV" - - name: Test with pytest (Anthropic, Hyperlight, Ollama, MCP integration) + - name: Test with pytest (Anthropic, Hyperlight, Mistral, Ollama, MCP integration) run: > uv run pytest --import-mode=importlib packages/anthropic/tests packages/hyperlight/tests + packages/mistral/tests packages/ollama/tests packages/core/tests/core/test_mcp.py packages/hosting-mcp/tests diff --git a/python/.env.example b/python/.env.example index 26fed3fb1c1..6811c4673ee 100644 --- a/python/.env.example +++ b/python/.env.example @@ -46,7 +46,9 @@ OLLAMA_ENDPOINT="" OLLAMA_MODEL="" # Mistral AI MISTRAL_API_KEY="" +MISTRAL_CHAT_MODEL="" MISTRAL_EMBEDDING_MODEL="" +MISTRAL_SERVER_URL="" # Observability (instrumentation is enabled by default; set "ENABLE_INSTRUMENTATION" to "false" to opt out) ENABLE_SENSITIVE_DATA=true OTEL_EXPORTER_OTLP_ENDPOINT="http://localhost:4317/" diff --git a/python/packages/core/agent_framework/mistral/__init__.py b/python/packages/core/agent_framework/mistral/__init__.py index c4d742a75f1..db583d2ada3 100644 --- a/python/packages/core/agent_framework/mistral/__init__.py +++ b/python/packages/core/agent_framework/mistral/__init__.py @@ -9,9 +9,13 @@ from typing import Any _IMPORTS: dict[str, tuple[str, str]] = { + "MistralChatClient": ("agent_framework_mistral", "agent-framework-mistral"), + "MistralChatOptions": ("agent_framework_mistral", "agent-framework-mistral"), "MistralEmbeddingClient": ("agent_framework_mistral", "agent-framework-mistral"), "MistralEmbeddingOptions": ("agent_framework_mistral", "agent-framework-mistral"), "MistralEmbeddingSettings": ("agent_framework_mistral", "agent-framework-mistral"), + "MistralSettings": ("agent_framework_mistral", "agent-framework-mistral"), + "RawMistralChatClient": ("agent_framework_mistral", "agent-framework-mistral"), } diff --git a/python/packages/core/agent_framework/mistral/__init__.pyi b/python/packages/core/agent_framework/mistral/__init__.pyi index 0868d21e02d..a42afaf73df 100644 --- a/python/packages/core/agent_framework/mistral/__init__.pyi +++ b/python/packages/core/agent_framework/mistral/__init__.pyi @@ -1,9 +1,21 @@ # Copyright (c) Microsoft. All rights reserved. -from agent_framework_mistral import MistralEmbeddingClient, MistralEmbeddingOptions, MistralEmbeddingSettings +from agent_framework_mistral import ( + MistralChatClient, + MistralChatOptions, + MistralEmbeddingClient, + MistralEmbeddingOptions, + MistralEmbeddingSettings, + MistralSettings, + RawMistralChatClient, +) __all__ = [ + "MistralChatClient", + "MistralChatOptions", "MistralEmbeddingClient", "MistralEmbeddingOptions", "MistralEmbeddingSettings", + "MistralSettings", + "RawMistralChatClient", ] diff --git a/python/packages/core/tests/core/test_mistral_namespace.py b/python/packages/core/tests/core/test_mistral_namespace.py index 1f6cbb5a6cb..7bff7a28e29 100644 --- a/python/packages/core/tests/core/test_mistral_namespace.py +++ b/python/packages/core/tests/core/test_mistral_namespace.py @@ -11,9 +11,13 @@ def test_mistral_namespace_dir_lists_lazy_exports() -> None: names = dir(mistral) for expected in ( + "MistralChatClient", + "MistralChatOptions", "MistralEmbeddingClient", "MistralEmbeddingOptions", "MistralEmbeddingSettings", + "MistralSettings", + "RawMistralChatClient", ): assert expected in names diff --git a/python/packages/mistral/AGENTS.md b/python/packages/mistral/AGENTS.md index 6edd27d6ac6..92e837ba7aa 100644 --- a/python/packages/mistral/AGENTS.md +++ b/python/packages/mistral/AGENTS.md @@ -1,26 +1,51 @@ # Mistral Package (agent-framework-mistral) -Integration with Mistral AI for embedding generation. +Integration with Mistral AI for chat completions and embedding generation. + +## Implementation Notes + +- Talks to the Mistral REST API directly over `httpx`; the official `mistralai` SDK is not used + because its pinned OpenTelemetry requirements conflict with the rest of the framework. ## Main Classes +- **`MistralChatClient`** - Chat client for Mistral AI models with function invocation, middleware, and telemetry +- **`RawMistralChatClient`** - Chat client without the batteries-included layers +- **`MistralChatOptions`** - Options TypedDict for Mistral-specific chat parameters +- **`MistralSettings`** - TypedDict settings for Mistral chat configuration - **`MistralEmbeddingClient`** - Embedding client for Mistral AI models - **`MistralEmbeddingOptions`** - Options TypedDict for Mistral-specific embedding parameters - **`MistralEmbeddingSettings`** - TypedDict settings for Mistral configuration ## Usage +```python +from agent_framework import Agent +from agent_framework.mistral import MistralChatClient + +# Requires MISTRAL_API_KEY environment variable (or pass api_key= directly) +client = MistralChatClient(model="mistral-large-latest") +try: + agent = Agent(client=client) + result = await agent.run("Hello!") +finally: + await client.close() +``` + ```python from agent_framework.mistral import MistralEmbeddingClient # Requires MISTRAL_API_KEY environment variable (or pass api_key= directly) client = MistralEmbeddingClient(model="mistral-embed") -result = await client.get_embeddings(["Hello, world!"]) -print(result[0].vector) +try: + result = await client.get_embeddings(["Hello, world!"]) + print(result[0].vector) +finally: + await client.close() ``` ## Import Path ```python -from agent_framework.mistral import MistralEmbeddingClient +from agent_framework.mistral import MistralChatClient, MistralEmbeddingClient ``` diff --git a/python/packages/mistral/README.md b/python/packages/mistral/README.md index 0912bd421d1..ccf4bd3c7fd 100644 --- a/python/packages/mistral/README.md +++ b/python/packages/mistral/README.md @@ -8,7 +8,41 @@ pip install agent-framework-mistral --pre and see the [README](https://github.com/microsoft/agent-framework/tree/main/python/README.md) for more information. -See the [Mistral embedding sample](../../samples/02-agents/providers/mistral/mistral_embeddings.py) for a runnable example. +See the [Mistral agent sample](../../samples/02-agents/providers/mistral/mistral_agent_basic.py) and the +[Mistral embedding sample](../../samples/02-agents/providers/mistral/mistral_embeddings.py) for runnable examples. + +## Chat Client + +The `MistralChatClient` provides chat completions using Mistral AI models, with support for +streaming, function tools, and structured output. + +### Quick Start + +```python +from agent_framework import Agent, Message +from agent_framework.mistral import MistralChatClient + +# Using environment variables (MISTRAL_API_KEY, MISTRAL_CHAT_MODEL) +client = MistralChatClient() +try: + agent = Agent(client=client, instructions="You are a helpful assistant.") + result = await agent.run("Tell me a joke.") + + # Parameters can also be passed directly: + # MistralChatClient(model="mistral-large-latest", api_key="your-api-key") + response = await client.get_response([Message("user", ["Hello!"])]) + print(response.text) +finally: + await client.close() +``` + +### Configuration + +| Environment Variable | Description | +|---|---| +| `MISTRAL_API_KEY` | Your Mistral AI API key | +| `MISTRAL_CHAT_MODEL` | Chat model name (e.g., `mistral-large-latest`) | +| `MISTRAL_SERVER_URL` | Optional server URL override | ## Embedding Client @@ -22,17 +56,15 @@ from agent_framework.mistral import MistralEmbeddingClient # Using environment variables (MISTRAL_API_KEY, MISTRAL_EMBEDDING_MODEL) client = MistralEmbeddingClient() -# Or passing parameters directly -client = MistralEmbeddingClient( - model="mistral-embed", - api_key="your-api-key", -) - -# Generate embeddings -result = await client.get_embeddings(["Hello, world!", "How are you?"]) -for embedding in result: - print(f"Dimensions: {embedding.dimensions}") - print(f"Vector: {embedding.vector[:5]}...") +try: + # Parameters can also be passed directly: + # MistralEmbeddingClient(model="mistral-embed", api_key="your-api-key") + result = await client.get_embeddings(["Hello, world!", "How are you?"]) + for embedding in result: + print(f"Dimensions: {embedding.dimensions}") + print(f"Vector: {embedding.vector[:5]}...") +finally: + await client.close() ``` ### Configuration diff --git a/python/packages/mistral/agent_framework_mistral/__init__.py b/python/packages/mistral/agent_framework_mistral/__init__.py index 58d4677a82f..66032d1a501 100644 --- a/python/packages/mistral/agent_framework_mistral/__init__.py +++ b/python/packages/mistral/agent_framework_mistral/__init__.py @@ -2,6 +2,7 @@ import importlib.metadata +from ._chat_client import MistralChatClient, MistralChatOptions, MistralSettings, RawMistralChatClient from ._embedding_client import MistralEmbeddingClient, MistralEmbeddingOptions, MistralEmbeddingSettings try: @@ -10,8 +11,12 @@ __version__ = "0.0.0" # Fallback for development mode __all__ = [ + "MistralChatClient", + "MistralChatOptions", "MistralEmbeddingClient", "MistralEmbeddingOptions", "MistralEmbeddingSettings", + "MistralSettings", + "RawMistralChatClient", "__version__", ] diff --git a/python/packages/mistral/agent_framework_mistral/_chat_client.py b/python/packages/mistral/agent_framework_mistral/_chat_client.py new file mode 100644 index 00000000000..dbe0348511f --- /dev/null +++ b/python/packages/mistral/agent_framework_mistral/_chat_client.py @@ -0,0 +1,987 @@ +# Copyright (c) Microsoft. All rights reserved. + +from __future__ import annotations + +import hashlib +import json +import logging +import re +import sys +from collections.abc import AsyncIterable, Awaitable, Mapping, Sequence +from datetime import datetime, timezone +from typing import Any, ClassVar, Generic, Literal, cast + +import httpx +from agent_framework import ( + BaseChatClient, + ChatAndFunctionMiddlewareTypes, + ChatMiddlewareLayer, + ChatOptions, + ChatResponse, + ChatResponseUpdate, + Content, + FinishReasonLiteral, + FunctionInvocationConfiguration, + FunctionInvocationLayer, + FunctionTool, + Message, + ResponseStream, + UsageDetails, + validate_tool_mode, +) +from agent_framework._settings import SecretString, load_settings +from agent_framework._telemetry import get_user_agent +from agent_framework._types import prepend_instructions_to_messages +from agent_framework.exceptions import ( + ChatClientException, + ChatClientInvalidAuthException, + ChatClientInvalidRequestException, + ChatClientInvalidResponseException, +) +from agent_framework.observability import ChatTelemetryLayer +from pydantic import BaseModel + +if sys.version_info >= (3, 13): + from typing import TypeVar # pragma: no cover +else: + from typing_extensions import TypeVar # pragma: no cover + +if sys.version_info >= (3, 12): + from typing import override # pragma: no cover +else: + from typing_extensions import override # pragma: no cover + +if sys.version_info >= (3, 11): + from typing import TypedDict # pragma: no cover +else: + from typing_extensions import TypedDict # pragma: no cover + +logger = logging.getLogger("agent_framework.mistral") + +ResponseModelT = TypeVar("ResponseModelT", bound=BaseModel | None, default=None) + + +# region Options & Settings + + +class MistralChatOptions(ChatOptions[ResponseModelT], Generic[ResponseModelT], total=False): + """Mistral AI-specific chat options. + + Extends ``ChatOptions`` with Mistral-specific fields. Standard options are mapped to their + Mistral chat-completion equivalents; Mistral-specific fields are declared below. + + See: https://docs.mistral.ai/api/#tag/chat + + Inherited fields from ``ChatOptions``: + model: Model to use for this call (e.g. ``"mistral-large-latest"``). + temperature: Controls randomness. Higher values produce more varied output. + max_tokens: Maximum number of tokens to generate. + top_p: Nucleus sampling cutoff. + stop: One or more sequences that stop generation when encountered. + seed: Fixed seed for reproducible outputs, translates to ``random_seed``. + frequency_penalty: Reduces repetition by penalising frequent tokens. + presence_penalty: Reduces repetition by penalising tokens already present. + tools: Function tools the model may call. + tool_choice: How the model picks a tool. One of ``'auto'``, ``'none'``, or ``'required'``. + allow_multiple_tool_calls: Translates to ``parallel_tool_calls``. + response_format: Pydantic model type or JSON schema mapping for structured JSON output. + The response text is parsed and exposed via ``ChatResponse.value``. + instructions: Extra system-level instructions prepended to the system message. + metadata: Arbitrary key/value metadata attached to the request. + + Not supported, and passing these raises a type error: + - ``logit_bias`` + - ``store`` + - ``user`` + - ``conversation_id`` + """ + + safe_prompt: bool + """Whether to inject a safety prompt before all conversations.""" + + prompt_mode: str + """Toggle between reasoning mode and no system prompt (e.g. ``"reasoning"``).""" + + prediction: dict[str, Any] + """Predicted output to optimize response time when large parts of the response are known.""" + + n: int + """Number of completions to return.""" + + guardrails: list[dict[str, Any]] + """Guardrail configurations applied to the request.""" + + prompt_cache_key: str + """Cache key shared by requests with the same prompt prefix.""" + + reasoning_effort: Literal["none", "minimal", "low", "medium", "high", "xhigh"] + """Effort level for models that support reasoning.""" + + # Unsupported base options. Override with None to indicate not supported + logit_bias: None # type: ignore[misc] + """Not supported in the Mistral API.""" + + store: None # type: ignore[misc] + """Not supported in the Mistral API.""" + + user: None # type: ignore[misc] + """Not supported in the Mistral API.""" + + conversation_id: None # type: ignore[misc] + """Not supported in the Mistral API.""" + + +MistralChatOptionsT = TypeVar("MistralChatOptionsT", bound=TypedDict, default="MistralChatOptions", covariant=True) # type: ignore[valid-type] + + +class MistralSettings(TypedDict, total=False): + """Mistral AI chat settings. + + Fields: + api_key: Mistral API key. Resolved from ``MISTRAL_API_KEY``. + chat_model: Chat model name. Resolved from ``MISTRAL_CHAT_MODEL``. + server_url: Optional server URL override. Resolved from ``MISTRAL_SERVER_URL``. + """ + + api_key: SecretString | None + chat_model: str | None + server_url: str | None + + +# endregion + +_MISTRAL_API_BASE_URL = "https://api.mistral.ai" +_CHAT_COMPLETIONS_PATH = "/v1/chat/completions" +_DEFAULT_TIMEOUT_SECONDS = 60.0 +_SSE_DATA_PREFIX = "data:" +_SSE_DONE = "[DONE]" + +# Keys mapping to a different Mistral chat-completion parameter name +_OPTION_TRANSLATIONS: dict[str, str] = { + "seed": "random_seed", + "allow_multiple_tool_calls": "parallel_tool_calls", +} + +# Keys handled with dedicated logic, not via the generic passthrough +_OPTION_EXPLICIT_KEYS: frozenset[str] = frozenset( + { + "tools", + "tool_choice", + "response_format", + } +) + +# Keys consumed upstream and not forwarded to the Mistral API +_OPTION_CONSUMED_KEYS: frozenset[str] = frozenset( + { + "model", + "instructions", + } +) + +_OPTION_EXCLUDE_KEYS: frozenset[str] = _OPTION_EXPLICIT_KEYS | _OPTION_CONSUMED_KEYS + +_FINISH_REASON_MAP: dict[str, FinishReasonLiteral] = { + "stop": "stop", + "length": "length", + "model_length": "length", + "tool_calls": "tool_calls", +} + +# La Plateforme requires tool call IDs to be exactly 9 alphanumeric characters. +_MISTRAL_TOOL_CALL_ID_PATTERN = re.compile(r"^[a-zA-Z0-9]{9}$") + + +def _sanitize_tool_call_id(call_id: str) -> str: + """Return a Mistral-compatible tool call ID, deterministically derived when needed.""" + if _MISTRAL_TOOL_CALL_ID_PATTERN.match(call_id): + return call_id + return hashlib.sha256(call_id.encode("utf-8")).hexdigest()[:9] + + +def _tool_call_id_of(tool_call: Mapping[str, Any]) -> str: + """Return the wire tool call ID, treating null/"null" placeholders as missing.""" + call_id = tool_call.get("id") + if isinstance(call_id, str) and call_id and call_id != "null": + return call_id + return "" + + +class _StreamedToolCall: + """Accumulates the fragments of one streamed tool call. + + Mistral streams may fragment a tool call across chunks — the ID typically arrives + only on the first fragment, and the name and arguments may arrive in pieces. + """ + + __slots__ = ("arguments_dict", "arguments_text", "call_id", "name", "raw") + + def __init__(self) -> None: + self.call_id = "" + self.name = "" + self.arguments_text = "" + self.arguments_dict: dict[str, Any] | None = None + self.raw: Any = None + + def add(self, tool_call: Mapping[str, Any]) -> None: + if not self.call_id: + self.call_id = _tool_call_id_of(tool_call) + function: Mapping[str, Any] = tool_call.get("function") or {} + if name := function.get("name"): + self.name += name + arguments = function.get("arguments") + if isinstance(arguments, dict): + self.arguments_dict = {**(self.arguments_dict or {}), **cast("dict[str, Any]", arguments)} + elif isinstance(arguments, str): + self.arguments_text += arguments + self.raw = tool_call + + def to_content(self, fallback_id: str) -> Content: + arguments = self.arguments_dict if self.arguments_dict is not None else self.arguments_text + return Content.from_function_call( + call_id=self.call_id or fallback_id, + name=self.name, + arguments=arguments, + raw_representation=self.raw, + ) + + +class RawMistralChatClient( + BaseChatClient[MistralChatOptionsT], + Generic[MistralChatOptionsT], +): + """A raw Mistral AI chat client. + + Talks to the Mistral REST API directly over HTTP; the ``mistralai`` SDK is not required. + + Use this when you want full control over the request pipeline. For instance, to opt out of + telemetry, use custom middleware, or compose your own layers. If you want the full-featured + client with batteries included, use `MistralChatClient` instead. + """ + + OTEL_PROVIDER_NAME: ClassVar[str] = "mistralai" + + INJECTABLE: ClassVar[set[str]] = {"client"} + + def __init__( + self, + *, + model: str | None = None, + api_key: str | SecretString | None = None, + server_url: str | None = None, + client: httpx.AsyncClient | None = None, + additional_properties: dict[str, Any] | None = None, + env_file_path: str | None = None, + env_file_encoding: str | None = None, + ) -> None: + """Create a raw Mistral AI chat client. + + Keyword Args: + model: The Mistral chat model to use (e.g. "mistral-large-latest"). + Can also be set via environment variable ``MISTRAL_CHAT_MODEL``. + api_key: Mistral API key. Defaults to ``MISTRAL_API_KEY`` environment variable. + server_url: Optional server URL override. Defaults to ``MISTRAL_SERVER_URL`` + environment variable, or the Mistral default. + client: Optional pre-configured ``httpx.AsyncClient``. When provided, api_key is + not required and the client is expected to carry its own auth headers and + base URL. + additional_properties: Additional properties stored on the client instance. + env_file_path: Path to ``.env`` file for settings. + env_file_encoding: Encoding for ``.env`` file. + """ + mistral_settings = load_settings( + MistralSettings, + env_prefix="MISTRAL_", + required_fields=[] if client is not None else ["api_key"], + api_key=api_key, + chat_model=model, + server_url=server_url, + env_file_path=env_file_path, + env_file_encoding=env_file_encoding, + ) + + self.model = mistral_settings.get("chat_model") + self.server_url = mistral_settings.get("server_url") + self._owns_client = client is None + + if client is not None: + self.client = client + if self.server_url is None: + client_base_url = str(client.base_url).rstrip("/") + self.server_url = client_base_url or None + else: + resolved_api_key: SecretString = mistral_settings["api_key"] # type: ignore[assignment] + self.client = httpx.AsyncClient( + base_url=self.server_url or _MISTRAL_API_BASE_URL, + headers={ + "Authorization": f"Bearer {resolved_api_key.get_secret_value()}", + "User-Agent": get_user_agent(), + "Accept": "application/json", + }, + timeout=_DEFAULT_TIMEOUT_SECONDS, + ) + + super().__init__(additional_properties=additional_properties) + + async def close(self) -> None: + """Close the internally created HTTP client.""" + if self._owns_client: + await self.client.aclose() + + @override + def service_url(self) -> str: + """Get the URL of the service.""" + return self.server_url or _MISTRAL_API_BASE_URL + + @override + def _inner_get_response( + self, + *, + messages: Sequence[Message], + options: Mapping[str, Any], + stream: bool = False, + **kwargs: Any, + ) -> Awaitable[ChatResponse] | ResponseStream[ChatResponseUpdate, ChatResponse]: + if stream: + + async def _stream() -> AsyncIterable[ChatResponseUpdate]: + validated = await self._validate_options(options) + request = self._prepare_request(messages, validated, **kwargs) + request["stream"] = True + pending_tool_calls: dict[tuple[int, int], _StreamedToolCall] = {} + last_chunk_id = "" + try: + async with self.client.stream("POST", _CHAT_COMPLETIONS_PATH, json=request) as response: + await self._raise_for_status(response) + async for line in response.aiter_lines(): + chunk = self._parse_sse_line(line) + if chunk is None: + continue + last_chunk_id = chunk.get("id") or last_chunk_id + yield self._parse_chunk(chunk, pending_tool_calls) + except ChatClientException: + raise + except Exception as ex: + raise ChatClientException( + f"Mistral streaming chat request failed: {ex}", + inner_exception=ex, + ) from ex + if pending_tool_calls: + # The stream ended without a finish chunk; emit the accumulated calls. + yield ChatResponseUpdate( + contents=self._flush_pending_tool_calls(pending_tool_calls, last_chunk_id), + role="assistant", + ) + + return self._build_response_stream(_stream(), response_format=options.get("response_format")) + + async def _get_response() -> ChatResponse: + validated = await self._validate_options(options) + request = self._prepare_request(messages, validated, **kwargs) + try: + response = await self.client.post(_CHAT_COMPLETIONS_PATH, json=request) + await self._raise_for_status(response) + except ChatClientException: + raise + except Exception as ex: + raise ChatClientException(f"Mistral chat request failed: {ex}", inner_exception=ex) from ex + try: + raw_payload = response.json() + if not isinstance(raw_payload, Mapping): + raise ChatClientInvalidResponseException("Mistral chat response must be a JSON object.") + payload = cast("Mapping[str, Any]", raw_payload) + return self._parse_response(payload, response_format=validated.get("response_format")) + except ChatClientException: + raise + except Exception as ex: + raise ChatClientInvalidResponseException( + f"Mistral chat response was invalid: {ex}", + inner_exception=ex, + ) from ex + + return _get_response() + + @staticmethod + async def _raise_for_status(response: httpx.Response) -> None: + if response.status_code < 400: + return + body = (await response.aread()).decode("utf-8", errors="replace") + message = f"Mistral chat request failed with status {response.status_code}: {body[:2000]}" + if response.status_code in (401, 403): + raise ChatClientInvalidAuthException(message) + if response.status_code < 500: + raise ChatClientInvalidRequestException(message) + raise ChatClientException(message) + + @staticmethod + def _parse_sse_line(line: str) -> dict[str, Any] | None: + """Parse one server-sent-events line into a completion chunk, or None to skip.""" + line = line.strip() + if not line.startswith(_SSE_DATA_PREFIX): + return None + data = line[len(_SSE_DATA_PREFIX) :].strip() + if not data or data == _SSE_DONE: + return None + try: + parsed = json.loads(data) + except json.JSONDecodeError as ex: + raise ChatClientInvalidResponseException( + "Mistral streaming chat response contained malformed SSE data.", + inner_exception=ex, + ) from ex + if not isinstance(parsed, dict): + raise ChatClientInvalidResponseException("Mistral streaming chat SSE data must be a JSON object.") + return cast("dict[str, Any]", parsed) + + # region Request preparation + + def _prepare_request( + self, messages: Sequence[Message], options: Mapping[str, Any], **kwargs: Any + ) -> dict[str, Any]: + """Build the JSON body for a Mistral chat-completion request. + + Args: + messages: The conversation history as framework Message objects. + options: Validated and normalized chat options. + kwargs: Additional keyword arguments merged into the request body. + + Returns: + The request body for ``POST /v1/chat/completions``. + + Raises: + ValueError: If no model is set on the options or the client instance. + """ + model = options.get("model") or self.model + if not model: + raise ValueError( + "Mistral model is required. Set via model parameter or MISTRAL_CHAT_MODEL environment variable." + ) + + if instructions := options.get("instructions"): + messages = prepend_instructions_to_messages(list(messages), instructions, role="system") + + request: dict[str, Any] = { + "model": model, + "messages": self._prepare_mistral_messages(messages), + } + + for key, value in options.items(): + if key in _OPTION_EXCLUDE_KEYS or value is None: + continue + request[_OPTION_TRANSLATIONS.get(key, key)] = value + + if tools := self._prepare_tools(options.get("tools")): + request["tools"] = tools + if (tool_choice := self._prepare_tool_choice(options.get("tool_choice"))) is not None: + request["tool_choice"] = tool_choice + if (response_format := self._prepare_response_format(options.get("response_format"))) is not None: + request["response_format"] = response_format + + request.update(kwargs) + return request + + def _prepare_mistral_messages(self, messages: Sequence[Message]) -> list[dict[str, Any]]: + mistral_messages: list[dict[str, Any]] = [] + for message in messages: + match message.role: + case "system": + if message.text: + mistral_messages.append({"role": "system", "content": message.text}) + case "user": + mistral_messages.append(self._format_user_message(message)) + case "assistant": + mistral_messages.append(self._format_assistant_message(message)) + case "tool": + mistral_messages.extend(self._format_tool_messages(message)) + case _: + logger.debug("Skipping unsupported message role for Mistral: %s", message.role) + return mistral_messages + + def _format_user_message(self, message: Message) -> dict[str, Any]: + chunks: list[dict[str, Any]] = [] + text_only = True + for content in message.contents: + match content.type: + case "text": + chunks.append({"type": "text", "text": content.text or ""}) + case "data" | "uri": + chunk = self._convert_data_or_uri_content(content) + if chunk is not None: + chunks.append(chunk) + text_only = False + case _: + logger.debug("Skipping unsupported user content type for Mistral: %s", content.type) + + if text_only: + return {"role": "user", "content": message.text} + return {"role": "user", "content": chunks} + + def _convert_data_or_uri_content(self, content: Content) -> dict[str, Any] | None: + """Convert a ``data`` or ``uri`` Content to a Mistral content chunk. + + Images become ``image_url`` chunks (data URIs are passed through as-is). + PDF documents referenced by external URI become ``document_url`` chunks. + """ + uri = content.uri + if not uri: + logger.warning("Skipping %s content for Mistral: missing uri", content.type) + return None + + if content.has_top_level_media_type("image"): + return {"type": "image_url", "image_url": uri} + + if content.type == "uri" and content.media_type == "application/pdf": + return {"type": "document_url", "document_url": uri} + + logger.warning( + "Skipping unsupported %s content for Mistral: media_type=%s", + content.type, + content.media_type, + ) + return None + + def _format_assistant_message(self, message: Message) -> dict[str, Any]: + tool_calls: list[dict[str, Any]] = [] + for content in message.contents: + if content.type == "function_call": + arguments = content.arguments if isinstance(content.arguments, (str, Mapping)) else "{}" + if isinstance(arguments, Mapping): + arguments = dict(arguments) + tool_calls.append( + { + "id": _sanitize_tool_call_id(content.call_id or ""), + "type": "function", + "function": {"name": content.name or "", "arguments": arguments}, + } + ) + formatted: dict[str, Any] = {"role": "assistant", "content": message.text or None} + if tool_calls: + formatted["tool_calls"] = tool_calls + return formatted + + def _format_tool_messages(self, message: Message) -> list[dict[str, Any]]: + tool_messages: list[dict[str, Any]] = [] + for content in message.contents: + if content.type != "function_result": + continue + if content.items: + text_parts = [c.text or "" for c in content.items if c.type == "text"] + if any(c.type in ("data", "uri") for c in content.items): + logger.warning( + "Mistral does not support rich content (images, audio) in tool results. " + "Rich content items will be omitted." + ) + result_text = "\n".join(text_parts) + else: + result_text = self._result_to_text(content.result) + tool_message: dict[str, Any] = { + "role": "tool", + "content": result_text, + "tool_call_id": _sanitize_tool_call_id(content.call_id or ""), + } + if name := getattr(content, "name", None): + tool_message["name"] = name + tool_messages.append(tool_message) + return tool_messages + + @staticmethod + def _result_to_text(result: Any) -> str: + if result is None: + return "" + if isinstance(result, str): + return result + try: + return json.dumps(result) + except (TypeError, ValueError): + return str(result) + + def _prepare_tools(self, tools: Sequence[Any] | None) -> list[Any] | None: + """Translate the framework tool list into Mistral API tool definitions. + + ``FunctionTool`` instances are translated to Mistral function definitions; plain + mappings are passed through unchanged. + """ + if not tools: + return None + prepared: list[Any] = [] + for tool in tools: + if isinstance(tool, FunctionTool): + prepared.append( + { + "type": "function", + "function": { + "name": tool.name, + "description": tool.description or "", + "parameters": tool.parameters(), + }, + } + ) + else: + prepared.append(tool) + return prepared or None + + def _prepare_tool_choice(self, tool_choice: Any) -> Any | None: + """Build the Mistral ``tool_choice`` value from the framework ``tool_choice`` option.""" + tool_mode = validate_tool_mode(tool_choice) + if not tool_mode: + return None + + match tool_mode.get("mode"): + case "auto": + if "allowed_tools" in tool_mode: + logger.warning("Mistral does not support restricting auto tool choice to specific tools.") + return "auto" + case "none": + return "none" + case "required": + if name := tool_mode.get("required_function_name"): + return {"type": "function", "function": {"name": name}} + return "required" + case unknown_mode: + logger.warning("Unsupported tool_choice mode for Mistral: %s", unknown_mode) + return None + + def _prepare_response_format(self, response_format: Any) -> dict[str, Any] | None: + """Build a Mistral ``response_format`` object from the framework option. + + Supports Pydantic model types, raw JSON schema mappings, response-format envelopes + (``{"type": "json_object"}`` / ``{"type": "json_schema", "json_schema": {...}}``), + and the string ``"json"``. + """ + if response_format is None: + return None + + if isinstance(response_format, type) and issubclass(response_format, BaseModel): + return { + "type": "json_schema", + "json_schema": { + "name": response_format.__name__, + "schema": response_format.model_json_schema(), + "strict": True, + }, + } + + if isinstance(response_format, str): + if response_format in ("json", "json_object"): + return {"type": "json_object"} + logger.warning("Unsupported response_format string for Mistral: %s", response_format) + return None + + if isinstance(response_format, Mapping): + mapping: dict[str, Any] = dict(cast("Mapping[str, Any]", response_format)) + format_type = mapping.get("type") + if format_type == "json_object": + return {"type": "json_object"} + if format_type == "json_schema": + json_schema: dict[str, Any] = dict(mapping.get("json_schema") or {}) + return { + "type": "json_schema", + "json_schema": { + "name": json_schema.get("name", "response"), + "schema": json_schema.get("schema") or json_schema.get("schema_definition") or {}, + "strict": json_schema.get("strict"), + }, + } + # A raw JSON schema mapping + return { + "type": "json_schema", + "json_schema": { + "name": str(mapping.get("title", "response")), + "schema": mapping, + "strict": True, + }, + } + + type_name = type(cast(object, response_format)).__name__ + logger.warning("Unsupported response_format for Mistral: %s", type_name) + return None + + # endregion + + # region Response parsing + + def _parse_response( + self, + response: Mapping[str, Any], + *, + response_format: Any | None = None, + ) -> ChatResponse: + """Convert a Mistral chat-completion response payload to a framework ChatResponse.""" + response_id = response.get("id") + messages: list[Message] = [] + finish_reason: FinishReasonLiteral | None = None + choices = cast("Sequence[Mapping[str, Any]]", response.get("choices") or ()) + for choice in choices: + message: Mapping[str, Any] = choice.get("message") or {} + contents = self._parse_message_contents(message, fallback_id_prefix=response_id or "") + messages.append(Message(role="assistant", contents=contents, raw_representation=choice)) + if (reason := choice.get("finish_reason")) and finish_reason is None: + finish_reason = _FINISH_REASON_MAP.get(str(reason)) + if not messages: + messages.append(Message(role="assistant", contents=[])) + return ChatResponse( + response_id=response_id, + messages=messages, + usage_details=self._parse_usage(response.get("usage")), + model=response.get("model") or self.model, + created_at=self._format_created_at(response.get("created")), + finish_reason=finish_reason, + response_format=response_format, + raw_representation=response, + ) + + def _parse_chunk( + self, + chunk: Mapping[str, Any], + pending_tool_calls: dict[tuple[int, int], _StreamedToolCall], + ) -> ChatResponseUpdate: + """Convert a Mistral streaming completion chunk to a framework ChatResponseUpdate. + + Tool-call fragments are accumulated in ``pending_tool_calls`` (keyed by choice and + tool-call index) and emitted as complete calls on the chunk that carries the finish + reason. Fragments cannot be emitted as they arrive: the ID and name may be split + across chunks and fragments of parallel calls may interleave, which the framework's + adjacency-based merging cannot reassemble. + """ + chunk_id = chunk.get("id") or "" + contents: list[Content] = [] + finish_reason: FinishReasonLiteral | None = None + choices = cast("Sequence[Mapping[str, Any]]", chunk.get("choices") or ()) + for choice in choices: + delta: Mapping[str, Any] = choice.get("delta") or {} + contents.extend(self._parse_content_chunks(delta)) + self._accumulate_tool_call_fragments(choice, chunk_id, pending_tool_calls, contents) + if (reason := choice.get("finish_reason")) and finish_reason is None: + finish_reason = _FINISH_REASON_MAP.get(str(reason)) + if finish_reason and pending_tool_calls: + contents.extend(self._flush_pending_tool_calls(pending_tool_calls, chunk_id)) + if usage := self._parse_usage(chunk.get("usage")): + contents.append(Content.from_usage(usage_details=usage, raw_representation=chunk)) + return ChatResponseUpdate( + contents=contents, + role="assistant", + response_id=chunk.get("id"), + model=chunk.get("model"), + created_at=self._format_created_at(chunk.get("created")), + finish_reason=finish_reason, + raw_representation=chunk, + ) + + def _accumulate_tool_call_fragments( + self, + choice: Mapping[str, Any], + chunk_id: str, + pending_tool_calls: dict[tuple[int, int], _StreamedToolCall], + contents: list[Content], + ) -> None: + delta: Mapping[str, Any] = choice.get("delta") or {} + tool_calls: Sequence[Mapping[str, Any]] = delta.get("tool_calls") or [] + if not tool_calls: + return + + indexes = [tool_call.get("index") for tool_call in tool_calls] + if len(tool_calls) > 1 and len(set(indexes)) != len(indexes): + # Parallel complete calls whose indexes don't distinguish them (the API may omit + # them entirely). They cannot be fragments, so emit directly with positionally + # distinct IDs. + for position, tool_call in enumerate(tool_calls): + call = _StreamedToolCall() + call.add(tool_call) + contents.append(call.to_content(f"{chunk_id}:{position}")) + return + + choice_index = choice.get("index") or 0 + for tool_call in tool_calls: + key = (choice_index, tool_call.get("index") or 0) + entry = pending_tool_calls.get(key) + incoming_id = _tool_call_id_of(tool_call) + if entry is not None and entry.call_id and incoming_id and incoming_id != entry.call_id: + # A new call reusing the index; the accumulated one is complete. + contents.append(entry.to_content(f"{chunk_id}:{key[1]}")) + entry = None + if entry is None: + entry = pending_tool_calls[key] = _StreamedToolCall() + entry.add(tool_call) + + @staticmethod + def _flush_pending_tool_calls( + pending_tool_calls: dict[tuple[int, int], _StreamedToolCall], + fallback_id_prefix: str, + ) -> list[Content]: + contents = [ + entry.to_content(f"{fallback_id_prefix}:{index}") for (_, index), entry in pending_tool_calls.items() + ] + pending_tool_calls.clear() + return contents + + def _parse_message_contents( + self, + message: Mapping[str, Any], + *, + fallback_id_prefix: str = "", + ) -> list[Content]: + contents = self._parse_content_chunks(message) + tool_calls = cast("Sequence[Mapping[str, Any]]", message.get("tool_calls") or ()) + for position, tool_call in enumerate(tool_calls): + function: Mapping[str, Any] = tool_call.get("function") or {} + arguments = function.get("arguments") + if isinstance(arguments, str): + normalized_arguments: str | dict[str, Any] = arguments + elif isinstance(arguments, dict): + normalized_arguments = cast("dict[str, Any]", arguments) + else: + normalized_arguments = str(cast(object, arguments)) + contents.append( + Content.from_function_call( + call_id=self._resolve_tool_call_id(tool_call, position, fallback_id_prefix), + name=function.get("name") or "", + arguments=normalized_arguments, + raw_representation=tool_call, + ) + ) + return contents + + def _parse_content_chunks(self, message: Mapping[str, Any]) -> list[Content]: + contents: list[Content] = [] + content = message.get("content") + if isinstance(content, str): + if content: + contents.append(Content.from_text(text=content)) + elif content: + for chunk in cast("Sequence[Mapping[str, Any]]", content): + chunk_type = chunk.get("type") + if chunk_type == "text": + if text := chunk.get("text"): + contents.append(Content.from_text(text=text, raw_representation=chunk)) + elif chunk_type == "thinking": + if reasoning := self._thinking_to_text(chunk): + contents.append(Content.from_text_reasoning(text=reasoning, raw_representation=chunk)) + else: + logger.debug("Skipping unsupported response chunk from Mistral: %s", chunk_type) + return contents + + @staticmethod + def _resolve_tool_call_id(tool_call: Mapping[str, Any], position: int, fallback_id_prefix: str) -> str: + """Return the tool call ID, synthesizing a distinct one when the API omits it.""" + if call_id := _tool_call_id_of(tool_call): + return call_id + # The index only disambiguates when non-zero (the API may omit it); otherwise position does. + index = tool_call.get("index") or position + return f"{fallback_id_prefix}:{index}" + + @staticmethod + def _format_created_at(created: Any) -> str | None: + if not isinstance(created, (int, float)): + return None + return datetime.fromtimestamp(created, tz=timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.%fZ") + + @staticmethod + def _thinking_to_text(chunk: Mapping[str, Any]) -> str: + thinking = chunk.get("thinking") + if isinstance(thinking, str): + return thinking + if isinstance(thinking, Sequence): + return "".join( + part.get("text") or "" + for part in cast("Sequence[Mapping[str, Any]]", thinking) + if isinstance(part, Mapping) + ) + return "" + + def _parse_usage(self, usage: Mapping[str, Any] | None) -> UsageDetails | None: + if not usage: + return None + details: UsageDetails = {} + if (value := usage.get("prompt_tokens")) is not None: + details["input_token_count"] = value + if (value := usage.get("completion_tokens")) is not None: + details["output_token_count"] = value + if (value := usage.get("total_tokens")) is not None: + details["total_token_count"] = value + return details or None + + # endregion + + +class MistralChatClient( + FunctionInvocationLayer[MistralChatOptionsT], + ChatMiddlewareLayer[MistralChatOptionsT], + ChatTelemetryLayer[MistralChatOptionsT], + RawMistralChatClient[MistralChatOptionsT], + Generic[MistralChatOptionsT], +): + """Mistral AI chat client with function invocation, middleware, and telemetry support. + + This is the recommended client for most use cases. It builds on ``RawMistralChatClient`` + and adds: + + - **Function invocation**: automatically calls ``FunctionTool`` implementations and feeds + results back to the model until it produces a final text response. + - **Middleware**: a composable chain for cross-cutting concerns (logging, retries, etc.). + - **Telemetry**: OpenTelemetry traces and metrics emitted for every request. + + Use ``RawMistralChatClient`` instead when you need full control over the request pipeline + and want to opt out of one or more of these layers. + + Examples: + .. code-block:: python + + from agent_framework_mistral import MistralChatClient + + # Using environment variables + # Set MISTRAL_API_KEY=your-key + # Set MISTRAL_CHAT_MODEL=mistral-large-latest + client = MistralChatClient() + + # Or passing parameters directly + client = MistralChatClient( + model="mistral-large-latest", + api_key="your-api-key", + ) + + response = await client.get_response("Hello!") + print(response.text) + await client.close() + """ + + def __init__( + self, + *, + model: str | None = None, + api_key: str | SecretString | None = None, + server_url: str | None = None, + client: httpx.AsyncClient | None = None, + additional_properties: dict[str, Any] | None = None, + middleware: Sequence[ChatAndFunctionMiddlewareTypes] | None = None, + function_invocation_configuration: FunctionInvocationConfiguration | None = None, + env_file_path: str | None = None, + env_file_encoding: str | None = None, + ) -> None: + """Create a Mistral AI chat client. + + Keyword Args: + model: The Mistral chat model to use (e.g. "mistral-large-latest"). + Can also be set via environment variable ``MISTRAL_CHAT_MODEL``. + api_key: Mistral API key. Defaults to ``MISTRAL_API_KEY`` environment variable. + server_url: Optional server URL override. Defaults to ``MISTRAL_SERVER_URL`` + environment variable, or the Mistral default. + client: Optional pre-configured ``httpx.AsyncClient``. When provided, api_key is + not required and the client is expected to carry its own auth headers and + base URL. + additional_properties: Additional properties stored on the client instance. + middleware: Optional middleware chain applied to every call. + function_invocation_configuration: Optional configuration for the function invocation loop. + env_file_path: Path to ``.env`` file for settings. + env_file_encoding: Encoding for ``.env`` file. + """ + super().__init__( + model=model, + api_key=api_key, + server_url=server_url, + client=client, + additional_properties=additional_properties, + middleware=middleware, + function_invocation_configuration=function_invocation_configuration, + env_file_path=env_file_path, + env_file_encoding=env_file_encoding, + ) diff --git a/python/packages/mistral/agent_framework_mistral/_embedding_client.py b/python/packages/mistral/agent_framework_mistral/_embedding_client.py index 4f6af9e4aa4..dd6b24f04e0 100644 --- a/python/packages/mistral/agent_framework_mistral/_embedding_client.py +++ b/python/packages/mistral/agent_framework_mistral/_embedding_client.py @@ -4,10 +4,10 @@ import logging import sys -from collections.abc import Sequence -from importlib import import_module -from typing import Any, ClassVar, Generic, TypedDict +from collections.abc import Mapping, Sequence +from typing import Any, ClassVar, Generic, TypedDict, cast +import httpx from agent_framework import ( BaseEmbeddingClient, Embedding, @@ -17,29 +17,15 @@ load_settings, ) from agent_framework._settings import SecretString -from agent_framework._telemetry import mark_feature_used +from agent_framework._telemetry import get_user_agent +from agent_framework.exceptions import ( + IntegrationException, + IntegrationInvalidAuthException, + IntegrationInvalidRequestException, + IntegrationInvalidResponseException, +) from agent_framework.observability import EmbeddingTelemetryLayer -from ._feature_usage import FeatureIndex - - -def _load_mistral_client_class() -> Any: - try: - mistral_class = getattr(import_module("mistralai.client"), "Mistral", None) - except ModuleNotFoundError as exc: - if exc.name != "mistralai.client": - raise - mistral_class = None - - if mistral_class is None: - mistral_class = getattr(import_module("mistralai"), "Mistral", None) - if mistral_class is None: - raise ImportError("The installed mistralai package does not expose the Mistral client class.") - return mistral_class - - -Mistral: Any = _load_mistral_client_class() - if sys.version_info >= (3, 13): from typing import TypeVar # pragma: no cover else: @@ -48,6 +34,10 @@ def _load_mistral_client_class() -> Any: logger = logging.getLogger("agent_framework.mistral") +_MISTRAL_API_BASE_URL = "https://api.mistral.ai" +_EMBEDDINGS_PATH = "/v1/embeddings" +_DEFAULT_TIMEOUT_SECONDS = 60.0 + class MistralEmbeddingOptions(EmbeddingGenerationOptions, total=False): """Mistral AI-specific embedding options. @@ -94,13 +84,16 @@ class RawMistralEmbeddingClient( ): """Raw Mistral AI embedding client without telemetry. + Talks to the Mistral REST API directly over HTTP; the ``mistralai`` SDK is not required. + Keyword Args: model: The Mistral embedding model (e.g. "mistral-embed"). Can also be set via environment variable ``MISTRAL_EMBEDDING_MODEL``. api_key: Mistral API key. Defaults to ``MISTRAL_API_KEY`` environment variable. server_url: Optional server URL override. Defaults to ``MISTRAL_SERVER_URL`` environment variable, or the Mistral default. - client: Optional pre-configured ``Mistral`` client instance. + client: Optional pre-configured ``httpx.AsyncClient``. When provided, api_key is + not required and the client is expected to carry its own auth headers and base URL. additional_properties: Additional properties stored on the client instance. env_file_path: Path to ``.env`` file for settings. env_file_encoding: Encoding for ``.env`` file. @@ -114,16 +107,17 @@ def __init__( model: str | None = None, api_key: str | SecretString | None = None, server_url: str | None = None, - client: Any | None = None, + client: httpx.AsyncClient | None = None, additional_properties: dict[str, Any] | None = None, env_file_path: str | None = None, env_file_encoding: str | None = None, ) -> None: """Initialize a raw Mistral AI embedding client.""" + required_fields = ["embedding_model"] if client is not None else ["embedding_model", "api_key"] mistral_settings = load_settings( MistralEmbeddingSettings, env_prefix="MISTRAL_", - required_fields=["embedding_model", "api_key"], + required_fields=required_fields, api_key=str(api_key) if isinstance(api_key, SecretString) else api_key, embedding_model=model, server_url=server_url, @@ -132,23 +126,36 @@ def __init__( ) self.model: str = mistral_settings["embedding_model"] # type: ignore[assignment] - resolved_api_key: str = mistral_settings["api_key"] # type: ignore[assignment] - resolved_server_url = mistral_settings.get("server_url") + self.server_url = mistral_settings.get("server_url") + self._owns_client = client is None if client is not None: self.client = client + if self.server_url is None: + client_base_url = str(client.base_url).rstrip("/") + self.server_url = client_base_url or None else: - client_kwargs: dict[str, Any] = {"api_key": resolved_api_key} - if resolved_server_url: - client_kwargs["server_url"] = resolved_server_url - self.client = Mistral(**client_kwargs) + resolved_api_key: str = mistral_settings["api_key"] # type: ignore[assignment] + self.client = httpx.AsyncClient( + base_url=self.server_url or _MISTRAL_API_BASE_URL, + headers={ + "Authorization": f"Bearer {resolved_api_key}", + "User-Agent": get_user_agent(), + "Accept": "application/json", + }, + timeout=_DEFAULT_TIMEOUT_SECONDS, + ) - self.server_url = resolved_server_url super().__init__(additional_properties=additional_properties) + async def close(self) -> None: + """Close the internally created HTTP client.""" + if self._owns_client: + await self.client.aclose() + def service_url(self) -> str: """Get the URL of the service.""" - return self.server_url or "https://api.mistral.ai" + return self.server_url or _MISTRAL_API_BASE_URL async def get_embeddings( self, @@ -167,6 +174,10 @@ async def get_embeddings( Raises: ValueError: If model is not provided or values is empty. + IntegrationInvalidAuthException: If Mistral rejects the configured credentials. + IntegrationInvalidRequestException: If Mistral rejects the request. + IntegrationInvalidResponseException: If Mistral returns an invalid response. + IntegrationException: If the request fails for another reason. """ if not values: return GeneratedEmbeddings([], options=options) @@ -176,34 +187,60 @@ async def get_embeddings( if not model: raise ValueError("model is required") - kwargs: dict[str, Any] = {"model": model, "inputs": list(values)} + request: dict[str, Any] = {"model": model, "input": list(values)} if "dimensions" in opts: - kwargs["output_dimension"] = opts["dimensions"] - - mark_feature_used(FeatureIndex.MISTRAL) - response = await self.client.embeddings.create_async(**kwargs) + request["output_dimension"] = opts["dimensions"] - embeddings: list[Embedding[list[float]]] = [] - if response and response.data: - items = sorted(response.data, key=lambda d: d.index if d.index is not None else 0) + try: + response = await self.client.post(_EMBEDDINGS_PATH, json=request) + if response.status_code >= 400: + message = ( + f"Mistral embeddings request failed with status {response.status_code}: {response.text[:2000]}" + ) + if response.status_code in (401, 403): + raise IntegrationInvalidAuthException(message) + if response.status_code < 500: + raise IntegrationInvalidRequestException(message) + raise IntegrationException(message) + except IntegrationException: + raise + except Exception as ex: + raise IntegrationException(f"Mistral embeddings request failed: {ex}", inner_exception=ex) from ex + + try: + raw_payload = response.json() + if not isinstance(raw_payload, Mapping): + raise IntegrationInvalidResponseException("Mistral embeddings response must be a JSON object.") + payload = cast("Mapping[str, Any]", raw_payload) + embeddings: list[Embedding[list[float]]] = [] + data = cast("Sequence[Mapping[str, Any]]", payload.get("data") or ()) + items = sorted(data, key=lambda item: item.get("index") or 0) for item in items: - vector = list(item.embedding) if item.embedding else [] + vector = [float(v) for v in cast("Sequence[float]", item.get("embedding") or ())] embeddings.append( Embedding( vector=vector, dimensions=len(vector), - model=response.model or model, + model=payload.get("model") or model, ) ) - usage_dict: UsageDetails | None = None - if response and response.usage: - usage_dict = { - "input_token_count": response.usage.prompt_tokens, - "total_token_count": response.usage.total_tokens, - } + usage_dict: UsageDetails | None = None + if usage := payload.get("usage"): + usage_dict = {} + if (value := usage.get("prompt_tokens")) is not None: + usage_dict["input_token_count"] = value + if (value := usage.get("total_tokens")) is not None: + usage_dict["total_token_count"] = value - return GeneratedEmbeddings(embeddings, options=options, usage=usage_dict) + return GeneratedEmbeddings(embeddings, options=options, usage=usage_dict or None) + except IntegrationException: + raise + except Exception as ex: + raise IntegrationInvalidResponseException( + f"Mistral embeddings response was invalid: {ex}", + inner_exception=ex, + ) from ex class MistralEmbeddingClient( @@ -219,7 +256,7 @@ class MistralEmbeddingClient( api_key: Mistral API key. Defaults to ``MISTRAL_API_KEY`` environment variable. server_url: Optional server URL override. Defaults to ``MISTRAL_SERVER_URL`` environment variable, or the Mistral default. - client: Optional pre-configured ``Mistral`` client instance. + client: Optional pre-configured ``httpx.AsyncClient``. otel_provider_name: Optional telemetry provider name override. env_file_path: Path to ``.env`` file for settings. env_file_encoding: Encoding for ``.env`` file. @@ -243,6 +280,7 @@ class MistralEmbeddingClient( # Generate embeddings result = await client.get_embeddings(["Hello, world!"]) print(result[0].vector) + await client.close() """ OTEL_PROVIDER_NAME: ClassVar[str] = "mistralai" @@ -253,7 +291,7 @@ def __init__( model: str | None = None, api_key: str | SecretString | None = None, server_url: str | None = None, - client: Any | None = None, + client: httpx.AsyncClient | None = None, otel_provider_name: str | None = None, additional_properties: dict[str, Any] | None = None, env_file_path: str | None = None, diff --git a/python/packages/mistral/pyproject.toml b/python/packages/mistral/pyproject.toml index af839f8d859..00b7fb3a1ea 100644 --- a/python/packages/mistral/pyproject.toml +++ b/python/packages/mistral/pyproject.toml @@ -24,8 +24,9 @@ classifiers = [ ] dependencies = [ "agent-framework-core>=1.13.0,<2", - # Mistral 1.x retains the embeddings API without the OpenTelemetry semantic-conventions cap in 2.x. - "mistralai>=1.8.1,<3", + # Talks to the Mistral REST API directly; the mistralai SDK is not used because its + # pinned OpenTelemetry requirements conflict with the rest of the framework. + "httpx>=0.23.1,<1", ] [tool.uv] diff --git a/python/packages/mistral/tests/mistral/test_mistral_chat_client.py b/python/packages/mistral/tests/mistral/test_mistral_chat_client.py new file mode 100644 index 00000000000..db9006c59f6 --- /dev/null +++ b/python/packages/mistral/tests/mistral/test_mistral_chat_client.py @@ -0,0 +1,978 @@ +# Copyright (c) Microsoft. All rights reserved. + +import json +import logging +import os +from collections.abc import AsyncIterator, Sequence +from typing import Any + +import httpx +import pytest +from agent_framework import Agent, ChatResponse, Content, Message, tool +from agent_framework.exceptions import ( + ChatClientException, + ChatClientInvalidAuthException, + ChatClientInvalidRequestException, + ChatClientInvalidResponseException, +) +from pydantic import BaseModel + +import agent_framework_mistral._chat_client as chat_client_module +from agent_framework_mistral import MistralChatClient, MistralChatOptions +from agent_framework_mistral._chat_client import _sanitize_tool_call_id # pyright: ignore[reportPrivateUsage] + +# region: Helpers + + +def make_response_payload( + content: Any = None, + tool_calls: list[dict[str, Any]] | None = None, + finish_reason: str = "stop", + usage: dict[str, Any] | None = None, + choices: list[dict[str, Any]] | None = None, +) -> dict[str, Any]: + if choices is None: + message: dict[str, Any] = {"role": "assistant", "content": content} + if tool_calls is not None: + message["tool_calls"] = tool_calls + choices = [{"index": 0, "finish_reason": finish_reason, "message": message}] + return { + "id": "resp-id", + "object": "chat.completion", + "model": "mistral-small-latest", + "created": 1722249600, + "usage": usage or {"prompt_tokens": 5, "completion_tokens": 7, "total_tokens": 12}, + "choices": choices, + } + + +def make_chunk_payload( + content: Any = None, + tool_calls: list[dict[str, Any]] | None = None, + finish_reason: str | None = None, + usage: dict[str, Any] | None = None, +) -> dict[str, Any]: + delta: dict[str, Any] = {"role": "assistant", "content": content} + if tool_calls is not None: + delta["tool_calls"] = tool_calls + return { + "id": "chunk-id", + "model": "mistral-small-latest", + "created": 1722249600, + "usage": usage, + "choices": [{"index": 0, "finish_reason": finish_reason, "delta": delta}], + } + + +def tool_call_payload( + name: str, + arguments: Any, + call_id: str | None = None, + index: int | None = None, +) -> dict[str, Any]: + payload: dict[str, Any] = {"type": "function", "function": {"name": name, "arguments": arguments}} + if call_id is not None: + payload["id"] = call_id + if index is not None: + payload["index"] = index + return payload + + +def sse_bytes(*chunks: dict[str, Any]) -> bytes: + body = b"".join(f"data: {json.dumps(chunk)}\n\n".encode() for chunk in chunks) + return body + b"data: [DONE]\n\n" + + +class MockMistral: + """Routes requests to a queue of responses and records request bodies.""" + + def __init__(self, responses: Sequence[httpx.Response]) -> None: + self._responses = list(responses) + self.requests: list[dict[str, Any]] = [] + + def handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(json.loads(request.content)) + return self._responses.pop(0) + + @property + def last_request(self) -> dict[str, Any]: + return self.requests[-1] + + +def make_client(*responses: httpx.Response) -> tuple[MistralChatClient, MockMistral]: + server = MockMistral(responses) + http_client = httpx.AsyncClient( + base_url="https://api.mistral.ai", + transport=httpx.MockTransport(server.handler), + ) + client = MistralChatClient(model="mistral-small-latest", client=http_client) + return client, server + + +def json_response(payload: Any) -> httpx.Response: + return httpx.Response(200, json=payload) + + +def stream_response(*chunks: dict[str, Any]) -> httpx.Response: + return httpx.Response(200, content=sse_bytes(*chunks), headers={"content-type": "text/event-stream"}) + + +# region: Construction + + +def test_mistral_chat_construction_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("MISTRAL_CHAT_MODEL", "mistral-large-latest") + monkeypatch.setenv("MISTRAL_API_KEY", "test-key") + client = MistralChatClient() + assert client.model == "mistral-large-latest" + + +def test_mistral_chat_construction_with_params() -> None: + client = MistralChatClient(model="mistral-large-latest", api_key="test-key") + assert client.model == "mistral-large-latest" + assert client.client.headers["Authorization"] == "Bearer test-key" + + +def test_mistral_chat_construction_with_server_url() -> None: + client = MistralChatClient( + model="mistral-large-latest", + api_key="test-key", + server_url="https://custom.mistral.ai", + ) + assert client.service_url() == "https://custom.mistral.ai" + assert str(client.client.base_url) == "https://custom.mistral.ai" + + +def test_mistral_chat_construction_with_client_needs_no_api_key(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("MISTRAL_API_KEY", raising=False) + http_client = httpx.AsyncClient(base_url="https://api.mistral.ai") + client = MistralChatClient(model="mistral-large-latest", client=http_client) + assert client.client is http_client + + +def test_mistral_chat_construction_missing_api_key_raises(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("MISTRAL_API_KEY", raising=False) + from agent_framework.exceptions import SettingNotFoundError + + with pytest.raises(SettingNotFoundError): + MistralChatClient(model="mistral-large-latest") + + +def test_mistral_chat_service_url_default() -> None: + client = MistralChatClient(model="mistral-large-latest", api_key="test-key") + assert client.service_url() == "https://api.mistral.ai" + + +async def test_mistral_chat_close_only_closes_owned_client() -> None: + owned = MistralChatClient(model="mistral-large-latest", api_key="test-key") + await owned.close() + assert owned.client.is_closed + + http_client = httpx.AsyncClient(base_url="https://custom.mistral.ai") + injected = MistralChatClient(model="mistral-large-latest", client=http_client) + assert injected.service_url() == "https://custom.mistral.ai" + + await injected.close() + + assert not http_client.is_closed + await http_client.aclose() + + +async def test_mistral_chat_missing_model_raises_at_request(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("MISTRAL_CHAT_MODEL", raising=False) + http_client = httpx.AsyncClient(base_url="https://api.mistral.ai") + client = MistralChatClient(client=http_client, api_key="test-key") + with pytest.raises(ValueError, match="Mistral model is required"): + await client.get_response([Message("user", ["hi"])]) + + +# region: Request preparation + + +async def test_get_response_basic() -> None: + client, server = make_client(json_response(make_response_payload(content="hello"))) + + response = await client.get_response([Message("user", ["hi"])]) + + assert isinstance(response, ChatResponse) + assert response.text == "hello" + assert response.finish_reason == "stop" + assert response.usage_details == { + "input_token_count": 5, + "output_token_count": 7, + "total_token_count": 12, + } + assert server.last_request["model"] == "mistral-small-latest" + assert server.last_request["messages"] == [{"role": "user", "content": "hi"}] + + +@pytest.mark.parametrize( + ("status_code", "expected_exception"), + [ + (401, ChatClientInvalidAuthException), + (400, ChatClientInvalidRequestException), + (500, ChatClientException), + ], +) +async def test_get_response_http_error_wrapped( + status_code: int, + expected_exception: type[ChatClientException], +) -> None: + client, _ = make_client(httpx.Response(status_code, json={"message": "request failed"})) + + with pytest.raises(expected_exception, match=f"status {status_code}"): + await client.get_response([Message("user", ["hi"])]) + + +async def test_get_response_network_error_wrapped() -> None: + def raise_connect_error(request: httpx.Request) -> httpx.Response: + raise httpx.ConnectError("offline", request=request) + + http_client = httpx.AsyncClient( + base_url="https://api.mistral.ai", + transport=httpx.MockTransport(raise_connect_error), + ) + client = MistralChatClient(model="mistral-small-latest", client=http_client) + + with pytest.raises(ChatClientException, match="Mistral chat request failed"): + await client.get_response([Message("user", ["hi"])]) + + +@pytest.mark.parametrize( + ("response", "message"), + [ + (httpx.Response(200, content=b"{"), "response was invalid"), + (json_response([]), "must be a JSON object"), + (json_response({"choices": ["not-an-object"]}), "response was invalid"), + ], +) +async def test_get_response_invalid_payload_wrapped(response: httpx.Response, message: str) -> None: + client, _ = make_client(response) + + with pytest.raises(ChatClientInvalidResponseException, match=message): + await client.get_response([Message("user", ["hi"])]) + + +async def test_get_response_option_mapping() -> None: + client, server = make_client(json_response(make_response_payload(content="ok"))) + + options: MistralChatOptions = { + "temperature": 0.5, + "max_tokens": 100, + "seed": 42, + "allow_multiple_tool_calls": False, + "safe_prompt": True, + "stop": ["END"], + "n": 2, + "guardrails": [{"name": "test-guardrail"}], + "prompt_cache_key": "shared-prefix", + "reasoning_effort": "high", + } + await client.get_response([Message("user", ["hi"])], options=options) + + request = server.last_request + assert request["temperature"] == 0.5 + assert request["max_tokens"] == 100 + assert request["random_seed"] == 42 + assert request["parallel_tool_calls"] is False + assert request["safe_prompt"] is True + assert request["stop"] == ["END"] + assert request["n"] == 2 + assert request["guardrails"] == [{"name": "test-guardrail"}] + assert request["prompt_cache_key"] == "shared-prefix" + assert request["reasoning_effort"] == "high" + assert "seed" not in request + assert "allow_multiple_tool_calls" not in request + + +async def test_get_response_instructions_prepended_as_system_message() -> None: + client, server = make_client(json_response(make_response_payload(content="ok"))) + + await client.get_response([Message("user", ["hi"])], options={"instructions": "Be brief."}) + + assert server.last_request["messages"][0] == {"role": "system", "content": "Be brief."} + assert "instructions" not in server.last_request + + +async def test_get_response_model_override() -> None: + client, server = make_client(json_response(make_response_payload(content="ok"))) + + await client.get_response([Message("user", ["hi"])], options={"model": "mistral-large-latest"}) + + assert server.last_request["model"] == "mistral-large-latest" + + +async def test_message_conversion_roles() -> None: + client, server = make_client(json_response(make_response_payload(content="ok"))) + + messages = [ + Message("system", ["You are helpful."]), + Message("user", ["Question?"]), + Message( + "assistant", + [ + Content.from_text(text="Let me check."), + Content.from_function_call(call_id="call123AB", name="lookup", arguments='{"q": "x"}'), + ], + ), + Message("tool", [Content.from_function_result(call_id="call123AB", result="42")]), + ] + await client.get_response(messages) + + sent = server.last_request["messages"] + assert sent[0] == {"role": "system", "content": "You are helpful."} + assert sent[1] == {"role": "user", "content": "Question?"} + assert sent[2]["role"] == "assistant" + assert sent[2]["content"] == "Let me check." + assert sent[2]["tool_calls"] == [ + {"id": "call123AB", "type": "function", "function": {"name": "lookup", "arguments": '{"q": "x"}'}} + ] + assert sent[3]["role"] == "tool" + assert sent[3]["tool_call_id"] == "call123AB" + assert sent[3]["content"] == "42" + + +async def test_message_conversion_image_content() -> None: + client, server = make_client(json_response(make_response_payload(content="ok"))) + + messages = [ + Message( + "user", + [ + Content.from_text(text="What is this?"), + Content.from_uri(uri="https://example.com/image.png", media_type="image/png"), + ], + ), + ] + await client.get_response(messages) + + chunks = server.last_request["messages"][0]["content"] + assert chunks[0] == {"type": "text", "text": "What is this?"} + assert chunks[1] == {"type": "image_url", "image_url": "https://example.com/image.png"} + + +def test_message_conversion_edge_cases(caplog: pytest.LogCaptureFixture) -> None: + caplog.set_level(logging.DEBUG, logger="agent_framework.mistral") + client, _ = make_client() + + messages = client._prepare_mistral_messages( # pyright: ignore[reportPrivateUsage] + [ + Message("developer", ["ignored"]), + Message("user", [Content.from_error(message="ignored")]), + ] + ) + assert messages == [{"role": "user", "content": ""}] + + assert ( + client._convert_data_or_uri_content( # pyright: ignore[reportPrivateUsage] + Content("uri", media_type="image/png") + ) + is None + ) + assert client._convert_data_or_uri_content( # pyright: ignore[reportPrivateUsage] + Content.from_uri(uri="https://example.com/file.pdf", media_type="application/pdf") + ) == {"type": "document_url", "document_url": "https://example.com/file.pdf"} + assert ( + client._convert_data_or_uri_content( # pyright: ignore[reportPrivateUsage] + Content.from_uri(uri="https://example.com/audio.mp3", media_type="audio/mpeg") + ) + is None + ) + assert "Skipping unsupported message role" in caplog.text + assert "Skipping unsupported user content type" in caplog.text + + +def test_assistant_and_tool_message_edge_cases(caplog: pytest.LogCaptureFixture) -> None: + client, _ = make_client() + assistant = client._format_assistant_message( # pyright: ignore[reportPrivateUsage] + Message( + "assistant", + [Content.from_function_call(call_id="call", name="lookup", arguments={"query": "x"})], + ) + ) + assert assistant["tool_calls"][0]["function"]["arguments"] == {"query": "x"} + + rich_result = Content.from_function_result( + call_id="call", + result=[ + Content.from_text("text result"), + Content.from_uri(uri="https://example.com/image.png", media_type="image/png"), + ], + ) + named_result = Content("function_result", call_id="call", name="lookup", result=None) + tool_messages = client._format_tool_messages( # pyright: ignore[reportPrivateUsage] + Message("tool", [Content.from_text("ignored"), rich_result, named_result]) + ) + assert tool_messages[0]["content"] == "text result" + assert tool_messages[1]["content"] == "" + assert tool_messages[1]["name"] == "lookup" + assert "Rich content items will be omitted" in caplog.text + + +def test_result_to_text_variants() -> None: + client, _ = make_client() + assert client._result_to_text(None) == "" # pyright: ignore[reportPrivateUsage] + assert client._result_to_text("result") == "result" # pyright: ignore[reportPrivateUsage] + assert client._result_to_text({"value": 42}) == '{"value": 42}' # pyright: ignore[reportPrivateUsage] + assert "object" in client._result_to_text(object()) # pyright: ignore[reportPrivateUsage] + + +def test_sanitize_tool_call_id() -> None: + assert _sanitize_tool_call_id("abc123XYZ") == "abc123XYZ" + sanitized = _sanitize_tool_call_id("call_abc-123-too-long") + assert len(sanitized) == 9 + assert sanitized.isalnum() + assert sanitized == _sanitize_tool_call_id("call_abc-123-too-long") + + +async def test_tools_and_tool_choice() -> None: + client, server = make_client(json_response(make_response_payload(content="ok"))) + + @tool(approval_mode="never_require") + def get_weather(location: str) -> str: + """Get the weather.""" + return "sunny" + + await client.get_response( + [Message("user", ["hi"])], + options={"tools": [get_weather], "tool_choice": "auto"}, + ) + + request = server.last_request + assert request["tool_choice"] == "auto" + assert len(request["tools"]) == 1 + assert request["tools"][0]["type"] == "function" + assert request["tools"][0]["function"]["name"] == "get_weather" + + +async def test_tool_choice_required_function() -> None: + client, server = make_client(json_response(make_response_payload(content="ok"))) + + @tool(approval_mode="never_require") + def get_weather(location: str) -> str: + """Get the weather.""" + return "sunny" + + await client.get_response( + [Message("user", ["hi"])], + options={ + "tools": [get_weather], + "tool_choice": {"mode": "required", "required_function_name": "get_weather"}, + }, + ) + + assert server.last_request["tool_choice"] == {"type": "function", "function": {"name": "get_weather"}} + + +def test_tool_preparation_edge_cases( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + client, _ = make_client() + native_tool = {"type": "web_search"} + assert client._prepare_tools([native_tool]) == [native_tool] # pyright: ignore[reportPrivateUsage] + assert ( + client._prepare_tool_choice( # pyright: ignore[reportPrivateUsage] + {"mode": "auto", "allowed_tools": ["lookup"]} + ) + == "auto" + ) + assert client._prepare_tool_choice("none") == "none" # pyright: ignore[reportPrivateUsage] + assert client._prepare_tool_choice("required") == "required" # pyright: ignore[reportPrivateUsage] + + monkeypatch.setattr(chat_client_module, "validate_tool_mode", lambda _: {"mode": "unsupported"}) + assert client._prepare_tool_choice("auto") is None # pyright: ignore[reportPrivateUsage] + assert "Unsupported tool_choice mode" in caplog.text + + +async def test_response_format_pydantic_model() -> None: + client, server = make_client(json_response(make_response_payload(content='{"answer": "42"}'))) + + class Answer(BaseModel): + answer: str + + response = await client.get_response([Message("user", ["hi"])], options={"response_format": Answer}) + + response_format = server.last_request["response_format"] + assert response_format["type"] == "json_schema" + assert response_format["json_schema"]["name"] == "Answer" + assert response_format["json_schema"]["schema"] == Answer.model_json_schema() + assert response_format["json_schema"]["strict"] is True + assert response.value is not None + assert response.value.answer == "42" + + +async def test_response_format_json_object() -> None: + client, server = make_client(json_response(make_response_payload(content="{}"))) + + await client.get_response([Message("user", ["hi"])], options={"response_format": {"type": "json_object"}}) + + assert server.last_request["response_format"] == {"type": "json_object"} + + +def test_response_format_edge_cases(caplog: pytest.LogCaptureFixture) -> None: + client, _ = make_client() + assert client._prepare_response_format("json") == {"type": "json_object"} # pyright: ignore[reportPrivateUsage] + assert client._prepare_response_format("yaml") is None # pyright: ignore[reportPrivateUsage] + assert client._prepare_response_format( # pyright: ignore[reportPrivateUsage] + { + "type": "json_schema", + "json_schema": { + "name": "Answer", + "schema_definition": {"type": "object"}, + "strict": False, + }, + } + ) == { + "type": "json_schema", + "json_schema": { + "name": "Answer", + "schema": {"type": "object"}, + "strict": False, + }, + } + raw_schema = {"title": "Answer", "type": "object"} + assert client._prepare_response_format(raw_schema) == { # pyright: ignore[reportPrivateUsage] + "type": "json_schema", + "json_schema": {"name": "Answer", "schema": raw_schema, "strict": True}, + } + assert client._prepare_response_format(object()) is None # pyright: ignore[reportPrivateUsage] + assert "Unsupported response_format" in caplog.text + + +# region: Response parsing + + +async def test_parse_tool_calls() -> None: + client, _ = make_client( + json_response( + make_response_payload( + tool_calls=[tool_call_payload("get_weather", '{"location": "Paris"}', call_id="abc123XYZ")], + finish_reason="tool_calls", + ) + ) + ) + + response = await client.get_response([Message("user", ["hi"])]) + + assert response.finish_reason == "tool_calls" + calls = [c for c in response.messages[0].contents if c.type == "function_call"] + assert len(calls) == 1 + assert calls[0].call_id == "abc123XYZ" + assert calls[0].name == "get_weather" + assert calls[0].parse_arguments() == {"location": "Paris"} + + +async def test_parse_tool_calls_null_ids_stay_distinct() -> None: + """The Mistral API omits tool call IDs on some paths (serialized as null or "null").""" + client, _ = make_client( + json_response( + make_response_payload( + tool_calls=[ + tool_call_payload("get_weather", '{"location": "Paris"}', call_id="null"), + tool_call_payload("get_time", '{"tz": "CET"}', call_id="null"), + ], + finish_reason="tool_calls", + ) + ) + ) + + response = await client.get_response([Message("user", ["hi"])]) + + calls = [c for c in response.messages[0].contents if c.type == "function_call"] + assert len(calls) == 2 + assert calls[0].call_id != calls[1].call_id + assert "null" not in (calls[0].call_id, calls[1].call_id) + assert calls[0].parse_arguments() == {"location": "Paris"} + assert calls[1].parse_arguments() == {"tz": "CET"} + + +async def test_parse_multiple_choices() -> None: + payload = make_response_payload( + choices=[ + {"index": i, "finish_reason": "stop", "message": {"role": "assistant", "content": f"choice {i}"}} + for i in range(2) + ] + ) + client, _ = make_client(json_response(payload)) + + response = await client.get_response([Message("user", ["hi"])]) + + assert len(response.messages) == 2 + assert response.messages[0].text == "choice 0" + assert response.messages[1].text == "choice 1" + + +async def test_parse_empty_choices_returns_empty_assistant_message() -> None: + client, _ = make_client(json_response(make_response_payload(choices=[]))) + + response = await client.get_response([Message("user", ["hi"])]) + + assert len(response.messages) == 1 + assert response.messages[0].role == "assistant" + assert response.messages[0].contents == [] + + +async def test_parse_thinking_chunks() -> None: + content = [ + {"type": "thinking", "thinking": [{"type": "text", "text": "reasoning..."}]}, + {"type": "text", "text": "answer"}, + ] + client, _ = make_client(json_response(make_response_payload(content=content))) + + response = await client.get_response([Message("user", ["hi"])]) + + contents = response.messages[0].contents + assert contents[0].type == "text_reasoning" + assert contents[0].text == "reasoning..." + assert response.text == "answer" + + +def test_response_content_edge_cases() -> None: + client, _ = make_client() + contents = client._parse_message_contents( # pyright: ignore[reportPrivateUsage] + { + "content": [ + {"type": "thinking", "thinking": "reasoning"}, + {"type": "unsupported"}, + ], + "tool_calls": [ + tool_call_payload("mapping", {"value": 1}, call_id="abc123XYZ"), + tool_call_payload("missing", None, call_id="def456UVW"), + ], + } + ) + calls = [content for content in contents if content.type == "function_call"] + assert contents[0].text == "reasoning" + assert calls[0].arguments == {"value": 1} + assert calls[1].arguments == "None" + assert client._format_created_at("invalid") is None # pyright: ignore[reportPrivateUsage] + assert client._thinking_to_text({"thinking": object()}) == "" # pyright: ignore[reportPrivateUsage] + + +async def test_parse_finish_reason_model_length() -> None: + client, _ = make_client(json_response(make_response_payload(content="x", finish_reason="model_length"))) + + response = await client.get_response([Message("user", ["hi"])]) + + assert response.finish_reason == "length" + + +async def test_function_invocation_loop() -> None: + client, server = make_client( + json_response( + make_response_payload( + tool_calls=[tool_call_payload("get_weather", '{"location": "Paris"}', call_id="abc123XYZ")], + finish_reason="tool_calls", + ) + ), + json_response(make_response_payload(content="It is sunny in Paris.")), + ) + + @tool(approval_mode="never_require") + def get_weather(location: str) -> str: + """Get the weather.""" + return f"sunny in {location}" + + response = await client.get_response( + [Message("user", ["Weather in Paris?"])], + options={"tools": [get_weather]}, + ) + + assert response.text == "It is sunny in Paris." + assert len(server.requests) == 2 + assert any(m["role"] == "tool" for m in server.requests[1]["messages"]) + + +# region: Streaming + + +async def test_streaming_response() -> None: + client, server = make_client( + stream_response( + make_chunk_payload(content="Hel"), + make_chunk_payload(content="lo"), + make_chunk_payload( + finish_reason="stop", + usage={"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5}, + ), + ) + ) + + stream = client.get_response([Message("user", ["hi"])], stream=True) + updates = [update async for update in stream] + assert [u.text for u in updates] == ["Hel", "lo", ""] + + response = await stream.get_final_response() + assert response.text == "Hello" + assert response.finish_reason == "stop" + assert response.usage_details == { + "input_token_count": 3, + "output_token_count": 2, + "total_token_count": 5, + } + assert server.last_request["stream"] is True + + +async def test_streaming_tool_calls() -> None: + client, _ = make_client( + stream_response( + make_chunk_payload( + tool_calls=[tool_call_payload("get_weather", '{"location": "Paris"}', call_id="abc123XYZ")], + finish_reason="tool_calls", + ), + ) + ) + + stream = client.get_response([Message("user", ["hi"])], stream=True) + updates = [update async for update in stream] + + calls = [c for u in updates for c in u.contents if c.type == "function_call"] + assert len(calls) == 1 + assert calls[0].name == "get_weather" + + +async def test_streaming_fragmented_tool_call_coalesces() -> None: + """Real streams carry the ID and name only on the first fragment; later fragments carry argument pieces.""" + client, _ = make_client( + stream_response( + make_chunk_payload( + tool_calls=[tool_call_payload("get_weather", '{"loc', call_id="abc123XYZ", index=0)], + ), + make_chunk_payload( + tool_calls=[tool_call_payload("", 'ation": "Paris"}', index=0)], + finish_reason="tool_calls", + ), + ) + ) + + stream = client.get_response([Message("user", ["hi"])], stream=True) + async for _ in stream: + pass + response = await stream.get_final_response() + + calls = [c for c in response.messages[0].contents if c.type == "function_call"] + assert len(calls) == 1 + assert calls[0].call_id == "abc123XYZ" + assert calls[0].name == "get_weather" + assert calls[0].parse_arguments() == {"location": "Paris"} + + +async def test_streaming_fragmented_function_name_coalesces() -> None: + """The function name itself may be split across fragments.""" + client, _ = make_client( + stream_response( + make_chunk_payload(tool_calls=[tool_call_payload("get_", "", call_id="abc123XYZ", index=0)]), + make_chunk_payload( + tool_calls=[tool_call_payload("weather", '{"location": "Paris"}', index=0)], + finish_reason="tool_calls", + ), + ) + ) + + stream = client.get_response([Message("user", ["hi"])], stream=True) + async for _ in stream: + pass + response = await stream.get_final_response() + + calls = [c for c in response.messages[0].contents if c.type == "function_call"] + assert len(calls) == 1 + assert calls[0].call_id == "abc123XYZ" + assert calls[0].name == "get_weather" + assert calls[0].parse_arguments() == {"location": "Paris"} + + +async def test_streaming_interleaved_parallel_tool_calls() -> None: + """Fragments of two parallel calls may interleave across chunks.""" + client, _ = make_client( + stream_response( + make_chunk_payload( + tool_calls=[ + tool_call_payload("get_weather", '{"loc', call_id="abc123XYZ", index=0), + tool_call_payload("get_time", '{"t', call_id="def456UVW", index=1), + ], + ), + make_chunk_payload( + tool_calls=[ + tool_call_payload("", 'ation": "Paris"}', index=0), + tool_call_payload("", 'z": "CET"}', index=1), + ], + finish_reason="tool_calls", + ), + ) + ) + + stream = client.get_response([Message("user", ["hi"])], stream=True) + async for _ in stream: + pass + response = await stream.get_final_response() + + calls = [c for c in response.messages[0].contents if c.type == "function_call"] + assert len(calls) == 2 + by_id = {c.call_id: c for c in calls} + assert by_id["abc123XYZ"].name == "get_weather" + assert by_id["abc123XYZ"].parse_arguments() == {"location": "Paris"} + assert by_id["def456UVW"].name == "get_time" + assert by_id["def456UVW"].parse_arguments() == {"tz": "CET"} + + +async def test_streaming_parallel_calls_without_indexes() -> None: + client, _ = make_client( + stream_response( + make_chunk_payload( + tool_calls=[ + tool_call_payload("get_weather", {"location": "Paris"}, call_id="abc123XYZ"), + tool_call_payload("get_time", {"tz": "CET"}, call_id="def456UVW"), + ], + finish_reason="tool_calls", + ) + ) + ) + + stream = client.get_response([Message("user", ["hi"])], stream=True) + response = await stream.get_final_response() + + calls = [content for content in response.messages[0].contents if content.type == "function_call"] + assert [call.call_id for call in calls] == ["abc123XYZ", "def456UVW"] + assert calls[0].arguments == {"location": "Paris"} + + +async def test_streaming_reused_index_flushes_previous_call() -> None: + client, _ = make_client( + stream_response( + make_chunk_payload(tool_calls=[tool_call_payload("first", '{"value": 1}', call_id="abc123XYZ", index=0)]), + make_chunk_payload( + tool_calls=[tool_call_payload("second", '{"value": 2}', call_id="def456UVW", index=0)], + finish_reason="tool_calls", + ), + ) + ) + + stream = client.get_response([Message("user", ["hi"])], stream=True) + response = await stream.get_final_response() + + calls = [content for content in response.messages[0].contents if content.type == "function_call"] + assert [call.name for call in calls] == ["first", "second"] + + +async def test_streaming_mid_stream_error_wrapped() -> None: + """Exceptions raised while iterating the stream surface as ChatClientException.""" + + class ExplodingStream(httpx.AsyncByteStream): + async def __aiter__(self) -> AsyncIterator[bytes]: + yield f"data: {json.dumps(make_chunk_payload(content='partial'))}\n\n".encode() + raise ConnectionError("connection dropped") + + client, _ = make_client( + httpx.Response(200, stream=ExplodingStream(), headers={"content-type": "text/event-stream"}) + ) + + stream = client.get_response([Message("user", ["hi"])], stream=True) + with pytest.raises(ChatClientException, match="Mistral streaming chat request failed"): + async for _ in stream: + pass + + +async def test_streaming_http_error_wrapped() -> None: + client, _ = make_client(httpx.Response(429, json={"message": "rate limited"})) + + stream = client.get_response([Message("user", ["hi"])], stream=True) + with pytest.raises(ChatClientException, match="status 429"): + async for _ in stream: + pass + + +def test_parse_sse_line_variants() -> None: + payload = make_chunk_payload(content="hello") + assert MistralChatClient._parse_sse_line(f"data:{json.dumps(payload)}") == payload # pyright: ignore[reportPrivateUsage] + assert MistralChatClient._parse_sse_line("") is None # pyright: ignore[reportPrivateUsage] + assert MistralChatClient._parse_sse_line("event: message") is None # pyright: ignore[reportPrivateUsage] + assert MistralChatClient._parse_sse_line("data:") is None # pyright: ignore[reportPrivateUsage] + assert MistralChatClient._parse_sse_line("data: [DONE]") is None # pyright: ignore[reportPrivateUsage] + + with pytest.raises(ChatClientInvalidResponseException, match="malformed SSE"): + MistralChatClient._parse_sse_line("data: {") # pyright: ignore[reportPrivateUsage] + with pytest.raises(ChatClientInvalidResponseException, match="must be a JSON object"): + MistralChatClient._parse_sse_line("data: []") # pyright: ignore[reportPrivateUsage] + + +async def test_streaming_tool_call_flushed_without_finish_chunk() -> None: + """A stream that ends without a finish chunk still emits accumulated calls.""" + client, _ = make_client( + stream_response( + make_chunk_payload( + tool_calls=[tool_call_payload("get_weather", '{"location": "Paris"}', call_id="abc123XYZ", index=0)], + ), + ) + ) + + stream = client.get_response([Message("user", ["hi"])], stream=True) + async for _ in stream: + pass + response = await stream.get_final_response() + + calls = [c for c in response.messages[0].contents if c.type == "function_call"] + assert len(calls) == 1 + assert calls[0].call_id == "abc123XYZ" + assert calls[0].parse_arguments() == {"location": "Paris"} + + +# region: Integration Tests + +skip_if_mistral_chat_integration_tests_disabled = pytest.mark.skipif( + os.getenv("MISTRAL_CHAT_MODEL", "") in ("", "test-model") or os.getenv("MISTRAL_API_KEY", "") == "", + reason="No real Mistral chat model or API key provided; skipping integration tests.", +) + + +@pytest.mark.flaky +@pytest.mark.integration +@skip_if_mistral_chat_integration_tests_disabled +async def test_mistral_chat_integration_basic() -> None: + client = MistralChatClient() + try: + response = await client.get_response([Message("user", ["Reply with exactly the word: hello"])]) + + assert response.text + assert response.usage_details is not None + finally: + await client.close() + + +@pytest.mark.flaky +@pytest.mark.integration +@skip_if_mistral_chat_integration_tests_disabled +async def test_mistral_chat_integration_streaming() -> None: + client = MistralChatClient() + try: + stream = client.get_response([Message("user", ["Count from 1 to 5."])], stream=True) + updates = [update async for update in stream] + + assert updates + response = await stream.get_final_response() + assert response.text + finally: + await client.close() + + +@pytest.mark.flaky +@pytest.mark.integration +@skip_if_mistral_chat_integration_tests_disabled +async def test_mistral_chat_integration_agent_with_tool() -> None: + @tool(approval_mode="never_require") + def get_secret_word() -> str: + """Get the secret word.""" + return "pineapple" + + client = MistralChatClient() + agent = Agent( + client=client, + instructions="Use the get_secret_word tool and reply with its result.", + tools=get_secret_word, + ) + try: + result = await agent.run("What is the secret word?") + + assert "pineapple" in result.text.lower() + finally: + await client.close() diff --git a/python/packages/mistral/tests/mistral/test_mistral_embedding_client.py b/python/packages/mistral/tests/mistral/test_mistral_embedding_client.py index 09f47973a5f..ccee2de6cea 100644 --- a/python/packages/mistral/tests/mistral/test_mistral_embedding_client.py +++ b/python/packages/mistral/tests/mistral/test_mistral_embedding_client.py @@ -1,88 +1,94 @@ # Copyright (c) Microsoft. All rights reserved. -import inspect +import json import os -from types import SimpleNamespace -from unittest.mock import AsyncMock, MagicMock, patch, sentinel +from collections.abc import Sequence +from typing import Any +import httpx import pytest from agent_framework import Embedding, GeneratedEmbeddings +from agent_framework.exceptions import ( + IntegrationException, + IntegrationInvalidAuthException, + IntegrationInvalidRequestException, + IntegrationInvalidResponseException, +) from agent_framework_mistral import MistralEmbeddingClient, MistralEmbeddingOptions -from agent_framework_mistral._embedding_client import _load_mistral_client_class # pyright: ignore[reportPrivateUsage] -from agent_framework_mistral._feature_usage import FeatureIndex # region: Unit Tests +def make_embeddings_payload( + vectors: Sequence[Sequence[float]], + model: str = "mistral-embed", + usage: dict[str, Any] | None = None, +) -> dict[str, Any]: + return { + "object": "list", + "model": model, + "data": [{"object": "embedding", "index": i, "embedding": list(vector)} for i, vector in enumerate(vectors)], + "usage": usage if usage is not None else {"prompt_tokens": 10, "total_tokens": 10}, + } + + +class MockMistral: + def __init__(self, responses: Sequence[httpx.Response]) -> None: + self._responses = list(responses) + self.requests: list[dict[str, Any]] = [] + + def handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(json.loads(request.content)) + return self._responses.pop(0) + + @property + def last_request(self) -> dict[str, Any]: + return self.requests[-1] + + +def make_client(*responses: httpx.Response) -> tuple[MistralEmbeddingClient, MockMistral]: + server = MockMistral(responses) + http_client = httpx.AsyncClient( + base_url="https://api.mistral.ai", + transport=httpx.MockTransport(server.handler), + ) + client = MistralEmbeddingClient(model="mistral-embed", client=http_client) + return client, server + + def test_mistral_embedding_construction(monkeypatch: pytest.MonkeyPatch) -> None: """Test construction with environment variables.""" monkeypatch.setenv("MISTRAL_EMBEDDING_MODEL", "mistral-embed") monkeypatch.setenv("MISTRAL_API_KEY", "test-key") - with patch("agent_framework_mistral._embedding_client.Mistral") as mock_cls: - mock_cls.return_value = MagicMock() - client = MistralEmbeddingClient() - assert client.model == "mistral-embed" + client = MistralEmbeddingClient() + assert client.model == "mistral-embed" def test_mistral_embedding_construction_with_params() -> None: """Test construction with explicit parameters.""" - with patch("agent_framework_mistral._embedding_client.Mistral") as mock_cls: - mock_cls.return_value = MagicMock() - client = MistralEmbeddingClient( - model="mistral-embed", - api_key="test-key", - ) - assert client.model == "mistral-embed" - mock_cls.assert_called_once_with(api_key="test-key") + client = MistralEmbeddingClient(model="mistral-embed", api_key="test-key") + assert client.model == "mistral-embed" + assert client.client.headers["Authorization"] == "Bearer test-key" def test_mistral_embedding_construction_with_server_url() -> None: """Test construction with custom server URL.""" - with patch("agent_framework_mistral._embedding_client.Mistral") as mock_cls: - mock_cls.return_value = MagicMock() - client = MistralEmbeddingClient( - model="mistral-embed", - api_key="test-key", - server_url="https://custom.mistral.ai", - ) - assert client.model == "mistral-embed" - assert client.server_url == "https://custom.mistral.ai" - mock_cls.assert_called_once_with( - api_key="test-key", - server_url="https://custom.mistral.ai", - ) + client = MistralEmbeddingClient( + model="mistral-embed", + api_key="test-key", + server_url="https://custom.mistral.ai", + ) + assert client.model == "mistral-embed" + assert client.server_url == "https://custom.mistral.ai" + assert str(client.client.base_url) == "https://custom.mistral.ai" def test_mistral_embedding_construction_with_client() -> None: """Test construction with a pre-configured client.""" - mock_client = MagicMock() - with patch("agent_framework_mistral._embedding_client.Mistral"): - client = MistralEmbeddingClient( - model="mistral-embed", - api_key="test-key", - client=mock_client, - ) - assert client.client is mock_client - - -def test_mistral_client_import_falls_back_when_client_module_is_missing() -> None: - """Test Mistral 1.x layouts that expose the client only from the package root.""" - - def import_mistral_module(name: str) -> object: - if name == "mistralai.client": - raise ModuleNotFoundError(name="mistralai.client") - return SimpleNamespace(Mistral=sentinel.mistral_class) - - with patch("agent_framework_mistral._embedding_client.import_module", side_effect=import_mistral_module): - assert _load_mistral_client_class() is sentinel.mistral_class - - -def test_mistral_sdk_supports_output_dimension() -> None: - """Test that the supported SDK range includes the dimensions parameter.""" - client = MistralEmbeddingClient(model="mistral-embed", api_key="test-key") - - assert "output_dimension" in inspect.signature(client.client.embeddings.create_async).parameters + http_client = httpx.AsyncClient(base_url="https://api.mistral.ai") + client = MistralEmbeddingClient(model="mistral-embed", client=http_client) + assert client.client is http_client def test_mistral_embedding_construction_missing_model_raises(monkeypatch: pytest.MonkeyPatch) -> None: @@ -107,163 +113,158 @@ def test_mistral_embedding_construction_missing_api_key_raises(monkeypatch: pyte def test_mistral_embedding_service_url() -> None: """Test service_url returns the correct URL.""" - with patch("agent_framework_mistral._embedding_client.Mistral") as mock_cls: - mock_cls.return_value = MagicMock() - client = MistralEmbeddingClient( - model="mistral-embed", - api_key="test-key", - ) - assert client.service_url() == "https://api.mistral.ai" + client = MistralEmbeddingClient(model="mistral-embed", api_key="test-key") + assert client.service_url() == "https://api.mistral.ai" def test_mistral_embedding_service_url_custom() -> None: """Test service_url returns custom URL when set.""" - with patch("agent_framework_mistral._embedding_client.Mistral") as mock_cls: - mock_cls.return_value = MagicMock() - client = MistralEmbeddingClient( - model="mistral-embed", - api_key="test-key", - server_url="https://custom.mistral.ai", - ) - assert client.service_url() == "https://custom.mistral.ai" + client = MistralEmbeddingClient( + model="mistral-embed", + api_key="test-key", + server_url="https://custom.mistral.ai", + ) + assert client.service_url() == "https://custom.mistral.ai" + + +async def test_mistral_embedding_close_only_closes_owned_client() -> None: + owned = MistralEmbeddingClient(model="mistral-embed", api_key="test-key") + await owned.close() + assert owned.client.is_closed + + http_client = httpx.AsyncClient(base_url="https://custom.mistral.ai") + injected = MistralEmbeddingClient(model="mistral-embed", client=http_client) + assert injected.service_url() == "https://custom.mistral.ai" + + await injected.close() + + assert not http_client.is_closed + await http_client.aclose() async def test_mistral_embedding_get_embeddings() -> None: """Test generating embeddings via the Mistral API.""" - mock_response = MagicMock() - mock_response.data = [ - MagicMock(embedding=[0.1, 0.2, 0.3], index=0, object="embedding"), - MagicMock(embedding=[0.4, 0.5, 0.6], index=1, object="embedding"), - ] - mock_response.model = "mistral-embed" - mock_response.usage = MagicMock(prompt_tokens=10, total_tokens=10) - - with ( - patch("agent_framework_mistral._embedding_client.Mistral") as mock_cls, - patch("agent_framework_mistral._embedding_client.mark_feature_used") as mark_feature_used, - ): - mock_client = MagicMock() - mock_client.embeddings = MagicMock() - mock_client.embeddings.create_async = AsyncMock(return_value=mock_response) - mock_cls.return_value = mock_client - - client = MistralEmbeddingClient(model="mistral-embed", api_key="test-key") - result = await client.get_embeddings(["hello", "world"]) - - mark_feature_used.assert_called_once_with(FeatureIndex.MISTRAL) - assert isinstance(result, GeneratedEmbeddings) - assert len(result) == 2 - assert result[0].vector == [0.1, 0.2, 0.3] - assert result[1].vector == [0.4, 0.5, 0.6] - assert result[0].model == "mistral-embed" - assert result.usage == {"input_token_count": 10, "total_token_count": 10} - - mock_client.embeddings.create_async.assert_called_once_with( - model="mistral-embed", - inputs=["hello", "world"], - ) + client, server = make_client(httpx.Response(200, json=make_embeddings_payload([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]))) + + result = await client.get_embeddings(["hello", "world"]) + + assert isinstance(result, GeneratedEmbeddings) + assert len(result) == 2 + assert result[0].vector == [0.1, 0.2, 0.3] + assert result[1].vector == [0.4, 0.5, 0.6] + assert result[0].model == "mistral-embed" + assert result.usage == {"input_token_count": 10, "total_token_count": 10} + assert server.last_request == {"model": "mistral-embed", "input": ["hello", "world"]} async def test_mistral_embedding_get_embeddings_empty_input() -> None: """Test generating embeddings with empty input.""" - with patch("agent_framework_mistral._embedding_client.Mistral") as mock_cls: - mock_client = MagicMock() - mock_cls.return_value = mock_client + client, server = make_client() - client = MistralEmbeddingClient(model="mistral-embed", api_key="test-key") - result = await client.get_embeddings([]) + result = await client.get_embeddings([]) - assert isinstance(result, GeneratedEmbeddings) - assert len(result) == 0 + assert isinstance(result, GeneratedEmbeddings) + assert len(result) == 0 + assert server.requests == [] async def test_mistral_embedding_get_embeddings_with_dimensions() -> None: """Test generating embeddings with custom dimensions option.""" - mock_response = MagicMock() - mock_response.data = [ - MagicMock(embedding=[0.1, 0.2], index=0, object="embedding"), - ] - mock_response.model = "mistral-embed" - mock_response.usage = MagicMock(prompt_tokens=5, total_tokens=5) - - with patch("agent_framework_mistral._embedding_client.Mistral") as mock_cls: - mock_client = MagicMock() - mock_client.embeddings = MagicMock() - mock_client.embeddings.create_async = AsyncMock(return_value=mock_response) - mock_cls.return_value = mock_client - - client = MistralEmbeddingClient(model="mistral-embed", api_key="test-key") - options: MistralEmbeddingOptions = {"dimensions": 512} - result = await client.get_embeddings(["hello"], options=options) - - assert len(result) == 1 - mock_client.embeddings.create_async.assert_called_once_with( - model="mistral-embed", - inputs=["hello"], - output_dimension=512, - ) + client, server = make_client( + httpx.Response(200, json=make_embeddings_payload([[0.1, 0.2]], usage={"prompt_tokens": 5, "total_tokens": 5})) + ) + + options: MistralEmbeddingOptions = {"dimensions": 512} + result = await client.get_embeddings(["hello"], options=options) + + assert len(result) == 1 + assert server.last_request == {"model": "mistral-embed", "input": ["hello"], "output_dimension": 512} async def test_mistral_embedding_get_embeddings_no_model_raises() -> None: """Test that missing model at call time raises ValueError.""" - with patch("agent_framework_mistral._embedding_client.Mistral") as mock_cls: - mock_client = MagicMock() - mock_cls.return_value = mock_client + client, _ = make_client() + client.model = None # type: ignore[assignment] # ty: ignore[invalid-assignment] - client = MistralEmbeddingClient(model="mistral-embed", api_key="test-key") - client.model = None # type: ignore[assignment] # ty: ignore[invalid-assignment] - - with pytest.raises(ValueError, match="model is required"): - await client.get_embeddings(["hello"]) + with pytest.raises(ValueError, match="model is required"): + await client.get_embeddings(["hello"]) async def test_mistral_embedding_get_embeddings_model_override() -> None: """Test that model can be overridden via options.""" - mock_response = MagicMock() - mock_response.data = [ - MagicMock(embedding=[0.1, 0.2, 0.3], index=0, object="embedding"), - ] - mock_response.model = "custom-embed" - mock_response.usage = MagicMock(prompt_tokens=5, total_tokens=5) - - with patch("agent_framework_mistral._embedding_client.Mistral") as mock_cls: - mock_client = MagicMock() - mock_client.embeddings = MagicMock() - mock_client.embeddings.create_async = AsyncMock(return_value=mock_response) - mock_cls.return_value = mock_client - - client = MistralEmbeddingClient(model="mistral-embed", api_key="test-key") - options: MistralEmbeddingOptions = {"model": "custom-embed"} - result = await client.get_embeddings(["hello"], options=options) - - assert len(result) == 1 - assert result[0].model == "custom-embed" - mock_client.embeddings.create_async.assert_called_once_with( - model="custom-embed", - inputs=["hello"], + client, server = make_client( + httpx.Response( + 200, + json=make_embeddings_payload( + [[0.1, 0.2, 0.3]], model="custom-embed", usage={"prompt_tokens": 5, "total_tokens": 5} + ), ) + ) + + options: MistralEmbeddingOptions = {"model": "custom-embed"} + result = await client.get_embeddings(["hello"], options=options) + + assert len(result) == 1 + assert result[0].model == "custom-embed" + assert server.last_request == {"model": "custom-embed", "input": ["hello"]} async def test_mistral_embedding_get_embeddings_no_usage() -> None: """Test handling response without usage information.""" - mock_response = MagicMock() - mock_response.data = [ - MagicMock(embedding=[0.1, 0.2, 0.3], index=0, object="embedding"), - ] - mock_response.model = "mistral-embed" - mock_response.usage = None + client, _ = make_client(httpx.Response(200, json=make_embeddings_payload([[0.1, 0.2, 0.3]], usage={}))) - with patch("agent_framework_mistral._embedding_client.Mistral") as mock_cls: - mock_client = MagicMock() - mock_client.embeddings = MagicMock() - mock_client.embeddings.create_async = AsyncMock(return_value=mock_response) - mock_cls.return_value = mock_client + result = await client.get_embeddings(["hello"]) - client = MistralEmbeddingClient(model="mistral-embed", api_key="test-key") - result = await client.get_embeddings(["hello"]) + assert len(result) == 1 + assert result.usage is None - assert len(result) == 1 - assert result.usage is None + +@pytest.mark.parametrize( + ("status_code", "expected_exception"), + [ + (401, IntegrationInvalidAuthException), + (400, IntegrationInvalidRequestException), + (500, IntegrationException), + ], +) +async def test_mistral_embedding_http_error_wrapped( + status_code: int, + expected_exception: type[IntegrationException], +) -> None: + """Test that HTTP errors surface with the appropriate integration exception.""" + client, _ = make_client(httpx.Response(status_code, json={"message": "request failed"})) + + with pytest.raises(expected_exception, match=f"status {status_code}"): + await client.get_embeddings(["hello"]) + + +async def test_mistral_embedding_network_error_wrapped() -> None: + def raise_connect_error(request: httpx.Request) -> httpx.Response: + raise httpx.ConnectError("offline", request=request) + + http_client = httpx.AsyncClient( + base_url="https://api.mistral.ai", + transport=httpx.MockTransport(raise_connect_error), + ) + client = MistralEmbeddingClient(model="mistral-embed", client=http_client) + + with pytest.raises(IntegrationException, match="Mistral embeddings request failed"): + await client.get_embeddings(["hello"]) + + +@pytest.mark.parametrize( + ("response", "message"), + [ + (httpx.Response(200, content=b"{"), "response was invalid"), + (httpx.Response(200, json=[]), "must be a JSON object"), + (httpx.Response(200, json={"data": ["not-an-object"]}), "response was invalid"), + ], +) +async def test_mistral_embedding_invalid_payload_wrapped(response: httpx.Response, message: str) -> None: + client, _ = make_client(response) + + with pytest.raises(IntegrationInvalidResponseException, match=message): + await client.get_embeddings(["hello"]) # region: Integration Tests @@ -280,15 +281,18 @@ async def test_mistral_embedding_get_embeddings_no_usage() -> None: async def test_mistral_embedding_integration() -> None: """Integration test for Mistral AI embedding client.""" client = MistralEmbeddingClient() - result = await client.get_embeddings(["Hello, world!", "How are you?"]) + try: + result = await client.get_embeddings(["Hello, world!", "How are you?"]) - assert isinstance(result, GeneratedEmbeddings) - assert len(result) == 2 - for embedding in result: - assert isinstance(embedding, Embedding) - assert isinstance(embedding.vector, list) - assert len(embedding.vector) > 0 - assert all(isinstance(v, float) for v in embedding.vector) - assert result.usage is not None - assert result.usage["input_token_count"] is not None - assert result.usage["input_token_count"] > 0 + assert isinstance(result, GeneratedEmbeddings) + assert len(result) == 2 + for embedding in result: + assert isinstance(embedding, Embedding) + assert isinstance(embedding.vector, list) + assert len(embedding.vector) > 0 + assert all(isinstance(v, float) for v in embedding.vector) + assert result.usage is not None + assert result.usage["input_token_count"] is not None + assert result.usage["input_token_count"] > 0 + finally: + await client.close() diff --git a/python/samples/02-agents/providers/mistral/README.md b/python/samples/02-agents/providers/mistral/README.md index b58ecf6ef26..b47858c1a4e 100644 --- a/python/samples/02-agents/providers/mistral/README.md +++ b/python/samples/02-agents/providers/mistral/README.md @@ -1,15 +1,17 @@ -# Mistral AI Embedding Examples +# Mistral AI Examples -This folder contains examples demonstrating how to use Mistral AI embedding models with the Agent Framework. +This folder contains examples demonstrating how to use Mistral AI models with the Agent Framework. ## Examples | File | Description | |------|-------------| +| [`mistral_agent_basic.py`](mistral_agent_basic.py) | Basic agent with tool usage using the Mistral AI chat client. | | [`mistral_embeddings.py`](mistral_embeddings.py) | Basic embedding generation with the Mistral AI embedding client. | ## Environment Variables - `MISTRAL_API_KEY`: Your Mistral AI API key +- `MISTRAL_CHAT_MODEL`: Chat model name (e.g., `mistral-small-latest`) - `MISTRAL_EMBEDDING_MODEL`: Embedding model name (e.g., `mistral-embed`) - `MISTRAL_SERVER_URL` (optional): Server URL override for custom deployments diff --git a/python/samples/02-agents/providers/mistral/mistral_agent_basic.py b/python/samples/02-agents/providers/mistral/mistral_agent_basic.py new file mode 100644 index 00000000000..4c604e0ef8a --- /dev/null +++ b/python/samples/02-agents/providers/mistral/mistral_agent_basic.py @@ -0,0 +1,99 @@ +# Copyright (c) Microsoft. All rights reserved. +# /// script +# requires-python = ">=3.10" +# dependencies = ["agent-framework-mistral"] +# /// + +import asyncio +from datetime import datetime +from zoneinfo import ZoneInfo + +from agent_framework import Agent, tool +from agent_framework.mistral import MistralChatClient +from dotenv import load_dotenv + +"""Demonstrates a Mistral AI agent with basic tool usage. + +Requires ``MISTRAL_API_KEY`` and ``MISTRAL_CHAT_MODEL`` environment variables +(e.g. MISTRAL_CHAT_MODEL=mistral-small-latest). +""" + +# Load environment variables from the local .env file. +load_dotenv() + + +# NOTE: approval_mode="never_require" is for sample brevity. Use "always_require" in production; +# see samples/02-agents/tools/function_tool_with_approval.py +# and samples/02-agents/tools/function_tool_with_approval_and_sessions.py. +@tool(approval_mode="never_require") +def get_time(timezone: str) -> str: + """Get the current time in an IANA timezone (e.g. 'America/Los_Angeles').""" + now = datetime.now(ZoneInfo(timezone)) + return f"The current time in {timezone} is {now.strftime('%I:%M %p')}." + + +async def non_streaming_example() -> None: + """Example of non-streaming response (get the complete result at once).""" + print("=== Non-streaming Response Example ===") + + client = MistralChatClient() + agent = Agent( + client=client, + name="TimeAgent", + instructions="You are a helpful time agent, answer in one sentence.", + tools=get_time, + ) + + query = "What time is it in Seattle? Use a tool call" + print(f"User: {query}") + try: + result = await agent.run(query) + print(f"Result: {result}\n") + finally: + await client.close() + + +async def streaming_example() -> None: + """Example of streaming response (get results as they are generated).""" + print("=== Streaming Response Example ===") + + client = MistralChatClient() + agent = Agent( + client=client, + name="TimeAgent", + instructions="You are a helpful time agent, answer in one sentence.", + tools=get_time, + ) + query = "What time is it in San Francisco? Use a tool call" + print(f"User: {query}") + print("Agent: ", end="", flush=True) + try: + async for chunk in agent.run(query, stream=True): + if chunk.text: + print(chunk.text, end="", flush=True) + print("\n") + finally: + await client.close() + + +async def main() -> None: + print("=== Basic Mistral Chat Client Agent Example ===") + + await non_streaming_example() + await streaming_example() + + +if __name__ == "__main__": + asyncio.run(main()) + +""" +Sample output: +=== Basic Mistral Chat Client Agent Example === +=== Non-streaming Response Example === +User: What time is it in Seattle? Use a tool call +Result: The current time in Seattle is 10:30 AM. + +=== Streaming Response Example === +User: What time is it in San Francisco? Use a tool call +Agent: The current time in San Francisco is 10:30 AM. +""" diff --git a/python/samples/02-agents/providers/mistral/mistral_embeddings.py b/python/samples/02-agents/providers/mistral/mistral_embeddings.py index 20419ded62d..b2468d7859c 100644 --- a/python/samples/02-agents/providers/mistral/mistral_embeddings.py +++ b/python/samples/02-agents/providers/mistral/mistral_embeddings.py @@ -21,23 +21,26 @@ async def basic_embedding_example() -> None: """Generate embeddings for a list of texts.""" print("=== Basic Embedding Generation ===") - # 1. Create the embedding client (uses MISTRAL_API_KEY and MISTRAL_EMBEDDING_MODEL env vars). + # 1. Create the embedding client using environment-based configuration. client = MistralEmbeddingClient() # 2. Generate embeddings for multiple texts. texts = ["Hello, world!", "How are you?", "Agent Framework with Mistral AI"] - result = await client.get_embeddings(texts) + try: + result = await client.get_embeddings(texts) - # 3. Print results. - print(f"Generated {len(result)} embeddings") - for i, embedding in enumerate(result): - print(f" Text {i + 1}: dimensions={embedding.dimensions}, vector={embedding.vector[:5]}...") + # 3. Print the generated vectors and usage metadata. + print(f"Generated {len(result)} embeddings") + for i, embedding in enumerate(result): + print(f" Text {i + 1}: dimensions={embedding.dimensions}, vector={embedding.vector[:5]}...") - if result.usage: - print( - f" Usage: {result.usage['input_token_count']} input tokens, " - f"{result.usage['total_token_count']} total tokens" - ) + if result.usage: + print( + f" Usage: {result.usage['input_token_count']} input tokens, " + f"{result.usage['total_token_count']} total tokens" + ) + finally: + await client.close() async def embedding_with_options_example() -> None: @@ -46,14 +49,16 @@ async def embedding_with_options_example() -> None: from agent_framework.mistral import MistralEmbeddingOptions - client = MistralEmbeddingClient() + # Only some models support a custom output dimension (e.g. codestral-embed; mistral-embed does not). + client = MistralEmbeddingClient(model="codestral-embed") - # Request a specific output dimension (model must support it). options: MistralEmbeddingOptions = {"dimensions": 256} - result = await client.get_embeddings(["Dimensionality reduction example"], options=options) - - print(f" Dimensions: {result[0].dimensions}") - print(f" Vector (first 5): {result[0].vector[:5]}...") + try: + result = await client.get_embeddings(["Dimensionality reduction example"], options=options) + print(f" Dimensions: {result[0].dimensions}") + print(f" Vector (first 5): {result[0].vector[:5]}...") + finally: + await client.close() async def main() -> None: diff --git a/python/samples/README.md b/python/samples/README.md index 6b7bc4a1737..ba5d4eedb3d 100644 --- a/python/samples/README.md +++ b/python/samples/README.md @@ -148,6 +148,10 @@ variable. | `agent-framework-github-copilot` | `GitHubCopilotAgent` | `GITHUB_COPILOT_TIMEOUT` | `60` | | `agent-framework-github-copilot` | `GitHubCopilotAgent` | `GITHUB_COPILOT_LOG_LEVEL` | `info` | | `agent-framework-mem0` | `agent_framework_mem0 package import` | `MEM0_TELEMETRY` | `false` | +| `agent-framework-mistral` | `MistralChatClient / MistralEmbeddingClient` | `MISTRAL_API_KEY` | `your-api-key` | +| `agent-framework-mistral` | `MistralChatClient` | `MISTRAL_CHAT_MODEL` | `mistral-small-latest` | +| `agent-framework-mistral` | `MistralEmbeddingClient` | `MISTRAL_EMBEDDING_MODEL` | `mistral-embed` | +| `agent-framework-mistral` | `MistralChatClient / MistralEmbeddingClient` | `MISTRAL_SERVER_URL` | `https://api.mistral.ai` | | `agent-framework-ollama` | `OllamaChatClient` | `OLLAMA_HOST` | `http://localhost:11434` | | `agent-framework-ollama` | `OllamaChatClient` | `OLLAMA_MODEL` | `llama3.1:8b` | | `agent-framework-openai` | `OpenAIChatClient / OpenAIChatCompletionClient / OpenAIEmbeddingClient` | `OPENAI_API_KEY` | `sk-proj-...` | diff --git a/python/uv.lock b/python/uv.lock index 3807b2b29f9..f810d2aa56d 100644 --- a/python/uv.lock +++ b/python/uv.lock @@ -887,13 +887,13 @@ version = "1.0.0b260730" source = { editable = "packages/mistral" } dependencies = [ { name = "agent-framework-core" }, - { name = "mistralai" }, + { name = "httpx" }, ] [package.metadata] requires-dist = [ { name = "agent-framework-core", editable = "packages/core" }, - { name = "mistralai", specifier = ">=1.8.1,<3" }, + { name = "httpx", specifier = ">=0.23.1,<1" }, ] [[package]] @@ -2361,15 +2361,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/de/15/545e2b6cf2e3be84bc1ed85613edd75b8aea69807a71c26f4ca6a9258e82/email_validator-2.3.0-py3-none-any.whl", hash = "sha256:80f13f623413e6b197ae73bb10bf4eb0908faf509ad8362c5edeb0be7fd450b4", size = 35604, upload-time = "2025-08-26T13:09:05.858Z" }, ] -[[package]] -name = "eval-type-backport" -version = "0.4.0" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/1c/15/273a4baf8248d6d76220723c3caf039d283774b31a7c46ba686120145b76/eval_type_backport-0.4.0.tar.gz", hash = "sha256:8397d25e6524c2e67b9576bb0636be27dea2192017711220c534ec2de921e9b0", size = 10260, upload-time = "2026-06-02T13:22:06.059Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/50/a7/bb99bf5e6f78736ddb53480f2c3ff3702ffe2196a7c5e1661c03081d398e/eval_type_backport-0.4.0-py3-none-any.whl", hash = "sha256:ad5e2a8db71b6696a56eafb938b0f5a337d3217f256b8e158b469422b4772b20", size = 6432, upload-time = "2026-06-02T13:22:04.827Z" }, -] - [[package]] name = "execnet" version = "2.1.2" @@ -3309,15 +3300,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/cb/b1/3846dd7f199d53cb17f49cba7e651e9ce294d8497c8c150530ed11865bb8/iniconfig-2.3.0-py3-none-any.whl", hash = "sha256:f631c04d2c48c52b84d0d0549c99ff3859c98df65b3101406327ecc7d53fbf12", size = 7484, upload-time = "2025-10-18T21:55:41.639Z" }, ] -[[package]] -name = "invoke" -version = "2.2.1" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/de/bd/b461d3424a24c80490313fd77feeb666ca4f6a28c7e72713e3d9095719b4/invoke-2.2.1.tar.gz", hash = "sha256:515bf49b4a48932b79b024590348da22f39c4942dff991ad1fb8b8baea1be707", size = 304762, upload-time = "2025-10-11T00:36:35.172Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/32/4b/b99e37f88336009971405cbb7630610322ed6fbfa31e1d7ab3fbf3049a2d/invoke-2.2.1-py3-none-any.whl", hash = "sha256:2413bc441b376e5cd3f55bb5d364f973ad8bdd7bf87e53c79de3c11bf3feecc8", size = 160287, upload-time = "2025-10-11T00:36:33.703Z" }, -] - [[package]] name = "isodate" version = "0.7.2" @@ -4081,27 +4063,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/ad/8d/9a0d65e6e683778adf54c617e42ef008060d00e08f74255c45d409bfe9b5/microsoft_opentelemetry-1.3.5-py3-none-any.whl", hash = "sha256:36bf6eb0e90d358f5a1886eb907b2974a91f66b105c1bf3f25851b2a2f3bd327", size = 208881, upload-time = "2026-07-01T19:13:16.159Z" }, ] -[[package]] -name = "mistralai" -version = "1.12.4" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "eval-type-backport" }, - { name = "httpx" }, - { name = "invoke" }, - { name = "opentelemetry-api" }, - { name = "opentelemetry-exporter-otlp-proto-http" }, - { name = "opentelemetry-sdk" }, - { name = "pydantic" }, - { name = "python-dateutil" }, - { name = "pyyaml" }, - { name = "typing-inspection" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/aa/12/c3476c53e907255b5f485f085ba50dd9a84b40fe662e9a888d6ded26fa7b/mistralai-1.12.4.tar.gz", hash = "sha256:e52b53bab58025dcd208eeac13e3c3df5778d4112eeca1f08124096c7738929f", size = 243129, upload-time = "2026-02-20T17:55:13.73Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/c9/f9/98d825105c450b9c67c27026caa374112b7e466c18331601d02ca278a01b/mistralai-1.12.4-py3-none-any.whl", hash = "sha256:7b69fcbc306436491ad3377fbdead527c9f3a0ce145ec029bf04c6308ff2cca6", size = 509321, upload-time = "2026-02-20T17:55:15.27Z" }, -] - [[package]] name = "ml-dtypes" version = "0.5.4" From 7683e83118fa0c4cf1a9bd9154ea65525f38eb1b Mon Sep 17 00:00:00 2001 From: NekoPunch Date: Wed, 29 Jul 2026 10:22:42 -0700 Subject: [PATCH 2/8] refactor(python): simplify Mistral client per review Drop the streamed tool-call accumulator and multi-choice parsing in favor of the framework's built-in fragment merging, mark n unsupported, omit unset strict from json_schema, and leave CI secret wiring to maintainers. --- .../workflows/integration-tests-manual.yml | 1 - .../workflows/python-integration-tests.yml | 5 - .github/workflows/python-merge-tests.yml | 3 - python/packages/mistral/README.md | 10 +- .../agent_framework_mistral/_chat_client.py | 180 +++--------------- .../tests/mistral/test_mistral_chat_client.py | 115 ++--------- .../providers/mistral/mistral_agent_basic.py | 12 +- 7 files changed, 54 insertions(+), 272 deletions(-) diff --git a/.github/workflows/integration-tests-manual.yml b/.github/workflows/integration-tests-manual.yml index 209853d7749..f43e7bbbba6 100644 --- a/.github/workflows/integration-tests-manual.yml +++ b/.github/workflows/integration-tests-manual.yml @@ -132,5 +132,4 @@ jobs: AZURE_TENANT_ID: ${{ secrets.AZURE_TENANT_ID }} AZURE_SUBSCRIPTION_ID: ${{ secrets.AZURE_SUBSCRIPTION_ID }} FOUNDRY_MODELS_API_KEY: ${{ secrets.FOUNDRY_MODELS_API_KEY }} - MISTRAL_API_KEY: ${{ secrets.MISTRAL_API_KEY }} OPENAI__APIKEY: ${{ secrets.OPENAI__APIKEY }} diff --git a/.github/workflows/python-integration-tests.yml b/.github/workflows/python-integration-tests.yml index 62052fd9a7b..e9cff7d1252 100644 --- a/.github/workflows/python-integration-tests.yml +++ b/.github/workflows/python-integration-tests.yml @@ -27,8 +27,6 @@ on: required: true FOUNDRY_MODELS_API_KEY: required: false - MISTRAL_API_KEY: - required: false OPENAI__APIKEY: required: true @@ -174,9 +172,6 @@ jobs: ANTHROPIC_API_KEY: ${{ secrets.ANTHROPIC_API_KEY }} ANTHROPIC_CHAT_MODEL: ${{ vars.ANTHROPIC_CHAT_MODEL_ID }} LOCAL_MCP_URL: ${{ vars.LOCAL_MCP__URL }} - MISTRAL_API_KEY: ${{ secrets.MISTRAL_API_KEY }} - MISTRAL_CHAT_MODEL: ${{ vars.MISTRAL_CHAT_MODEL_ID }} - MISTRAL_EMBEDDING_MODEL: ${{ vars.MISTRAL_EMBEDDING_MODEL_ID }} OLLAMA_MODEL: qwen2.5:1.5b OLLAMA_EMBEDDING_MODEL: nomic-embed-text defaults: diff --git a/.github/workflows/python-merge-tests.yml b/.github/workflows/python-merge-tests.yml index 3113ca1664a..9207e381bf3 100644 --- a/.github/workflows/python-merge-tests.yml +++ b/.github/workflows/python-merge-tests.yml @@ -279,9 +279,6 @@ jobs: ANTHROPIC_API_KEY: ${{ secrets.ANTHROPIC_API_KEY }} ANTHROPIC_CHAT_MODEL: ${{ vars.ANTHROPIC_CHAT_MODEL_ID }} LOCAL_MCP_URL: ${{ vars.LOCAL_MCP__URL }} - MISTRAL_API_KEY: ${{ secrets.MISTRAL_API_KEY }} - MISTRAL_CHAT_MODEL: ${{ vars.MISTRAL_CHAT_MODEL_ID }} - MISTRAL_EMBEDDING_MODEL: ${{ vars.MISTRAL_EMBEDDING_MODEL_ID }} OLLAMA_MODEL: qwen2.5:1.5b OLLAMA_EMBEDDING_MODEL: nomic-embed-text defaults: diff --git a/python/packages/mistral/README.md b/python/packages/mistral/README.md index ccf4bd3c7fd..8c48d1b5ea8 100644 --- a/python/packages/mistral/README.md +++ b/python/packages/mistral/README.md @@ -19,18 +19,16 @@ streaming, function tools, and structured output. ### Quick Start ```python -from agent_framework import Agent, Message +from agent_framework import Agent from agent_framework.mistral import MistralChatClient # Using environment variables (MISTRAL_API_KEY, MISTRAL_CHAT_MODEL) +# Parameters can also be passed directly: +# MistralChatClient(model="mistral-large-latest", api_key="your-api-key") client = MistralChatClient() try: agent = Agent(client=client, instructions="You are a helpful assistant.") - result = await agent.run("Tell me a joke.") - - # Parameters can also be passed directly: - # MistralChatClient(model="mistral-large-latest", api_key="your-api-key") - response = await client.get_response([Message("user", ["Hello!"])]) + response = await agent.run("Hello!") print(response.text) finally: await client.close() diff --git a/python/packages/mistral/agent_framework_mistral/_chat_client.py b/python/packages/mistral/agent_framework_mistral/_chat_client.py index dbe0348511f..a15d0064863 100644 --- a/python/packages/mistral/agent_framework_mistral/_chat_client.py +++ b/python/packages/mistral/agent_framework_mistral/_chat_client.py @@ -94,6 +94,7 @@ class MistralChatOptions(ChatOptions[ResponseModelT], Generic[ResponseModelT], t - ``store`` - ``user`` - ``conversation_id`` + - ``n`` """ safe_prompt: bool @@ -105,12 +106,12 @@ class MistralChatOptions(ChatOptions[ResponseModelT], Generic[ResponseModelT], t prediction: dict[str, Any] """Predicted output to optimize response time when large parts of the response are known.""" - n: int - """Number of completions to return.""" - guardrails: list[dict[str, Any]] """Guardrail configurations applied to the request.""" + n: None + """Not supported. The framework expects a single completion per request.""" + prompt_cache_key: str """Cache key shared by requests with the same prompt prefix.""" @@ -207,45 +208,6 @@ def _tool_call_id_of(tool_call: Mapping[str, Any]) -> str: return "" -class _StreamedToolCall: - """Accumulates the fragments of one streamed tool call. - - Mistral streams may fragment a tool call across chunks — the ID typically arrives - only on the first fragment, and the name and arguments may arrive in pieces. - """ - - __slots__ = ("arguments_dict", "arguments_text", "call_id", "name", "raw") - - def __init__(self) -> None: - self.call_id = "" - self.name = "" - self.arguments_text = "" - self.arguments_dict: dict[str, Any] | None = None - self.raw: Any = None - - def add(self, tool_call: Mapping[str, Any]) -> None: - if not self.call_id: - self.call_id = _tool_call_id_of(tool_call) - function: Mapping[str, Any] = tool_call.get("function") or {} - if name := function.get("name"): - self.name += name - arguments = function.get("arguments") - if isinstance(arguments, dict): - self.arguments_dict = {**(self.arguments_dict or {}), **cast("dict[str, Any]", arguments)} - elif isinstance(arguments, str): - self.arguments_text += arguments - self.raw = tool_call - - def to_content(self, fallback_id: str) -> Content: - arguments = self.arguments_dict if self.arguments_dict is not None else self.arguments_text - return Content.from_function_call( - call_id=self.call_id or fallback_id, - name=self.name, - arguments=arguments, - raw_representation=self.raw, - ) - - class RawMistralChatClient( BaseChatClient[MistralChatOptionsT], Generic[MistralChatOptionsT], @@ -348,17 +310,13 @@ async def _stream() -> AsyncIterable[ChatResponseUpdate]: validated = await self._validate_options(options) request = self._prepare_request(messages, validated, **kwargs) request["stream"] = True - pending_tool_calls: dict[tuple[int, int], _StreamedToolCall] = {} - last_chunk_id = "" try: async with self.client.stream("POST", _CHAT_COMPLETIONS_PATH, json=request) as response: await self._raise_for_status(response) async for line in response.aiter_lines(): chunk = self._parse_sse_line(line) - if chunk is None: - continue - last_chunk_id = chunk.get("id") or last_chunk_id - yield self._parse_chunk(chunk, pending_tool_calls) + if chunk is not None: + yield self._parse_chunk(chunk) except ChatClientException: raise except Exception as ex: @@ -366,12 +324,6 @@ async def _stream() -> AsyncIterable[ChatResponseUpdate]: f"Mistral streaming chat request failed: {ex}", inner_exception=ex, ) from ex - if pending_tool_calls: - # The stream ended without a finish chunk; emit the accumulated calls. - yield ChatResponseUpdate( - contents=self._flush_pending_tool_calls(pending_tool_calls, last_chunk_id), - role="assistant", - ) return self._build_response_stream(_stream(), response_format=options.get("response_format")) @@ -674,14 +626,13 @@ def _prepare_response_format(self, response_format: Any) -> dict[str, Any] | Non return {"type": "json_object"} if format_type == "json_schema": json_schema: dict[str, Any] = dict(mapping.get("json_schema") or {}) - return { - "type": "json_schema", - "json_schema": { - "name": json_schema.get("name", "response"), - "schema": json_schema.get("schema") or json_schema.get("schema_definition") or {}, - "strict": json_schema.get("strict"), - }, + prepared_schema: dict[str, Any] = { + "name": json_schema.get("name", "response"), + "schema": json_schema.get("schema") or json_schema.get("schema_definition") or {}, } + if (strict := json_schema.get("strict")) is not None: + prepared_schema["strict"] = strict + return {"type": "json_schema", "json_schema": prepared_schema} # A raw JSON schema mapping return { "type": "json_schema", @@ -707,21 +658,16 @@ def _parse_response( response_format: Any | None = None, ) -> ChatResponse: """Convert a Mistral chat-completion response payload to a framework ChatResponse.""" - response_id = response.get("id") - messages: list[Message] = [] - finish_reason: FinishReasonLiteral | None = None choices = cast("Sequence[Mapping[str, Any]]", response.get("choices") or ()) - for choice in choices: - message: Mapping[str, Any] = choice.get("message") or {} - contents = self._parse_message_contents(message, fallback_id_prefix=response_id or "") - messages.append(Message(role="assistant", contents=contents, raw_representation=choice)) - if (reason := choice.get("finish_reason")) and finish_reason is None: - finish_reason = _FINISH_REASON_MAP.get(str(reason)) - if not messages: - messages.append(Message(role="assistant", contents=[])) + choice: Mapping[str, Any] = choices[0] if choices else {} + message: Mapping[str, Any] = choice.get("message") or {} + contents = self._parse_message_contents(message) + finish_reason: FinishReasonLiteral | None = None + if reason := choice.get("finish_reason"): + finish_reason = _FINISH_REASON_MAP.get(str(reason)) return ChatResponse( - response_id=response_id, - messages=messages, + response_id=response.get("id"), + messages=[Message(role="assistant", contents=contents, raw_representation=choice or None)], usage_details=self._parse_usage(response.get("usage")), model=response.get("model") or self.model, created_at=self._format_created_at(response.get("created")), @@ -730,31 +676,20 @@ def _parse_response( raw_representation=response, ) - def _parse_chunk( - self, - chunk: Mapping[str, Any], - pending_tool_calls: dict[tuple[int, int], _StreamedToolCall], - ) -> ChatResponseUpdate: + def _parse_chunk(self, chunk: Mapping[str, Any]) -> ChatResponseUpdate: """Convert a Mistral streaming completion chunk to a framework ChatResponseUpdate. - Tool-call fragments are accumulated in ``pending_tool_calls`` (keyed by choice and - tool-call index) and emitted as complete calls on the chunk that carries the finish - reason. Fragments cannot be emitted as they arrive: the ID and name may be split - across chunks and fragments of parallel calls may interleave, which the framework's - adjacency-based merging cannot reassemble. + Continuation fragments of a streamed tool call carry an empty ``call_id``; the + framework merges them into the preceding call when building the final response. """ - chunk_id = chunk.get("id") or "" contents: list[Content] = [] finish_reason: FinishReasonLiteral | None = None choices = cast("Sequence[Mapping[str, Any]]", chunk.get("choices") or ()) for choice in choices: delta: Mapping[str, Any] = choice.get("delta") or {} - contents.extend(self._parse_content_chunks(delta)) - self._accumulate_tool_call_fragments(choice, chunk_id, pending_tool_calls, contents) + contents.extend(self._parse_message_contents(delta)) if (reason := choice.get("finish_reason")) and finish_reason is None: finish_reason = _FINISH_REASON_MAP.get(str(reason)) - if finish_reason and pending_tool_calls: - contents.extend(self._flush_pending_tool_calls(pending_tool_calls, chunk_id)) if usage := self._parse_usage(chunk.get("usage")): contents.append(Content.from_usage(usage_details=usage, raw_representation=chunk)) return ChatResponseUpdate( @@ -767,62 +702,10 @@ def _parse_chunk( raw_representation=chunk, ) - def _accumulate_tool_call_fragments( - self, - choice: Mapping[str, Any], - chunk_id: str, - pending_tool_calls: dict[tuple[int, int], _StreamedToolCall], - contents: list[Content], - ) -> None: - delta: Mapping[str, Any] = choice.get("delta") or {} - tool_calls: Sequence[Mapping[str, Any]] = delta.get("tool_calls") or [] - if not tool_calls: - return - - indexes = [tool_call.get("index") for tool_call in tool_calls] - if len(tool_calls) > 1 and len(set(indexes)) != len(indexes): - # Parallel complete calls whose indexes don't distinguish them (the API may omit - # them entirely). They cannot be fragments, so emit directly with positionally - # distinct IDs. - for position, tool_call in enumerate(tool_calls): - call = _StreamedToolCall() - call.add(tool_call) - contents.append(call.to_content(f"{chunk_id}:{position}")) - return - - choice_index = choice.get("index") or 0 - for tool_call in tool_calls: - key = (choice_index, tool_call.get("index") or 0) - entry = pending_tool_calls.get(key) - incoming_id = _tool_call_id_of(tool_call) - if entry is not None and entry.call_id and incoming_id and incoming_id != entry.call_id: - # A new call reusing the index; the accumulated one is complete. - contents.append(entry.to_content(f"{chunk_id}:{key[1]}")) - entry = None - if entry is None: - entry = pending_tool_calls[key] = _StreamedToolCall() - entry.add(tool_call) - - @staticmethod - def _flush_pending_tool_calls( - pending_tool_calls: dict[tuple[int, int], _StreamedToolCall], - fallback_id_prefix: str, - ) -> list[Content]: - contents = [ - entry.to_content(f"{fallback_id_prefix}:{index}") for (_, index), entry in pending_tool_calls.items() - ] - pending_tool_calls.clear() - return contents - - def _parse_message_contents( - self, - message: Mapping[str, Any], - *, - fallback_id_prefix: str = "", - ) -> list[Content]: + def _parse_message_contents(self, message: Mapping[str, Any]) -> list[Content]: contents = self._parse_content_chunks(message) tool_calls = cast("Sequence[Mapping[str, Any]]", message.get("tool_calls") or ()) - for position, tool_call in enumerate(tool_calls): + for tool_call in tool_calls: function: Mapping[str, Any] = tool_call.get("function") or {} arguments = function.get("arguments") if isinstance(arguments, str): @@ -833,7 +716,7 @@ def _parse_message_contents( normalized_arguments = str(cast(object, arguments)) contents.append( Content.from_function_call( - call_id=self._resolve_tool_call_id(tool_call, position, fallback_id_prefix), + call_id=_tool_call_id_of(tool_call), name=function.get("name") or "", arguments=normalized_arguments, raw_representation=tool_call, @@ -860,15 +743,6 @@ def _parse_content_chunks(self, message: Mapping[str, Any]) -> list[Content]: logger.debug("Skipping unsupported response chunk from Mistral: %s", chunk_type) return contents - @staticmethod - def _resolve_tool_call_id(tool_call: Mapping[str, Any], position: int, fallback_id_prefix: str) -> str: - """Return the tool call ID, synthesizing a distinct one when the API omits it.""" - if call_id := _tool_call_id_of(tool_call): - return call_id - # The index only disambiguates when non-zero (the API may omit it); otherwise position does. - index = tool_call.get("index") or position - return f"{fallback_id_prefix}:{index}" - @staticmethod def _format_created_at(created: Any) -> str | None: if not isinstance(created, (int, float)): diff --git a/python/packages/mistral/tests/mistral/test_mistral_chat_client.py b/python/packages/mistral/tests/mistral/test_mistral_chat_client.py index db9006c59f6..c7f66ebe65f 100644 --- a/python/packages/mistral/tests/mistral/test_mistral_chat_client.py +++ b/python/packages/mistral/tests/mistral/test_mistral_chat_client.py @@ -510,6 +510,23 @@ async def test_response_format_json_object() -> None: assert server.last_request["response_format"] == {"type": "json_object"} +async def test_response_format_json_schema_omits_unset_strict() -> None: + client, server = make_client(json_response(make_response_payload(content="{}"))) + + await client.get_response( + [Message("user", ["hi"])], + options={ + "response_format": { + "type": "json_schema", + "json_schema": {"name": "answer", "schema": {"type": "object"}}, + } + }, + ) + + json_schema = server.last_request["response_format"]["json_schema"] + assert "strict" not in json_schema + + def test_response_format_edge_cases(caplog: pytest.LogCaptureFixture) -> None: client, _ = make_client() assert client._prepare_response_format("json") == {"type": "json_object"} # pyright: ignore[reportPrivateUsage] @@ -563,46 +580,6 @@ async def test_parse_tool_calls() -> None: assert calls[0].parse_arguments() == {"location": "Paris"} -async def test_parse_tool_calls_null_ids_stay_distinct() -> None: - """The Mistral API omits tool call IDs on some paths (serialized as null or "null").""" - client, _ = make_client( - json_response( - make_response_payload( - tool_calls=[ - tool_call_payload("get_weather", '{"location": "Paris"}', call_id="null"), - tool_call_payload("get_time", '{"tz": "CET"}', call_id="null"), - ], - finish_reason="tool_calls", - ) - ) - ) - - response = await client.get_response([Message("user", ["hi"])]) - - calls = [c for c in response.messages[0].contents if c.type == "function_call"] - assert len(calls) == 2 - assert calls[0].call_id != calls[1].call_id - assert "null" not in (calls[0].call_id, calls[1].call_id) - assert calls[0].parse_arguments() == {"location": "Paris"} - assert calls[1].parse_arguments() == {"tz": "CET"} - - -async def test_parse_multiple_choices() -> None: - payload = make_response_payload( - choices=[ - {"index": i, "finish_reason": "stop", "message": {"role": "assistant", "content": f"choice {i}"}} - for i in range(2) - ] - ) - client, _ = make_client(json_response(payload)) - - response = await client.get_response([Message("user", ["hi"])]) - - assert len(response.messages) == 2 - assert response.messages[0].text == "choice 0" - assert response.messages[1].text == "choice 1" - - async def test_parse_empty_choices_returns_empty_assistant_message() -> None: client, _ = make_client(json_response(make_response_payload(choices=[]))) @@ -758,64 +735,6 @@ async def test_streaming_fragmented_tool_call_coalesces() -> None: assert calls[0].parse_arguments() == {"location": "Paris"} -async def test_streaming_fragmented_function_name_coalesces() -> None: - """The function name itself may be split across fragments.""" - client, _ = make_client( - stream_response( - make_chunk_payload(tool_calls=[tool_call_payload("get_", "", call_id="abc123XYZ", index=0)]), - make_chunk_payload( - tool_calls=[tool_call_payload("weather", '{"location": "Paris"}', index=0)], - finish_reason="tool_calls", - ), - ) - ) - - stream = client.get_response([Message("user", ["hi"])], stream=True) - async for _ in stream: - pass - response = await stream.get_final_response() - - calls = [c for c in response.messages[0].contents if c.type == "function_call"] - assert len(calls) == 1 - assert calls[0].call_id == "abc123XYZ" - assert calls[0].name == "get_weather" - assert calls[0].parse_arguments() == {"location": "Paris"} - - -async def test_streaming_interleaved_parallel_tool_calls() -> None: - """Fragments of two parallel calls may interleave across chunks.""" - client, _ = make_client( - stream_response( - make_chunk_payload( - tool_calls=[ - tool_call_payload("get_weather", '{"loc', call_id="abc123XYZ", index=0), - tool_call_payload("get_time", '{"t', call_id="def456UVW", index=1), - ], - ), - make_chunk_payload( - tool_calls=[ - tool_call_payload("", 'ation": "Paris"}', index=0), - tool_call_payload("", 'z": "CET"}', index=1), - ], - finish_reason="tool_calls", - ), - ) - ) - - stream = client.get_response([Message("user", ["hi"])], stream=True) - async for _ in stream: - pass - response = await stream.get_final_response() - - calls = [c for c in response.messages[0].contents if c.type == "function_call"] - assert len(calls) == 2 - by_id = {c.call_id: c for c in calls} - assert by_id["abc123XYZ"].name == "get_weather" - assert by_id["abc123XYZ"].parse_arguments() == {"location": "Paris"} - assert by_id["def456UVW"].name == "get_time" - assert by_id["def456UVW"].parse_arguments() == {"tz": "CET"} - - async def test_streaming_parallel_calls_without_indexes() -> None: client, _ = make_client( stream_response( diff --git a/python/samples/02-agents/providers/mistral/mistral_agent_basic.py b/python/samples/02-agents/providers/mistral/mistral_agent_basic.py index 4c604e0ef8a..b0ecc63b6cb 100644 --- a/python/samples/02-agents/providers/mistral/mistral_agent_basic.py +++ b/python/samples/02-agents/providers/mistral/mistral_agent_basic.py @@ -4,6 +4,12 @@ # dependencies = ["agent-framework-mistral"] # /// +"""Demonstrates a Mistral AI agent with basic tool usage. + +Requires ``MISTRAL_API_KEY`` and ``MISTRAL_CHAT_MODEL`` environment variables +(e.g. MISTRAL_CHAT_MODEL=mistral-small-latest). +""" + import asyncio from datetime import datetime from zoneinfo import ZoneInfo @@ -12,12 +18,6 @@ from agent_framework.mistral import MistralChatClient from dotenv import load_dotenv -"""Demonstrates a Mistral AI agent with basic tool usage. - -Requires ``MISTRAL_API_KEY`` and ``MISTRAL_CHAT_MODEL`` environment variables -(e.g. MISTRAL_CHAT_MODEL=mistral-small-latest). -""" - # Load environment variables from the local .env file. load_dotenv() From 686bf40deceb1da32ad1771ea308fc13b83bfa45 Mon Sep 17 00:00:00 2001 From: NekoPunch Date: Thu, 30 Jul 2026 01:48:22 -0700 Subject: [PATCH 3/8] test(python): drop n forwarding assertion n is typed as unsupported on MistralChatOptions; the option-mapping test still passed n, failing pyrefly/ty/zuban/mypy in CI. --- .../packages/mistral/tests/mistral/test_mistral_chat_client.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/python/packages/mistral/tests/mistral/test_mistral_chat_client.py b/python/packages/mistral/tests/mistral/test_mistral_chat_client.py index c7f66ebe65f..ba286e31a2e 100644 --- a/python/packages/mistral/tests/mistral/test_mistral_chat_client.py +++ b/python/packages/mistral/tests/mistral/test_mistral_chat_client.py @@ -263,7 +263,6 @@ async def test_get_response_option_mapping() -> None: "allow_multiple_tool_calls": False, "safe_prompt": True, "stop": ["END"], - "n": 2, "guardrails": [{"name": "test-guardrail"}], "prompt_cache_key": "shared-prefix", "reasoning_effort": "high", @@ -277,7 +276,7 @@ async def test_get_response_option_mapping() -> None: assert request["parallel_tool_calls"] is False assert request["safe_prompt"] is True assert request["stop"] == ["END"] - assert request["n"] == 2 + assert "n" not in request assert request["guardrails"] == [{"name": "test-guardrail"}] assert request["prompt_cache_key"] == "shared-prefix" assert request["reasoning_effort"] == "high" From f999f304f58eedb114c12c77e858672cfff68a63 Mon Sep 17 00:00:00 2001 From: NekoPunch Date: Thu, 30 Jul 2026 02:11:36 -0700 Subject: [PATCH 4/8] refactor(python): drop n from MistralChatOptions n is not part of the base ChatOptions, so removing the key rejects it without an explicit None override. --- .../packages/mistral/agent_framework_mistral/_chat_client.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/python/packages/mistral/agent_framework_mistral/_chat_client.py b/python/packages/mistral/agent_framework_mistral/_chat_client.py index a15d0064863..8635ec62204 100644 --- a/python/packages/mistral/agent_framework_mistral/_chat_client.py +++ b/python/packages/mistral/agent_framework_mistral/_chat_client.py @@ -94,7 +94,6 @@ class MistralChatOptions(ChatOptions[ResponseModelT], Generic[ResponseModelT], t - ``store`` - ``user`` - ``conversation_id`` - - ``n`` """ safe_prompt: bool @@ -109,9 +108,6 @@ class MistralChatOptions(ChatOptions[ResponseModelT], Generic[ResponseModelT], t guardrails: list[dict[str, Any]] """Guardrail configurations applied to the request.""" - n: None - """Not supported. The framework expects a single completion per request.""" - prompt_cache_key: str """Cache key shared by requests with the same prompt prefix.""" From 1865140e2f1b6da30179d2055509255797205c3d Mon Sep 17 00:00:00 2001 From: NekoPunch Date: Thu, 30 Jul 2026 09:51:45 -0700 Subject: [PATCH 5/8] feat(python): mark Mistral feature usage Both clients flip the shared FeatureIndex.MISTRAL bit before each request, matching the feature-usage telemetry other providers emit. --- .../agent_framework_mistral/_chat_client.py | 6 +++++- .../agent_framework_mistral/_embedding_client.py | 5 ++++- .../tests/mistral/test_mistral_chat_client.py | 14 ++++++++++++++ .../mistral/test_mistral_embedding_client.py | 15 +++++++++++++++ 4 files changed, 38 insertions(+), 2 deletions(-) diff --git a/python/packages/mistral/agent_framework_mistral/_chat_client.py b/python/packages/mistral/agent_framework_mistral/_chat_client.py index 8635ec62204..c3dba5540d4 100644 --- a/python/packages/mistral/agent_framework_mistral/_chat_client.py +++ b/python/packages/mistral/agent_framework_mistral/_chat_client.py @@ -30,7 +30,7 @@ validate_tool_mode, ) from agent_framework._settings import SecretString, load_settings -from agent_framework._telemetry import get_user_agent +from agent_framework._telemetry import get_user_agent, mark_feature_used from agent_framework._types import prepend_instructions_to_messages from agent_framework.exceptions import ( ChatClientException, @@ -41,6 +41,8 @@ from agent_framework.observability import ChatTelemetryLayer from pydantic import BaseModel +from ._feature_usage import FeatureIndex + if sys.version_info >= (3, 13): from typing import TypeVar # pragma: no cover else: @@ -306,6 +308,7 @@ async def _stream() -> AsyncIterable[ChatResponseUpdate]: validated = await self._validate_options(options) request = self._prepare_request(messages, validated, **kwargs) request["stream"] = True + mark_feature_used(FeatureIndex.MISTRAL) try: async with self.client.stream("POST", _CHAT_COMPLETIONS_PATH, json=request) as response: await self._raise_for_status(response) @@ -326,6 +329,7 @@ async def _stream() -> AsyncIterable[ChatResponseUpdate]: async def _get_response() -> ChatResponse: validated = await self._validate_options(options) request = self._prepare_request(messages, validated, **kwargs) + mark_feature_used(FeatureIndex.MISTRAL) try: response = await self.client.post(_CHAT_COMPLETIONS_PATH, json=request) await self._raise_for_status(response) diff --git a/python/packages/mistral/agent_framework_mistral/_embedding_client.py b/python/packages/mistral/agent_framework_mistral/_embedding_client.py index dd6b24f04e0..63f462867f5 100644 --- a/python/packages/mistral/agent_framework_mistral/_embedding_client.py +++ b/python/packages/mistral/agent_framework_mistral/_embedding_client.py @@ -17,7 +17,7 @@ load_settings, ) from agent_framework._settings import SecretString -from agent_framework._telemetry import get_user_agent +from agent_framework._telemetry import get_user_agent, mark_feature_used from agent_framework.exceptions import ( IntegrationException, IntegrationInvalidAuthException, @@ -26,6 +26,8 @@ ) from agent_framework.observability import EmbeddingTelemetryLayer +from ._feature_usage import FeatureIndex + if sys.version_info >= (3, 13): from typing import TypeVar # pragma: no cover else: @@ -191,6 +193,7 @@ async def get_embeddings( if "dimensions" in opts: request["output_dimension"] = opts["dimensions"] + mark_feature_used(FeatureIndex.MISTRAL) try: response = await self.client.post(_EMBEDDINGS_PATH, json=request) if response.status_code >= 400: diff --git a/python/packages/mistral/tests/mistral/test_mistral_chat_client.py b/python/packages/mistral/tests/mistral/test_mistral_chat_client.py index ba286e31a2e..50492180425 100644 --- a/python/packages/mistral/tests/mistral/test_mistral_chat_client.py +++ b/python/packages/mistral/tests/mistral/test_mistral_chat_client.py @@ -189,6 +189,20 @@ async def test_mistral_chat_missing_model_raises_at_request(monkeypatch: pytest. # region: Request preparation +async def test_get_response_marks_feature_used(monkeypatch: pytest.MonkeyPatch) -> None: + from unittest.mock import MagicMock + + from agent_framework_mistral._feature_usage import FeatureIndex + + mark = MagicMock() + monkeypatch.setattr(chat_client_module, "mark_feature_used", mark) + client, _ = make_client(json_response(make_response_payload(content="ok"))) + + await client.get_response([Message("user", ["hi"])]) + + mark.assert_called_once_with(FeatureIndex.MISTRAL) + + async def test_get_response_basic() -> None: client, server = make_client(json_response(make_response_payload(content="hello"))) diff --git a/python/packages/mistral/tests/mistral/test_mistral_embedding_client.py b/python/packages/mistral/tests/mistral/test_mistral_embedding_client.py index ccee2de6cea..6aea0924574 100644 --- a/python/packages/mistral/tests/mistral/test_mistral_embedding_client.py +++ b/python/packages/mistral/tests/mistral/test_mistral_embedding_client.py @@ -142,6 +142,21 @@ async def test_mistral_embedding_close_only_closes_owned_client() -> None: await http_client.aclose() +async def test_mistral_embedding_marks_feature_used(monkeypatch: pytest.MonkeyPatch) -> None: + from unittest.mock import MagicMock + + import agent_framework_mistral._embedding_client as embedding_client_module + from agent_framework_mistral._feature_usage import FeatureIndex + + mark = MagicMock() + monkeypatch.setattr(embedding_client_module, "mark_feature_used", mark) + client, _ = make_client(httpx.Response(200, json=make_embeddings_payload([[0.1, 0.2]]))) + + await client.get_embeddings(["hello"]) + + mark.assert_called_once_with(FeatureIndex.MISTRAL) + + async def test_mistral_embedding_get_embeddings() -> None: """Test generating embeddings via the Mistral API.""" client, server = make_client(httpx.Response(200, json=make_embeddings_payload([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]))) From 7301acb28fd4aba3a2e0fd18fae32d71df0a4732 Mon Sep 17 00:00:00 2001 From: NekoPunch Date: Sun, 2 Aug 2026 21:37:16 -0700 Subject: [PATCH 6/8] fix(python): key streamed tool calls by index Mistral omits the tool call id on continuation fragments, and the framework only coalesces empty-id fragments into the immediately preceding call, so interleaved parallel calls merged into the wrong call with corrupted arguments. Accumulate fragments per (choice, index) and emit each call only once complete. --- .../agent_framework_mistral/_chat_client.py | 129 ++++++++++++++---- .../tests/mistral/test_mistral_chat_client.py | 34 +++++ 2 files changed, 139 insertions(+), 24 deletions(-) diff --git a/python/packages/mistral/agent_framework_mistral/_chat_client.py b/python/packages/mistral/agent_framework_mistral/_chat_client.py index c3dba5540d4..18f8ecd17fb 100644 --- a/python/packages/mistral/agent_framework_mistral/_chat_client.py +++ b/python/packages/mistral/agent_framework_mistral/_chat_client.py @@ -206,6 +206,95 @@ def _tool_call_id_of(tool_call: Mapping[str, Any]) -> str: return "" +def _function_call_content(tool_call: Mapping[str, Any]) -> Content: + function: Mapping[str, Any] = tool_call.get("function") or {} + arguments = function.get("arguments") + if isinstance(arguments, str): + normalized_arguments: str | dict[str, Any] = arguments + elif isinstance(arguments, dict): + normalized_arguments = cast("dict[str, Any]", arguments) + else: + normalized_arguments = str(cast(object, arguments)) + return Content.from_function_call( + call_id=_tool_call_id_of(tool_call), + name=function.get("name") or "", + arguments=normalized_arguments, + raw_representation=tool_call, + ) + + +class _StreamedToolCalls: + """Correlates streamed tool-call fragments by ``(choice index, tool-call index)``. + + Mistral may interleave fragments of parallel calls and omit ``id`` on + continuations, so a call is only emitted once it is complete: when its + choice finishes, when its index is reused by a new call, or at stream end. + """ + + def __init__(self) -> None: + self._pending: dict[tuple[int, int | str], dict[str, Any]] = {} + self._auto_key_count = 0 + + def add(self, choice_index: int, fragment: Mapping[str, Any]) -> list[Content]: + """Fold a fragment into its pending call; returns calls completed by an index reuse.""" + flushed: list[Content] = [] + key = self._key_for(choice_index, fragment) + pending = self._pending.get(key) + if pending is not None: + fragment_id = _tool_call_id_of(fragment) + if fragment_id and (pending_id := _tool_call_id_of(pending)) and pending_id != fragment_id: + flushed.append(_function_call_content(self._pending.pop(key))) + pending = None + if pending is None: + self._pending[key] = {**fragment, "function": dict(fragment.get("function") or {})} + else: + self._merge(pending, fragment) + return flushed + + def flush_choice(self, choice_index: int) -> list[Content]: + keys = [key for key in self._pending if key[0] == choice_index] + return [_function_call_content(self._pending.pop(key)) for key in keys] + + def flush_all(self) -> list[Content]: + contents = [_function_call_content(pending) for pending in self._pending.values()] + self._pending.clear() + return contents + + def _key_for(self, choice_index: int, fragment: Mapping[str, Any]) -> tuple[int, int | str]: + index = fragment.get("index") + if isinstance(index, int): + return (choice_index, index) + if fragment_id := _tool_call_id_of(fragment): + for key, pending in self._pending.items(): + if key[0] == choice_index and _tool_call_id_of(pending) == fragment_id: + return key + else: + for key in reversed(self._pending): + if key[0] == choice_index: + return key + self._auto_key_count += 1 + return (choice_index, f"auto-{self._auto_key_count}") + + @staticmethod + def _merge(pending: dict[str, Any], fragment: Mapping[str, Any]) -> None: + if fragment_id := _tool_call_id_of(fragment): + pending["id"] = fragment_id + function: Mapping[str, Any] = fragment.get("function") or {} + pending_function: dict[str, Any] = pending["function"] + if (name := function.get("name")) and not pending_function.get("name"): + pending_function["name"] = name + new_arguments = function.get("arguments") + old_arguments = pending_function.get("arguments") + if new_arguments is None: + return + if isinstance(old_arguments, str) and isinstance(new_arguments, str): + pending_function["arguments"] = old_arguments + new_arguments + elif isinstance(old_arguments, dict) and isinstance(new_arguments, dict): + cast("dict[str, Any]", old_arguments).update(cast("dict[str, Any]", new_arguments)) + else: + pending_function["arguments"] = new_arguments + + class RawMistralChatClient( BaseChatClient[MistralChatOptionsT], Generic[MistralChatOptionsT], @@ -309,13 +398,16 @@ async def _stream() -> AsyncIterable[ChatResponseUpdate]: request = self._prepare_request(messages, validated, **kwargs) request["stream"] = True mark_feature_used(FeatureIndex.MISTRAL) + tool_calls = _StreamedToolCalls() try: async with self.client.stream("POST", _CHAT_COMPLETIONS_PATH, json=request) as response: await self._raise_for_status(response) async for line in response.aiter_lines(): chunk = self._parse_sse_line(line) if chunk is not None: - yield self._parse_chunk(chunk) + yield self._parse_chunk(chunk, tool_calls) + if remaining := tool_calls.flush_all(): + yield ChatResponseUpdate(contents=remaining, role="assistant") except ChatClientException: raise except Exception as ex: @@ -676,20 +768,25 @@ def _parse_response( raw_representation=response, ) - def _parse_chunk(self, chunk: Mapping[str, Any]) -> ChatResponseUpdate: + def _parse_chunk(self, chunk: Mapping[str, Any], tool_calls: _StreamedToolCalls) -> ChatResponseUpdate: """Convert a Mistral streaming completion chunk to a framework ChatResponseUpdate. - Continuation fragments of a streamed tool call carry an empty ``call_id``; the - framework merges them into the preceding call when building the final response. + Tool-call fragments are folded into ``tool_calls`` keyed by (choice, index) and + emitted as complete calls when their choice finishes. """ contents: list[Content] = [] finish_reason: FinishReasonLiteral | None = None choices = cast("Sequence[Mapping[str, Any]]", chunk.get("choices") or ()) for choice in choices: + choice_index = index if isinstance(index := choice.get("index"), int) else 0 delta: Mapping[str, Any] = choice.get("delta") or {} - contents.extend(self._parse_message_contents(delta)) - if (reason := choice.get("finish_reason")) and finish_reason is None: - finish_reason = _FINISH_REASON_MAP.get(str(reason)) + contents.extend(self._parse_content_chunks(delta)) + for fragment in cast("Sequence[Mapping[str, Any]]", delta.get("tool_calls") or ()): + contents.extend(tool_calls.add(choice_index, fragment)) + if reason := choice.get("finish_reason"): + contents.extend(tool_calls.flush_choice(choice_index)) + if finish_reason is None: + finish_reason = _FINISH_REASON_MAP.get(str(reason)) if usage := self._parse_usage(chunk.get("usage")): contents.append(Content.from_usage(usage_details=usage, raw_representation=chunk)) return ChatResponseUpdate( @@ -705,23 +802,7 @@ def _parse_chunk(self, chunk: Mapping[str, Any]) -> ChatResponseUpdate: def _parse_message_contents(self, message: Mapping[str, Any]) -> list[Content]: contents = self._parse_content_chunks(message) tool_calls = cast("Sequence[Mapping[str, Any]]", message.get("tool_calls") or ()) - for tool_call in tool_calls: - function: Mapping[str, Any] = tool_call.get("function") or {} - arguments = function.get("arguments") - if isinstance(arguments, str): - normalized_arguments: str | dict[str, Any] = arguments - elif isinstance(arguments, dict): - normalized_arguments = cast("dict[str, Any]", arguments) - else: - normalized_arguments = str(cast(object, arguments)) - contents.append( - Content.from_function_call( - call_id=_tool_call_id_of(tool_call), - name=function.get("name") or "", - arguments=normalized_arguments, - raw_representation=tool_call, - ) - ) + contents.extend(_function_call_content(tool_call) for tool_call in tool_calls) return contents def _parse_content_chunks(self, message: Mapping[str, Any]) -> list[Content]: diff --git a/python/packages/mistral/tests/mistral/test_mistral_chat_client.py b/python/packages/mistral/tests/mistral/test_mistral_chat_client.py index 50492180425..5249ea1d9d4 100644 --- a/python/packages/mistral/tests/mistral/test_mistral_chat_client.py +++ b/python/packages/mistral/tests/mistral/test_mistral_chat_client.py @@ -748,6 +748,40 @@ async def test_streaming_fragmented_tool_call_coalesces() -> None: assert calls[0].parse_arguments() == {"location": "Paris"} +async def test_streaming_interleaved_parallel_tool_calls() -> None: + """Continuation fragments without IDs must merge into the call with the same index, not the preceding one.""" + client, _ = make_client( + stream_response( + make_chunk_payload( + tool_calls=[tool_call_payload("get_weather", '{"loc', call_id="abc123XYZ", index=0)], + ), + make_chunk_payload( + tool_calls=[tool_call_payload("get_time", '{"tz', call_id="def456UVW", index=1)], + ), + make_chunk_payload( + tool_calls=[tool_call_payload("", 'ation": "Paris"}', index=0)], + ), + make_chunk_payload( + tool_calls=[tool_call_payload("", '": "CET"}', index=1)], + finish_reason="tool_calls", + ), + ) + ) + + stream = client.get_response([Message("user", ["hi"])], stream=True) + async for _ in stream: + pass + response = await stream.get_final_response() + + calls = [c for c in response.messages[0].contents if c.type == "function_call"] + assert len(calls) == 2 + by_id = {c.call_id: c for c in calls} + assert by_id["abc123XYZ"].name == "get_weather" + assert by_id["abc123XYZ"].parse_arguments() == {"location": "Paris"} + assert by_id["def456UVW"].name == "get_time" + assert by_id["def456UVW"].parse_arguments() == {"tz": "CET"} + + async def test_streaming_parallel_calls_without_indexes() -> None: client, _ = make_client( stream_response( From 89448cf94d8235229a7743761e8d96d82b024ec5 Mon Sep 17 00:00:00 2001 From: NekoPunch Date: Sun, 2 Aug 2026 21:37:24 -0700 Subject: [PATCH 7/8] fix(python): restore Mistral SDK client injection Dropping the mistralai dependency turned the embedding client's client= parameter into a breaking change for injected SDK clients. Add http_client= for httpx.AsyncClient and keep client= working: httpx goes to the REST path, a duck-typed mistralai.Mistral goes through the legacy SDK path with a DeprecationWarning until the next major release. --- .../_embedding_client.py | 107 ++++++++++++++++-- .../mistral/test_mistral_embedding_client.py | 62 +++++++++- 2 files changed, 152 insertions(+), 17 deletions(-) diff --git a/python/packages/mistral/agent_framework_mistral/_embedding_client.py b/python/packages/mistral/agent_framework_mistral/_embedding_client.py index 63f462867f5..b347f870f49 100644 --- a/python/packages/mistral/agent_framework_mistral/_embedding_client.py +++ b/python/packages/mistral/agent_framework_mistral/_embedding_client.py @@ -4,6 +4,7 @@ import logging import sys +import warnings from collections.abc import Mapping, Sequence from typing import Any, ClassVar, Generic, TypedDict, cast @@ -41,6 +42,35 @@ _DEFAULT_TIMEOUT_SECONDS = 60.0 +def _resolve_injected_clients( + http_client: httpx.AsyncClient | None, + client: Any | None, +) -> tuple[httpx.AsyncClient | None, Any | None]: + """Split the deprecated ``client`` parameter into REST and legacy-SDK forms. + + Returns ``(http_client, sdk_client)``; at most one is set. The SDK form is + duck-typed on ``.embeddings`` so the ``mistralai`` dependency stays optional. + """ + if client is None: + return http_client, None + warnings.warn( + "The 'client' parameter is deprecated; pass an httpx.AsyncClient as 'http_client' instead. " + "Support for injected mistralai.Mistral clients will be removed in the next major release.", + DeprecationWarning, + stacklevel=3, + ) + if http_client is not None: + raise ValueError("Provide either 'http_client' or the deprecated 'client' parameter, not both.") + if isinstance(client, httpx.AsyncClient): + return client, None + if hasattr(client, "embeddings"): + return None, client + raise TypeError( + "The 'client' parameter accepts an httpx.AsyncClient or a mistralai.Mistral instance; " + f"got {type(client).__name__}." + ) + + class MistralEmbeddingOptions(EmbeddingGenerationOptions, total=False): """Mistral AI-specific embedding options. @@ -94,14 +124,17 @@ class RawMistralEmbeddingClient( api_key: Mistral API key. Defaults to ``MISTRAL_API_KEY`` environment variable. server_url: Optional server URL override. Defaults to ``MISTRAL_SERVER_URL`` environment variable, or the Mistral default. - client: Optional pre-configured ``httpx.AsyncClient``. When provided, api_key is + http_client: Optional pre-configured ``httpx.AsyncClient``. When provided, api_key is not required and the client is expected to carry its own auth headers and base URL. + client: Deprecated. Accepts an ``httpx.AsyncClient`` (treated as ``http_client``) or a + ``mistralai.Mistral`` instance, which keeps working through the legacy SDK path + until the next major release. additional_properties: Additional properties stored on the client instance. env_file_path: Path to ``.env`` file for settings. env_file_encoding: Encoding for ``.env`` file. """ - INJECTABLE: ClassVar[set[str]] = {"client"} + INJECTABLE: ClassVar[set[str]] = {"http_client", "client"} def __init__( self, @@ -109,13 +142,16 @@ def __init__( model: str | None = None, api_key: str | SecretString | None = None, server_url: str | None = None, - client: httpx.AsyncClient | None = None, + http_client: httpx.AsyncClient | None = None, + client: Any | None = None, additional_properties: dict[str, Any] | None = None, env_file_path: str | None = None, env_file_encoding: str | None = None, ) -> None: """Initialize a raw Mistral AI embedding client.""" - required_fields = ["embedding_model"] if client is not None else ["embedding_model", "api_key"] + http_client, sdk_client = _resolve_injected_clients(http_client, client) + injected = http_client is not None or sdk_client is not None + required_fields = ["embedding_model"] if injected else ["embedding_model", "api_key"] mistral_settings = load_settings( MistralEmbeddingSettings, env_prefix="MISTRAL_", @@ -129,12 +165,16 @@ def __init__( self.model: str = mistral_settings["embedding_model"] # type: ignore[assignment] self.server_url = mistral_settings.get("server_url") - self._owns_client = client is None - - if client is not None: - self.client = client + self._owns_client = not injected + self._sdk_client = sdk_client + self.client: Any + + if sdk_client is not None: + self.client = sdk_client + elif http_client is not None: + self.client = http_client if self.server_url is None: - client_base_url = str(client.base_url).rstrip("/") + client_base_url = str(http_client.base_url).rstrip("/") self.server_url = client_base_url or None else: resolved_api_key: str = mistral_settings["api_key"] # type: ignore[assignment] @@ -189,11 +229,14 @@ async def get_embeddings( if not model: raise ValueError("model is required") + mark_feature_used(FeatureIndex.MISTRAL) + if self._sdk_client is not None: + return await self._get_embeddings_sdk(self._sdk_client, model, values, opts, options) + request: dict[str, Any] = {"model": model, "input": list(values)} if "dimensions" in opts: request["output_dimension"] = opts["dimensions"] - mark_feature_used(FeatureIndex.MISTRAL) try: response = await self.client.post(_EMBEDDINGS_PATH, json=request) if response.status_code >= 400: @@ -245,6 +288,43 @@ async def get_embeddings( inner_exception=ex, ) from ex + async def _get_embeddings_sdk( + self, + sdk_client: Any, + model: str, + values: Sequence[str], + opts: Mapping[str, Any], + options: MistralEmbeddingOptionsT | None, + ) -> GeneratedEmbeddings[list[float], MistralEmbeddingOptionsT]: + """Legacy path for injected mistralai.Mistral clients; removed in the next major release.""" + kwargs: dict[str, Any] = {"model": model, "inputs": list(values)} + if "dimensions" in opts: + kwargs["output_dimension"] = opts["dimensions"] + + response = await sdk_client.embeddings.create_async(**kwargs) + + embeddings: list[Embedding[list[float]]] = [] + if response and response.data: + items = sorted(response.data, key=lambda d: d.index if d.index is not None else 0) + for item in items: + vector = list(item.embedding) if item.embedding else [] + embeddings.append( + Embedding( + vector=vector, + dimensions=len(vector), + model=response.model or model, + ) + ) + + usage_dict: UsageDetails | None = None + if response and response.usage: + usage_dict = { + "input_token_count": response.usage.prompt_tokens, + "total_token_count": response.usage.total_tokens, + } + + return GeneratedEmbeddings(embeddings, options=options, usage=usage_dict) + class MistralEmbeddingClient( EmbeddingTelemetryLayer[str, list[float], MistralEmbeddingOptionsT], @@ -259,7 +339,8 @@ class MistralEmbeddingClient( api_key: Mistral API key. Defaults to ``MISTRAL_API_KEY`` environment variable. server_url: Optional server URL override. Defaults to ``MISTRAL_SERVER_URL`` environment variable, or the Mistral default. - client: Optional pre-configured ``httpx.AsyncClient``. + http_client: Optional pre-configured ``httpx.AsyncClient``. + client: Deprecated. Accepts an ``httpx.AsyncClient`` or a ``mistralai.Mistral`` instance. otel_provider_name: Optional telemetry provider name override. env_file_path: Path to ``.env`` file for settings. env_file_encoding: Encoding for ``.env`` file. @@ -294,7 +375,8 @@ def __init__( model: str | None = None, api_key: str | SecretString | None = None, server_url: str | None = None, - client: httpx.AsyncClient | None = None, + http_client: httpx.AsyncClient | None = None, + client: Any | None = None, otel_provider_name: str | None = None, additional_properties: dict[str, Any] | None = None, env_file_path: str | None = None, @@ -305,6 +387,7 @@ def __init__( model=model, api_key=api_key, server_url=server_url, + http_client=http_client, client=client, additional_properties=additional_properties, otel_provider_name=otel_provider_name, diff --git a/python/packages/mistral/tests/mistral/test_mistral_embedding_client.py b/python/packages/mistral/tests/mistral/test_mistral_embedding_client.py index 6aea0924574..1eee61cfe4d 100644 --- a/python/packages/mistral/tests/mistral/test_mistral_embedding_client.py +++ b/python/packages/mistral/tests/mistral/test_mistral_embedding_client.py @@ -3,6 +3,7 @@ import json import os from collections.abc import Sequence +from types import SimpleNamespace from typing import Any import httpx @@ -53,7 +54,7 @@ def make_client(*responses: httpx.Response) -> tuple[MistralEmbeddingClient, Moc base_url="https://api.mistral.ai", transport=httpx.MockTransport(server.handler), ) - client = MistralEmbeddingClient(model="mistral-embed", client=http_client) + client = MistralEmbeddingClient(model="mistral-embed", http_client=http_client) return client, server @@ -84,13 +85,64 @@ def test_mistral_embedding_construction_with_server_url() -> None: assert str(client.client.base_url) == "https://custom.mistral.ai" -def test_mistral_embedding_construction_with_client() -> None: +def test_mistral_embedding_construction_with_http_client() -> None: """Test construction with a pre-configured client.""" http_client = httpx.AsyncClient(base_url="https://api.mistral.ai") - client = MistralEmbeddingClient(model="mistral-embed", client=http_client) + client = MistralEmbeddingClient(model="mistral-embed", http_client=http_client) assert client.client is http_client +def test_mistral_embedding_deprecated_client_param_accepts_httpx() -> None: + http_client = httpx.AsyncClient(base_url="https://api.mistral.ai") + with pytest.deprecated_call(): + client = MistralEmbeddingClient(model="mistral-embed", client=http_client) + assert client.client is http_client + + +class FakeMistralSDK: + """Duck-typed stand-in for a mistralai.Mistral client.""" + + def __init__(self, vectors: Sequence[Sequence[float]] = ((0.1, 0.2),)) -> None: + self.requests: list[dict[str, Any]] = [] + self._vectors = vectors + self.embeddings = SimpleNamespace(create_async=self._create_async) + + async def _create_async(self, **kwargs: Any) -> Any: + self.requests.append(kwargs) + return SimpleNamespace( + model="mistral-embed", + data=[SimpleNamespace(index=i, embedding=list(v)) for i, v in enumerate(self._vectors)], + usage=SimpleNamespace(prompt_tokens=3, total_tokens=3), + ) + + +async def test_mistral_embedding_deprecated_client_param_accepts_sdk_client() -> None: + """An injected mistralai.Mistral keeps working through the legacy SDK path.""" + sdk = FakeMistralSDK() + with pytest.deprecated_call(): + client = MistralEmbeddingClient(model="mistral-embed", client=sdk) + + result = await client.get_embeddings(["hello"], options=MistralEmbeddingOptions(dimensions=2)) + + assert [e.vector for e in result] == [[0.1, 0.2]] + assert result.usage == {"input_token_count": 3, "total_token_count": 3} + assert sdk.requests == [{"model": "mistral-embed", "inputs": ["hello"], "output_dimension": 2}] + + +def test_mistral_embedding_deprecated_client_param_rejects_unknown_client() -> None: + class NotAClient: + pass + + with pytest.deprecated_call(), pytest.raises(TypeError, match="httpx.AsyncClient"): + MistralEmbeddingClient(model="mistral-embed", client=NotAClient()) + + +def test_mistral_embedding_client_and_http_client_conflict() -> None: + http_client = httpx.AsyncClient(base_url="https://api.mistral.ai") + with pytest.deprecated_call(), pytest.raises(ValueError, match="not both"): + MistralEmbeddingClient(model="mistral-embed", http_client=http_client, client=http_client) + + def test_mistral_embedding_construction_missing_model_raises(monkeypatch: pytest.MonkeyPatch) -> None: """Test that missing model raises an error.""" monkeypatch.delenv("MISTRAL_EMBEDDING_MODEL", raising=False) @@ -133,7 +185,7 @@ async def test_mistral_embedding_close_only_closes_owned_client() -> None: assert owned.client.is_closed http_client = httpx.AsyncClient(base_url="https://custom.mistral.ai") - injected = MistralEmbeddingClient(model="mistral-embed", client=http_client) + injected = MistralEmbeddingClient(model="mistral-embed", http_client=http_client) assert injected.service_url() == "https://custom.mistral.ai" await injected.close() @@ -261,7 +313,7 @@ def raise_connect_error(request: httpx.Request) -> httpx.Response: base_url="https://api.mistral.ai", transport=httpx.MockTransport(raise_connect_error), ) - client = MistralEmbeddingClient(model="mistral-embed", client=http_client) + client = MistralEmbeddingClient(model="mistral-embed", http_client=http_client) with pytest.raises(IntegrationException, match="Mistral embeddings request failed"): await client.get_embeddings(["hello"]) From f467931f6d8d70a6d99543ab8797de969b2ca881 Mon Sep 17 00:00:00 2001 From: NekoPunch Date: Sun, 2 Aug 2026 23:18:26 -0700 Subject: [PATCH 8/8] chore(python): tidy Mistral sample header --- .../providers/mistral/mistral_agent_basic.py | 16 ++++++---------- 1 file changed, 6 insertions(+), 10 deletions(-) diff --git a/python/samples/02-agents/providers/mistral/mistral_agent_basic.py b/python/samples/02-agents/providers/mistral/mistral_agent_basic.py index b0ecc63b6cb..ff66a285ba6 100644 --- a/python/samples/02-agents/providers/mistral/mistral_agent_basic.py +++ b/python/samples/02-agents/providers/mistral/mistral_agent_basic.py @@ -1,14 +1,4 @@ # Copyright (c) Microsoft. All rights reserved. -# /// script -# requires-python = ">=3.10" -# dependencies = ["agent-framework-mistral"] -# /// - -"""Demonstrates a Mistral AI agent with basic tool usage. - -Requires ``MISTRAL_API_KEY`` and ``MISTRAL_CHAT_MODEL`` environment variables -(e.g. MISTRAL_CHAT_MODEL=mistral-small-latest). -""" import asyncio from datetime import datetime @@ -21,6 +11,12 @@ # Load environment variables from the local .env file. load_dotenv() +"""Demonstrates a Mistral AI agent with basic tool usage. + +Requires ``MISTRAL_API_KEY`` and ``MISTRAL_CHAT_MODEL`` environment variables +(e.g. MISTRAL_CHAT_MODEL=mistral-small-latest). +""" + # NOTE: approval_mode="never_require" is for sample brevity. Use "always_require" in production; # see samples/02-agents/tools/function_tool_with_approval.py