diff --git a/pyproject.toml b/pyproject.toml index 985350fb5..ee23052fb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -39,12 +39,12 @@ mistral = [ anthropic = [] gemini = [ - "google-genai>=1.51.0", + "google-genai>=1.70.0", "google-cloud-storage", ] vertexai = [ - "google-genai>=1.51.0", + "google-genai>=1.70.0", "google-cloud-storage", ] diff --git a/src/any_llm/any_llm.py b/src/any_llm/any_llm.py index 1d34f145c..5598a9839 100644 --- a/src/any_llm/any_llm.py +++ b/src/any_llm/any_llm.py @@ -54,6 +54,7 @@ from any_llm.utils.structured_output import ( build_parsed_message, is_structured_output_type, + normalize_output_config, parse_json_content, parse_responses_output, ) @@ -132,6 +133,9 @@ class AnyLLM(ABC): SUPPORTS_MESSAGES: bool = True """Anthropic Messages API (all providers support it via conversion)""" + SUPPORTS_MESSAGES_STRUCTURED_OUTPUT_STREAMING: bool = False + """Whether Messages structured output can be streamed by this provider.""" + PROMPT_CACHE_KEY_SUPPORT: Literal["unsupported", "supported", "passthrough"] = "unsupported" """Whether prompt_cache_key is supported, forwarded to a router, or rejected.""" @@ -960,7 +964,8 @@ async def amessages( ``output_config``. Either a Pydantic ``BaseModel``/dataclass **type** (typed ``parsed_output``) or a raw Anthropic ``output_config`` **dict** for non-Pydantic JSON schemas (``parsed_output`` holds the parsed JSON). The call returns - Anthropic's ``ParsedMessage``. Not supported with ``stream=True``. + Anthropic's ``ParsedMessage`` for non-streaming requests. Providers with native + support can stream schema-constrained Messages events instead. timeout: Per-request timeout in seconds, passed through to the provider's client/SDK. An explicit ``None`` is treated the same as omitting it (the provider's default applies), so it cannot request an unbounded timeout. Providers that have no @@ -973,12 +978,13 @@ async def amessages( iterator of MessageStreamEvent (if streaming). Raises: - ValueError: If `output_format` is combined with `stream=True`. + ValueError: If `output_format` is combined with `stream=True` for a provider that + does not support streaming structured output. NotImplementedError: If `container`, `context_management`, or `betas` is used with a provider that has no native Anthropic Messages API. """ - if output_format is not None and stream: + if output_format is not None and stream and not self.SUPPORTS_MESSAGES_STRUCTURED_OUTPUT_STREAMING: msg = "stream is not supported for output_format" raise ValueError(msg) @@ -1014,6 +1020,11 @@ async def amessages( # case); for the raw-dict case and for all bridged providers it returns a MessageResponse, # so build the same ParsedMessage shape from the response's JSON text here. if output_format is not None and isinstance(result, MessageResponse): + if isinstance(output_format, dict): + format_config = normalize_output_config(output_format).get("format") + schema = format_config.get("schema") if isinstance(format_config, dict) else None + if not isinstance(schema, dict) or not schema: + return result return build_parsed_message(result, output_format) return result diff --git a/src/any_llm/api.py b/src/any_llm/api.py index 86996881c..96f2a2ca6 100644 --- a/src/any_llm/api.py +++ b/src/any_llm/api.py @@ -636,8 +636,8 @@ def messages( output_format: Structured output, mirroring Anthropic's ``messages.parse``/``output_config``. Either a Pydantic ``BaseModel``/dataclass **type** (typed ``parsed_output``) or a raw Anthropic ``output_config`` **dict** for non-Pydantic JSON schemas (``parsed_output`` - holds the parsed JSON). The call returns Anthropic's ``ParsedMessage``. Not supported - with streaming. + holds the parsed JSON). Non-streaming calls return Anthropic's ``ParsedMessage``; + providers with native support can stream schema-constrained Messages events. timeout: Per-request timeout in seconds, passed through to the provider's client/SDK. An explicit ``None`` is treated the same as omitting it (the provider's default applies), so it cannot request an unbounded timeout. Providers that have no @@ -744,8 +744,8 @@ async def amessages( output_format: Structured output, mirroring Anthropic's ``messages.parse``/``output_config``. Either a Pydantic ``BaseModel``/dataclass **type** (typed ``parsed_output``) or a raw Anthropic ``output_config`` **dict** for non-Pydantic JSON schemas (``parsed_output`` - holds the parsed JSON). The call returns Anthropic's ``ParsedMessage``. Not supported - with streaming. + holds the parsed JSON). Non-streaming calls return Anthropic's ``ParsedMessage``; + providers with native support can stream schema-constrained Messages events. timeout: Per-request timeout in seconds, passed through to the provider's client/SDK. An explicit ``None`` is treated the same as omitting it (the provider's default applies), so it cannot request an unbounded timeout. Providers that have no diff --git a/src/any_llm/providers/anthropic/base.py b/src/any_llm/providers/anthropic/base.py index 32a00b454..5d9d2f10b 100644 --- a/src/any_llm/providers/anthropic/base.py +++ b/src/any_llm/providers/anthropic/base.py @@ -236,6 +236,7 @@ class BaseAnthropicProvider(AnyLLM, ABC): SUPPORTS_LIST_MODELS = False SUPPORTS_BATCH = True SUPPORTS_RERANK = False + SUPPORTS_MESSAGES_STRUCTURED_OUTPUT_STREAMING = True # The Anthropic SDK accepts a per-request `timeout` on messages.create, so it forwards unchanged. TIMEOUT_SUPPORT = "native" @@ -322,7 +323,8 @@ async def _amessages( (which drives the GA ``output_config`` primitive) and returns the SDK's ``ParsedMessage`` unchanged. When it is a raw ``output_config`` dict, passes it straight to native ``messages.create(output_config=...)`` and returns a ``MessageResponse`` (the base layer - then builds the matching ``ParsedMessage`` from its JSON text). + then builds the matching ``ParsedMessage`` from its JSON text). Streaming requests use + ``messages.stream`` with the matching typed or raw output configuration. """ header_betas = _pop_anthropic_beta_header(kwargs) betas = _messages_betas(params, header_betas) @@ -335,6 +337,14 @@ async def _amessages( if betas: native_kwargs["betas"] = betas native_kwargs.update(kwargs) + if params.stream: + if is_structured_output_type(params.output_format): + native_kwargs["output_format"] = params.output_format + else: + native_kwargs["output_config"] = normalize_output_config( + cast("dict[str, Any]", params.output_format) + ) + return self._stream_messages_async(use_beta=use_beta, **native_kwargs) if is_structured_output_type(params.output_format): with _translating_nonstreaming_guard(self, params.max_tokens): parsed = await messages_resource.parse(output_format=params.output_format, **native_kwargs) diff --git a/src/any_llm/providers/anthropic/utils.py b/src/any_llm/providers/anthropic/utils.py index d7cee8f9f..cfdcdbec0 100644 --- a/src/any_llm/providers/anthropic/utils.py +++ b/src/any_llm/providers/anthropic/utils.py @@ -307,7 +307,6 @@ def _create_openai_chunk_from_anthropic_chunk(chunk: Any, model_id: str) -> Chat delta["extra_content"] = {"anthropic": {"stop_details": stop_details}} elif isinstance(chunk, MessageStopEvent): - finish_reason = None if hasattr(chunk, "message") and chunk.message.usage: anthropic_usage = chunk.message.usage cache_read = anthropic_usage.cache_read_input_tokens or 0 @@ -319,6 +318,9 @@ def _create_openai_chunk_from_anthropic_chunk(chunk: Any, model_id: str) -> Chat "total_tokens": total_prompt_tokens + anthropic_usage.output_tokens, "prompt_tokens_details": PromptTokensDetails(cached_tokens=cache_read) if cache_read else None, } + # The stop event carries no delta or finish_reason, only usage. Leave choices + # empty so it matches the trailing usage-only chunk OpenAI-compatible providers emit. + return ChatCompletionChunk.model_validate(chunk_dict) choice = { "index": 0, diff --git a/src/any_llm/providers/bedrock/utils.py b/src/any_llm/providers/bedrock/utils.py index f57935c50..8bf7b0729 100644 --- a/src/any_llm/providers/bedrock/utils.py +++ b/src/any_llm/providers/bedrock/utils.py @@ -36,6 +36,23 @@ # Titan, Llama, ...) have no equivalent verified mechanism, so they keep raising. _STRUCTURED_OUTPUT_TOOL_NAME = "any_llm_structured_output" +_FinishReason = Literal["stop", "length", "tool_calls", "content_filter", "function_call"] + +# The Converse API reports nine stop reasons (botocore's bedrock-runtime service model, shape +# "StopReason"). OpenAI has no counterpart for "stop_sequence", "malformed_model_output" and +# "malformed_tool_use", so those fall through to the "stop" default along with any reason a +# future service model adds. The rest do have one, and without it a guardrail block or a +# context overflow looks like a normal completion to callers, including the structured-output +# guard in any_llm.py that is supposed to raise ContentFilterFinishReasonError. +BEDROCK_STOP_REASON_TO_FINISH_REASON: dict[str, _FinishReason] = { + "end_turn": "stop", + "max_tokens": "length", + "model_context_window_exceeded": "length", + "tool_use": "tool_calls", + "content_filtered": "content_filter", + "guardrail_intervened": "content_filter", +} + REASONING_EFFORT_TO_THINKING_BUDGETS = { "minimal": 1024, "low": 2048, @@ -46,6 +63,13 @@ } +def _map_stop_reason(stop_reason: Any) -> _FinishReason: + """Map a Converse API stopReason onto the OpenAI finish_reason vocabulary.""" + if not isinstance(stop_reason, str): + return "stop" + return BEDROCK_STOP_REASON_TO_FINISH_REASON.get(stop_reason, "stop") + + def _is_anthropic_model(model_id: str) -> bool: """Return True if the Bedrock model id refers to an Anthropic Claude model. @@ -521,8 +545,7 @@ def _convert_response(response: dict[str, Any]) -> ChatCompletion: ) content = "".join(content_parts) - stop_reason = response.get("stopReason") - finish_reason: Literal["stop", "length"] = "length" if stop_reason == "max_tokens" else "stop" + finish_reason = _map_stop_reason(response.get("stopReason")) message = ChatCompletionMessage( role="assistant", @@ -534,9 +557,7 @@ def _convert_response(response: dict[str, Any]) -> ChatCompletion: choices_out.append( Choice( index=0, - finish_reason=cast( - "Literal['stop', 'length', 'tool_calls', 'content_filter', 'function_call']", finish_reason - ), + finish_reason=finish_reason, message=message, ) ) @@ -570,7 +591,7 @@ def _create_openai_chunk_from_aws_chunk( content: str | None = None reasoning_content: str | None = None - finish_reason: Literal["stop", "length", "tool_calls"] | None = None + finish_reason: _FinishReason | None = None tool_call: ChoiceDeltaToolCall | None = None usage: CompletionUsage | None = None @@ -616,13 +637,7 @@ def _create_openai_chunk_from_aws_chunk( ), ) elif "messageStop" in chunk: - stop_reason = chunk["messageStop"]["stopReason"] - if stop_reason == "max_tokens": - finish_reason = "length" - elif stop_reason == "tool_use": - finish_reason = "tool_calls" - else: - finish_reason = "stop" + finish_reason = _map_stop_reason(chunk["messageStop"]["stopReason"]) elif "messageStart" in chunk: content = "" elif "metadata" in chunk: diff --git a/src/any_llm/providers/deepseek/deepseek.py b/src/any_llm/providers/deepseek/deepseek.py index c61e77e4f..ac0e8bf8a 100644 --- a/src/any_llm/providers/deepseek/deepseek.py +++ b/src/any_llm/providers/deepseek/deepseek.py @@ -1,8 +1,10 @@ +import re from collections.abc import AsyncIterator from typing import Any from typing_extensions import override +from any_llm.exceptions import InvalidRequestError from any_llm.providers.deepseek.utils import ( _inject_cached_tokens, _inject_cached_tokens_chunk, @@ -10,13 +12,38 @@ _preprocess_messages, ) from any_llm.providers.openai.base import BaseOpenAIProvider -from any_llm.types.completion import ChatCompletion, ChatCompletionChunk, CompletionParams +from any_llm.types.completion import ChatCompletion, ChatCompletionChunk, CompletionParams, ReasoningEffort -# The two legacy API model names being discontinued 2026-07-24 in favor of deepseek-v4-flash / -# deepseek-v4-pro. See https://api-docs.deepseek.com/updates/#date-2026-04-24. They hard-code -# their own thinking behavior (non-thinking / thinking respectively) and don't need, and may not -# accept, the `thinking` request toggle added below for the new model family. -_LEGACY_MODEL_IDS = frozenset({"deepseek-chat", "deepseek-reasoner"}) +# Each entry maps an any-llm effort to DeepSeek's top-level effort and thinking toggle. +# DeepSeek Chat accepts low, high, and max, maps medium and xhigh to high, and does not +# accept OpenAI's minimal value. +# https://api-docs.deepseek.com/guides/thinking_mode/ +_REASONING_CONTROLS: dict[ReasoningEffort | None, tuple[str | None, str | None]] = { + None: (None, None), + "auto": (None, None), + "none": (None, "disabled"), + "low": ("low", "enabled"), + "medium": ("high", "enabled"), + "high": ("high", "enabled"), + "xhigh": ("high", "enabled"), + "max": ("max", "enabled"), +} + +# These normalized fields are absent from the current DeepSeek Chat schema, except for the two +# penalty fields, which the API marks deprecated and ineffective. They remain accepted by +# any_llm's shared interface for compatibility but are not sent to DeepSeek. +# https://api-docs.deepseek.com/api/create-chat-completion +_UNSUPPORTED_DEEPSEEK_FIELDS = frozenset( + { + "frequency_penalty", + "logit_bias", + "n", + "parallel_tool_calls", + "presence_penalty", + "seed", + "service_tier", + } +) class DeepseekProvider(BaseOpenAIProvider): @@ -36,25 +63,43 @@ class DeepseekProvider(BaseOpenAIProvider): def _convert_completion_params(params: CompletionParams, **kwargs: Any) -> dict[str, Any]: """DeepSeek only accepts ``max_tokens``, not ``max_completion_tokens``. - Also maps ``reasoning_effort`` to DeepSeek's ``thinking`` toggle for the V4 model - family. DeepSeek's V4 models default to thinking mode ENABLED when the toggle is - omitted from the request (see https://api-docs.deepseek.com/guides/thinking_mode), so - any_llm explicitly defaults it to disabled here -- matching the legacy ``deepseek-chat`` - behavior -- unless the caller opts in via ``reasoning_effort``. A caller-supplied - ``extra_body`` override is respected and not clobbered. + DeepSeek's V4 models default to enabled thinking with high effort, so ``None`` and the + normalized ``auto`` sentinel leave both controls absent. An explicit ``none`` uses the + provider's thinking toggle. Caller-supplied ``extra_body`` values take precedence. """ converted_params = BaseOpenAIProvider._convert_completion_params(params, **kwargs) if "max_completion_tokens" in converted_params: converted_params["max_tokens"] = converted_params.pop("max_completion_tokens") - if params.model_id not in _LEGACY_MODEL_IDS: - # ``"auto"`` means "no explicit reasoning requested" (BaseOpenAIProvider._acompletion - # normalizes it to this provider's default before we run), so it is treated the same - # as ``None``/``"none"`` here -- matching every other provider's converter and keeping - # this self-contained even if called directly with ``"auto"``. - thinking_disabled = params.reasoning_effort in (None, "none", "auto") - extra_body = converted_params.setdefault("extra_body", {}) - extra_body.setdefault("thinking", {"type": "disabled" if thinking_disabled else "enabled"}) + user_id = converted_params.pop("user", None) + for field in _UNSUPPORTED_DEEPSEEK_FIELDS: + converted_params.pop(field, None) + + converted_params.pop("reasoning_effort", None) + controls = _REASONING_CONTROLS.get(params.reasoning_effort) + if controls is None: + msg = f"reasoning_effort {params.reasoning_effort!r} is not supported by DeepSeek Chat" + raise InvalidRequestError(msg, provider_name=DeepseekProvider.PROVIDER_NAME) + reasoning_effort, thinking_type = controls + if reasoning_effort is not None: + converted_params["reasoning_effort"] = reasoning_effort + thinking = {"type": thinking_type} if thinking_type is not None else None + + if user_id is not None or thinking is not None: + extra_body = dict(converted_params.get("extra_body") or {}) + converted_params["extra_body"] = extra_body + if user_id is not None and "user_id" not in extra_body: + # DeepSeek's user_id contract is stricter than any-llm's shared user field. + # https://api-docs.deepseek.com/quick_start/rate_limit/#setting-user_id + if re.fullmatch(r"[a-zA-Z0-9_-]{1,512}", user_id) is None: + msg = ( + "DeepSeek user_id must contain only ASCII letters, digits, underscores, or hyphens " + "and be between 1 and 512 characters" + ) + raise InvalidRequestError(msg, provider_name=DeepseekProvider.PROVIDER_NAME) + extra_body["user_id"] = user_id + if thinking is not None: + extra_body.setdefault("thinking", thinking) return converted_params @staticmethod diff --git a/src/any_llm/providers/deepseek/utils.py b/src/any_llm/providers/deepseek/utils.py index 05177f913..3ca26b5d9 100644 --- a/src/any_llm/providers/deepseek/utils.py +++ b/src/any_llm/providers/deepseek/utils.py @@ -44,15 +44,13 @@ def _convert_structured_type_to_deepseek_json( return modified_messages -def _reinject_reasoning_content(messages: list[dict[str, Any]]) -> list[dict[str, Any]]: - """Restore ``reasoning_content`` on replayed assistant tool-call turns. +def _reinject_reasoning_content(messages: list[dict[str, Any]], *, replay_reasoning: bool) -> list[dict[str, Any]]: + """Restore ``reasoning_content`` when the current request carries tools. - DeepSeek's thinking mode requires ``reasoning_content`` to be passed back verbatim on any - assistant turn that performed a tool call, or the API returns a 400. any_llm's shared - message serialization (``AnyLLM.acompletion``) strips the normalized ``reasoning`` field - before replaying a ``ChatCompletionMessage`` back as a request message, so we restore it - here from the ``extra_content["deepseek"]`` side-channel populated by - ``_inject_reasoning_extra_content`` when the response was first received. + DeepSeek requires all previous assistant reasoning to be replayed when ``tools`` is present, + including turns that did not call a tool. any_llm's shared message serialization strips the + normalized ``reasoning`` field, so this restores it from the provider side-channel populated + by ``_inject_reasoning_extra_content``. Reference: https://api-docs.deepseek.com/guides/thinking_mode#tool-calls @@ -64,7 +62,7 @@ def _reinject_reasoning_content(messages: list[dict[str, Any]]) -> list[dict[str for message in messages: extra_content = message.get("extra_content") cleaned = {k: v for k, v in message.items() if k != "extra_content"} if extra_content is not None else message - if message.get("role") == "assistant" and message.get("tool_calls") is not None: + if replay_reasoning and message.get("role") == "assistant": deepseek_extra = extra_content.get("deepseek") if isinstance(extra_content, dict) else None if isinstance(deepseek_extra, dict) and isinstance(deepseek_extra.get("reasoning_content"), str): result.append({**cleaned, "reasoning_content": deepseek_extra["reasoning_content"]}) @@ -81,7 +79,7 @@ def _preprocess_messages(params: CompletionParams) -> CompletionParams: params.response_format = {"type": "json_object"} params.messages = modified_messages - params.messages = _reinject_reasoning_content(params.messages) + params.messages = _reinject_reasoning_content(params.messages, replay_reasoning=params.tools is not None) return params diff --git a/src/any_llm/providers/gemini/base.py b/src/any_llm/providers/gemini/base.py index 7d49f0892..e23c06829 100644 --- a/src/any_llm/providers/gemini/base.py +++ b/src/any_llm/providers/gemini/base.py @@ -20,6 +20,7 @@ CreateEmbeddingResponse, Function, Reasoning, + ReasoningEffort, ) from any_llm.utils.structured_output import get_json_schema, is_structured_output_type @@ -52,7 +53,7 @@ from any_llm.types.model import Model REASONING_EFFORT_TO_THINKING_BUDGETS = { - "minimal": 256, + "minimal": 1024, "low": 1024, "medium": 8192, "high": 24576, @@ -68,17 +69,118 @@ "max": types.ThinkingLevel.HIGH, } _SUPPORTED_BATCH_ENDPOINTS = frozenset({"/v1/chat/completions"}) -_THINKING_LEVEL_MIN_GEMINI_VERSION = (3, 5) +_ALL_THINKING_LEVELS = frozenset(REASONING_EFFORT_TO_THINKING_LEVELS.values()) +# Known model capabilities differ within the Gemini 3 family, so these exceptions refine rather than replace +# the permissive version routing used for custom, dated, and newly released model IDs. +# Source: https://ai.google.dev/gemini-api/docs/generate-content/thinking#thinking-levels +_THINKING_LEVELS_BY_MODEL = { + "gemini-3.8-flash": frozenset({types.ThinkingLevel.LOW, types.ThinkingLevel.MEDIUM, types.ThinkingLevel.HIGH}), + "gemini-3.7-flash": frozenset({types.ThinkingLevel.LOW, types.ThinkingLevel.MEDIUM, types.ThinkingLevel.HIGH}), + "gemini-3.6-flash": _ALL_THINKING_LEVELS, + "gemini-3.5-flash": _ALL_THINKING_LEVELS, + "gemini-3.5-flash-lite": _ALL_THINKING_LEVELS, + "gemini-3.1-flash-lite": _ALL_THINKING_LEVELS, + "gemini-3.1-pro-preview": frozenset( + {types.ThinkingLevel.LOW, types.ThinkingLevel.MEDIUM, types.ThinkingLevel.HIGH} + ), + "gemini-3.1-flash-image": frozenset({types.ThinkingLevel.MINIMAL, types.ThinkingLevel.HIGH}), + "gemini-3.1-flash-lite-image": frozenset({types.ThinkingLevel.MINIMAL, types.ThinkingLevel.HIGH}), + "gemini-3-flash-preview": _ALL_THINKING_LEVELS, +} +_MAX_THINKING_BUDGET_BY_MODEL = { + "gemini-2.5-pro": 32768, + "gemini-2.5-flash": 24576, + "gemini-2.5-flash-lite": 24576, +} _GEMINI_VERSION_PATTERN = re.compile(r"(?:^|/)gemini-(\d+)(?:\.(\d+))?") def _uses_thinking_level(model_id: str) -> bool: - """Gemini 3.5 and newer reject `thinking_budget` and expect `thinking_level` instead.""" + """Route Gemini 3 and newer to `thinking_level`, including unlisted model IDs.""" match = _GEMINI_VERSION_PATTERN.search(model_id.lower()) if match is None: return False - major, minor = int(match.group(1)), int(match.group(2) or 0) - return (major, minor) >= _THINKING_LEVEL_MIN_GEMINI_VERSION + return int(match.group(1)) >= 3 + + +def _matches_known_model(model_name: str, known_model: str) -> bool: + version_suffix = model_name.removeprefix(known_model) + return model_name == known_model or ( + model_name.startswith(known_model) and re.fullmatch(r"(?:-\d+)+", version_suffix) is not None + ) + + +def _known_thinking_levels(model_name: str) -> frozenset[types.ThinkingLevel] | None: + for known_model, supported_levels in _THINKING_LEVELS_BY_MODEL.items(): + if _matches_known_model(model_name, known_model): + return supported_levels + return None + + +def _known_max_thinking_budget(model_name: str) -> int | None: + for known_model, max_budget in _MAX_THINKING_BUDGET_BY_MODEL.items(): + if _matches_known_model(model_name, known_model): + return max_budget + return None + + +def _convert_reasoning_effort( + model_id: str, + reasoning_effort: ReasoningEffort | None, + provider_name: str, +) -> types.ThinkingConfig | None: + if reasoning_effort is None or reasoning_effort == "auto": + return None + + parameter_name = "reasoning_effort" + model_name = model_id.rsplit("/", maxsplit=1)[-1].lower() + supported_levels = _known_thinking_levels(model_name) + additional_message = f"'{reasoning_effort}' is not available for model '{model_id}'." + if reasoning_effort == "none": + if _uses_thinking_level(model_id) or _matches_known_model(model_name, "gemini-2.5-pro"): + raise UnsupportedParameterError(parameter_name, provider_name, additional_message) + return types.ThinkingConfig(thinking_budget=0) + + if supported_levels is not None or _uses_thinking_level(model_id): + thinking_level = REASONING_EFFORT_TO_THINKING_LEVELS.get(reasoning_effort) + # Google's OpenAI compatibility contract maps `minimal` to `low` for Gemini 3.1 Pro. + # Source: https://ai.google.dev/gemini-api/docs/openai#thinking + if _matches_known_model(model_name, "gemini-3.1-pro-preview") and reasoning_effort == "minimal": + thinking_level = types.ThinkingLevel.LOW + if thinking_level is None or (supported_levels is not None and thinking_level not in supported_levels): + raise UnsupportedParameterError(parameter_name, provider_name, additional_message) + return types.ThinkingConfig(include_thoughts=True, thinking_level=thinking_level) + + thinking_budget = REASONING_EFFORT_TO_THINKING_BUDGETS.get(reasoning_effort) + if thinking_budget is None: + raise UnsupportedParameterError(parameter_name, provider_name, additional_message) + max_budget = _known_max_thinking_budget(model_name) + if max_budget is not None: + thinking_budget = min(thinking_budget, max_budget) + return types.ThinkingConfig(include_thoughts=True, thinking_budget=thinking_budget) + + +def _convert_response_format(response_format: dict[str, Any] | type | None) -> dict[str, Any]: + if is_structured_output_type(response_format): + schema = get_json_schema(response_format) + schema_key = "response_json_schema" if _has_additional_properties(schema) else "response_schema" + return {"response_mime_type": "application/json", schema_key: schema} + if not isinstance(response_format, dict): + return {} + + response_type = response_format.get("type") + if response_type == "json_schema": + return { + "response_mime_type": "application/json", + "response_json_schema": response_format["json_schema"]["schema"], + } + if response_type == "json_object": + return {"response_mime_type": "application/json"} + if response_type in (None, "text"): + return {} + + msg = f"Unsupported response_format type: {response_type}" + raise ValueError(msg) class GoogleProvider(AnyLLM): @@ -140,61 +242,36 @@ def _convert_completion_params(params: CompletionParams, **kwargs: Any) -> dict[ error_message = "parallel_tool_calls" raise UnsupportedParameterError(error_message, provider_name) - if params.frequency_penalty is not None: - kwargs["frequency_penalty"] = params.frequency_penalty - if params.max_tokens is not None: - kwargs["max_output_tokens"] = params.max_tokens - if params.presence_penalty is not None: - kwargs["presence_penalty"] = params.presence_penalty - if params.reasoning_effort != "auto": - if params.reasoning_effort is None or params.reasoning_effort == "none": - kwargs["thinking_config"] = types.ThinkingConfig(include_thoughts=False) - elif _uses_thinking_level(params.model_id): - kwargs["thinking_config"] = types.ThinkingConfig( - include_thoughts=True, thinking_level=REASONING_EFFORT_TO_THINKING_LEVELS[params.reasoning_effort] - ) - else: - kwargs["thinking_config"] = types.ThinkingConfig( - include_thoughts=True, thinking_budget=REASONING_EFFORT_TO_THINKING_BUDGETS[params.reasoning_effort] + kwargs.update( + { + option_name: option_value + for option_name, option_value in ( + ("frequency_penalty", params.frequency_penalty), + ("max_output_tokens", params.max_tokens), + ("presence_penalty", params.presence_penalty), + ("seed", params.seed), + ("service_tier", params.service_tier), + ("temperature", params.temperature), + ("top_p", params.top_p), ) - if params.seed is not None: - kwargs["seed"] = params.seed - if params.service_tier is not None: - kwargs["service_tier"] = params.service_tier - if params.temperature is not None: - kwargs["temperature"] = params.temperature + if option_value is not None + } + ) + + thinking_config = _convert_reasoning_effort(params.model_id, params.reasoning_effort, provider_name) + if thinking_config is not None: + kwargs["thinking_config"] = thinking_config if params.tools is not None: kwargs["tools"] = _convert_tool_spec(params.tools, provider_name) if params.tool_choice is not None: kwargs["tool_config"] = _convert_tool_choice(params.tool_choice, provider_name) - if params.top_p is not None: - kwargs["top_p"] = params.top_p if params.stop is not None: if isinstance(params.stop, str): kwargs["stop_sequences"] = [params.stop] else: kwargs["stop_sequences"] = params.stop - response_format = params.response_format - if is_structured_output_type(response_format): - kwargs["response_mime_type"] = "application/json" - schema = get_json_schema(response_format) - if _has_additional_properties(schema): - kwargs["response_json_schema"] = schema - else: - kwargs["response_schema"] = schema - elif isinstance(response_format, dict): - response_type = response_format.get("type") - if response_type == "json_schema": - kwargs["response_mime_type"] = "application/json" - kwargs["response_json_schema"] = response_format["json_schema"]["schema"] - elif response_type == "json_object": - kwargs["response_mime_type"] = "application/json" - elif response_type == "text": - pass - else: - msg = f"Unsupported response_format type: {response_type}" - raise ValueError(msg) + kwargs.update(_convert_response_format(params.response_format)) formatted_messages, system_instruction = _convert_messages(params.messages, provider_name=provider_name) if system_instruction: @@ -242,6 +319,8 @@ def _convert_completion_response(response: Any) -> ChatCompletion: tool_calls=tool_calls, reasoning=Reasoning(content=reasoning_content) if reasoning_content else None, extra_content=message_dict.get("extra_content"), + images=message_dict.get("images"), + audio=message_dict.get("audio"), ) from typing import Literal diff --git a/src/any_llm/providers/gemini/utils.py b/src/any_llm/providers/gemini/utils.py index 270e80f48..e4116d3d3 100644 --- a/src/any_llm/providers/gemini/utils.py +++ b/src/any_llm/providers/gemini/utils.py @@ -16,6 +16,7 @@ from any_llm.types.completion import ( ChatCompletionChunk, ChoiceDelta, + ChoiceDeltaAudio, ChoiceDeltaToolCall, ChoiceDeltaToolCallFunction, ChunkChoice, @@ -23,6 +24,7 @@ CompletionUsage, CreateEmbeddingResponse, Embedding, + ImageContent, PromptTokensDetails, Reasoning, Usage, @@ -425,6 +427,86 @@ def _thought_signature_extra_content(part: types.Part) -> dict[str, Any] | None: return None +def _inline_data_image(part: types.Part) -> dict[str, Any] | None: + """Convert Gemini inline data into an OpenAI-compatible data URL.""" + blob = part.inline_data + if ( + blob is None + or not isinstance(blob.data, bytes) + or not blob.data + or not isinstance(blob.mime_type, str) + or not blob.mime_type.startswith("image/") + ): + return None + return { + "type": "image_url", + "image_url": {"url": f"data:{blob.mime_type};base64,{base64.b64encode(blob.data).decode('ascii')}"}, + } + + +def _inline_audio_blob(part: types.Part) -> types.Blob | None: + """Return the part's non-empty audio blob, if it has one.""" + blob = part.inline_data + if ( + blob is None + or not isinstance(blob.data, bytes) + or not blob.data + or not isinstance(blob.mime_type, str) + or not blob.mime_type.startswith("audio/") + ): + return None + return blob + + +def _wav_from_pcm(pcm: bytes, mime_type: str) -> bytes: + """Wrap headerless 16-bit mono PCM in a WAV header.""" + rate = 24000 + for parameter in mime_type.split(";"): + parameter = parameter.strip() + if parameter.startswith("rate="): + with suppress(ValueError): + parsed_rate = int(parameter.removeprefix("rate=")) + if parsed_rate > 0: + rate = parsed_rate + break + return ( + b"RIFF" + + (36 + len(pcm)).to_bytes(4, "little") + + b"WAVE" + + b"fmt " + + (16).to_bytes(4, "little") + + (1).to_bytes(2, "little") + + (1).to_bytes(2, "little") + + rate.to_bytes(4, "little") + + (rate * 2).to_bytes(4, "little") + + (2).to_bytes(2, "little") + + (16).to_bytes(2, "little") + + b"data" + + len(pcm).to_bytes(4, "little") + + pcm + ) + + +def _inline_data_audio(blobs: list[types.Blob], transcript: str, *, playable: bool) -> dict[str, Any] | None: + """Convert Gemini audio blobs into an OpenAI-compatible audio object. + + Complete audio/L16 responses are wrapped as WAV so the sample rate from the MIME + type is retained. Streaming chunks stay raw because their total length is unknown. + """ + if not blobs: + return None + data = b"".join(cast("bytes", blob.data) for blob in blobs) + mime_type = cast("str", blobs[0].mime_type) + if playable and mime_type.startswith("audio/L16"): + data = _wav_from_pcm(data, mime_type) + return { + "id": "google_genai_audio", + "data": base64.b64encode(data).decode("ascii"), + "expires_at": 0, + "transcript": transcript, + } + + _FINISH_REASON_MAP: dict[types.FinishReason, Literal["stop", "length", "content_filter"]] = { types.FinishReason.STOP: "stop", types.FinishReason.MAX_TOKENS: "length", @@ -492,15 +574,17 @@ def _convert_response_to_response_dict(response: types.GenerateContentResponse) reasoning = None tool_calls_list: list[dict[str, Any]] = [] text_content = None + images: list[dict[str, Any]] = [] + audio_blobs: list[types.Blob] = [] # Gemini 3 signs the last non-function-call part of a text answer. It rides message.extra_content, # the same spelling Google's OpenAI-compatible endpoint uses. message_extra_content = None parts = candidate.content.parts if candidate.content else None for part in parts or []: - if getattr(part, "thought", None): + if part.thought: reasoning = (reasoning or "") + (part.text or "") - elif function_call := getattr(part, "function_call", None): + elif function_call := part.function_call: args_dict = {} if args := getattr(function_call, "args", None): for key, value in args.items(): @@ -521,14 +605,20 @@ def _convert_response_to_response_dict(response: types.GenerateContentResponse) tool_calls_list.append(tool_call_dict) else: + if image := _inline_data_image(part): + images.append(image) + if audio_blob := _inline_audio_blob(part): + audio_blobs.append(audio_blob) if part.text: text_content = (text_content or "") + part.text message_extra_content = _thought_signature_extra_content(part) or message_extra_content + audio = _inline_data_audio(audio_blobs, text_content or "", playable=True) + # Truncated or filtered responses produce a choice even without content or tool # calls, e.g. a thinking model that spent the whole max_output_tokens budget on # reasoning, so callers see the terminal reason instead of an empty choices list. - if tool_calls_list or text_content or mapped_finish_reason in ("length", "content_filter"): + if tool_calls_list or text_content or images or audio or mapped_finish_reason in ("length", "content_filter"): choices.append( { "message": { @@ -536,6 +626,8 @@ def _convert_response_to_response_dict(response: types.GenerateContentResponse) "content": text_content, "reasoning": reasoning or None, "tool_calls": tool_calls_list or None, + "images": images or None, + "audio": audio, "extra_content": message_extra_content, "refusal": _GEMINI_CONTENT_FILTER_REFUSAL if mapped_finish_reason == "content_filter" else None, }, @@ -613,6 +705,8 @@ def _create_openai_chunk_from_google_chunk( reasoning_content = "" tool_calls_list: list[ChoiceDeltaToolCall] = [] message_extra_content = None + images: list[dict[str, Any]] = [] + audio_blobs: list[types.Blob] = [] # Content can be absent on terminal chunks, e.g. when the response is truncated or # filtered before any part is produced; the finish reason must still be surfaced. @@ -647,9 +741,17 @@ def _create_openai_chunk_from_google_chunk( ) ) else: + if image := _inline_data_image(part): + images.append(image) + if audio_blob := _inline_audio_blob(part): + audio_blobs.append(audio_blob) content += part.text or "" # the signed final part may carry empty text message_extra_content = _thought_signature_extra_content(part) or message_extra_content + audio = None + if converted_audio := _inline_data_audio(audio_blobs, content, playable=False): + audio = ChoiceDeltaAudio(data=converted_audio["data"], transcript=converted_audio["transcript"] or None) + # Unmapped reasons stay None so non-final chunks are not forced to a terminal reason. mapped_finish_reason = _map_finish_reason(candidate.finish_reason) if candidate else None prompt_was_blocked = _prompt_was_blocked(response) @@ -664,6 +766,8 @@ def _create_openai_chunk_from_google_chunk( reasoning=Reasoning(content=reasoning_content) if reasoning_content else None, tool_calls=tool_calls_list or None, extra_content=message_extra_content, + images=cast("list[ImageContent] | None", images or None), + audio=audio, ) choice = ChunkChoice( diff --git a/src/any_llm/providers/ollama/utils.py b/src/any_llm/providers/ollama/utils.py index 94bef0494..70d74043b 100644 --- a/src/any_llm/providers/ollama/utils.py +++ b/src/any_llm/providers/ollama/utils.py @@ -29,6 +29,22 @@ from any_llm.types.model import Model +def _map_ollama_done_reason( + done_reason: str | None, +) -> Literal["stop", "length", "tool_calls", "content_filter", "function_call"]: + """Normalize a terminal reason, using stop when Ollama has no OpenAI equivalent. + + Ollama's load/unload responses complete a model lifecycle operation without + generating text. Missing and unknown reasons use the same stop fallback as + other providers; they do not imply truncation, filtering, or a tool call. + """ + match done_reason: + case "stop" | "length" | "tool_calls" | "content_filter" | "function_call": + return done_reason + case _: + return "stop" + + def _convert_tool_calls(tool_calls: list[dict[str, Any]]) -> list[dict[str, Any]]: """Convert OpenAI tool calls to the shape Ollama expects. @@ -130,9 +146,10 @@ def _create_openai_chunk_from_ollama_chunk(ollama_chunk: OllamaChatResponse) -> choice = ChunkChoice( index=0, delta=delta, - finish_reason=cast( - "Literal['stop', 'length', 'tool_calls', 'content_filter', 'function_call'] | None", - ollama_chunk.done_reason, + finish_reason=( + _map_ollama_done_reason(ollama_chunk.done_reason) + if ollama_chunk.done or ollama_chunk.done_reason is not None + else None ), ) @@ -217,7 +234,7 @@ def _create_chat_completion_from_ollama_response(response: OllamaChatResponse) - reasoning=Reasoning(content=response_message.thinking) if response_message.thinking else None, ) - finish_reason: Any = "tool_calls" if openai_tool_calls else response.done_reason + finish_reason = "tool_calls" if openai_tool_calls else _map_ollama_done_reason(response.done_reason) choice = Choice(index=0, finish_reason=finish_reason, message=message) diff --git a/src/any_llm/providers/otari/otari.py b/src/any_llm/providers/otari/otari.py index 8ca31fec9..6168d2a66 100644 --- a/src/any_llm/providers/otari/otari.py +++ b/src/any_llm/providers/otari/otari.py @@ -5,6 +5,7 @@ import os from typing import TYPE_CHECKING, Any, TypedDict, cast +from anthropic import transform_schema from pydantic import BaseModel from typing_extensions import override @@ -36,6 +37,7 @@ build_responses_text_format, get_json_schema, is_structured_output_type, + normalize_output_config, parse_json_content, ) @@ -179,6 +181,7 @@ class OtariProvider(BaseOpenAIProvider): SUPPORTS_AUDIO_TRANSCRIPTION = True SUPPORTS_AUDIO_SPEECH = True SUPPORTS_RERANK = True + SUPPORTS_MESSAGES_STRUCTURED_OUTPUT_STREAMING = True otari_client: Any @@ -356,21 +359,16 @@ async def _amessages( parameter_name = "container" raise UnsupportedParameterError(parameter_name, self.PROVIDER_NAME) - if params.output_format is not None: - # Structured output is handled by the base Messages<->Completions bridge, which - # routes output_format through otari's completion path. A follow-up could adopt - # otari's native /messages structured-output support directly. - if params.context_management is not None or params.betas: - msg = ( - "output_format cannot be combined with context_management or betas on otari: " - "structured output routes through the Completions bridge, which drops both. " - "Send them in separate requests until otari's native /messages structured " - "output is adopted." - ) - raise NotImplementedError(msg) - return await super()._amessages(params, **kwargs) - - api_kwargs = params.model_dump(exclude_none=True) + api_kwargs = params.model_dump(exclude_none=True, exclude={"output_format"}) + if is_structured_output_type(params.output_format): + api_kwargs["output_format"] = { + "format": { + "type": "json_schema", + "schema": transform_schema(get_json_schema(params.output_format)), + } + } + elif isinstance(params.output_format, dict): + api_kwargs["output_format"] = normalize_output_config(params.output_format) api_kwargs.update(kwargs) api_kwargs.pop("stream", None) diff --git a/src/any_llm/types/completion.py b/src/any_llm/types/completion.py index 2333813a6..0b98a6d20 100644 --- a/src/any_llm/types/completion.py +++ b/src/any_llm/types/completion.py @@ -73,6 +73,33 @@ class ChatCompletionMessageFunctionToolCall(OpenAIChatCompletionMessageFunctionT ChatCompletionMessageToolCall = ChatCompletionMessageFunctionToolCall | OpenAIChatCompletionMessageToolCall +class ImageURL(BaseModel): + """OpenAI-compatible URL for an image response part.""" + + url: str + + +class ImageContent(BaseModel): + """OpenAI-compatible image response content.""" + + type: Literal["image_url"] + image_url: ImageURL + + +class ChoiceDeltaAudio(BaseModel): + """Partial audio object emitted by a streaming chat completion. + + Providers may send the identifier, transcript, data, and expiration timestamp in + separate chunks, so every field is optional. The data field contains the + base64-encoded bytes for the current chunk. + """ + + id: str | None = None + data: str | None = None + transcript: str | None = None + expires_at: int | None = None + + class ChatCompletionMessage(OpenAIChatCompletionMessage): tool_calls: list[ChatCompletionMessageToolCall] | None = None # type: ignore[assignment] reasoning: Reasoning | None = None @@ -87,6 +114,8 @@ class ChatCompletionMessage(OpenAIChatCompletionMessage): {"anthropic": {"signature": ""}} """ + images: list[ImageContent] | None = None + class Choice(OpenAIChoice): message: ChatCompletionMessage @@ -133,6 +162,9 @@ class ChoiceDelta(OpenAIChoiceDelta): that arrives as part of a streaming delta rather than the final message. """ + images: list[ImageContent] | None = None + audio: ChoiceDeltaAudio | None = None + class ChunkChoice(OpenAIChunkChoice): delta: ChoiceDelta diff --git a/src/any_llm/types/messages.py b/src/any_llm/types/messages.py index 3ceb71375..2b4eb7bc0 100644 --- a/src/any_llm/types/messages.py +++ b/src/any_llm/types/messages.py @@ -236,12 +236,18 @@ class MessagesParams(BaseModel): Either a Pydantic ``BaseModel`` subclass or dataclass **type**, or a raw Anthropic ``output_config`` **dict** (e.g. ``{"format": {"type": "json_schema", "schema": {...}}}``) for non-Pydantic JSON schemas. The bare ``format`` object - (``{"type": "json_schema", "schema": {...}}``) is accepted as well. A type goes to native - ``messages.parse`` on Anthropic; a dict is passed through to native - ``messages.create(output_config=...)``. Other providers route either form through the - completion bridge. The result is Anthropic's ``ParsedMessage``: its ``parsed_output`` holds - the typed object for a type, or the parsed JSON (plain ``dict``/``list``) for a raw schema. - - A dict that carries no schema in either shape raises ``InvalidRequestError`` on the bridged - path rather than sending a ``response_format`` that constrains nothing. + (``{"type": "json_schema", "schema": {...}}``) is accepted as well. Anthropic sends a type + to native ``messages.parse`` and a dict to native + ``messages.create(output_config=...)``. Providers with native Messages structured-output + support, such as Anthropic and Otari, keep these requests on their Messages APIs; other + providers route schema-backed forms through the completion bridge. Non-streaming + schema-backed requests return Anthropic's ``ParsedMessage``: its ``parsed_output`` holds the + typed object for a type, or the parsed JSON (plain ``dict``/``list``) for a raw schema. + Providers with native streaming support return schema-constrained Messages events when + ``stream=True``. + + A dict that names a format but carries no usable schema raises ``InvalidRequestError`` on the + bridged path rather than sending a ``response_format`` that constrains nothing. A dict without + a format remains a provider-specific output configuration and returns a regular + ``MessageResponse``. """ diff --git a/src/any_llm/utils/exception_handler.py b/src/any_llm/utils/exception_handler.py index c3f1bb1c6..de564705d 100644 --- a/src/any_llm/utils/exception_handler.py +++ b/src/any_llm/utils/exception_handler.py @@ -142,9 +142,11 @@ def _extract_status_code(exception: Exception) -> int | None: response = getattr(exception, "response", None) if response is not None: - response_status = getattr(response, "status_code", None) - if isinstance(response_status, int): - return response_status + # httpx and requests spell it status_code; aiohttp spells it status + for name in ("status_code", "status"): + response_status = getattr(response, name, None) + if isinstance(response_status, int): + return response_status return None diff --git a/tests/conftest.py b/tests/conftest.py index 3e78eb185..83cee0229 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -108,7 +108,7 @@ def provider_model_map() -> dict[LLMProvider, str]: LLMProvider.VERTEXAI: "gemini-3-flash-preview", LLMProvider.MOONSHOT: "kimi-k3", LLMProvider.SAMBANOVA: "gpt-oss-120b", - LLMProvider.TOGETHER: "openai/gpt-oss-20b", + LLMProvider.TOGETHER: "Qwen/Qwen3.5-9B", LLMProvider.XAI: "grok-3-mini-latest", LLMProvider.INCEPTION: "mercury", LLMProvider.NEOSANTARA: "gemini-3-flash-preview", diff --git a/tests/integration/test_agent_loop.py b/tests/integration/test_agent_loop.py index 00eddc217..cdad46506 100644 --- a/tests/integration/test_agent_loop.py +++ b/tests/integration/test_agent_loop.py @@ -1,8 +1,9 @@ import inspect import json +import re import warnings from collections.abc import Callable -from typing import TYPE_CHECKING, Any +from typing import Any import httpx import pytest @@ -13,11 +14,9 @@ from any_llm import AnyLLM, LLMProvider from any_llm.exceptions import MissingApiKeyError +from any_llm.types.completion import ChatCompletionMessage from tests.constants import EXPECTED_PROVIDERS, LOCAL_PROVIDERS -if TYPE_CHECKING: - from any_llm.types.completion import ChatCompletion, ChatCompletionMessage - def get_current_date() -> str: """Get the current date and time.""" @@ -46,8 +45,82 @@ def _call_tool(tool_fn: Callable[..., str], args: dict[str, Any]) -> str: return tool_fn(**{name: value for name, value in args.items() if name in accepted}) +async def _run_agent_loop( + llm: AnyLLM, + model_id: str, + messages: list[dict[str, Any] | ChatCompletionMessage], + available_tools: dict[str, Callable[..., str]], + calls_complete: Callable[[list[tuple[str, dict[str, Any]]]], bool], + *, + include_tool_name: bool, + max_iterations: int = 5, +) -> tuple[ChatCompletionMessage, list[tuple[str, dict[str, Any]]]]: + """Execute tool calls until the required calls are complete and the model answers.""" + calls_made: list[tuple[str, dict[str, Any]]] = [] + + for _ in range(max_iterations): + result = await llm.acompletion( + model=model_id, + messages=messages, + tools=list(available_tools.values()), + ) + message = result.choices[0].message + tool_calls = message.tool_calls + + if not tool_calls: + assert calls_complete(calls_made), f"Model answered before making the required tool calls: {calls_made}" + return message, calls_made + + messages.append(message) + for tool_call in tool_calls: + assert isinstance(tool_call, OpenAIChatCompletionMessageFunctionToolCall), ( + f"Expected a function tool call, got: {tool_call}" + ) + tool_name = tool_call.function.name + assert tool_name in available_tools, f"Unknown tool: {tool_name}" + args = json.loads(tool_call.function.arguments) if tool_call.function.arguments else {} + calls_made.append((tool_name, args)) + + tool_message: dict[str, Any] = { + "role": "tool", + "content": _call_tool(available_tools[tool_name], args), + "tool_call_id": tool_call.id, + } + if include_tool_name: + tool_message["name"] = tool_name + messages.append(tool_message) + + error = f"Agent loop did not answer within {max_iterations} iterations; calls: {calls_made}" + raise AssertionError(error) + + +def _mentions_tool_result(content: str | None) -> bool: + """The weather tool returns 15C and sunny, so an answer built on it repeats one of them. + + ``15`` must not run into another digit, so ``150F`` does not count; ``15C``, ``15°C`` and + ``15 degrees`` all do. + """ + return content is not None and re.search(r"\b15(?!\d)|\bsunny\b", content, re.IGNORECASE) is not None + + +@pytest.mark.parametrize( + ("content", "expected"), + [ + (None, False), + ("", False), + ("It rains in Paris.", False), + ("It is 150F in Paris.", False), + ("It is 15C in Paris.", True), + ("It is 15°C in Paris.", True), + ("Sunny in London.", True), + ], +) +def test_mentions_tool_result(content: str | None, expected: bool) -> None: + assert _mentions_tool_result(content) is expected + + @pytest.mark.asyncio -async def test_agent_loop_parallel_tool_calls( +async def test_agent_loop_multiple_tool_calls( provider: LLMProvider, provider_model_map: dict[LLMProvider, str], provider_client_config: dict[LLMProvider, dict[str, Any]], @@ -69,40 +142,19 @@ async def test_agent_loop_parallel_tool_calls( } ] - result: ChatCompletion = await llm.acompletion( - model=model_id, - messages=messages, - tools=[get_weather], - ) - - tool_calls = result.choices[0].message.tool_calls - assert tool_calls is not None, f"Expected tool calls, got: {result.choices[0].message}" - - messages.append(result.choices[0].message) - - for tool_call in tool_calls: - if not isinstance(tool_call, OpenAIChatCompletionMessageFunctionToolCall): - continue - assert tool_call.function.name == "get_weather" - args = json.loads(tool_call.function.arguments) - tool_result = _call_tool(get_weather, args) - - messages.append( - { - "role": "tool", - "content": tool_result, - "tool_call_id": tool_call.id, - "name": tool_call.function.name, - } - ) - - second_result: ChatCompletion = await llm.acompletion( - model=model_id, - messages=messages, - tools=[get_weather], + def called_both_locations(calls: list[tuple[str, dict[str, Any]]]) -> bool: + locations = {args.get("location") for tool_name, args in calls if tool_name == "get_weather"} + return locations >= {"Paris", "London"} + + message, _ = await _run_agent_loop( + llm, + model_id, + messages, + {"get_weather": get_weather}, + called_both_locations, + include_tool_name=False, ) - - assert second_result.choices[0].message.content is not None or second_result.choices[0].message.tool_calls + assert _mentions_tool_result(message.content), f"Expected an answer from the tool results, got: {message}" except MissingApiKeyError: if provider in EXPECTED_PROVIDERS: @@ -115,7 +167,7 @@ async def test_agent_loop_parallel_tool_calls( @pytest.mark.asyncio -async def test_agent_loop_sequential_tool_calls( +async def test_agent_loop_multiple_tool_types( provider: LLMProvider, provider_model_map: dict[LLMProvider, str], provider_client_config: dict[LLMProvider, dict[str, Any]], @@ -137,52 +189,24 @@ async def test_agent_loop_sequential_tool_calls( } ] - tools = [get_current_date, get_weather] available_tools: dict[str, Callable[..., str]] = { "get_current_date": get_current_date, "get_weather": get_weather, } - max_iterations = 5 - iteration = 0 - - while iteration < max_iterations: - iteration += 1 - - result: ChatCompletion = await llm.acompletion( - model=model_id, - messages=messages, - tools=tools, - ) - - tool_calls = result.choices[0].message.tool_calls - - if tool_calls is None: - assert result.choices[0].message.content is not None - break - - messages.append(result.choices[0].message) - - for tool_call in tool_calls: - if not isinstance(tool_call, OpenAIChatCompletionMessageFunctionToolCall): - continue - tool_name = tool_call.function.name - assert tool_name in available_tools, f"Unknown tool: {tool_name}" - tool_fn = available_tools[tool_name] - - args = json.loads(tool_call.function.arguments) if tool_call.function.arguments else {} - tool_result = _call_tool(tool_fn, args) - - messages.append( - { - "role": "tool", - "content": tool_result, - "tool_call_id": tool_call.id, - "name": tool_call.function.name, - } - ) - - assert iteration <= max_iterations, "Agent loop did not complete within max iterations" + def called_both_tools(calls: list[tuple[str, dict[str, Any]]]) -> bool: + return {tool_name for tool_name, _ in calls} >= set(available_tools) + + # Callers may still send name on tool messages, so one loop keeps that shape on the wire. + message, _ = await _run_agent_loop( + llm, + model_id, + messages, + available_tools, + called_both_tools, + include_tool_name=True, + ) + assert _mentions_tool_result(message.content), f"Expected an answer from the tool results, got: {message}" except MissingApiKeyError: if provider in EXPECTED_PROVIDERS: diff --git a/tests/unit/providers/test_anthropic_messages.py b/tests/unit/providers/test_anthropic_messages.py index da35f100a..a6ff1600b 100644 --- a/tests/unit/providers/test_anthropic_messages.py +++ b/tests/unit/providers/test_anthropic_messages.py @@ -9,6 +9,7 @@ import httpx import pytest +from anthropic import transform_schema from anthropic.types import Message, TextBlock, ThinkingBlock, ToolUseBlock, Usage from anthropic.types.beta import BetaMCPToolUseBlock, BetaMessage, BetaThinkingBlock, BetaUsage from pydantic import BaseModel @@ -1188,6 +1189,195 @@ async def test_amessages_output_config_dict_passes_through_to_create() -> None: assert call_kwargs["max_tokens"] == 1024 +@pytest.mark.asyncio +async def test_amessages_output_config_dict_streams_with_anthropic_fields() -> None: + output_config = {"format": {"type": "json_schema", "schema": {"type": "object"}}} + stream_result = AsyncMock() + mock_client = Mock() + mock_client.beta.messages.create = AsyncMock( + return_value=_make_message(content=[TextBlock(type="text", text="{}")]) + ) + + provider = Mock(spec=BaseAnthropicProvider) + provider.client = mock_client + provider._stream_messages_async = Mock(return_value=stream_result) + provider._convert_native_message_to_response = BaseAnthropicProvider._convert_native_message_to_response + context_management = {"edits": [{"type": "compact_20260112"}]} + params = MessagesParams( + model="claude-3-5-sonnet", + messages=[{"role": "user", "content": "Capital of France?"}], + max_tokens=1024, + stream=True, + output_format=output_config, + context_management=context_management, + betas=["compact-2026-01-12"], + cache_control={"type": "ephemeral"}, + ) + + result = await BaseAnthropicProvider._amessages(provider, params) + + assert result is stream_result + provider._stream_messages_async.assert_called_once() + call_kwargs = provider._stream_messages_async.call_args.kwargs + assert call_kwargs["use_beta"] is True + assert call_kwargs["output_config"] == output_config + assert call_kwargs["context_management"] == context_management + assert call_kwargs["betas"] == ["compact-2026-01-12"] + assert call_kwargs["cache_control"] == {"type": "ephemeral"} + mock_client.beta.messages.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_amessages_bare_output_config_is_normalized_for_streaming() -> None: + output_format = {"type": "json_schema", "schema": {"type": "object"}} + stream_result = AsyncMock() + provider = Mock(spec=BaseAnthropicProvider) + provider.client = Mock() + provider._stream_messages_async = Mock(return_value=stream_result) + params = MessagesParams( + model="claude-3-5-sonnet", + messages=[{"role": "user", "content": "Capital of France?"}], + max_tokens=1024, + stream=True, + output_format=output_format, + ) + + result = await BaseAnthropicProvider._amessages(provider, params) + + assert result is stream_result + call_kwargs = provider._stream_messages_async.call_args.kwargs + assert call_kwargs["output_config"] == {"format": output_format} + + +@pytest.mark.asyncio +async def test_amessages_typed_output_format_streams_with_sdk_parser() -> None: + class City(BaseModel): + city: str + + stream_result = AsyncMock() + mock_client = Mock() + mock_client.messages.parse = AsyncMock( + return_value=_make_message(content=[TextBlock(type="text", text='{"city": "Paris"}')]) + ) + + provider = Mock(spec=BaseAnthropicProvider) + provider.client = mock_client + provider._stream_messages_async = Mock(return_value=stream_result) + params = MessagesParams( + model="claude-3-5-sonnet", + messages=[{"role": "user", "content": "Capital of France?"}], + max_tokens=1024, + stream=True, + output_format=City, + ) + + result = await BaseAnthropicProvider._amessages(provider, params) + + assert result is stream_result + call_kwargs = provider._stream_messages_async.call_args.kwargs + assert call_kwargs["use_beta"] is False + assert call_kwargs["output_format"] is City + mock_client.messages.parse.assert_not_called() + + +@pytest.mark.asyncio +async def test_anthropic_provider_allows_streaming_output_format() -> None: + class City(BaseModel): + city: str + + async def events() -> AsyncIterator[MessageStopEvent]: + yield MessageStopEvent(type="message_stop") + + with patch("any_llm.providers.anthropic.anthropic.AsyncAnthropic"): + provider = AnthropicProvider(api_key="test-key") + provider._stream_messages_async = Mock(return_value=events()) # type: ignore[method-assign] + + result = await provider.amessages( + model="claude-3-5-sonnet", + messages=[{"role": "user", "content": "Capital of France?"}], + max_tokens=1024, + stream=True, + output_format=City, + ) + collected = [event async for event in cast("AsyncIterator[MessageStopEvent]", result)] + + assert [event.type for event in collected] == ["message_stop"] + + +@pytest.mark.asyncio +async def test_amessages_typed_output_format_streams_through_sdk_transport() -> None: + class City(BaseModel): + city: str + + requests: list[httpx.Request] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + events = [ + ( + "message_start", + { + "type": "message_start", + "message": { + "id": "msg_structured", + "type": "message", + "role": "assistant", + "model": "claude-3-5-sonnet", + "stop_reason": None, + "stop_sequence": None, + "content": [], + "usage": {"input_tokens": 1, "output_tokens": 0}, + }, + }, + ), + ( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ), + ( + "content_block_delta", + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": '{"city":"Paris"}'}, + }, + ), + ("content_block_stop", {"type": "content_block_stop", "index": 0}), + ( + "message_delta", + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": 1}, + }, + ), + ("message_stop", {"type": "message_stop"}), + ] + body = "".join(f"event: {name}\ndata: {json.dumps(payload)}\n\n" for name, payload in events) + return httpx.Response(200, headers={"content-type": "text/event-stream"}, content=body.encode()) + + http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = AnthropicProvider(api_key="test-key", http_client=http_client) + try: + result = await provider.amessages( + model="claude-3-5-sonnet", + messages=[{"role": "user", "content": "Capital of France?"}], + max_tokens=1024, + stream=True, + output_format=City, + ) + collected = [event async for event in cast("AsyncIterator[Any]", result)] + finally: + await http_client.aclose() + + assert collected[-1].type == "message_stop" + assert len(requests) == 1 + request_body = json.loads(requests[0].content) + assert request_body["output_config"] == { + "format": {"type": "json_schema", "schema": transform_schema(City.model_json_schema())} + } + + @pytest.mark.asyncio async def test_amessages_cache_control_passthrough() -> None: """Test that cache_control is passed through to the API call.""" diff --git a/tests/unit/providers/test_anthropic_provider.py b/tests/unit/providers/test_anthropic_provider.py index 9fe8468e8..c9d396ad4 100644 --- a/tests/unit/providers/test_anthropic_provider.py +++ b/tests/unit/providers/test_anthropic_provider.py @@ -3,8 +3,8 @@ from collections.abc import AsyncIterator from contextlib import asynccontextmanager, contextmanager from datetime import UTC, datetime -from typing import Any, cast, get_args -from unittest.mock import AsyncMock, Mock, patch +from typing import Any, Self, cast, get_args +from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest from anthropic import transform_schema @@ -21,6 +21,7 @@ _convert_models_list, _convert_response_format, _convert_tool_spec, + _create_openai_chunk_from_anthropic_chunk, ) from any_llm.types.completion import ChatCompletionMessageFunctionToolCall, CompletionParams, ReasoningEffort @@ -930,7 +931,7 @@ def test_streaming_chunk_includes_cache_tokens_in_usage() -> None: assert result.usage.total_tokens == expected_total_tokens assert result.usage.prompt_tokens_details is not None assert result.usage.prompt_tokens_details.cached_tokens == 13332 - assert result.choices[0].finish_reason is None + assert result.choices == [] @pytest.mark.asyncio @@ -1000,7 +1001,7 @@ def test_streaming_chunk_without_cache_tokens() -> None: assert result.usage.completion_tokens == 50 assert result.usage.total_tokens == 150 assert result.usage.prompt_tokens_details is None - assert result.choices[0].finish_reason is None + assert result.choices == [] def test_streaming_tool_chunks_preserve_parallel_tool_index() -> None: @@ -1824,3 +1825,140 @@ def test_convert_response_non_datetime_created_at(created_at: Any) -> None: assert result.created == 0 assert result.choices[0].message.content == "hello" + + +def test_stream_trailing_usage_chunk_has_no_choices() -> None: + """Usage arrives on the message_stop chunk after finish_reason, with choices left empty like OpenAI.""" + from anthropic.types import ( + ContentBlockDeltaEvent, + ContentBlockStopEvent, + MessageDeltaEvent, + MessageDeltaUsage, + MessageStopEvent, + TextDelta, + Usage, + ) + from anthropic.types.raw_message_delta_event import Delta + + stop_event = MessageStopEvent(type="message_stop") + stop_event.message = MagicMock(usage=Usage(input_tokens=12, output_tokens=7)) # type: ignore[attr-defined] + events = [ + ContentBlockDeltaEvent(type="content_block_delta", index=0, delta=TextDelta(type="text_delta", text="hi")), + ContentBlockStopEvent(type="content_block_stop", index=0), + MessageDeltaEvent( + type="message_delta", + delta=Delta(stop_reason="end_turn", stop_sequence=None), + usage=MessageDeltaUsage(output_tokens=7), + ), + stop_event, + ] + + results = [_create_openai_chunk_from_anthropic_chunk(event, "claude-sonnet-4-5") for event in events] + + finish_index = next(i for i, r in enumerate(results) if r.choices and r.choices[0].finish_reason == "stop") + usage_chunks = [r for r in results if r.usage is not None] + assert len(usage_chunks) == 1 + assert results.index(usage_chunks[0]) > finish_index + assert usage_chunks[0].choices == [] + assert usage_chunks[0].usage is not None + assert usage_chunks[0].usage.total_tokens == 19 + + +def test_stream_message_stop_without_message_has_no_choices_or_usage() -> None: + """A raw message_stop event with no accumulated message yields neither choices nor usage.""" + from anthropic.types import MessageStopEvent + + result = _create_openai_chunk_from_anthropic_chunk(MessageStopEvent(type="message_stop"), "claude-sonnet-4-5") + + assert result.choices == [] + assert result.usage is None + + +@pytest.mark.asyncio +async def test_stream_usage_reaches_openai_style_consumer() -> None: + """A loop written against OpenAI's stream contract must get text, finish_reason and usage from Anthropic. + + OpenAI reports final usage on a trailing chunk with no choices, so callers read it with + ``if not chunk.choices``. The same loop has to work unchanged when the provider is Anthropic. + """ + from anthropic.lib.streaming import MessageStopEvent as StreamMessageStopEvent + from anthropic.types import ( + ContentBlockDeltaEvent, + ContentBlockStartEvent, + ContentBlockStopEvent, + MessageDeltaEvent, + MessageDeltaUsage, + MessageStartEvent, + TextBlock, + TextDelta, + Usage, + ) + from anthropic.types.raw_message_delta_event import Delta + + final_message = Message( + id="msg_1", + type="message", + role="assistant", + model="claude-sonnet-4-6", + content=[TextBlock(type="text", text="Hi!")], + stop_reason="end_turn", + usage=Usage(input_tokens=13, output_tokens=5), + ) + events: list[Any] = [ + MessageStartEvent( + type="message_start", + message=final_message.model_copy( + update={"content": [], "stop_reason": None, "usage": Usage(input_tokens=13, output_tokens=0)} + ), + ), + ContentBlockStartEvent(type="content_block_start", index=0, content_block=TextBlock(type="text", text="")), + ContentBlockDeltaEvent(type="content_block_delta", index=0, delta=TextDelta(type="text_delta", text="Hi")), + ContentBlockDeltaEvent(type="content_block_delta", index=0, delta=TextDelta(type="text_delta", text="!")), + ContentBlockStopEvent(type="content_block_stop", index=0), + MessageDeltaEvent( + type="message_delta", + delta=Delta(stop_reason="end_turn", stop_sequence=None), + usage=MessageDeltaUsage(output_tokens=5), + ), + StreamMessageStopEvent(type="message_stop", message=final_message), + ] + + class FakeMessageStream: + def __init__(self) -> None: + self._events = iter(events) + + async def __aenter__(self) -> Self: + return self + + async def __aexit__(self, *args: object) -> None: + pass + + def __aiter__(self) -> Self: + return self + + async def __anext__(self) -> Any: + try: + return next(self._events) + except StopIteration: + raise StopAsyncIteration from None + + provider = AnthropicProvider(api_key="sk-test") + with patch.object(provider.client.messages, "stream", return_value=FakeMessageStream()): + stream = await provider.acompletion( + model="claude-sonnet-4-6", messages=[{"role": "user", "content": "Say hi."}], stream=True + ) + text = "" + finish_reasons: list[str] = [] + usage = None + async for chunk in stream: + if chunk.choices: + text += chunk.choices[0].delta.content or "" + if chunk.choices[0].finish_reason: + finish_reasons.append(chunk.choices[0].finish_reason) + else: + usage = chunk.usage + + assert text == "Hi!" + assert finish_reasons == ["stop"] + assert usage is not None + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (13, 5, 18) diff --git a/tests/unit/providers/test_aws_provider.py b/tests/unit/providers/test_aws_provider.py index b518e1ec9..4fd8f72b5 100644 --- a/tests/unit/providers/test_aws_provider.py +++ b/tests/unit/providers/test_aws_provider.py @@ -13,7 +13,12 @@ from botocore.tokens import ScopedEnvTokenProvider from pydantic import BaseModel -from any_llm.exceptions import InvalidRequestError, MissingApiKeyError, UnsupportedParameterError +from any_llm.exceptions import ( + ContentFilterFinishReasonError, + InvalidRequestError, + MissingApiKeyError, + UnsupportedParameterError, +) from any_llm.providers.bedrock import BedrockProvider from any_llm.providers.bedrock.utils import ( _STRUCTURED_OUTPUT_TOOL_NAME, @@ -1584,3 +1589,90 @@ def test_convert_messages_merges_consecutive_user_messages() -> None: {"toolResult": {"toolUseId": "t1", "content": [{"text": "here"}]}}, {"text": "what is in it"}, ] + + +def _bedrock_stop_reasons() -> tuple[str, ...]: + """Every stopReason the Converse API can return, read from the installed botocore service model. + + Reading the enum out of the service model rather than hardcoding it means a botocore upgrade + that adds a stop reason fails the parametrized tests below instead of silently defaulting the + new reason to "stop". + """ + service_model = botocore.session.Session().get_service_model("bedrock-runtime") # type: ignore[no-untyped-call] + return tuple(service_model.shape_for("StopReason").enum) + + +_EXPECTED_FINISH_REASONS = { + "end_turn": "stop", + "stop_sequence": "stop", + "malformed_model_output": "stop", + "malformed_tool_use": "stop", + "max_tokens": "length", + "model_context_window_exceeded": "length", + "tool_use": "tool_calls", + "content_filtered": "content_filter", + "guardrail_intervened": "content_filter", +} + + +@pytest.mark.parametrize("stop_reason", _bedrock_stop_reasons()) +def test_convert_response_maps_every_bedrock_stop_reason(stop_reason: str) -> None: + """Every stopReason the Converse API can return needs an explicit OpenAI finish_reason. + + An unmapped one falls back to "stop", which tells callers the model answered normally when it + was actually blocked by a guardrail or ran out of context. + """ + assert stop_reason in _EXPECTED_FINISH_REASONS, ( + f"New Bedrock stop reason {stop_reason!r} needs a finish_reason mapping." + ) + + response: dict[str, Any] = { + "output": {"message": {"content": [{"text": "Hello!"}]}}, + "stopReason": stop_reason, + } + + result = _convert_response(response) + + assert result.choices[0].finish_reason == _EXPECTED_FINISH_REASONS[stop_reason] + + +@pytest.mark.parametrize("stop_reason", _bedrock_stop_reasons()) +def test_streaming_chunk_maps_every_bedrock_stop_reason(stop_reason: str) -> None: + """The streaming path must agree with the non-streaming one on every stopReason.""" + result = _create_openai_chunk_from_aws_chunk({"messageStop": {"stopReason": stop_reason}}, "test-model") + + assert result is not None + assert result.choices[0].finish_reason == _EXPECTED_FINISH_REASONS[stop_reason] + + +def test_convert_response_without_stop_reason_finishes_as_stop() -> None: + """A response with no stopReason at all still needs a valid finish_reason.""" + response: dict[str, Any] = {"output": {"message": {"content": [{"text": "Hello!"}]}}} + + result = _convert_response(response) + + assert result.choices[0].finish_reason == "stop" + + +@pytest.mark.asyncio +async def test_guardrail_blocked_structured_output_raises_content_filter_error() -> None: + """A guardrail block on a structured-output call must reach the caller as a typed error. + + While guardrail_intervened mapped to "stop", the content_filter guard in + AnyLLM.acompletion never fired and the guardrail's blocked-message prose was handed to the + JSON parser instead, so callers saw a pydantic ValidationError from an unrelated layer. + """ + custom_client = Mock() + custom_client.converse.return_value = { + "output": {"message": {"content": [{"text": "Sorry, I cannot answer that."}]}}, + "stopReason": "guardrail_intervened", + } + + provider = BedrockProvider(client=custom_client) + + with pytest.raises(ContentFilterFinishReasonError): + await provider.acompletion( + model="us.anthropic.claude-sonnet-4-20250514-v1:0", + messages=[{"role": "user", "content": "Hello"}], + response_format=_City, + ) diff --git a/tests/unit/providers/test_deepseek_provider.py b/tests/unit/providers/test_deepseek_provider.py index 834f7ba9c..63bbf9aeb 100644 --- a/tests/unit/providers/test_deepseek_provider.py +++ b/tests/unit/providers/test_deepseek_provider.py @@ -1,14 +1,17 @@ import dataclasses +import json from typing import Any +import httpx import pytest from openai.types.chat.chat_completion import ChatCompletion as OpenAIChatCompletion from openai.types.chat.chat_completion_chunk import ChatCompletionChunk as OpenAIChatCompletionChunk from pydantic import BaseModel +from any_llm.exceptions import InvalidRequestError from any_llm.providers.deepseek.deepseek import DeepseekProvider from any_llm.providers.deepseek.utils import _preprocess_messages, _reinject_reasoning_content -from any_llm.types.completion import CompletionParams +from any_llm.types.completion import CompletionParams, ReasoningEffort class PersonResponseFormat(BaseModel): @@ -251,19 +254,20 @@ def test_deepseek_no_max_tokens_when_neither_set() -> None: assert "max_completion_tokens" not in result -def test_deepseek_thinking_disabled_by_default_for_v4_model() -> None: - """V4 model ids without reasoning_effort should default to thinking disabled.""" +def test_deepseek_preserves_provider_thinking_default() -> None: + """No explicit effort leaves DeepSeek's enabled/high provider default in control.""" params = CompletionParams( model_id="deepseek-v4-flash", messages=[{"role": "user", "content": "hi"}], reasoning_effort=None, ) result = DeepseekProvider._convert_completion_params(params) - assert result["extra_body"]["thinking"] == {"type": "disabled"} + assert "reasoning_effort" not in result + assert "extra_body" not in result def test_deepseek_thinking_disabled_for_none_reasoning_effort_value() -> None: - """reasoning_effort="none" is also treated as "no reasoning requested".""" + """An explicit none uses DeepSeek's thinking toggle, not an invalid wire effort.""" params = CompletionParams( model_id="deepseek-v4-pro", messages=[{"role": "user", "content": "hi"}], @@ -271,52 +275,197 @@ def test_deepseek_thinking_disabled_for_none_reasoning_effort_value() -> None: ) result = DeepseekProvider._convert_completion_params(params) assert result["extra_body"]["thinking"] == {"type": "disabled"} + assert "reasoning_effort" not in result -def test_deepseek_thinking_disabled_for_auto_reasoning_effort_value() -> None: - """reasoning_effort="auto" (the default) is treated as "no reasoning requested".""" +def test_deepseek_auto_preserves_provider_thinking_default() -> None: + """The normalized auto sentinel does not override DeepSeek's provider default.""" params = CompletionParams( model_id="deepseek-v4-flash", messages=[{"role": "user", "content": "hi"}], reasoning_effort="auto", ) result = DeepseekProvider._convert_completion_params(params) - assert result["extra_body"]["thinking"] == {"type": "disabled"} + assert "reasoning_effort" not in result + assert "extra_body" not in result -def test_deepseek_thinking_enabled_when_reasoning_effort_set() -> None: - """An explicit reasoning_effort should enable thinking mode and be passed through.""" +@pytest.mark.parametrize( + ("reasoning_effort", "expected_effort"), + [("low", "low"), ("medium", "high"), ("high", "high"), ("xhigh", "high"), ("max", "max")], +) +def test_deepseek_maps_current_reasoning_efforts(reasoning_effort: ReasoningEffort, expected_effort: str) -> None: + """Normalized efforts map to DeepSeek's current low, high, and max wire values.""" params = CompletionParams( model_id="deepseek-v4-flash", messages=[{"role": "user", "content": "hi"}], - reasoning_effort="high", + reasoning_effort=reasoning_effort, ) result = DeepseekProvider._convert_completion_params(params) assert result["extra_body"]["thinking"] == {"type": "enabled"} - assert result["reasoning_effort"] == "high" + assert result["reasoning_effort"] == expected_effort + + +def test_deepseek_rejects_unsupported_minimal_reasoning_effort() -> None: + """DeepSeek Chat does not document OpenAI's minimal effort.""" + params = CompletionParams( + model_id="deepseek-v4-pro", + messages=[{"role": "user", "content": "hi"}], + reasoning_effort="minimal", + ) + + with pytest.raises(InvalidRequestError, match="minimal"): + DeepseekProvider._convert_completion_params(params) + + +def test_deepseek_thinking_respects_explicit_extra_body_override() -> None: + """Caller-supplied DeepSeek fields take precedence over normalized controls.""" + params = CompletionParams( + model_id="deepseek-v4-flash", + messages=[{"role": "user", "content": "hi"}], + reasoning_effort="max", + user="ignored.invalid-user", + ) + result = DeepseekProvider._convert_completion_params( + params, + extra_body={"thinking": {"type": "disabled"}, "user_id": "caller-user"}, + ) + assert result["reasoning_effort"] == "max" + assert result["extra_body"] == {"thinking": {"type": "disabled"}, "user_id": "caller-user"} -def test_deepseek_thinking_untouched_for_legacy_model_ids() -> None: - """Legacy deepseek-chat/deepseek-reasoner ids must not get the thinking toggle injected.""" - for model_id in ("deepseek-chat", "deepseek-reasoner"): +def test_deepseek_does_not_mutate_reused_extra_body_when_mapping_user() -> None: + extra_body = {"custom": "value"} + results: list[dict[str, Any]] = [] + + for user in ("alice", "bob"): params = CompletionParams( - model_id=model_id, + model_id="deepseek-v4-flash", messages=[{"role": "user", "content": "hi"}], - reasoning_effort="high", + reasoning_effort="low", + user=user, ) - result = DeepseekProvider._convert_completion_params(params) - assert "extra_body" not in result + results.append(DeepseekProvider._convert_completion_params(params, extra_body=extra_body)) -def test_deepseek_thinking_respects_explicit_extra_body_override() -> None: - """A caller-supplied extra_body/thinking value should not be clobbered by the default.""" + assert extra_body == {"custom": "value"} + + assert results[0]["extra_body"] == { + "custom": "value", + "thinking": {"type": "enabled"}, + "user_id": "alice", + } + assert results[1]["extra_body"] == { + "custom": "value", + "thinking": {"type": "enabled"}, + "user_id": "bob", + } + assert all(result["extra_body"] is not extra_body for result in results) + + +@pytest.mark.asyncio +async def test_deepseek_emits_current_chat_wire_contract() -> None: + requests: list[httpx.Request] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + request.read() + requests.append(request) + return httpx.Response( + 200, + json={ + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 1, + "model": "deepseek-v4-pro", + "choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "ok"}}], + }, + ) + + user_id = "a" * 512 + provider = DeepseekProvider( + api_key="test-key", + http_client=httpx.AsyncClient(transport=httpx.MockTransport(handle_request)), + max_retries=0, + ) + try: + await provider.acompletion( + model="deepseek-v4-pro", + messages=[ + {"role": "user", "content": "Use the lookup tool"}, + { + "role": "assistant", + "content": "", + "extra_content": {"deepseek": {"reasoning_content": "I should use lookup."}}, + }, + ], + tools=[{"type": "function", "function": {"name": "lookup", "parameters": {"type": "object"}}}], + reasoning_effort="medium", + user=user_id, + n=2, + frequency_penalty=0.5, + presence_penalty=0.5, + seed=7, + parallel_tool_calls=False, + logit_bias={"123": 1.0}, + service_tier="priority", + ) + finally: + await provider.client.close() + + assert len(requests) == 1 + request = requests[0] + body = json.loads(request.content) + assert request.url.path == "/chat/completions" + assert body["model"] == "deepseek-v4-pro" + assert body["reasoning_effort"] == "high" + assert body["thinking"] == {"type": "enabled"} + assert body["user_id"] == user_id + assert body["messages"][1]["reasoning_content"] == "I should use lookup." + assert "extra_content" not in body["messages"][1] + for field in ( + "user", + "n", + "frequency_penalty", + "presence_penalty", + "seed", + "parallel_tool_calls", + "logit_bias", + "service_tier", + ): + assert field not in body + + +@pytest.mark.parametrize("user_id", ["", "account.42", "a" * 513]) +def test_deepseek_rejects_invalid_user_id(user_id: str) -> None: params = CompletionParams( model_id="deepseek-v4-flash", messages=[{"role": "user", "content": "hi"}], - reasoning_effort=None, + user=user_id, ) - result = DeepseekProvider._convert_completion_params(params, extra_body={"thinking": {"type": "enabled"}}) - assert result["extra_body"]["thinking"] == {"type": "enabled"} + + with pytest.raises(InvalidRequestError, match="user_id"): + DeepseekProvider._convert_completion_params(params) + + +@pytest.mark.parametrize( + ("params_kwargs", "expected_extra_body"), + [ + ({"user": "account_42"}, {"user_id": "account_42"}), + ({"reasoning_effort": "low"}, {"thinking": {"type": "enabled"}}), + ], +) +def test_deepseek_treats_explicit_none_extra_body_as_absent( + params_kwargs: dict[str, Any], expected_extra_body: dict[str, Any] +) -> None: + params = CompletionParams( + model_id="deepseek-v4-flash", + messages=[{"role": "user", "content": "hi"}], + **params_kwargs, + ) + + result = DeepseekProvider._convert_completion_params(params, extra_body=None) + + assert result["extra_body"] == expected_extra_body def test_convert_completion_response_stashes_reasoning_into_extra_content() -> None: @@ -389,7 +538,7 @@ def test_reinject_reasoning_content_on_tool_call_message() -> None: {"role": "tool", "tool_call_id": "call_1", "content": "Sunny"}, ] - result = _reinject_reasoning_content(messages) + result = _reinject_reasoning_content(messages, replay_reasoning=True) assert result[1]["reasoning_content"] == "I should call get_weather." # The any_llm-internal extra_content must not be forwarded to DeepSeek's API. @@ -399,8 +548,8 @@ def test_reinject_reasoning_content_on_tool_call_message() -> None: assert "extra_content" in messages[1] -def test_reinject_reasoning_content_skips_non_tool_call_message() -> None: - """Assistant messages without tool_calls should be left untouched.""" +def test_reinject_reasoning_content_includes_assistant_turn_without_tool_call() -> None: + """A request carrying tools replays reasoning from every previous assistant turn.""" messages: list[dict[str, Any]] = [ {"role": "user", "content": "hi"}, { @@ -410,13 +559,65 @@ def test_reinject_reasoning_content_skips_non_tool_call_message() -> None: }, ] - result = _reinject_reasoning_content(messages) + result = _reinject_reasoning_content(messages, replay_reasoning=True) - assert "reasoning_content" not in result[1] - # extra_content is any_llm-internal and is stripped even when reasoning is not reinjected. + assert result[1]["reasoning_content"] == "greeting" + # extra_content remains internal and is stripped after its reasoning is reinjected. assert "extra_content" not in result[1] +def test_preprocess_messages_replays_all_assistant_reasoning_when_tools_are_present() -> None: + params = CompletionParams( + model_id="deepseek-v4-pro", + messages=[ + { + "role": "assistant", + "content": "hello", + "extra_content": {"deepseek": {"reasoning_content": "greeting"}}, + } + ], + tools=[{"type": "function", "function": {"name": "lookup"}}], + ) + + processed = _preprocess_messages(params) + + assert processed.messages[0]["reasoning_content"] == "greeting" + assert "extra_content" not in processed.messages[0] + + +def test_preprocess_messages_omits_assistant_reasoning_without_tools() -> None: + params = CompletionParams( + model_id="deepseek-v4-pro", + messages=[ + { + "role": "assistant", + "content": "hello", + "extra_content": {"deepseek": {"reasoning_content": "greeting"}}, + } + ], + ) + + processed = _preprocess_messages(params) + + assert "reasoning_content" not in processed.messages[0] + assert "extra_content" not in processed.messages[0] + + +def test_reinject_reasoning_content_omits_reasoning_without_tools() -> None: + messages: list[dict[str, Any]] = [ + { + "role": "assistant", + "content": "hello", + "extra_content": {"deepseek": {"reasoning_content": "greeting"}}, + } + ] + + result = _reinject_reasoning_content(messages, replay_reasoning=False) + + assert "reasoning_content" not in result[0] + assert "extra_content" not in result[0] + + def test_reinject_reasoning_content_handles_missing_extra_content() -> None: """A tool-call message with no extra_content should pass through unchanged.""" messages: list[dict[str, Any]] = [ @@ -429,6 +630,6 @@ def test_reinject_reasoning_content_handles_missing_extra_content() -> None: }, ] - result = _reinject_reasoning_content(messages) + result = _reinject_reasoning_content(messages, replay_reasoning=True) assert "reasoning_content" not in result[0] diff --git a/tests/unit/providers/test_gemini_provider.py b/tests/unit/providers/test_gemini_provider.py index 58be7d63a..97b48b341 100644 --- a/tests/unit/providers/test_gemini_provider.py +++ b/tests/unit/providers/test_gemini_provider.py @@ -2,9 +2,10 @@ import json from collections.abc import AsyncIterator from contextlib import contextmanager -from typing import Any, get_args +from typing import Any, cast from unittest.mock import AsyncMock, Mock, patch +import httpx import pytest from google.genai import types from pydantic import BaseModel, ConfigDict @@ -16,7 +17,7 @@ UnsupportedParameterError, ) from any_llm.providers.gemini import GeminiProvider -from any_llm.providers.gemini.base import REASONING_EFFORT_TO_THINKING_BUDGETS, GoogleProvider +from any_llm.providers.gemini.base import GoogleProvider, _convert_reasoning_effort from any_llm.providers.gemini.utils import ( _convert_messages, _convert_response_to_response_dict, @@ -816,45 +817,47 @@ async def test_completion_inside_agent_loop(agent_loop_messages: list[dict[str, assert contents[2].role == "function" -@pytest.mark.parametrize("reasoning_effort", [None, *get_args(ReasoningEffort)]) -@pytest.mark.asyncio -async def test_completion_with_custom_reasoning_effort(reasoning_effort: ReasoningEffort | None) -> None: - api_key = "test-api-key" - model = "model-id" - messages = [{"role": "user", "content": "Hello"}] - - with mock_gemini_provider() as mock_genai: - provider = GeminiProvider(api_key=api_key) - await provider._acompletion( - CompletionParams(model_id=model, messages=messages, reasoning_effort=reasoning_effort) - ) - - _, call_kwargs = mock_genai.return_value.aio.models.generate_content.call_args - thinking_config = call_kwargs["config"].thinking_config - - if reasoning_effort == "auto": - assert thinking_config is None - elif reasoning_effort is None or reasoning_effort == "none": - assert thinking_config == types.ThinkingConfig(include_thoughts=False) - else: - assert thinking_config == types.ThinkingConfig( - include_thoughts=True, thinking_budget=REASONING_EFFORT_TO_THINKING_BUDGETS[reasoning_effort] - ) - - @pytest.mark.parametrize( - ("model_id", "reasoning_effort", "expected_level"), + ("model_id", "reasoning_effort", "expected"), [ - ("gemini-3.5-flash", "xhigh", types.ThinkingLevel.HIGH), - ("gemini-3.5-flash", "max", types.ThinkingLevel.HIGH), - ("gemini-3.5-pro", "low", types.ThinkingLevel.LOW), - ("models/gemini-3.5-flash", "medium", types.ThinkingLevel.MEDIUM), - ("gemini-3.10-flash", "minimal", types.ThinkingLevel.MINIMAL), - ("gemini-4-pro", "high", types.ThinkingLevel.HIGH), + ("gemini-3.8-flash", "low", {"includeThoughts": True, "thinkingLevel": "LOW"}), + ("gemini-3.7-flash", "medium", {"includeThoughts": True, "thinkingLevel": "MEDIUM"}), + ("gemini-3.6-flash", "minimal", {"includeThoughts": True, "thinkingLevel": "MINIMAL"}), + ("gemini-3.5-flash", "high", {"includeThoughts": True, "thinkingLevel": "HIGH"}), + ("gemini-3.5-flash-lite", "minimal", {"includeThoughts": True, "thinkingLevel": "MINIMAL"}), + ("gemini-3.1-flash-lite", "medium", {"includeThoughts": True, "thinkingLevel": "MEDIUM"}), + ("models/gemini-3.1-pro-preview", "minimal", {"includeThoughts": True, "thinkingLevel": "LOW"}), + ("gemini-3.1-flash-image", "minimal", {"includeThoughts": True, "thinkingLevel": "MINIMAL"}), + ("gemini-3.1-flash-lite-image", "high", {"includeThoughts": True, "thinkingLevel": "HIGH"}), + ("gemini-3-flash-preview", "minimal", {"includeThoughts": True, "thinkingLevel": "MINIMAL"}), + ("gemini-2.5-flash", "none", {"thinkingBudget": 0}), + ("gemini-2.5-flash", "minimal", {"includeThoughts": True, "thinkingBudget": 1024}), + ("gemini-2.5-flash", "low", {"includeThoughts": True, "thinkingBudget": 1024}), + ("gemini-2.5-flash-lite", "none", {"thinkingBudget": 0}), + ("gemini-2.5-flash-lite", "minimal", {"includeThoughts": True, "thinkingBudget": 1024}), + ("gemini-2.5-flash-lite", "medium", {"includeThoughts": True, "thinkingBudget": 8192}), + ("gemini-2.5-pro", "minimal", {"includeThoughts": True, "thinkingBudget": 1024}), + ("gemini-2.5-pro", "high", {"includeThoughts": True, "thinkingBudget": 24576}), + ("gemini-2.5-pro", "xhigh", {"includeThoughts": True, "thinkingBudget": 32768}), + ("gemini-2.5-flash", "max", {"includeThoughts": True, "thinkingBudget": 24576}), + ("gemini-2.5-flash-lite", "xhigh", {"includeThoughts": True, "thinkingBudget": 24576}), + ("gemini-3.8-flash", "xhigh", {"includeThoughts": True, "thinkingLevel": "HIGH"}), + ("-001", "minimal", {"includeThoughts": True, "thinkingBudget": 1024}), + ("custom-gemini-model", "minimal", {"includeThoughts": True, "thinkingBudget": 1024}), + ("gemini-3.1-custom", "high", {"includeThoughts": True, "thinkingLevel": "HIGH"}), + ("models/gemini-3.10-flash-preview-202609", "max", {"includeThoughts": True, "thinkingLevel": "HIGH"}), + ("gemini-3.8-flash-custom", "minimal", {"includeThoughts": True, "thinkingLevel": "MINIMAL"}), + ( + "projects/p/locations/l/publishers/google/models/gemini-3.8-flash-001", + "low", + {"includeThoughts": True, "thinkingLevel": "LOW"}, + ), ], ) -def test_new_gemini_models_use_thinking_level( - model_id: str, reasoning_effort: ReasoningEffort, expected_level: types.ThinkingLevel +def test_gemini_reasoning_effort_matches_documented_thinking_config( + model_id: str, + reasoning_effort: ReasoningEffort, + expected: dict[str, object], ) -> None: result = GoogleProvider._convert_completion_params( CompletionParams( @@ -863,29 +866,119 @@ def test_new_gemini_models_use_thinking_level( provider_name="gemini", ) - assert result["config"].thinking_config == types.ThinkingConfig( - include_thoughts=True, thinking_level=expected_level - ) + config = result["config"].model_dump(by_alias=True, exclude_none=True) + assert config["thinkingConfig"] == expected @pytest.mark.parametrize( - "model_id", + ("model_id", "reasoning_effort"), [ - "gemini-3.0-flash", - "gemini-3.4-flash", - "gemini-3-pro-preview", - "gemini-2.5-flash", - "gemini-pro", - "projects/p/locations/l/publishers/google/models/gemini-3-pro", + ("gemini-3.8-flash", "minimal"), + ("gemini-3.1-flash-image", "low"), + ("gemini-3.1-flash-lite-image", "low"), + ("gemini-3.8-flash-001", "minimal"), + ("gemini-3.8-flash", "none"), + ("gemini-3.1-flash-image", "none"), + ("gemini-2.5-pro", "none"), ], ) -def test_older_gemini_models_keep_thinking_budget(model_id: str) -> None: +def test_gemini_rejects_undocumented_reasoning_effort( + model_id: str, + reasoning_effort: ReasoningEffort, +) -> None: + with pytest.raises(UnsupportedParameterError) as exc_info: + GoogleProvider._convert_completion_params( + CompletionParams( + model_id=model_id, + messages=[{"role": "user", "content": "Hello"}], + reasoning_effort=reasoning_effort, + ), + provider_name="gemini", + ) + + assert str(exc_info.value) == ( + "[gemini] 'reasoning_effort' is not supported for gemini.\n" + f"'{reasoning_effort}' is not available for model '{model_id}'." + ) + + +def test_gemini_invalid_reasoning_effort_error_identifies_model_and_effort() -> None: + reasoning_effort = cast("ReasoningEffort", "invalid") + + with pytest.raises(UnsupportedParameterError) as exc_info: + _convert_reasoning_effort("custom-gemini-model", reasoning_effort, "gemini") + + assert str(exc_info.value) == ( + "[gemini] 'reasoning_effort' is not supported for gemini.\n" + "'invalid' is not available for model 'custom-gemini-model'." + ) + + +@pytest.mark.parametrize( + ("reasoning_effort", "expected"), + [(None, None), ("auto", None)], +) +def test_gemini_preserves_default_thinking_config_wire_behavior( + reasoning_effort: ReasoningEffort | None, expected: dict[str, object] | None +) -> None: result = GoogleProvider._convert_completion_params( - CompletionParams(model_id=model_id, messages=[{"role": "user", "content": "Hello"}], reasoning_effort="high"), + CompletionParams( + model_id="gemini-3.8-flash", + messages=[{"role": "user", "content": "Hello"}], + reasoning_effort=reasoning_effort, + ), provider_name="gemini", ) - assert result["config"].thinking_config == types.ThinkingConfig(include_thoughts=True, thinking_budget=24576) + config = result["config"].model_dump(by_alias=True, exclude_none=True) + assert config.get("thinkingConfig") == expected + + +@pytest.mark.parametrize( + ("model_id", "reasoning_effort", "expected"), + [ + ("gemini-3.8-flash", "high", {"include_thoughts": True, "thinking_level": "HIGH"}), + ("gemini-2.5-flash", "none", {"thinking_budget": 0}), + ], +) +@pytest.mark.asyncio +async def test_gemini_reasoning_effort_reaches_official_sdk_wire( + model_id: str, + reasoning_effort: ReasoningEffort, + expected: dict[str, object], +) -> None: + requests: list[dict[str, object]] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + requests.append(json.loads(request.content)) + return httpx.Response( + 200, + json={ + "candidates": [ + { + "content": {"parts": [{"text": "ok"}], "role": "model"}, + "finishReason": "STOP", + } + ] + }, + ) + + provider = GeminiProvider( + api_key="test-api-key", + http_options=types.HttpOptions( + async_client_args={"transport": httpx.MockTransport(handler)}, + ), + ) + await provider._acompletion( + CompletionParams( + model_id=model_id, + messages=[{"role": "user", "content": "Hello"}], + reasoning_effort=reasoning_effort, + ) + ) + await provider.client.aio.aclose() + + assert requests[0]["generationConfig"] == {"thinkingConfig": expected} @pytest.mark.asyncio @@ -1192,6 +1285,168 @@ def test_convert_response_skips_parts_without_text_or_function_call() -> None: message = response_dict["choices"][0]["message"] assert message["content"] == "Described." assert message["tool_calls"] is None + assert message["images"][0]["image_url"]["url"] == "data:image/png;base64,iVBORw==" + + +def test_convert_response_emits_choice_for_image_only_response() -> None: + response = _make_gemini_response( + [types.Part(inline_data=types.Blob(mime_type="image/png", data=b"\x89PNG"))], + types.FinishReason.STOP, + ) + + response_dict = _convert_response_to_response_dict(response) + + assert len(response_dict["choices"]) == 1 + message = response_dict["choices"][0]["message"] + assert message["content"] is None + assert message["images"][0]["image_url"]["url"] == "data:image/png;base64,iVBORw==" + + +@pytest.mark.parametrize( + "blob", + [ + types.Blob(mime_type="image/png", data=b""), + types.Blob(mime_type="image/png", data=None), + types.Blob(mime_type=None, data=b"data"), + types.Blob(mime_type="application/pdf", data=b"data"), + ], +) +def test_convert_response_skips_inline_data_without_image_payload(blob: types.Blob) -> None: + response = _make_gemini_response([types.Part(inline_data=blob)], types.FinishReason.STOP) + + response_dict = _convert_response_to_response_dict(response) + + assert response_dict["choices"] == [] + + +def test_streaming_completion_with_inline_image() -> None: + response = _make_gemini_response( + [types.Part(inline_data=types.Blob(mime_type="image/png", data=b"\x89PNG"))], + types.FinishReason.STOP, + ) + + chunk = _create_openai_chunk_from_google_chunk(response) + + images = chunk.choices[0].delta.images + assert images is not None + assert images[0].image_url.url == "data:image/png;base64,iVBORw==" + + +def test_convert_response_preserves_inline_audio() -> None: + response = _make_gemini_response( + [ + types.Part(text="Here is the audio."), + types.Part(inline_data=types.Blob(mime_type="audio/wav", data=b"WAVE")), + ], + types.FinishReason.STOP, + ) + + result = GoogleProvider._convert_completion_response( + (_convert_response_to_response_dict(response), "gemini-2.5-flash") + ) + + assert result.choices[0].message.audio is not None + assert result.choices[0].message.audio.data == "V0FWRQ==" + assert result.choices[0].message.audio.transcript == "Here is the audio." + + +def test_convert_response_emits_choice_for_audio_only_response() -> None: + response = _make_gemini_response( + [types.Part(inline_data=types.Blob(mime_type="audio/wav", data=b"WAVE"))], + types.FinishReason.STOP, + ) + + response_dict = _convert_response_to_response_dict(response) + + assert len(response_dict["choices"]) == 1 + message = response_dict["choices"][0]["message"] + assert message["content"] is None + assert message["audio"]["data"] == "V0FWRQ==" + + +def test_streaming_completion_with_inline_audio() -> None: + response = _make_gemini_response( + [ + types.Part(text="Here is the audio."), + types.Part(inline_data=types.Blob(mime_type="audio/wav", data=b"WAVE")), + ], + types.FinishReason.STOP, + ) + + chunk = _create_openai_chunk_from_google_chunk(response) + + assert chunk.choices[0].delta.audio is not None + assert chunk.choices[0].delta.audio.data == "V0FWRQ==" + assert chunk.choices[0].delta.audio.transcript == "Here is the audio." + + +def test_streaming_completion_emits_choice_for_audio_only_response() -> None: + response = _make_gemini_response( + [types.Part(inline_data=types.Blob(mime_type="audio/wav", data=b"WAVE"))], + types.FinishReason.STOP, + ) + + chunk = _create_openai_chunk_from_google_chunk(response) + + assert chunk.choices[0].delta.audio is not None + assert chunk.choices[0].delta.audio.data == "V0FWRQ==" + + +def _wav_fields(wav: bytes) -> tuple[bytes, int, int, bytes]: + """Return the RIFF tag, sample rate, data length, and payload from a WAV file.""" + return wav[:4], int.from_bytes(wav[24:28], "little"), int.from_bytes(wav[40:44], "little"), wav[44:] + + +@pytest.mark.parametrize("rate", [24000, 16000]) +def test_convert_response_wraps_pcm_audio_as_wav(rate: int) -> None: + """Complete Gemini TTS responses expose all PCM parts as a playable WAV.""" + mime_type = f"audio/L16;codec=pcm;rate={rate}" + response = _make_gemini_response( + [ + types.Part(inline_data=types.Blob(mime_type=mime_type, data=b"\x01\x02")), + types.Part(inline_data=types.Blob(mime_type=mime_type, data=b"\x03\x04")), + ], + types.FinishReason.STOP, + ) + + result = GoogleProvider._convert_completion_response( + (_convert_response_to_response_dict(response), "gemini-2.5-flash-preview-tts") + ) + + audio = result.choices[0].message.audio + assert audio is not None + assert _wav_fields(base64.b64decode(audio.data)) == (b"RIFF", rate, 4, b"\x01\x02\x03\x04") + assert result.choices[0].message.content is None + assert result.choices[0].finish_reason == "stop" + + +def test_streaming_completion_keeps_pcm_audio_raw() -> None: + """Streaming Gemini TTS parts stay raw and are joined in their original order.""" + mime_type = "audio/L16;codec=pcm;rate=24000" + response = _make_gemini_response( + [ + types.Part(inline_data=types.Blob(mime_type=mime_type, data=b"\x01\x02")), + types.Part(inline_data=types.Blob(mime_type=mime_type, data=b"\x03\x04")), + ], + types.FinishReason.STOP, + ) + + chunk = _create_openai_chunk_from_google_chunk(response) + + audio = chunk.choices[0].delta.audio + assert audio is not None + assert audio.data is not None + assert base64.b64decode(audio.data) == b"\x01\x02\x03\x04" + + +def test_convert_response_skips_empty_audio_blob() -> None: + """Empty inline audio must not create an otherwise empty completion choice.""" + response = _make_gemini_response( + [types.Part(inline_data=types.Blob(mime_type="audio/L16;codec=pcm;rate=24000", data=b""))], + types.FinishReason.STOP, + ) + + assert _convert_response_to_response_dict(response)["choices"] == [] def test_convert_response_emits_choice_for_filtered_response_without_content() -> None: @@ -1241,6 +1496,20 @@ def test_google_provider_preserves_prompt_block_as_refusal() -> None: assert result.choices[0].message.refusal == "Response blocked by Gemini content filtering." +def test_google_provider_preserves_images_on_completion_message() -> None: + response_dict = _convert_response_to_response_dict( + _make_gemini_response( + [types.Part(inline_data=types.Blob(mime_type="image/png", data=b"\x89PNG"))], + types.FinishReason.STOP, + ) + ) + + result = GoogleProvider._convert_completion_response((response_dict, "gemini-2.5-flash-image")) + + assert result.choices[0].message.images is not None + assert result.choices[0].message.images[0].image_url.url == "data:image/png;base64,iVBORw==" + + def test_convert_response_without_content_and_terminal_reason_has_no_choices() -> None: response_dict = _convert_response_to_response_dict(_make_gemini_response(None, types.FinishReason.STOP)) diff --git a/tests/unit/providers/test_ollama_provider.py b/tests/unit/providers/test_ollama_provider.py index 7dbb033e9..c64741c5a 100644 --- a/tests/unit/providers/test_ollama_provider.py +++ b/tests/unit/providers/test_ollama_provider.py @@ -1,8 +1,10 @@ +import json import logging from collections.abc import AsyncIterator from typing import Any from unittest.mock import AsyncMock, Mock, patch +import httpx import pytest from ollama import ChatResponse as OllamaChatResponse from ollama import Message as OllamaMessage @@ -12,7 +14,7 @@ _create_chat_completion_from_ollama_response, _create_openai_chunk_from_ollama_chunk, ) -from any_llm.types.completion import CompletionParams +from any_llm.types.completion import ChatCompletion, CompletionParams @pytest.mark.asyncio @@ -458,6 +460,7 @@ def _make_chunk(name: str, arguments: dict[str, str]) -> Mock: chunk.message = message chunk.created_at = None chunk.model = "llama3.1" + chunk.done = False chunk.done_reason = None chunk.prompt_eval_count = None chunk.eval_count = None @@ -780,3 +783,92 @@ async def test_tool_call_arguments_are_converted_to_a_mapping(arguments: Any, ex ) assert sent[1]["tool_calls"] == [{"function": {"name": "get_weather", "arguments": expected}}] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize( + ("reason_fields", "expected"), + [ + ({}, "stop"), + ({"done_reason": None}, "stop"), + ({"done_reason": "load"}, "stop"), + ({"done_reason": "unload"}, "stop"), + ({"done_reason": "stop"}, "stop"), + ({"done_reason": "length"}, "length"), + ({"done_reason": "future_reason"}, "stop"), + ({"done_reason": "tool_calls"}, "tool_calls"), + ({"done_reason": "content_filter"}, "content_filter"), + ({"done_reason": "function_call"}, "function_call"), + ], +) +async def test_completion_normalizes_sdk_done_reason( + stream: bool, reason_fields: dict[str, str | None], expected: str +) -> None: + """Exercise JSON and NDJSON decoding through the real Ollama SDK and provider.""" + terminal = { + "model": "llama3.2", + "created_at": "2024-09-12T21:17:29.110811Z", + "message": {"role": "assistant", "content": ""}, + "done": True, + **reason_fields, + } + + def handle(request: httpx.Request) -> httpx.Response: + assert request.url.path == "/api/chat" + assert json.loads(request.content)["stream"] is stream + if stream: + partial = {**terminal, "done": False, "message": {"role": "assistant", "content": "Hello"}} + partial.pop("done_reason", None) + return httpx.Response( + 200, + content="\n".join(json.dumps(item) for item in [partial, terminal]) + "\n", + headers={"content-type": "application/x-ndjson"}, + ) + return httpx.Response(200, json=terminal) + + provider = OllamaProvider(api_base="http://ollama.test", transport=httpx.MockTransport(handle)) + try: + result = await provider.acompletion( + model="llama3.2", messages=[{"role": "user", "content": "Hello"}], stream=stream + ) + if stream: + assert isinstance(result, AsyncIterator) + chunks = [chunk async for chunk in result] + assert len(chunks) == 2 + assert chunks[0].choices[0].delta.content == "Hello" + assert chunks[0].choices[0].finish_reason is None + assert chunks[1].choices[0].delta.content == "" + assert chunks[1].choices[0].finish_reason == expected + else: + assert isinstance(result, ChatCompletion) + assert result.choices[0].message.content == "" + assert result.choices[0].finish_reason == expected + finally: + await provider.client._client.aclose() + + +@pytest.mark.parametrize("done_reason", [None, "load", "unload", "stop", "length"]) +def test_completion_normalization_preserves_tool_call_precedence(done_reason: str | None) -> None: + response = OllamaChatResponse( + model="llama3.2", + created_at="2024-09-12T21:17:29.110811Z", + done=True, + done_reason=done_reason, + message=OllamaMessage( + role="assistant", + content="", + tool_calls=[ + OllamaMessage.ToolCall( + function=OllamaMessage.ToolCall.Function(name="get_weather", arguments={"city": "Paris"}) + ) + ], + ), + ) + completion = _create_chat_completion_from_ollama_response(response) + assert completion.choices[0].finish_reason == "tool_calls" + calls = completion.choices[0].message.tool_calls + assert calls is not None + assert calls[0].type == "function" + assert calls[0].function.name == "get_weather" + assert json.loads(calls[0].function.arguments) == {"city": "Paris"} diff --git a/tests/unit/providers/test_openai_utils.py b/tests/unit/providers/test_openai_utils.py index d0db0b6a4..f58c8ea5f 100644 --- a/tests/unit/providers/test_openai_utils.py +++ b/tests/unit/providers/test_openai_utils.py @@ -166,3 +166,30 @@ def test_convert_chunk_response_with_nonstandard_service_tier() -> None: ) result = BaseOpenAIProvider._convert_completion_chunk_response(openai_chunk) assert result.service_tier == "standard" + + +@pytest.mark.parametrize( + "audio_piece", + [ + {"id": "audio_6a98e62e", "transcript": "Sure"}, + {"data": "AAAB//8CAP=="}, + {"expires_at": 1788408898}, + ], +) +def test_convert_chunk_response_keeps_partial_audio_delta(audio_piece: dict[str, object]) -> None: + """OpenAI streams delta audio fields in separate chunks.""" + openai_chunk = OpenAIChatCompletionChunk.model_validate( + { + "id": "chatcmpl-audio", + "object": "chat.completion.chunk", + "created": 0, + "model": "gpt-audio-mini", + "choices": [{"index": 0, "delta": {"role": "assistant", "audio": audio_piece}, "finish_reason": None}], + } + ) + + result = BaseOpenAIProvider._convert_completion_chunk_response(openai_chunk) + + audio = result.choices[0].delta.audio + assert audio is not None + assert audio.model_dump(exclude_none=True) == audio_piece diff --git a/tests/unit/providers/test_otari_provider.py b/tests/unit/providers/test_otari_provider.py index d2450ac22..22a35114a 100644 --- a/tests/unit/providers/test_otari_provider.py +++ b/tests/unit/providers/test_otari_provider.py @@ -7,6 +7,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from anthropic import transform_schema from pydantic import BaseModel from any_llm.exceptions import BatchNotCompleteError, UnsupportedParameterError @@ -958,72 +959,151 @@ async def test_otari_amessages_streaming_forwards_anthropic_beta_params() -> Non @pytest.mark.asyncio -async def test_otari_amessages_output_format_falls_back_to_bridge() -> None: - """With output_format set, otari delegates to the Completions bridge, not /messages.""" - completion = ChatCompletion.model_validate( - { - "id": "chatcmpl-test", - "object": "chat.completion", - "created": 0, - "model": "claude-sonnet-4-5", - "choices": [ - {"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": '{"city": "Paris"}'}} - ], - } - ) - +async def test_otari_amessages_streaming_output_format_uses_native_messages() -> None: class City(BaseModel): city: str client = _mock_otari_client() + client.with_response_metadata.message.return_value = _MockMetadataStream([{"type": "message_stop"}]) provider = _build_provider(client) - provider._acompletion = AsyncMock(return_value=completion) # type: ignore[method-assign] + context_management = {"edits": [{"type": "compact_20260112"}]} + betas = ["compact-2026-01-12"] - params = MessagesParams( + result = await provider.amessages( model="claude-sonnet-4-5", messages=[{"role": "user", "content": "Capital of France?"}], max_tokens=100, output_format=City, + stream=True, + context_management=context_management, + betas=betas, + cache_control={"type": "ephemeral"}, ) + assert not isinstance(result, (MessageResponse, ParsedMessage, ParsedBetaMessage)) + collected = [event async for event in result] - result = await provider._amessages(params) - - assert isinstance(result, MessageResponse) - provider._acompletion.assert_called_once() - client.message.assert_not_called() + assert len(collected) == 1 + assert isinstance(collected[0], MessageStopEvent) + call_kwargs = client.message.call_args.kwargs + assert call_kwargs["stream"] is True + assert call_kwargs["output_format"] == { + "format": {"type": "json_schema", "schema": transform_schema(City.model_json_schema())} + } + assert call_kwargs["context_management"] == context_management + assert call_kwargs["betas"] == betas + assert call_kwargs["cache_control"] == {"type": "ephemeral"} @pytest.mark.asyncio -@pytest.mark.parametrize( - "beta_params", - [ - {"context_management": {"edits": [{"type": "compact_20260112"}]}}, - {"betas": ["compact-2026-01-12"]}, - ], -) -async def test_otari_amessages_output_format_with_beta_params_raises(beta_params: dict[str, Any]) -> None: - """output_format routes through the bridge, which cannot carry context_management or betas.""" +async def test_otari_amessages_streaming_bare_output_format_uses_native_messages() -> None: + client = _mock_otari_client() + client.with_response_metadata.message.return_value = _MockMetadataStream([{"type": "message_stop"}]) + provider = _build_provider(client) + provider._acompletion = AsyncMock() # type: ignore[method-assign] + output_format = {"type": "json_schema", "schema": {"type": "object"}} + context_management = {"edits": [{"type": "compact_20260112"}]} + betas = ["compact-2026-01-12"] + + result = await provider.amessages( + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "Capital of France?"}], + max_tokens=100, + output_format=output_format, + stream=True, + context_management=context_management, + betas=betas, + cache_control={"type": "ephemeral"}, + ) + assert not isinstance(result, (MessageResponse, ParsedMessage, ParsedBetaMessage)) + collected = [event async for event in result] + assert len(collected) == 1 + assert isinstance(collected[0], MessageStopEvent) + call_kwargs = client.message.call_args.kwargs + assert call_kwargs["stream"] is True + assert call_kwargs["output_format"] == {"format": output_format} + assert call_kwargs["context_management"] == context_management + assert call_kwargs["betas"] == betas + assert call_kwargs["cache_control"] == {"type": "ephemeral"} + provider._acompletion.assert_not_called() + client.completion.assert_not_called() + + +@pytest.mark.asyncio +async def test_otari_amessages_output_format_uses_native_messages_with_anthropic_fields() -> None: class City(BaseModel): city: str client = _mock_otari_client() + payload = _message_response_payload() + payload["content"] = [{"type": "text", "text": '{"city": "Paris"}'}] + client.message.return_value = SimpleNamespace(data=payload, request_id=None) provider = _build_provider(client) provider._acompletion = AsyncMock() # type: ignore[method-assign] + context_management = {"edits": [{"type": "compact_20260112"}]} + betas = ["compact-2026-01-12"] params = MessagesParams( model="claude-sonnet-4-5", messages=[{"role": "user", "content": "Capital of France?"}], max_tokens=100, output_format=City, - **beta_params, + context_management=context_management, + betas=betas, + cache_control={"type": "ephemeral"}, ) - with pytest.raises(NotImplementedError, match="output_format cannot be combined"): - await provider._amessages(params) + result = await provider._amessages(params) + assert isinstance(result, MessageResponse) provider._acompletion.assert_not_called() - client.message.assert_not_called() + call_kwargs = client.message.call_args.kwargs + assert call_kwargs["output_format"] == { + "format": {"type": "json_schema", "schema": transform_schema(City.model_json_schema())} + } + assert call_kwargs["context_management"] == context_management + assert call_kwargs["betas"] == betas + assert call_kwargs["cache_control"] == {"type": "ephemeral"} + + +@pytest.mark.asyncio +async def test_otari_amessages_bare_output_format_uses_native_messages() -> None: + client = _mock_otari_client() + client.message.return_value = SimpleNamespace(data=_message_response_payload(), request_id=None) + provider = _build_provider(client) + output_format = {"type": "json_schema", "schema": {"type": "object"}} + + params = MessagesParams( + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "Capital of France?"}], + max_tokens=100, + output_format=output_format, + ) + + await provider._amessages(params) + + assert client.message.call_args.kwargs["output_format"] == {"format": output_format} + client.completion.assert_not_called() + + +@pytest.mark.asyncio +async def test_otari_amessages_schema_less_output_config_returns_plain_message() -> None: + client = _mock_otari_client() + client.message.return_value = SimpleNamespace(data=_message_response_payload(), request_id=None) + provider = _build_provider(client) + + result = await provider.amessages( + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "Hello"}], + max_tokens=100, + output_format={"effort": "high"}, + ) + + assert isinstance(result, MessageResponse) + block = result.content[0] + assert isinstance(block, TextBlock) + assert block.text == "hi" + assert client.message.call_args.kwargs["output_format"] == {"effort": "high"} @pytest.mark.asyncio diff --git a/tests/unit/test_agent_loop_helpers.py b/tests/unit/test_agent_loop_helpers.py index 2f5b2e1d7..0d99b2c7f 100644 --- a/tests/unit/test_agent_loop_helpers.py +++ b/tests/unit/test_agent_loop_helpers.py @@ -1,6 +1,54 @@ +import json +from typing import TYPE_CHECKING, Any +from unittest.mock import AsyncMock, MagicMock + import pytest -from tests.integration.test_agent_loop import _call_tool, get_current_date, get_weather +from any_llm import AnyLLM +from any_llm.types.completion import ChatCompletion, ChatCompletionMessage +from tests.integration.test_agent_loop import ( + _call_tool, + _run_agent_loop, + get_current_date, + get_weather, +) + +if TYPE_CHECKING: + from collections.abc import Callable + + +def _completion( + *, + content: str | None = None, + tool_calls: list[tuple[str, dict[str, Any]]] | None = None, +) -> ChatCompletion: + serialized_tool_calls = [ + { + "id": f"call-{index}", + "type": "function", + "function": {"name": name, "arguments": json.dumps(arguments)}, + } + for index, (name, arguments) in enumerate(tool_calls or []) + ] + return ChatCompletion.model_validate( + { + "id": "test-completion", + "object": "chat.completion", + "created": 0, + "model": "test-model", + "choices": [ + { + "index": 0, + "finish_reason": "tool_calls" if serialized_tool_calls else "stop", + "message": { + "role": "assistant", + "content": content, + "tool_calls": serialized_tool_calls if tool_calls is not None else None, + }, + } + ], + } + ) def test_call_tool_ignores_spurious_model_arguments_for_zero_arg_tool() -> None: @@ -13,3 +61,114 @@ def test_call_tool_preserves_declared_tool_arguments() -> None: """Preserve arguments declared by a parameterized tool while filtering extras.""" with pytest.warns(UserWarning, match="Ignoring unexpected arguments for get_weather: result"): assert "Paris" in _call_tool(get_weather, {"location": "Paris", "result": "unexpected"}) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("empty_tool_calls", [None, []]) +async def test_run_agent_loop_continues_sequential_calls_before_answering( + empty_tool_calls: list[tuple[str, dict[str, Any]]] | None, +) -> None: + client = MagicMock(spec=AnyLLM) + client.acompletion = AsyncMock( + side_effect=[ + _completion(tool_calls=[("get_weather", {"location": "Paris"})]), + _completion(tool_calls=[("get_weather", {"location": "London"})]), + _completion(content="Paris and London are sunny at 15C.", tool_calls=empty_tool_calls), + ] + ) + messages: list[dict[str, Any] | ChatCompletionMessage] = [{"role": "user", "content": "Weather?"}] + + def called_both_locations(calls: list[tuple[str, dict[str, Any]]]) -> bool: + return {arguments.get("location") for _, arguments in calls} >= {"Paris", "London"} + + message, calls = await _run_agent_loop( + client, + "test-model", + messages, + {"get_weather": get_weather}, + called_both_locations, + include_tool_name=False, + ) + + assert message.content == "Paris and London are sunny at 15C." + assert calls == [ + ("get_weather", {"location": "Paris"}), + ("get_weather", {"location": "London"}), + ] + assert client.acompletion.await_count == 3 + assert all("tool_choice" not in call.kwargs for call in client.acompletion.await_args_list) + tool_messages = [item for item in messages if isinstance(item, dict) and item.get("role") == "tool"] + assert all("name" not in tool_message for tool_message in tool_messages) + + +@pytest.mark.asyncio +async def test_run_agent_loop_can_include_tool_names() -> None: + client = MagicMock(spec=AnyLLM) + client.acompletion = AsyncMock( + side_effect=[ + _completion( + tool_calls=[ + ("get_current_date", {}), + ("get_weather", {"location": "Paris"}), + ] + ), + _completion(content="Paris is sunny at 15C."), + ] + ) + messages: list[dict[str, Any] | ChatCompletionMessage] = [{"role": "user", "content": "Weather?"}] + available_tools: dict[str, Callable[..., str]] = { + "get_current_date": get_current_date, + "get_weather": get_weather, + } + + message, _ = await _run_agent_loop( + client, + "test-model", + messages, + available_tools, + lambda calls: {name for name, _ in calls} >= set(available_tools), + include_tool_name=True, + ) + + assert message.content == "Paris is sunny at 15C." + assert client.acompletion.await_count == 2 + tool_messages = [item for item in messages if isinstance(item, dict) and item.get("role") == "tool"] + assert [tool_message["name"] for tool_message in tool_messages] == ["get_current_date", "get_weather"] + + +@pytest.mark.asyncio +async def test_run_agent_loop_rejects_answer_before_required_calls() -> None: + client = MagicMock(spec=AnyLLM) + client.acompletion = AsyncMock(side_effect=[_completion(content="No tools needed.")]) + + with pytest.raises(AssertionError, match="answered before making the required tool calls"): + await _run_agent_loop( + client, + "test-model", + [{"role": "user", "content": "Weather?"}], + {"get_weather": get_weather}, + bool, + include_tool_name=False, + ) + + +@pytest.mark.asyncio +async def test_run_agent_loop_rejects_repeated_calls_at_iteration_limit() -> None: + client = MagicMock(spec=AnyLLM) + client.acompletion = AsyncMock( + side_effect=[ + _completion(tool_calls=[("get_weather", {"location": "Paris"})]), + _completion(tool_calls=[("get_weather", {"location": "Paris"})]), + ] + ) + + with pytest.raises(AssertionError, match="did not answer within 2 iterations"): + await _run_agent_loop( + client, + "test-model", + [{"role": "user", "content": "Weather?"}], + {"get_weather": get_weather}, + lambda calls: any(arguments.get("location") == "London" for _, arguments in calls), + include_tool_name=False, + max_iterations=2, + ) diff --git a/tests/unit/test_exception_handler.py b/tests/unit/test_exception_handler.py index ab157457b..27ccc5116 100644 --- a/tests/unit/test_exception_handler.py +++ b/tests/unit/test_exception_handler.py @@ -175,6 +175,14 @@ def test_convert_exception_carries_status_code_from_response() -> None: assert result.status_code == 400 +def test_convert_exception_carries_status_from_aiohttp_style_response() -> None: + """aiohttp responses spell it ``status``, which google-genai attaches on its async path.""" + error = _ResponseStatusError(400, "Invalid request") + del error.response.status_code # type: ignore[attr-defined] + error.response.status = 400 # type: ignore[attr-defined] + assert convert_exception(error, "openai").status_code == 400 + + def test_convert_exception_prefers_status_code_attribute_over_response() -> None: error = _ResponseStatusError(500, "Invalid request") error.status_code = 400 # type: ignore[attr-defined]