diff --git a/.github/workflows/python-integration-tests.yml b/.github/workflows/python-integration-tests.yml index 3308f294dc2..2104a15b0c0 100644 --- a/.github/workflows/python-integration-tests.yml +++ b/.github/workflows/python-integration-tests.yml @@ -232,11 +232,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 7db1c4a641d..bd984d8bd7f 100644 --- a/.github/workflows/python-merge-tests.yml +++ b/.github/workflows/python-merge-tests.yml @@ -68,6 +68,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' @@ -339,11 +340,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..8c48d1b5ea8 100644 --- a/python/packages/mistral/README.md +++ b/python/packages/mistral/README.md @@ -8,7 +8,39 @@ 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 +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.") + response = await agent.run("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 +54,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..c3dba5540d4 --- /dev/null +++ b/python/packages/mistral/agent_framework_mistral/_chat_client.py @@ -0,0 +1,861 @@ +# 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, mark_feature_used +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 + +from ._feature_usage import FeatureIndex + +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.""" + + 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 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 + 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) + async for line in response.aiter_lines(): + chunk = self._parse_sse_line(line) + if chunk is not None: + yield self._parse_chunk(chunk) + except ChatClientException: + raise + except Exception as ex: + raise ChatClientException( + f"Mistral streaming chat request failed: {ex}", + inner_exception=ex, + ) from ex + + 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) + mark_feature_used(FeatureIndex.MISTRAL) + 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 {}) + 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", + "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.""" + choices = cast("Sequence[Mapping[str, Any]]", response.get("choices") or ()) + 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.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")), + finish_reason=finish_reason, + response_format=response_format, + raw_representation=response, + ) + + def _parse_chunk(self, chunk: Mapping[str, Any]) -> 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. + """ + 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_message_contents(delta)) + if (reason := choice.get("finish_reason")) and 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( + 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 _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, + ) + ) + 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 _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..63f462867f5 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,17 @@ load_settings, ) from agent_framework._settings import SecretString -from agent_framework._telemetry import mark_feature_used +from agent_framework._telemetry import get_user_agent, mark_feature_used +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 +36,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 +86,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 +109,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 +128,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 +176,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 +189,61 @@ 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"] + request["output_dimension"] = opts["dimensions"] mark_feature_used(FeatureIndex.MISTRAL) - response = await self.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) + 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 +259,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 +283,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 +294,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..50492180425 --- /dev/null +++ b/python/packages/mistral/tests/mistral/test_mistral_chat_client.py @@ -0,0 +1,910 @@ +# 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_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"))) + + 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"], + "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 "n" not in request + 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"} + + +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] + 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_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_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..6aea0924574 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,173 @@ 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_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.""" - 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 +296,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..b0ecc63b6cb --- /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"] +# /// + +"""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 + +from agent_framework import Agent, tool +from agent_framework.mistral import MistralChatClient +from dotenv import load_dotenv + +# 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 19324c85e97..c2dff69219c 100644 --- a/python/samples/README.md +++ b/python/samples/README.md @@ -147,6 +147,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 12fba89b2a2..2d2cc349dd7 100644 --- a/python/uv.lock +++ b/python/uv.lock @@ -905,13 +905,13 @@ version = "1.0.0b260730" source = { editable = "packages/mistral" } dependencies = [ { name = "agent-framework-core", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "mistralai", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "httpx", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, ] [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]] @@ -2379,15 +2379,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" @@ -3327,15 +3318,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" @@ -4099,27 +4081,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", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "httpx", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "invoke", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "opentelemetry-api", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "opentelemetry-exporter-otlp-proto-http", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "opentelemetry-sdk", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "pydantic", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "python-dateutil", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "pyyaml", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "typing-inspection", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, -] -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"