diff --git a/README.md b/README.md index 5938de75..d84e5866 100644 --- a/README.md +++ b/README.md @@ -12,12 +12,14 @@ uv add ai ``` AI Gateway API-key usage works with the base package. Direct providers that -use an OpenAI-compatible or Anthropic-compatible adapter load the corresponding -official SDK lazily. Vercel OIDC for AI Gateway also uses an optional extra: +use an OpenAI-compatible, Anthropic-compatible, or Google adapter load the +corresponding official SDK lazily. Vercel OIDC for AI Gateway also uses an +optional extra: ```bash uv add "ai[openai]" # OpenAI-compatible providers uv add "ai[anthropic]" # Anthropic-compatible providers +uv add "ai[google]" # Google Gemini uv add "ai[vercel]" # Vercel OIDC for AI Gateway ``` @@ -72,12 +74,14 @@ model = ai.get_model("openai/gpt-5.4") # provider omitted: defaults to gateway model = ai.get_model("gateway:openai/gpt-5.4") model = ai.get_model("openai:gpt-5.4") model = ai.get_model("anthropic:claude-sonnet-4-6") +model = ai.get_model("google:gemini-2.5-flash") ``` Provider IDs without a `provider:` prefix route through AI Gateway by default. Direct OpenAI-compatible providers, including `openai:` and compatible models.dev provider IDs, require `ai[openai]`. Direct Anthropic-compatible -providers require `ai[anthropic]`. +providers require `ai[anthropic]`. The direct Google provider requires +`ai[google]`. Structured output: diff --git a/docs/ai-python/content/docs/basics/providers.mdx b/docs/ai-python/content/docs/basics/providers.mdx index 8fe5b172..f56bafd0 100644 --- a/docs/ai-python/content/docs/basics/providers.mdx +++ b/docs/ai-python/content/docs/basics/providers.mdx @@ -41,6 +41,7 @@ Providers read their provider-specific keys: export AI_GATEWAY_API_KEY="your_access_token_here" export OPENAI_API_KEY="your_access_token_here" export ANTHROPIC_API_KEY="your_access_token_here" +export GEMINI_API_KEY="your_access_token_here" ``` ## Use an explicit provider diff --git a/docs/ai-python/content/docs/reference/ai.providers/google/index.mdx b/docs/ai-python/content/docs/reference/ai.providers/google/index.mdx new file mode 100644 index 00000000..86980e2d --- /dev/null +++ b/docs/ai-python/content/docs/reference/ai.providers/google/index.mdx @@ -0,0 +1,53 @@ +--- +title: "ai.providers.google" +description: Reference for Google provider APIs. +type: reference +summary: Reference for ai.providers.google. +--- + +`ai.providers.google` contains the Google Gemini provider, protocol, and +provider-executed tools. + +```python +provider = ai.get_provider("google") +model = ai.Model(id="gemini-2.5-flash", provider=provider) +``` + +The optional upstream `google-genai` SDK loads lazily when the provider +creates or uses an SDK client. + +## GoogleProvider + +`GoogleProvider` implements `Provider` for the Google Gemini API. + +Default configuration for the `google` provider uses: + +- `GOOGLE_GENERATIVE_AI_API_KEY`, falling back to `GEMINI_API_KEY` and + `GOOGLE_API_KEY` +- `GOOGLE_GEMINI_BASE_URL` + +Pass `base_url`, `api_key`, or a custom client through `get_provider` when you +need explicit configuration. + +```python +provider = ai.get_provider( + "google", + base_url="https://gemini.example.com", +) +model = ai.Model(id="gemini-2.5-flash", provider=provider) +``` + +The provider supports model listing, probing, streaming, provider-executed +tools, and custom `google.genai.Client` instances. + +## GoogleGenerateContentProtocol + +`GoogleGenerateContentProtocol` translates SDK messages and params to the +Google generateContent API wire format. + +The provider uses this protocol by default. Use it directly only when you need +a protocol override. + +## Tools + +Provider-executed Google tools are documented on the child `tools` page. diff --git a/docs/ai-python/content/docs/reference/ai.providers/google/meta.json b/docs/ai-python/content/docs/reference/ai.providers/google/meta.json new file mode 100644 index 00000000..6c750497 --- /dev/null +++ b/docs/ai-python/content/docs/reference/ai.providers/google/meta.json @@ -0,0 +1,7 @@ +{ + "title": "google", + "description": "Reference for Google provider APIs.", + "pages": [ + "tools" + ] +} diff --git a/docs/ai-python/content/docs/reference/ai.providers/google/tools.mdx b/docs/ai-python/content/docs/reference/ai.providers/google/tools.mdx new file mode 100644 index 00000000..01b991f1 --- /dev/null +++ b/docs/ai-python/content/docs/reference/ai.providers/google/tools.mdx @@ -0,0 +1,14 @@ +--- +title: "tools" +description: Reference for Google provider-executed tools. +type: reference +summary: Reference for Google provider tool helpers. +--- + +Google tool helpers create provider-executed `ai.Tool` declarations. + +## Tools + +- `google_search` +- `url_context` +- `code_execution` diff --git a/docs/ai-python/content/docs/reference/ai.providers/index.mdx b/docs/ai-python/content/docs/reference/ai.providers/index.mdx index 446fc68f..b917921f 100644 --- a/docs/ai-python/content/docs/reference/ai.providers/index.mdx +++ b/docs/ai-python/content/docs/reference/ai.providers/index.mdx @@ -16,6 +16,7 @@ provider-specific namespaces. - `ProviderProtocol` - `OpenAICompatibleProvider` - `AnthropicCompatibleProvider` +- `GoogleProvider` - `GatewayProvider` - `history_utils` diff --git a/docs/ai-python/content/docs/reference/ai.providers/meta.json b/docs/ai-python/content/docs/reference/ai.providers/meta.json index fdac8296..812c5b59 100644 --- a/docs/ai-python/content/docs/reference/ai.providers/meta.json +++ b/docs/ai-python/content/docs/reference/ai.providers/meta.json @@ -4,6 +4,7 @@ "pages": [ "ai-gateway", "openai", - "anthropic" + "anthropic", + "google" ] } diff --git a/docs/ai-python/content/docs/reference/ai/get-provider.mdx b/docs/ai-python/content/docs/reference/ai/get-provider.mdx index f6de5162..8832c353 100644 --- a/docs/ai-python/content/docs/reference/ai/get-provider.mdx +++ b/docs/ai-python/content/docs/reference/ai/get-provider.mdx @@ -16,5 +16,5 @@ provider = ai.get_provider( ) ``` -Known providers include AI Gateway, OpenAI-compatible providers, and -Anthropic-compatible providers. +Known providers include AI Gateway, OpenAI-compatible providers, +Anthropic-compatible providers, and the Google Gemini provider. diff --git a/pyproject.toml b/pyproject.toml index c8204c64..6506a0b8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -35,6 +35,7 @@ dependencies = [ [project.optional-dependencies] anthropic = ["anthropic>=0.83.0"] +google = ["google-genai>=2.0.0"] mcp = ["mcp>=1.18.0"] openai = ["openai>=2.14.0"] otel = [ @@ -65,6 +66,7 @@ bump = true [dependency-groups] dev = [ "anthropic>=0.83.0", + "google-genai>=2.0.0", "mcp>=1.18.0", "python-dotenv>=1.2.1", "pytest>=8.0", diff --git a/src/ai/providers/__init__.py b/src/ai/providers/__init__.py index 2af8d02f..74488624 100644 --- a/src/ai/providers/__init__.py +++ b/src/ai/providers/__init__.py @@ -4,11 +4,13 @@ from .ai_gateway import GatewayProvider from .anthropic import AnthropicCompatibleProvider from .base import Provider, ProviderProtocol, get_provider +from .google import GoogleProvider from .openai import OpenAICompatibleProvider __all__ = [ "AnthropicCompatibleProvider", "GatewayProvider", + "GoogleProvider", "OpenAICompatibleProvider", "Provider", "ProviderProtocol", diff --git a/src/ai/providers/google/__init__.py b/src/ai/providers/google/__init__.py new file mode 100644 index 00000000..70cc7311 --- /dev/null +++ b/src/ai/providers/google/__init__.py @@ -0,0 +1,28 @@ +"""Google provider. + +Usage:: + + import ai + from ai.providers.google import tools as google_tools + + model = ai.get_model("google:gemini-2.5-flash") + provider = ai.get_provider("google", api_key="...") + model = ai.Model(id="gemini-2.5-flash", provider=provider) + ids = await ai.get_provider("google").list_models() + + # built-in tools + async with ai.stream( + model, msgs, + tools=[google_tools.google_search()], + ) as s: + ... + +The optional upstream ``google-genai`` SDK is loaded lazily when the +provider creates or uses an SDK client. +""" + +from . import tools +from .protocol import GoogleGenerateContentProtocol +from .provider import GoogleProvider + +__all__ = ["GoogleGenerateContentProtocol", "GoogleProvider", "tools"] diff --git a/src/ai/providers/google/_sdk.py b/src/ai/providers/google/_sdk.py new file mode 100644 index 00000000..83235014 --- /dev/null +++ b/src/ai/providers/google/_sdk.py @@ -0,0 +1,43 @@ +"""Lazy Google GenAI SDK imports.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Protocol, cast + +from .. import _optional + +if TYPE_CHECKING: + import google.genai as genai + from google.genai import errors as genai_errors + + +class GoogleSDK(Protocol): + Client: type[genai.Client] + + +class GoogleErrors(Protocol): + APIError: type[genai_errors.APIError] + ClientError: type[genai_errors.ClientError] + ServerError: type[genai_errors.ServerError] + + +def import_sdk(*, provider: str = "google") -> GoogleSDK: + return cast( + "GoogleSDK", + _optional.import_optional_sdk( + "google.genai", + provider=provider, + extra="google", + ), + ) + + +def import_errors(*, provider: str = "google") -> GoogleErrors: + return cast( + "GoogleErrors", + _optional.import_optional_sdk( + "google.genai.errors", + provider=provider, + extra="google", + ), + ) diff --git a/src/ai/providers/google/errors.py b/src/ai/providers/google/errors.py new file mode 100644 index 00000000..b757a35a --- /dev/null +++ b/src/ai/providers/google/errors.py @@ -0,0 +1,111 @@ +"""Google GenAI SDK error mapping.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any, cast + +import httpx + +from ... import errors as ai_errors + +if TYPE_CHECKING: + from google.genai import errors as genai_errors + + +def map_error( + exc: genai_errors.APIError, + *, + provider: str | None = None, + model_id: str | None = None, +) -> ai_errors.ProviderAPIError: + """Map a Google GenAI SDK exception to the public provider hierarchy.""" + status_code = exc.code if isinstance(exc.code, int) else None + if status_code == 404 and model_id is not None: + cls: type[ai_errors.ProviderAPIError] = ( + ai_errors.ProviderModelNotFoundError + ) + elif status_code is not None: + cls = ai_errors.http_status_to_provider_status_error_class(status_code) + else: + cls = ai_errors.ProviderAPIError + return _provider_error(cls, exc, provider=provider, model_id=model_id) + + +def _provider_error( + cls: type[ai_errors.ProviderAPIError], + exc: genai_errors.APIError, + *, + provider: str | None, + model_id: str | None, +) -> ai_errors.ProviderAPIError: + body = exc.details + if issubclass(cls, ai_errors.ProviderModelNotFoundError): + if model_id is None: # pragma: no cover - guarded by map_error + raise RuntimeError( + "model_id is required for ProviderModelNotFoundError" + ) + return cls( + _message(exc), + model_id=model_id, + provider=provider, + http_context=_http_context(exc), + body=body, + error_type=exc.status, + ) + return cls( + _message(exc), + provider=provider, + http_context=_http_context(exc), + body=body, + error_type=exc.status, + ) + + +def map_httpx_error( + exc: httpx.HTTPError, + *, + provider: str | None = None, +) -> ai_errors.ProviderAPIError: + """Map a raw httpx transport error to the public provider hierarchy. + + The Google GenAI SDK does not wrap transport failures, so connection + and timeout errors surface as bare httpx exceptions. + """ + cls: type[ai_errors.ProviderAPIError] = ( + ai_errors.ProviderTimeoutError + if isinstance(exc, httpx.TimeoutException) + else ai_errors.ProviderConnectionError + ) + return cls( + str(exc) or type(exc).__name__, + provider=provider, + is_retryable=True, + ) + + +def _http_context( + exc: genai_errors.APIError, +) -> ai_errors.HTTPErrorContext | None: + if not isinstance(exc.code, int): + return None + response = cast("Any", getattr(exc, "response", None)) + if not isinstance(response, httpx.Response): + return ai_errors.HTTPErrorContext(status_code=exc.code) + try: + request = response.request + except RuntimeError: + request = None + return ai_errors.HTTPErrorContext( + status_code=exc.code, + request=request, + response=response, + ) + + +def _message(exc: genai_errors.APIError) -> str: + if isinstance(exc.message, str) and exc.message: + return exc.message + return str(exc) + + +__all__ = ["map_error", "map_httpx_error"] diff --git a/src/ai/providers/google/protocol.py b/src/ai/providers/google/protocol.py new file mode 100644 index 00000000..33cce148 --- /dev/null +++ b/src/ai/providers/google/protocol.py @@ -0,0 +1,776 @@ +"""Google protocol — generateContent API. + +Message/tool conversion and streaming via the official ``google-genai`` +SDK. The Google provider owns the SDK client used by this protocol. +""" + +from __future__ import annotations + +import base64 +import json +from typing import TYPE_CHECKING, Any, Literal, cast + +import httpx + +from ... import errors as ai_errors +from ... import types +from ...models.core import params as params_ +from ...types import events +from .. import base, history_utils +from . import _sdk, errors + +if TYPE_CHECKING: + from collections.abc import AsyncGenerator, Sequence + + import google.genai as genai + import pydantic + + from ...models import core + +PROVIDER_NAME = "google" + +_CODE_EXECUTION_TOOL = "code_execution" + +# Models before Gemini 3 reject `thinking_level` and take a token budget +# instead; reasoning effort maps to these budgets there (-1 = automatic). +_THINKING_BUDGETS = { + "minimal": 128, + "low": 1024, + "medium": 8192, + "high": 24576, +} + + +def _provider_metadata(**values: Any) -> dict[str, Any]: + """Namespace metadata as ``{"google": {...}}``.""" + return {PROVIDER_NAME: {**values}} + + +def _google_metadata(pm: dict[str, Any] | None) -> dict[str, Any]: + """Read back the metadata written by :func:`_provider_metadata`.""" + meta = (pm or {}).get(PROVIDER_NAME) + return meta if isinstance(meta, dict) else {} + + +def _signature_from_metadata(pm: dict[str, Any] | None) -> bytes | None: + signature = _google_metadata(pm).get("thoughtSignature") + if isinstance(signature, str) and signature: + return base64.b64decode(signature) + return None + + +# --------------------------------------------------------------------------- +# Message / tool conversion — internal types → Google wire format +# --------------------------------------------------------------------------- + + +def _split_tools( + tools: Sequence[types.tools.Tool], +) -> tuple[list[types.tools.Tool], list[types.tools.Tool]]: + """Split ``tools`` into host-executed and provider-executed declarations.""" + custom: list[types.tools.Tool] = [] + builtin: list[types.tools.Tool] = [] + for t in tools: + if t.kind == "provider": + builtin.append(t) + else: + custom.append(t) + return custom, builtin + + +def _tools_to_google( + custom: Sequence[types.tools.Tool], + builtin: Sequence[types.tools.Tool], +) -> list[dict[str, Any]]: + """Convert tools to Google wire format. + + Host-executed tools become one ``function_declarations`` entry; + each provider-executed tool becomes its own entry keyed by the + tool id without the ``google.`` prefix (``google_search`` etc.). + """ + wire: list[dict[str, Any]] = [] + declarations: list[dict[str, Any]] = [] + for tool in custom: + spec = tool.spec + if spec is None: + raise TypeError(f"function tool {tool.name!r} has no spec") + declarations.append( + { + "name": tool.name, + "description": spec.description or "", + "parameters_json_schema": spec.params, + } + ) + if declarations: + wire.append({"function_declarations": declarations}) + for tool in builtin: + cfg = tool.tool_config + tool_id = cfg.id if cfg is not None else None + if cfg is None or tool_id is None or not tool_id.startswith("google."): + raise ValueError( + "GoogleModel does not support provider tool " + f"{tool_id or tool.name!r}" + ) + wire.append({tool_id.removeprefix("google."): {**cfg.args}}) + return wire + + +def _file_part_to_google(part: types.messages.FilePart) -> dict[str, Any]: + """Convert a :class:`FilePart` to a Google content part. + + Downloadable URLs map to ``file_data``; everything else is sent + inline as raw bytes. + """ + media_type = ( + "image/jpeg" if part.media_type == "image/*" else part.media_type + ) + if isinstance(part.data, str) and types.media.is_downloadable_url( + part.data + ): + return {"file_data": {"file_uri": part.data, "mime_type": media_type}} + return { + "inline_data": { + "data": base64.b64decode(types.media.data_to_base64(part.data)), + "mime_type": media_type, + } + } + + +def _tool_result_to_google( + part: types.messages.ToolResultPart, +) -> dict[str, Any]: + """Convert a tool result to a ``function_response`` response object. + + Google requires a JSON object: dicts pass through, everything else + is wrapped under ``"output"`` (or ``"error"`` for error results). + """ + value = part.get_model_input() + if isinstance(value, types.messages.ContentOutput): + texts: list[str] = [] + for item in value.value: + if isinstance(item, types.messages.FilePart): + raise ValueError( + "Google does not support file parts in tool results" + ) + texts.append(item.text) + value = "".join(texts) + if part.is_error: + return {"error": value if value is not None else ""} + if isinstance(value, dict): + return value + return {"output": value if value is not None else ""} + + +def _messages_to_google( + messages: list[types.messages.Message], +) -> tuple[str | None, list[dict[str, Any]]]: + """Convert internal messages to Google API format. + + Returns ``(system_instruction, contents)``. The system prompt is + extracted separately because the Google API takes it as a config + parameter. + """ + system_instruction: str | None = None + result: list[dict[str, Any]] = [] + + for msg in history_utils.repair(messages): + match msg.role: + case "system": + system_instruction = "".join( + p.text + for p in msg.parts + if isinstance(p, types.messages.TextPart) + ) + case "assistant": + parts: list[dict[str, Any]] = [] + for part in msg.parts: + match part: + case types.messages.ReasoningPart( + text=text, + provider_metadata=provider_metadata, + ): + # Thought summaries are not sent back; only + # signed thoughts round-trip. + signature = _signature_from_metadata( + provider_metadata + ) + if signature: + parts.append( + { + "text": text, + "thought": True, + "thought_signature": signature, + } + ) + case types.messages.TextPart( + text=text, + provider_metadata=provider_metadata, + ): + entry: dict[str, Any] = {"text": text} + signature = _signature_from_metadata( + provider_metadata + ) + if signature: + entry["thought_signature"] = signature + parts.append(entry) + case types.messages.ToolCallPart(): + call_entry: dict[str, Any] = { + "function_call": { + "id": part.tool_call_id, + "name": part.tool_name, + "args": json.loads(part.tool_args) + if part.tool_args + else {}, + } + } + signature = _signature_from_metadata( + part.provider_metadata + ) + if signature: + call_entry["thought_signature"] = signature + parts.append(call_entry) + case types.messages.BuiltinToolCallPart() if ( + part.tool_name == _CODE_EXECUTION_TOOL + ): + exec_entry: dict[str, Any] = { + "executable_code": json.loads(part.tool_args) + if part.tool_args + else {}, + } + signature = _signature_from_metadata( + part.provider_metadata + ) + if signature: + exec_entry["thought_signature"] = signature + parts.append(exec_entry) + case types.messages.BuiltinToolReturnPart() if ( + part.tool_name == _CODE_EXECUTION_TOOL + ): + result_entry: dict[str, Any] = { + "code_execution_result": part.result or {}, + } + signature = _signature_from_metadata( + part.provider_metadata + ) + if signature: + result_entry["thought_signature"] = signature + parts.append(result_entry) + if parts: + result.append({"role": "model", "parts": parts}) + + case "tool": + tool_parts: list[dict[str, Any]] = [] + for part in msg.parts: + if isinstance(part, types.messages.ToolResultPart): + tool_parts.append( + { + "function_response": { + "id": part.tool_call_id, + "name": part.tool_name, + "response": _tool_result_to_google(part), + } + } + ) + if tool_parts: + result.append({"role": "user", "parts": tool_parts}) + + case "user": + user_parts: list[dict[str, Any]] = [] + for p in msg.parts: + match p: + case types.messages.TextPart(text=text): + user_parts.append({"text": text}) + case types.messages.FilePart(): + user_parts.append(_file_part_to_google(p)) + if user_parts: + result.append({"role": "user", "parts": user_parts}) + + return system_instruction, _merge_consecutive_roles(result) + + +def _merge_consecutive_roles( + contents: list[dict[str, Any]], +) -> list[dict[str, Any]]: + """Merge consecutive contents that share the same role.""" + if not contents: + return contents + + merged: list[dict[str, Any]] = [contents[0]] + for content in contents[1:]: + if content["role"] == merged[-1]["role"]: + merged[-1]["parts"] = merged[-1]["parts"] + content["parts"] + else: + merged.append(content) + return merged + + +def _is_default(value: object) -> bool: + return isinstance(value, params_.ModelProviderDefault) + + +def _not_default(value: object) -> bool: + return not _is_default(value) + + +def _seed_value(seed: object) -> int | None: + if seed is None or seed == -1 or isinstance(seed, params_.RandomSeed): + return None + if isinstance(seed, int): + return seed + if isinstance(seed, params_.ModelProviderDefault): + return None + raise TypeError("seed must be an int, RANDOM, DEFAULT, or None") + + +def _http_options(config: dict[str, Any]) -> dict[str, Any]: + http_options = config.get("http_options") + if not isinstance(http_options, dict): + http_options = {} + config["http_options"] = http_options + return http_options + + +def _google_tool_choice( + tool_choice: params_.ToolChoiceMode + | params_.ToolRef + | params_.ToolSelection, +) -> dict[str, Any]: + if isinstance(tool_choice, params_.ToolChoiceMode): + match tool_choice: + case params_.ToolChoiceMode.AUTO: + return {"mode": "AUTO"} + case params_.ToolChoiceMode.REQUIRED: + return {"mode": "ANY"} + case params_.ToolChoiceMode.NONE: + return {"mode": "NONE"} + if isinstance(tool_choice, params_.ToolRef): + return {"mode": "ANY", "allowed_function_names": [str(tool_choice)]} + if tool_choice.mode in { + params_.ToolChoiceMode.AUTO, + params_.ToolChoiceMode.REQUIRED, + }: + return { + # VALIDATED lets the model choose between the allowed tools + # and plain text; ANY forces a tool call. + "mode": "ANY" + if tool_choice.mode == params_.ToolChoiceMode.REQUIRED + else "VALIDATED", + "allowed_function_names": sorted(str(t) for t in tool_choice.tools), + } + raise ValueError("Google does not support excluded tool subsets") + + +def _apply_sampling( + config: dict[str, Any], + request_params: params_.InferenceRequestParams, +) -> None: + sampling = request_params.sampling + if isinstance(sampling, params_.ModelProviderDefault): + return + for sampler in sampling.values(): + match sampler: + case params_.TemperatureSamplerParams(temperature=temperature): + if _not_default(temperature): + config["temperature"] = temperature + case params_.TopPSamplerParams(top_p=top_p): + if _not_default(top_p): + config["top_p"] = top_p + case params_.TopKSamplerParams(top_k=top_k): + if _not_default(top_k): + config["top_k"] = top_k + case params_.SeedSamplerParams(seed=seed): + if _not_default(seed): + seed_value = _seed_value(seed) + if seed_value is not None: + config["seed"] = seed_value + case params_.MinPSamplerParams(min_p=min_p): + if _not_default(min_p) and min_p is not None: + raise ValueError("Google does not support min_p") + case params_.RepetitionPenaltyParams() as repetition: + if ( + _not_default(repetition.frequency_penalty) + and repetition.frequency_penalty is not None + ): + config["frequency_penalty"] = repetition.frequency_penalty + if ( + _not_default(repetition.presence_penalty) + and repetition.presence_penalty is not None + ): + config["presence_penalty"] = repetition.presence_penalty + unsupported = { + "repetition_penalty": repetition.repetition_penalty, + "consideration_window": repetition.consideration_window, + } + for key, value in unsupported.items(): + if _not_default(value) and value is not None: + raise ValueError(f"Google does not support {key}") + + +def _apply_google_params( + config: dict[str, Any], + request_params: params_.InferenceRequestParams, + *, + model_id: str, + provider: str, +) -> None: + _ = provider + _apply_sampling(config, request_params) + + reasoning = request_params.reasoning + output = request_params.output + summary = params_.DEFAULT if output is None else output.reasoning_summary + thinking: dict[str, Any] = {} + if not isinstance(reasoning, params_.ModelProviderDefault) and _not_default( + reasoning.effort + ): + if reasoning.effort is None: + thinking["thinking_budget"] = 0 + elif model_id.startswith("gemini-3"): + thinking["thinking_level"] = reasoning.effort + else: + thinking["thinking_budget"] = _THINKING_BUDGETS.get( + cast("str", reasoning.effort), -1 + ) + if _not_default(summary): + # `reasoning_summary` controls whether thought summaries are + # surfaced in the response; it never turns thinking off (use + # `reasoning.effort=None` for that). + thinking["include_thoughts"] = summary is not None + if thinking: + config["thinking_config"] = thinking + + if request_params.tool_calling is not None: + tool_calling = request_params.tool_calling + if ( + _not_default(tool_calling.max_tool_calls) + and tool_calling.max_tool_calls is not None + ): + raise ValueError("Google does not support max_tool_calls") + if _not_default(tool_calling.parallel_tool_calls): + raise ValueError( + "Google does not support configuring parallel tool calls" + ) + config["tool_config"] = { + "function_calling_config": _google_tool_choice( + tool_calling.tool_choice + ) + } + + if request_params.provider_service is not None and _not_default( + request_params.provider_service.service_tier + ): + config["service_tier"] = request_params.provider_service.service_tier + + if request_params.metadata is not None: + raise ValueError("Google does not support request metadata") + if request_params.safety_identifier is not None: + raise ValueError("Google does not support safety identifiers") + if request_params.context_management is not None: + raise ValueError("Google does not support context management") + + if output is not None: + if output.max_tokens is not None: + config["max_output_tokens"] = output.max_tokens + if output.include is not None: + raise ValueError("Google does not support output include") + if ( + _not_default(output.text_verbosity) + and output.text_verbosity is not None + ): + raise ValueError("Google does not support text verbosity") + + if request_params.cache is not None: + cache = request_params.cache + if _not_default(cache.retention): + raise ValueError("Google does not support cache retention") + if cache.key is not None: + config["cached_content"] = cache.key + + if request_params.extra_headers is not None: + headers = { + key: value + for key, value in request_params.extra_headers.items() + if not isinstance(value, params_.Unset) + } + if headers: + _http_options(config)["headers"] = headers + if request_params.extra_query is not None: + raise ValueError("Google does not support extra query arguments") + if request_params.extra_body is not None: + _http_options(config)["extra_body"] = dict(request_params.extra_body) + + +# --------------------------------------------------------------------------- +# Public protocol function +# --------------------------------------------------------------------------- + + +async def stream( + sdk_client: genai.Client, + model: core.model.Model, + messages: list[types.messages.Message], + *, + tools: Sequence[types.tools.Tool] | None = None, + output_type: type[pydantic.BaseModel] | None = None, + params: params_.InferenceRequestParams | None = None, + provider: str, +) -> AsyncGenerator[events.Event]: + """Stream through the Google generateContent protocol using *sdk_client*. + + Yields :class:`~ai.types.events.Event` objects as the response streams in. + Pure delta emitter — the :class:`~ai.models.Stream` wrapper aggregates + parts into the final :class:`~ai.types.messages.Message`. + """ + genai_errors = _sdk.import_errors(provider=provider) + config: dict[str, Any] = {} + if params is not None: + if not isinstance(params, params_.InferenceRequestParams): + raise TypeError( + "google stream params must be InferenceRequestParams" + ) + _apply_google_params( + config, params, model_id=model.id, provider=provider + ) + system_instruction, contents = _messages_to_google(messages) + if system_instruction: + config["system_instruction"] = system_instruction + + custom_tools, builtin_tools = _split_tools(tools or ()) + wire_tools = _tools_to_google(custom_tools, builtin_tools) + if wire_tools: + config["tools"] = wire_tools + + if output_type is not None: + config["response_mime_type"] = "application/json" + config["response_json_schema"] = output_type.model_json_schema() + + # Google streams complete parts per chunk without explicit block + # boundaries; like the OpenAI chat-completions protocol we track one + # text and one reasoning block with fixed ids. + text_started = False + reasoning_started = False + text_signature: str | None = None + reasoning_signature: str | None = None + file_index = 0 + # Fallback pairing for code execution results when the server does + # not populate part ids (Vertex AI never does). + last_exec_tool_id = "" + + try: + sdk_stream = await sdk_client.aio.models.generate_content_stream( + model=model.id, + contents=contents, + config=cast("Any", config or None), + ) + yield events.StreamStart() + + usage_metadata: Any = None + finish_reason: str | None = None + async for chunk in sdk_stream: + feedback = chunk.prompt_feedback + if feedback is not None and feedback.block_reason: + raise ai_errors.ProviderResponseError( + "Google blocked the prompt: " + f"{feedback.block_reason_message or feedback.block_reason}", + provider=provider, + ) + if chunk.usage_metadata is not None: + usage_metadata = chunk.usage_metadata + candidate = chunk.candidates[0] if chunk.candidates else None + if candidate is not None and candidate.finish_reason is not None: + finish_reason = str(candidate.finish_reason.value) + content = candidate.content if candidate is not None else None + for part in (content.parts if content is not None else None) or []: + signature = ( + base64.b64encode(part.thought_signature).decode("ascii") + if part.thought_signature + else None + ) + if part.text is not None: + if part.thought: + if not reasoning_started: + reasoning_started = True + yield events.ReasoningStart(block_id="reasoning") + yield events.ReasoningDelta( + chunk=part.text, block_id="reasoning" + ) + if signature is not None: + reasoning_signature = signature + else: + if reasoning_started: + reasoning_started = False + yield events.ReasoningEnd( + block_id="reasoning", + provider_metadata=( + _provider_metadata( + thoughtSignature=reasoning_signature + ) + if reasoning_signature is not None + else None + ), + ) + if not text_started: + text_started = True + yield events.TextStart(block_id="text") + yield events.TextDelta(chunk=part.text, block_id="text") + if signature is not None: + text_signature = signature + elif part.function_call is not None: + fc = part.function_call + tool_id = fc.id or types.messages.generate_id("call") + yield events.ToolStart( + tool_call_id=tool_id, tool_name=fc.name or "" + ) + yield events.ToolDelta( + chunk=json.dumps(fc.args or {}, separators=(",", ":")), + tool_call_id=tool_id, + ) + yield events.ToolEnd( + tool_call_id=tool_id, + tool_call=types.messages.DUMMY_TOOL_CALL, + provider_metadata=( + _provider_metadata(thoughtSignature=signature) + if signature is not None + else None + ), + ) + elif part.executable_code is not None: + tool_id = part.executable_code.id or ( + types.messages.generate_id("call") + ) + last_exec_tool_id = tool_id + exec_args = json.dumps( + part.executable_code.model_dump( + mode="json", exclude_none=True + ), + separators=(",", ":"), + ) + exec_metadata = ( + _provider_metadata(thoughtSignature=signature) + if signature is not None + else _provider_metadata() + ) + yield events.BuiltinToolStart( + tool_call_id=tool_id, + tool_name=_CODE_EXECUTION_TOOL, + provider_metadata=exec_metadata, + ) + yield events.BuiltinToolDelta( + chunk=exec_args, tool_call_id=tool_id + ) + yield events.BuiltinToolEnd( + tool_call_id=tool_id, + tool_call=types.messages.BuiltinToolCallPart( + tool_call_id=tool_id, + tool_name=_CODE_EXECUTION_TOOL, + provider_metadata=exec_metadata, + ), + provider_metadata=exec_metadata, + ) + elif part.code_execution_result is not None: + result_payload = part.code_execution_result.model_dump( + mode="json", exclude_none=True + ) + tool_id = part.code_execution_result.id or last_exec_tool_id + yield events.BuiltinToolResult( + tool_call_id=tool_id, + result=types.messages.BuiltinToolReturnPart( + tool_call_id=tool_id, + tool_name=_CODE_EXECUTION_TOOL, + result=result_payload, + is_error=result_payload.get("outcome") + not in (None, "OUTCOME_OK"), + provider_metadata=( + _provider_metadata(thoughtSignature=signature) + if signature is not None + else _provider_metadata() + ), + ), + ) + elif part.inline_data is not None: + yield events.FileEvent( + block_id=f"file_{file_index}", + media_type=part.inline_data.mime_type or "", + data=part.inline_data.data or b"", + ) + file_index += 1 + + if reasoning_started: + yield events.ReasoningEnd( + block_id="reasoning", + provider_metadata=( + _provider_metadata(thoughtSignature=reasoning_signature) + if reasoning_signature is not None + else None + ), + ) + if text_started: + yield events.TextEnd( + block_id="text", + provider_metadata=( + _provider_metadata(thoughtSignature=text_signature) + if text_signature is not None + else None + ), + ) + + if usage_metadata is not None: + thoughts = usage_metadata.thoughts_token_count + usage = types.usage.Usage( + input_tokens=usage_metadata.prompt_token_count or 0, + output_tokens=( + (usage_metadata.candidates_token_count or 0) + + (thoughts or 0) + ), + reasoning_tokens=thoughts, + cache_read_tokens=usage_metadata.cached_content_token_count, + raw=usage_metadata.model_dump(mode="json", exclude_none=True) + or None, + ) + else: + usage = types.usage.Usage() + yield events.StreamEnd( + usage=usage, + provider_metadata=( + _provider_metadata(finishReason=finish_reason) + if finish_reason is not None + else None + ), + ) + except genai_errors.APIError as exc: + raise errors.map_error( + exc, + provider=provider, + model_id=model.id, + ) from exc + except httpx.HTTPError as exc: + raise errors.map_httpx_error(exc, provider=provider) from exc + + +class GoogleGenerateContentProtocol(base.ProviderProtocol[Any]): + """Google generateContent API protocol.""" + + protocol_class_id: Literal["google_generate_content"] = ( + "google_generate_content" + ) + + def stream( + self, + client: genai.Client, + model: core.model.Model, + messages: list[types.messages.Message], + *, + tools: Sequence[types.tools.Tool] | None = None, + output_type: type[pydantic.BaseModel] | None = None, + params: params_.InferenceRequestParams | None = None, + provider: str, + ) -> AsyncGenerator[events.Event]: + return stream( + client, + model, + messages, + tools=tools, + output_type=output_type, + params=params, + provider=provider, + ) diff --git a/src/ai/providers/google/provider.py b/src/ai/providers/google/provider.py new file mode 100644 index 00000000..eddc25c2 --- /dev/null +++ b/src/ai/providers/google/provider.py @@ -0,0 +1,245 @@ +"""Google provider.""" + +from __future__ import annotations + +import os +from typing import TYPE_CHECKING, Any, ClassVar, Literal, cast + +import httpx +import pydantic + +from ... import errors as ai_errors +from .. import base +from . import _sdk, errors +from . import protocol as protocol_module +from . import tools as tools_module + +if TYPE_CHECKING: + from collections.abc import AsyncGenerator, Mapping, Sequence + from types import ModuleType + + import google.genai as genai + import modelsdotdev + + from ...models.core import model as model_ + from ...models.core import params as params_ + from ...types import events + from ...types import messages as messages_ + from ...types import tools as tools_ + + GoogleClient = httpx.AsyncClient | genai.Client + GoogleSDKClient = genai.Client +else: + GoogleClient = Any + GoogleSDKClient = Any + +_BASE_URL = "https://generativelanguage.googleapis.com" +_BASE_URL_ENV = "GOOGLE_GEMINI_BASE_URL" +# Alternative API key envs: the first two come from models.dev, the +# last is the SDK's own canonical env var. The first is the primary, +# the rest are fallbacks checked by :attr:`api_key`. +_API_KEY_ENVS = ( + "GOOGLE_GENERATIVE_AI_API_KEY", + "GEMINI_API_KEY", + "GOOGLE_API_KEY", +) + + +class GoogleProvider(base.Provider[GoogleSDKClient]): + """Callable provider for the Google Gemini API.""" + + handles: ClassVar[tuple[str, ...]] = ("google", "@ai-sdk/google") + + provider_class_id: Literal["google"] = "google" + + _http_client: httpx.AsyncClient | None = pydantic.PrivateAttr(default=None) + _close_client_on_aclose: bool = pydantic.PrivateAttr(default=False) + _has_user_sdk_client: bool = pydantic.PrivateAttr(default=False) + + def model_post_init(self, __context: Any) -> None: + self._close_client_on_aclose = True + + def _set_runtime_client(self, client: GoogleClient | None) -> None: + google_sdk = None + if client is not None and not isinstance(client, httpx.AsyncClient): + google_sdk = _sdk.import_sdk(provider=self.name) + + if google_sdk is not None and isinstance(client, google_sdk.Client): + sdk_client = client + http_client = None + self._has_user_sdk_client = True + self._close_client_on_aclose = False + elif isinstance(client, httpx.AsyncClient) or client is None: + sdk_client = None + http_client = client + self._has_user_sdk_client = False + self._close_client_on_aclose = client is None + else: + raise TypeError( + "Google providers require an httpx.AsyncClient or " + "google.genai.Client" + ) + + self._http_client = http_client + if sdk_client is not None: + self._set_client(sdk_client) + + def _make_sdk_client( + self, + *, + http_client: httpx.AsyncClient | None = None, + ) -> GoogleSDKClient: + google_sdk = _sdk.import_sdk(provider=self.name) + http_options: dict[str, Any] = {"base_url": self.base_url} + if self.headers: + http_options["headers"] = dict(self.headers) + if http_client is not None: + http_options["httpx_async_client"] = http_client + return google_sdk.Client( + api_key=self.api_key or None, + http_options=cast("Any", http_options), + ) + + @property + def sdk_client(self) -> GoogleSDKClient: + """Provider SDK client used for Google API requests.""" + return self.client + + @property + def client(self) -> GoogleSDKClient: + """Lazily-created SDK client for Google API requests.""" + if self._client is None: + self._set_client( + self._make_sdk_client(http_client=self._http_client) + ) + return super().client + + @property + def api_key(self) -> str | None: + """API key configured directly or via one of the Gemini env vars.""" + api_key = super().api_key + if api_key is not None: + return api_key + for env in _API_KEY_ENVS: + value = self.env.get(env) or os.environ.get(env) + if value: + return value + return None + + def default_protocol(self) -> base.ProviderProtocol[GoogleSDKClient]: + """Return the default Google generateContent protocol.""" + return protocol_module.GoogleGenerateContentProtocol() + + def is_configured(self) -> bool: + if self._has_user_sdk_client: + return True + if not self.api_key: + return False + return all(self._config_value(env) for env in self.config_envs) + + async def aclose(self) -> None: + """Close the provider-owned SDK client, if any.""" + if self._close_client_on_aclose and self._client is not None: + await self.client.aio.aclose() + + def stream( + self, + model: model_.Model, + messages: list[messages_.Message], + *, + tools: Sequence[tools_.Tool] | None = None, + output_type: type[pydantic.BaseModel] | None = None, + params: params_.InferenceRequestParams | None = None, + ) -> AsyncGenerator[events.Event]: + """Stream via the Google generateContent protocol.""" + return super().stream( + model, + messages, + tools=tools, + output_type=output_type, + params=params, + ) + + @classmethod + def from_modelsdev_provider( + cls, + provider: modelsdotdev.Provider, + *, + model_provider_config: modelsdotdev.ModelProviderConfig | None = None, + base_url: str | None = None, + api_key: str | None = None, + headers: Mapping[str, str] | None = None, + env: Mapping[str, str] | None = None, + client: GoogleClient | None = None, + protocol: base.ProviderProtocol[Any] | None = None, + ) -> base.Provider[GoogleSDKClient]: + resolved_base_url = base_url or base.provider_base_url( + provider, + model_provider_config, + ) + if resolved_base_url is None: + resolved_base_url = _BASE_URL + api_key_env, config_envs = base.provider_config( + provider, model_provider_config + ) + provider_instance = cls( + name=provider.id, + default_base_url=resolved_base_url, + api_key_value=api_key, + api_key_env=api_key_env, + base_url_env=_BASE_URL_ENV if base_url is None else None, + # models.dev lists alternative API key envs alongside real + # config envs; the alternates are handled by `api_key`. + config_envs=tuple( + config_env + for config_env in config_envs + if config_env not in _API_KEY_ENVS + ), + headers=dict(headers or {}), + env=dict(env or {}), + protocol_override=protocol, + ) + provider_instance._set_runtime_client(client) + return provider_instance + + @property + def tools(self) -> ModuleType: + """The provider's built-in tool factories. + + Convenience accessor: ``google.tools.google_search()``. + """ + return tools_module + + async def list_models(self) -> list[str]: + """List available model IDs from the Google API.""" + genai_errors = _sdk.import_errors(provider=self.name) + try: + pager = await self.sdk_client.aio.models.list() + names = [sdk_model.name async for sdk_model in pager] + except genai_errors.APIError as exc: + raise errors.map_error(exc, provider=self.name) from exc + except httpx.HTTPError as exc: + raise errors.map_httpx_error(exc, provider=self.name) from exc + return sorted(name.removeprefix("models/") for name in names if name) + + async def probe(self, model: model_.Model) -> None: + """Raise unless credentials are valid and the model exists.""" + if not self.is_configured(): + raise ai_errors.ProviderNotConfiguredError( + f"provider {self.name!r} is not configured", + provider=self.name, + ) + genai_errors = _sdk.import_errors(provider=self.name) + try: + await self.sdk_client.aio.models.get(model=model.id) + except genai_errors.APIError as exc: + raise errors.map_error( + exc, + provider=self.name, + model_id=model.id, + ) from exc + except httpx.HTTPError as exc: + raise errors.map_httpx_error(exc, provider=self.name) from exc + + +__all__ = ["GoogleProvider"] diff --git a/src/ai/providers/google/tools.py b/src/ai/providers/google/tools.py new file mode 100644 index 00000000..d382a923 --- /dev/null +++ b/src/ai/providers/google/tools.py @@ -0,0 +1,32 @@ +"""Google provider-executed tools.""" + +from __future__ import annotations + +from ... import types + + +def google_search() -> types.tools.Tool: + return types.tools.Tool( + kind="provider", + name="google_search", + tool_config=types.tools.ToolConfig(id="google.google_search"), + ) + + +def url_context() -> types.tools.Tool: + return types.tools.Tool( + kind="provider", + name="url_context", + tool_config=types.tools.ToolConfig(id="google.url_context"), + ) + + +def code_execution() -> types.tools.Tool: + return types.tools.Tool( + kind="provider", + name="code_execution", + tool_config=types.tools.ToolConfig(id="google.code_execution"), + ) + + +__all__ = ["code_execution", "google_search", "url_context"] diff --git a/tests/models/test_resolution.py b/tests/models/test_resolution.py index a77a276a..adc21520 100644 --- a/tests/models/test_resolution.py +++ b/tests/models/test_resolution.py @@ -7,6 +7,7 @@ from ai import ConfigurationError, models from ai.providers.ai_gateway import GatewayV3Protocol from ai.providers.anthropic import AnthropicMessagesProtocol +from ai.providers.google import GoogleGenerateContentProtocol from ai.providers.openai import ( OpenAIChatCompletionsProtocol, OpenAIResponsesProtocol, @@ -211,14 +212,22 @@ def test_provider_from_id_rejects_unknown_provider() -> None: def test_provider_from_id_rejects_unsupported_provider_package() -> None: with pytest.raises(ai.UnsupportedProviderError) as exc_info: - ai.get_provider("google") + ai.get_provider("google-vertex") - assert exc_info.value.provider_id == "google" + assert exc_info.value.provider_id == "google-vertex" def test_get_rejects_unsupported_provider_package() -> None: with pytest.raises(ai.errors.UnsupportedProviderError): - models.get_model("google:gemini-2.5-pro") + models.get_model("google-vertex:gemini-2.5-pro") + + +def test_get_resolves_provider_qualified_google_model_id() -> None: + model = models.get_model("google:gemini-2.5-pro") + + assert model.id == "gemini-2.5-pro" + assert model.provider.name == "google" + assert isinstance(model.provider.protocol, GoogleGenerateContentProtocol) def test_get_rejects_empty_model_id() -> None: diff --git a/tests/providers/google/__init__.py b/tests/providers/google/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/providers/google/conftest.py b/tests/providers/google/conftest.py new file mode 100644 index 00000000..016c6e0c --- /dev/null +++ b/tests/providers/google/conftest.py @@ -0,0 +1,84 @@ +"""Shared fakes for the Google adapter tests. + +The adapter consumes ``google.genai.Client`` via +``client.aio.models.generate_content_stream(**kwargs)``. To exercise the +real adapter without hitting the network we build a tiny stand-in that +captures the kwargs and yields real SDK-typed response chunks. +""" + +from __future__ import annotations + +from types import SimpleNamespace +from typing import Any + +from google.genai import types as genai_types + + +def chunk( + parts: list[genai_types.Part] | None = None, + *, + usage: dict[str, Any] | None = None, + finish_reason: genai_types.FinishReason | None = None, + block_reason: genai_types.BlockedReason | None = None, +) -> genai_types.GenerateContentResponse: + """Build an SDK-typed streaming response chunk.""" + candidates = None + if parts is not None or finish_reason is not None: + candidates = [ + genai_types.Candidate( + content=( + genai_types.Content(role="model", parts=parts) + if parts is not None + else None + ), + finish_reason=finish_reason, + ) + ] + return genai_types.GenerateContentResponse( + candidates=candidates, + prompt_feedback=( + genai_types.GenerateContentResponsePromptFeedback( + block_reason=block_reason + ) + if block_reason is not None + else None + ), + usage_metadata=( + genai_types.GenerateContentResponseUsageMetadata(**usage) + if usage is not None + else None + ), + ) + + +class FakeAsyncModels: + def __init__( + self, + captured: dict[str, Any], + chunks: list[genai_types.GenerateContentResponse], + ) -> None: + self._captured = captured + self._chunks = chunks + + async def generate_content_stream(self, **kwargs: Any) -> Any: + self._captured.update(kwargs) + + async def _gen() -> Any: + for item in self._chunks: + yield item + + return _gen() + + +class FakeGoogleClient: + """Stand-in for ``google.genai.Client``.""" + + def __init__( + self, + captured: dict[str, Any] | None = None, + chunks: list[genai_types.GenerateContentResponse] | None = None, + ) -> None: + self.captured = captured if captured is not None else {} + self.aio = SimpleNamespace( + models=FakeAsyncModels(self.captured, chunks or []) + ) diff --git a/tests/providers/google/test_adapter.py b/tests/providers/google/test_adapter.py new file mode 100644 index 00000000..3965f5aa --- /dev/null +++ b/tests/providers/google/test_adapter.py @@ -0,0 +1,853 @@ +"""Tests for the Google adapter's request shaping. + +Focused on ``params`` translation, message/tool conversion, and SDK +error mapping. +""" + +from __future__ import annotations + +import base64 +from typing import Any, cast + +import httpx +import pydantic +import pytest +from google.genai import errors as genai_errors + +import ai +from ai.providers.google import protocol +from ai.providers.google import tools as google_tools +from ai.types import messages + +from .conftest import FakeGoogleClient + +_MODEL = ai.Model(id="gemini-2.5-flash", provider=ai.get_provider("google")) + + +class _RaisingAsyncModels: + def __init__(self, exc: Exception) -> None: + self._exc = exc + + async def generate_content_stream(self, **kwargs: Any) -> Any: + raise self._exc + + +class _RaisingGoogleClient: + def __init__(self, exc: Exception) -> None: + self.aio = type("Aio", (), {"models": _RaisingAsyncModels(exc)})() + + +async def _drain(stream: Any) -> None: + async for _ in stream: + pass + + +async def test_params_translate_to_config() -> None: + fake = FakeGoogleClient() + + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + [ai.user_message("Hi")], + params=ai.InferenceRequestParams( + sampling={ + ai.TemperatureSamplerParams: ai.TemperatureSamplerParams( + temperature=0.5 + ), + ai.TopPSamplerParams: ai.TopPSamplerParams(top_p=0.9), + ai.TopKSamplerParams: ai.TopKSamplerParams(top_k=40), + ai.SeedSamplerParams: ai.SeedSamplerParams(seed=123), + ai.RepetitionPenaltyParams: ai.RepetitionPenaltyParams( + frequency_penalty=0.1, presence_penalty=0.2 + ), + }, + reasoning=ai.ReasoningParams(effort="high"), + output=ai.OutputParams( + max_tokens=123, reasoning_summary="auto" + ), + tool_calling=ai.ToolCallingParams( + tool_choice=ai.ToolChoiceMode.REQUIRED, + ), + extra_headers={"x-goog-feature": "enabled"}, + extra_body={"future_option": {"enabled": True}}, + ), + provider="google", + ) + ) + + config = fake.captured["config"] + assert config["temperature"] == 0.5 + assert config["top_p"] == 0.9 + assert config["top_k"] == 40 + assert config["seed"] == 123 + assert config["frequency_penalty"] == 0.1 + assert config["presence_penalty"] == 0.2 + assert config["max_output_tokens"] == 123 + # Pre-Gemini-3 models take a token budget instead of a level. + assert config["thinking_config"] == { + "thinking_budget": 24576, + "include_thoughts": True, + } + assert config["tool_config"] == {"function_calling_config": {"mode": "ANY"}} + assert config["http_options"] == { + "headers": {"x-goog-feature": "enabled"}, + "extra_body": {"future_option": {"enabled": True}}, + } + + +async def test_reasoning_effort_maps_to_thinking_level_on_gemini_3() -> None: + fake = FakeGoogleClient() + + await _drain( + protocol.stream( + cast("Any", fake), + ai.Model(id="gemini-3-flash-preview", provider=_MODEL.provider), + [ai.user_message("Hi")], + params=ai.InferenceRequestParams( + reasoning=ai.ReasoningParams(effort="high") + ), + provider="google", + ) + ) + + assert fake.captured["config"]["thinking_config"] == { + "thinking_level": "high" + } + + +async def test_reasoning_disabled_maps_to_zero_budget() -> None: + fake = FakeGoogleClient() + + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + [ai.user_message("Hi")], + params=ai.InferenceRequestParams( + reasoning=ai.ReasoningParams(effort=None) + ), + provider="google", + ) + ) + + assert fake.captured["config"]["thinking_config"] == {"thinking_budget": 0} + + +async def test_random_seed_omitted_by_adapter() -> None: + fake = FakeGoogleClient() + + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + [ai.user_message("Hi")], + params=ai.InferenceRequestParams( + sampling={ai.SeedSamplerParams: ai.SeedSamplerParams(seed=-1)} + ), + provider="google", + ) + ) + + assert fake.captured["config"] is None + + +async def test_unsupported_params_rejected_by_adapter() -> None: + unsupported = [ + ( + ai.InferenceRequestParams( + sampling={ai.MinPSamplerParams: ai.MinPSamplerParams(min_p=0.1)} + ), + "min_p", + ), + (ai.InferenceRequestParams(metadata={"k": "v"}), "metadata"), + (ai.InferenceRequestParams(extra_query={"k": "v"}), "extra query"), + ( + ai.InferenceRequestParams( + tool_calling=ai.ToolCallingParams( + tool_choice=ai.ToolChoiceMode.AUTO, + parallel_tool_calls=False, + ) + ), + "parallel tool calls", + ), + ] + for params, match in unsupported: + with pytest.raises(ValueError, match=match): + await _drain( + protocol.stream( + cast("Any", FakeGoogleClient()), + _MODEL, + [ai.user_message("Hi")], + params=params, + provider="google", + ) + ) + + +async def test_system_prompt_becomes_system_instruction() -> None: + fake = FakeGoogleClient() + + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + [ai.system_message("Be brief."), ai.user_message("Hi")], + provider="google", + ) + ) + + assert fake.captured["config"]["system_instruction"] == "Be brief." + assert fake.captured["contents"] == [ + {"role": "user", "parts": [{"text": "Hi"}]} + ] + + +async def test_tool_round_trip_and_result_wrapping() -> None: + fake = FakeGoogleClient() + + convo = [ + ai.user_message("What's the weather?"), + messages.Message( + role="assistant", + parts=[ + messages.ToolCallPart( + tool_call_id="fc_1", + tool_name="get_weather", + tool_args='{"city":"Tokyo"}', + ) + ], + ), + messages.Message( + role="tool", + parts=[ + messages.ToolResultPart( + tool_call_id="fc_1", + tool_name="get_weather", + result="sunny", + ) + ], + ), + ] + + await _drain( + protocol.stream(cast("Any", fake), _MODEL, convo, provider="google") + ) + + assert fake.captured["contents"] == [ + {"role": "user", "parts": [{"text": "What's the weather?"}]}, + { + "role": "model", + "parts": [ + { + "function_call": { + "id": "fc_1", + "name": "get_weather", + "args": {"city": "Tokyo"}, + } + } + ], + }, + { + "role": "user", + "parts": [ + { + "function_response": { + "id": "fc_1", + "name": "get_weather", + "response": {"output": "sunny"}, + } + } + ], + }, + ] + + +async def test_dict_tool_result_passes_through_unwrapped() -> None: + fake = FakeGoogleClient() + + convo = [ + ai.user_message("Weather?"), + messages.Message( + role="assistant", + parts=[ + messages.ToolCallPart( + tool_call_id="fc_1", tool_name="get_weather", tool_args="{}" + ) + ], + ), + messages.Message( + role="tool", + parts=[ + messages.ToolResultPart( + tool_call_id="fc_1", + tool_name="get_weather", + result={"condition": "sunny", "temp_c": 30}, + ) + ], + ), + ai.user_message("Thanks"), + ] + + await _drain( + protocol.stream(cast("Any", fake), _MODEL, convo, provider="google") + ) + + # The dict result passes through unwrapped, and the trailing user + # message merges into the same user content as the tool response. + tool_content = fake.captured["contents"][-1] + assert tool_content["role"] == "user" + assert tool_content["parts"] == [ + { + "function_response": { + "id": "fc_1", + "name": "get_weather", + "response": {"condition": "sunny", "temp_c": 30}, + } + }, + {"text": "Thanks"}, + ] + + +async def test_thought_signature_round_trips_from_provider_metadata() -> None: + fake = FakeGoogleClient() + signature = base64.b64encode(b"sig-bytes").decode() + + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + [ + ai.assistant_message( + ai.thinking( + "hidden", + provider_metadata={ + "google": {"thoughtSignature": signature} + }, + ) + ), + ai.user_message("Hi"), + ], + provider="google", + ) + ) + + assert fake.captured["contents"][0] == { + "role": "model", + "parts": [ + { + "text": "hidden", + "thought": True, + "thought_signature": b"sig-bytes", + } + ], + } + + +async def test_text_part_signature_round_trips() -> None: + fake = FakeGoogleClient() + signature = base64.b64encode(b"sig-bytes").decode() + + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + [ + messages.Message( + role="assistant", + parts=[ + messages.TextPart( + text="answer", + provider_metadata={ + "google": {"thoughtSignature": signature} + }, + ) + ], + ), + ai.user_message("Hi"), + ], + provider="google", + ) + ) + + assert fake.captured["contents"][0] == { + "role": "model", + "parts": [{"text": "answer", "thought_signature": b"sig-bytes"}], + } + + +async def test_unsigned_reasoning_parts_are_dropped() -> None: + fake = FakeGoogleClient() + + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + [ + ai.assistant_message(ai.thinking("hidden"), "visible"), + ai.user_message("Hi"), + ], + provider="google", + ) + ) + + assert fake.captured["contents"][0] == { + "role": "model", + "parts": [{"text": "visible"}], + } + + +async def test_tools_translate_to_wire_format() -> None: + fake = FakeGoogleClient() + + tool = ai.types.tools.Tool( + kind="function", + name="get_weather", + spec=ai.types.tools.ToolSpec( + description="Get the weather", + params={"type": "object", "properties": {}}, + ), + ) + + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + [ai.user_message("Hi")], + tools=[tool, google_tools.google_search()], + provider="google", + ) + ) + + assert fake.captured["config"]["tools"] == [ + { + "function_declarations": [ + { + "name": "get_weather", + "description": "Get the weather", + "parameters_json_schema": { + "type": "object", + "properties": {}, + }, + } + ] + }, + {"google_search": {}}, + ] + + +async def test_foreign_provider_tool_rejected() -> None: + tool = ai.types.tools.Tool( + kind="provider", + name="web_search", + tool_config=ai.types.tools.ToolConfig( + id="anthropic.web_search_20260209" + ), + ) + + with pytest.raises(ValueError, match="provider tool"): + await _drain( + protocol.stream( + cast("Any", FakeGoogleClient()), + _MODEL, + [ai.user_message("Hi")], + tools=[tool], + provider="google", + ) + ) + + +async def test_output_type_sets_response_schema() -> None: + fake = FakeGoogleClient() + + class Weather(pydantic.BaseModel): + city: str + + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + [ai.user_message("Hi")], + output_type=Weather, + provider="google", + ) + ) + + config = fake.captured["config"] + assert config["response_mime_type"] == "application/json" + assert config["response_json_schema"] == Weather.model_json_schema() + + +async def test_sdk_errors_are_mapped_to_provider_hierarchy() -> None: + response = httpx.Response( + 429, + request=httpx.Request("POST", "https://google.test/v1beta/models"), + ) + sdk_error = genai_errors.APIError( + 429, + { + "error": { + "code": 429, + "message": "quota exceeded", + "status": "RESOURCE_EXHAUSTED", + } + }, + response, + ) + + with pytest.raises(ai.ProviderRateLimitError) as exc_info: + await _drain( + protocol.stream( + cast("Any", _RaisingGoogleClient(sdk_error)), + _MODEL, + [ai.user_message("Hi")], + provider="google", + ) + ) + + exc = exc_info.value + assert exc.provider == "google" + assert exc.http_context is not None + assert exc.http_context.status_code == 429 + assert exc.http_context.request is response.request + assert exc.http_context.response is response + assert exc.type == "RESOURCE_EXHAUSTED" + assert exc.__cause__ is sdk_error + + +async def test_transport_errors_are_mapped_to_provider_hierarchy() -> None: + for sdk_error, expected in [ + (httpx.ConnectError("connection refused"), ai.ProviderConnectionError), + (httpx.ConnectTimeout("timed out"), ai.ProviderTimeoutError), + ]: + with pytest.raises(expected) as exc_info: + await _drain( + protocol.stream( + cast("Any", _RaisingGoogleClient(sdk_error)), + _MODEL, + [ai.user_message("Hi")], + provider="google", + ) + ) + + exc = exc_info.value + assert exc.provider == "google" + assert exc.is_retryable is True + assert exc.__cause__ is sdk_error + + +async def test_model_404_is_mapped_to_model_not_found() -> None: + sdk_error = genai_errors.APIError( + 404, + {"error": {"code": 404, "message": "model not found"}}, + ) + + with pytest.raises(ai.ProviderModelNotFoundError) as exc_info: + await _drain( + protocol.stream( + cast("Any", _RaisingGoogleClient(sdk_error)), + _MODEL, + [ai.user_message("Hi")], + provider="google", + ) + ) + + assert exc_info.value.model_id == _MODEL.id + + +async def test_tool_choice_modes_map_to_function_calling_config() -> None: + cases: list[tuple[ai.ToolChoiceMode | ai.ToolRef, dict[str, Any]]] = [ + (ai.ToolChoiceMode.AUTO, {"mode": "AUTO"}), + (ai.ToolChoiceMode.NONE, {"mode": "NONE"}), + ( + ai.ToolRef("get_weather"), + {"mode": "ANY", "allowed_function_names": ["get_weather"]}, + ), + ] + for tool_choice, expected in cases: + fake = FakeGoogleClient() + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + [ai.user_message("Hi")], + params=ai.InferenceRequestParams( + tool_calling=ai.ToolCallingParams(tool_choice=tool_choice) + ), + provider="google", + ) + ) + + assert fake.captured["config"]["tool_config"] == { + "function_calling_config": expected + } + + +async def test_tool_selection_maps_auto_to_validated() -> None: + for mode, expected in [ + (ai.ToolChoiceMode.AUTO, "VALIDATED"), + (ai.ToolChoiceMode.REQUIRED, "ANY"), + ]: + fake = FakeGoogleClient() + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + [ai.user_message("Hi")], + params=ai.InferenceRequestParams( + tool_calling=ai.ToolCallingParams( + tool_choice=ai.ToolSelection( + tools=frozenset({ai.ToolRef("get_weather")}), + mode=mode, + ) + ) + ), + provider="google", + ) + ) + + assert fake.captured["config"]["tool_config"] == { + "function_calling_config": { + "mode": expected, + "allowed_function_names": ["get_weather"], + } + } + + +async def test_service_tier_maps_to_config() -> None: + fake = FakeGoogleClient() + + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + [ai.user_message("Hi")], + params=ai.InferenceRequestParams( + provider_service=ai.ProviderServiceParams(service_tier="flex") + ), + provider="google", + ) + ) + + assert fake.captured["config"]["service_tier"] == "flex" + + +async def test_builtin_code_execution_parts_round_trip() -> None: + """Built-in tool parts serialize back to wire with their signatures.""" + fake = FakeGoogleClient() + signature = base64.b64encode(b"sig-bytes").decode() + + call = messages.BuiltinToolCallPart( + tool_call_id="call_1", + tool_name="code_execution", + tool_args='{"code":"print(1)","language":"PYTHON"}', + provider_metadata={"google": {"thoughtSignature": signature}}, + ) + result = messages.BuiltinToolReturnPart( + tool_call_id="call_1", + tool_name="code_execution", + result={"outcome": "OUTCOME_OK", "output": "1\n"}, + provider_metadata={"google": {}}, + ) + convo = [ + ai.user_message("Compute 1"), + messages.Message(role="assistant", parts=[call, result]), + ai.user_message("Thanks"), + ] + + await _drain( + protocol.stream(cast("Any", fake), _MODEL, convo, provider="google") + ) + + model_content = next( + c for c in fake.captured["contents"] if c["role"] == "model" + ) + assert model_content["parts"] == [ + { + "executable_code": {"code": "print(1)", "language": "PYTHON"}, + "thought_signature": b"sig-bytes", + }, + { + "code_execution_result": { + "outcome": "OUTCOME_OK", + "output": "1\n", + } + }, + ] + + +async def test_tool_call_signature_round_trips() -> None: + fake = FakeGoogleClient() + signature = base64.b64encode(b"sig-bytes").decode() + + convo = [ + ai.user_message("Weather?"), + messages.Message( + role="assistant", + parts=[ + messages.ToolCallPart( + tool_call_id="fc_1", + tool_name="get_weather", + tool_args="{}", + provider_metadata={ + "google": {"thoughtSignature": signature} + }, + ) + ], + ), + messages.Message( + role="tool", + parts=[ + messages.ToolResultPart( + tool_call_id="fc_1", + tool_name="get_weather", + result="sunny", + ) + ], + ), + ] + + await _drain( + protocol.stream(cast("Any", fake), _MODEL, convo, provider="google") + ) + + model_content = next( + c for c in fake.captured["contents"] if c["role"] == "model" + ) + assert model_content["parts"] == [ + { + "function_call": { + "id": "fc_1", + "name": "get_weather", + "args": {}, + }, + "thought_signature": b"sig-bytes", + } + ] + + +async def test_file_parts_convert_to_inline_and_file_data() -> None: + fake = FakeGoogleClient() + + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + [ + messages.Message( + role="user", + parts=[ + messages.TextPart(text="What are these?"), + messages.FilePart( + data=b"\x89PNG", media_type="image/png" + ), + messages.FilePart( + data="https://example.com/doc.pdf", + media_type="application/pdf", + ), + ], + ) + ], + provider="google", + ) + ) + + (content,) = fake.captured["contents"] + assert content["parts"] == [ + {"text": "What are these?"}, + {"inline_data": {"data": b"\x89PNG", "mime_type": "image/png"}}, + { + "file_data": { + "file_uri": "https://example.com/doc.pdf", + "mime_type": "application/pdf", + } + }, + ] + + +async def test_multipart_tool_result_flattens_text_and_rejects_files() -> None: + def _convo(result: Any) -> list[messages.Message]: + return [ + ai.user_message("Go"), + messages.Message( + role="assistant", + parts=[ + messages.ToolCallPart( + tool_call_id="fc_1", tool_name="tool", tool_args="{}" + ) + ], + ), + messages.Message( + role="tool", + parts=[ + messages.ToolResultPart( + tool_call_id="fc_1", + tool_name="tool", + result=result, + result_kind="special", + ) + ], + ), + ] + + fake = FakeGoogleClient() + await _drain( + protocol.stream( + cast("Any", fake), + _MODEL, + _convo( + messages.ContentOutput( + value=[ + messages.TextPart(text="part one, "), + messages.TextPart(text="part two"), + ] + ) + ), + provider="google", + ) + ) + tool_content = fake.captured["contents"][-1] + assert tool_content["parts"][0]["function_response"]["response"] == { + "output": "part one, part two" + } + + with pytest.raises(ValueError, match="file parts in tool results"): + await _drain( + protocol.stream( + cast("Any", FakeGoogleClient()), + _MODEL, + _convo( + messages.ContentOutput( + value=[ + messages.FilePart( + data=b"\x89PNG", media_type="image/png" + ) + ] + ) + ), + provider="google", + ) + ) + + +async def test_messages_to_google_repairs_history() -> None: + """Conversion runs history_utils.repair: internal messages are dropped + and orphaned tool calls get a synthetic error result.""" + msgs = [ + messages.Message( + role="internal", + parts=[messages.TextPart(text="app-only")], + ), + messages.Message( + role="assistant", + parts=[ + messages.ToolCallPart( + tool_call_id="tc-1", tool_name="search", tool_args="{}" + ) + ], + ), + ] + _, wire = protocol._messages_to_google(msgs) + assert [m["role"] for m in wire] == ["model", "user"] + (tool_response,) = wire[1]["parts"] + assert "error" in tool_response["function_response"]["response"] diff --git a/tests/providers/google/test_provider.py b/tests/providers/google/test_provider.py new file mode 100644 index 00000000..5bda796c --- /dev/null +++ b/tests/providers/google/test_provider.py @@ -0,0 +1,167 @@ +from __future__ import annotations + +import importlib + +import google.genai +import httpx +import pytest + +import ai +from ai.providers.google import ( + GoogleGenerateContentProtocol, + GoogleProvider, +) + + +async def test_list_models_strips_prefix_and_sorts_ids() -> None: + captured_urls: list[str] = [] + captured_headers: dict[str, str] = {} + + def _handler(request: httpx.Request) -> httpx.Response: + captured_urls.append(str(request.url)) + captured_headers.update(dict(request.headers)) + return httpx.Response( + 200, + json={ + "models": [ + {"name": "models/gemini-z"}, + {"name": "models/gemini-a"}, + ] + }, + ) + + provider = ai.get_provider( + "google", + base_url="https://google.test", + api_key="sk-test", + headers={"X-Custom-Header": "example"}, + client=httpx.AsyncClient(transport=httpx.MockTransport(_handler)), + ) + + ids = await provider.list_models() + + assert captured_urls[0].startswith("https://google.test/") + assert "models" in captured_urls[0] + assert captured_headers["x-goog-api-key"] == "sk-test" + assert captured_headers["x-custom-header"] == "example" + assert ids == ["gemini-a", "gemini-z"] + + +async def test_probe_maps_404_to_model_not_found() -> None: + def _handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 404, + json={"error": {"code": 404, "message": "not found"}}, + ) + + provider = ai.get_provider( + "google", + base_url="https://google.test", + api_key="sk-test", + client=httpx.AsyncClient(transport=httpx.MockTransport(_handler)), + ) + model = ai.Model(id="gemini-nope", provider=provider) + + with pytest.raises(ai.ProviderModelNotFoundError) as exc_info: + await provider.probe(model) + + assert exc_info.value.model_id == "gemini-nope" + assert exc_info.value.provider == "google" + + +async def test_probe_unconfigured_raises_not_configured( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("GOOGLE_GENERATIVE_AI_API_KEY", raising=False) + monkeypatch.delenv("GEMINI_API_KEY", raising=False) + + provider = ai.get_provider("google") + model = ai.Model(id="gemini-2.5-flash", provider=provider) + + with pytest.raises(ai.ProviderNotConfiguredError): + await provider.probe(model) + + +async def test_get_provider_accepts_google_sdk_client() -> None: + sdk_client = google.genai.Client(api_key="sk-test") + provider = ai.get_provider("google", client=sdk_client) + + assert isinstance(provider, GoogleProvider) + assert provider.sdk_client is sdk_client + assert provider.is_configured() is True + + +def test_base_url_defaults_when_env_var_unset( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("GOOGLE_GEMINI_BASE_URL", raising=False) + assert ( + ai.get_provider("google").base_url + == "https://generativelanguage.googleapis.com" + ) + + +def test_base_url_reads_google_gemini_base_url_env( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("GOOGLE_GEMINI_BASE_URL", "https://proxy.example.com") + assert ai.get_provider("google").base_url == "https://proxy.example.com" + + +def test_api_key_env_fallbacks(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("GOOGLE_GENERATIVE_AI_API_KEY", raising=False) + monkeypatch.delenv("GEMINI_API_KEY", raising=False) + monkeypatch.delenv("GOOGLE_API_KEY", raising=False) + + assert ai.get_provider("google").is_configured() is False + + monkeypatch.setenv("GOOGLE_API_KEY", "sk-canonical") + provider = ai.get_provider("google") + assert provider.api_key == "sk-canonical" + assert provider.is_configured() is True + + monkeypatch.setenv("GEMINI_API_KEY", "sk-gemini") + assert ai.get_provider("google").api_key == "sk-gemini" + + monkeypatch.setenv("GOOGLE_GENERATIVE_AI_API_KEY", "sk-google") + assert ai.get_provider("google").api_key == "sk-google" + + +def test_get_provider_raises_installation_error_when_google_sdk_missing( + monkeypatch: pytest.MonkeyPatch, +) -> None: + real_import_module = importlib.import_module + + def _missing_google(name: str, package: str | None = None) -> object: + if name == "google.genai" or name.startswith("google.genai."): + raise ModuleNotFoundError(name=name) + return real_import_module(name, package) + + monkeypatch.setattr(importlib, "import_module", _missing_google) + + provider = ai.get_provider("google", api_key="sk-test") + + with pytest.raises(ai.InstallationError) as exc_info: + _ = provider.client + + assert "could not import `google`" in str(exc_info.value) + assert "required to use the google provider" in str(exc_info.value) + assert "ai[google]" in str(exc_info.value) + + +def test_get_provider_accepts_base_url_and_api_key() -> None: + provider = ai.get_provider( + "google", + base_url="https://custom.example.com", + api_key="sk-custom", + headers={"X-Custom-Header": "example"}, + ) + + model = ai.Model(id="custom-model", provider=provider) + assert repr(provider) == "google" + assert isinstance(provider.protocol, GoogleGenerateContentProtocol) + assert provider.base_url == "https://custom.example.com" + assert provider.api_key == "sk-custom" + assert provider.headers == {"X-Custom-Header": "example"} + assert provider.is_configured() is True + assert model.id == "custom-model" diff --git a/tests/providers/google/test_stream.py b/tests/providers/google/test_stream.py new file mode 100644 index 00000000..9d3b9326 --- /dev/null +++ b/tests/providers/google/test_stream.py @@ -0,0 +1,294 @@ +"""Tests for Google stream parsing. + +The adapter consumes ``GenerateContentResponse`` chunks and emits +framework events. Drained through :class:`models.Stream` to also +exercise event aggregation in ``core.api``. +""" + +from __future__ import annotations + +import base64 +from typing import Any, cast + +import pytest +from google.genai import types as genai_types + +import ai +from ai import models +from ai.providers.google import protocol +from ai.types import events, messages + +from .conftest import FakeGoogleClient, chunk + +_MODEL = ai.Model(id="gemini-2.5-flash", provider=ai.get_provider("google")) + + +async def _drain(chunks: list[Any]) -> models.Stream: + fake = FakeGoogleClient(chunks=chunks) + s = models.Stream( + protocol.stream( + cast("Any", fake), + _MODEL, + [ai.user_message("Hi")], + provider="google", + ) + ) + async for _ in s: + pass + return s + + +async def test_text_deltas_aggregate_into_one_block() -> None: + s = await _drain( + [ + chunk([genai_types.Part(text="Hello")]), + chunk([genai_types.Part(text=" world")]), + ] + ) + + (part,) = s.message.parts + assert isinstance(part, messages.TextPart) + assert part.text == "Hello world" + + +async def test_thought_parts_emit_reasoning_with_signature() -> None: + s = await _drain( + [ + chunk( + [ + genai_types.Part( + text="thinking...", + thought=True, + thought_signature=b"sig-bytes", + ), + genai_types.Part(text="Answer"), + ] + ), + ] + ) + + reasoning, text = s.message.parts + assert isinstance(reasoning, messages.ReasoningPart) + assert reasoning.text == "thinking..." + assert reasoning.provider_metadata == { + "google": {"thoughtSignature": base64.b64encode(b"sig-bytes").decode()} + } + assert isinstance(text, messages.TextPart) + assert text.text == "Answer" + + +async def test_function_call_emits_tool_events() -> None: + fake = FakeGoogleClient( + chunks=[ + chunk( + [ + genai_types.Part( + function_call=genai_types.FunctionCall( + id="fc_1", + name="get_weather", + args={"city": "Tokyo"}, + ) + ) + ] + ) + ] + ) + + seen: list[type] = [] + s = models.Stream( + protocol.stream( + cast("Any", fake), + _MODEL, + [ai.user_message("Hi")], + provider="google", + ) + ) + async for event in s: + seen.append(type(event)) + + assert seen == [ + events.StreamStart, + events.ToolStart, + events.ToolDelta, + events.ToolEnd, + events.StreamEnd, + ] + (call,) = s.message.tool_calls + assert call.tool_call_id == "fc_1" + assert call.tool_name == "get_weather" + assert call.tool_args == '{"city":"Tokyo"}' + + +async def test_function_call_without_id_gets_generated_id() -> None: + s = await _drain( + [ + chunk( + [ + genai_types.Part( + function_call=genai_types.FunctionCall( + name="get_weather", args={} + ) + ) + ] + ) + ] + ) + + (call,) = s.message.tool_calls + assert call.tool_call_id + + +async def test_code_execution_parts_emit_builtin_events() -> None: + s = await _drain( + [ + chunk( + [ + genai_types.Part( + executable_code=genai_types.ExecutableCode( + code="print(1)", + language=genai_types.Language.PYTHON, + ) + ), + genai_types.Part( + code_execution_result=genai_types.CodeExecutionResult( + outcome=genai_types.Outcome.OUTCOME_OK, + output="1\n", + ) + ), + ] + ) + ] + ) + + (call,) = s.message.builtin_tool_calls + assert call.tool_name == "code_execution" + assert call.tool_args == '{"code":"print(1)","language":"PYTHON"}' + + (ret,) = s.message.builtin_tool_returns + assert ret.tool_call_id == call.tool_call_id + assert ret.tool_name == "code_execution" + assert ret.result == {"outcome": "OUTCOME_OK", "output": "1\n"} + assert ret.is_error is False + + +async def test_code_execution_results_pair_by_id() -> None: + s = await _drain( + [ + chunk( + [ + genai_types.Part( + executable_code=genai_types.ExecutableCode( + code="print(1)", + language=genai_types.Language.PYTHON, + id="exec_1", + ) + ), + genai_types.Part( + executable_code=genai_types.ExecutableCode( + code="print(2)", + language=genai_types.Language.PYTHON, + id="exec_2", + ) + ), + genai_types.Part( + code_execution_result=genai_types.CodeExecutionResult( + outcome=genai_types.Outcome.OUTCOME_OK, + output="1\n", + id="exec_1", + ) + ), + genai_types.Part( + code_execution_result=genai_types.CodeExecutionResult( + outcome=genai_types.Outcome.OUTCOME_OK, + output="2\n", + id="exec_2", + ) + ), + ] + ) + ] + ) + + calls = s.message.builtin_tool_calls + assert [c.tool_call_id for c in calls] == ["exec_1", "exec_2"] + + returns = s.message.builtin_tool_returns + assert {r.tool_call_id: r.result["output"] for r in returns} == { + "exec_1": "1\n", + "exec_2": "2\n", + } + + +async def test_inline_data_emits_file_event() -> None: + fake = FakeGoogleClient( + chunks=[ + chunk( + [ + genai_types.Part( + inline_data=genai_types.Blob( + data=b"\x89PNG", mime_type="image/png" + ) + ) + ] + ) + ] + ) + + file_events = [] + s = models.Stream( + protocol.stream( + cast("Any", fake), + _MODEL, + [ai.user_message("Draw a cat")], + provider="google", + ) + ) + async for event in s: + if isinstance(event, events.FileEvent): + file_events.append(event) + + (file_event,) = file_events + assert file_event.media_type == "image/png" + assert file_event.data == b"\x89PNG" + + +async def test_blocked_prompt_raises_response_error() -> None: + with pytest.raises(ai.ProviderResponseError, match="blocked the prompt"): + await _drain([chunk(block_reason=genai_types.BlockedReason.SAFETY)]) + + +async def test_finish_reason_lands_in_provider_metadata() -> None: + s = await _drain( + [ + chunk( + [genai_types.Part(text="partial")], + finish_reason=genai_types.FinishReason.SAFETY, + ), + ] + ) + + assert s.message.provider_metadata == {"google": {"finishReason": "SAFETY"}} + + +async def test_usage_metadata_maps_to_usage() -> None: + s = await _drain( + [ + chunk([genai_types.Part(text="Hi")]), + chunk( + None, + usage={ + "prompt_token_count": 10, + "candidates_token_count": 5, + "thoughts_token_count": 3, + "cached_content_token_count": 2, + }, + ), + ] + ) + + usage = s.message.usage + assert usage is not None + assert usage.input_tokens == 10 + assert usage.output_tokens == 8 + assert usage.reasoning_tokens == 3 + assert usage.cache_read_tokens == 2 diff --git a/uv.lock b/uv.lock index 694be9d6..97685442 100644 --- a/uv.lock +++ b/uv.lock @@ -25,6 +25,9 @@ dependencies = [ anthropic = [ { name = "anthropic" }, ] +google = [ + { name = "google-genai" }, +] mcp = [ { name = "mcp" }, ] @@ -42,6 +45,7 @@ vercel = [ dev = [ { name = "anthropic" }, { name = "async-solipsism" }, + { name = "google-genai" }, { name = "mcp" }, { name = "mypy" }, { name = "openai" }, @@ -64,6 +68,7 @@ examples = [ [package.metadata] requires-dist = [ { name = "anthropic", marker = "extra == 'anthropic'", specifier = ">=0.83.0" }, + { name = "google-genai", marker = "extra == 'google'", specifier = ">=2.0.0" }, { name = "httpx", specifier = ">=0.28.1" }, { name = "mcp", marker = "extra == 'mcp'", specifier = ">=1.18.0" }, { name = "modelsdotdev", specifier = "==0.*" }, @@ -73,12 +78,13 @@ requires-dist = [ { name = "typing-extensions", specifier = ">=4.15.0" }, { name = "vercel", marker = "extra == 'vercel'", specifier = ">=0.5.9" }, ] -provides-extras = ["anthropic", "mcp", "openai", "otel", "vercel"] +provides-extras = ["anthropic", "google", "mcp", "openai", "otel", "vercel"] [package.metadata.requires-dev] dev = [ { name = "anthropic", specifier = ">=0.83.0" }, { name = "async-solipsism", specifier = ">=0.9" }, + { name = "google-genai", specifier = ">=2.0.0" }, { name = "mcp", specifier = ">=1.18.0" }, { name = "mypy", specifier = "~=2.1.0" }, { name = "openai", specifier = ">=2.14.0" }, @@ -486,6 +492,45 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/5a/ff/2e4eca3ade2c22fe1dea7043b8ee9dabe47753349eb1b56a202de8af6349/fastapi-0.136.1-py3-none-any.whl", hash = "sha256:a6e9d7eeada96c93a4d69cb03836b44fa34e2854accb7244a1ece36cd4781c3f", size = 117683, upload-time = "2026-04-23T16:49:42.437Z" }, ] +[[package]] +name = "google-auth" +version = "2.55.2" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "cryptography" }, + { name = "pyasn1-modules" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/79/b9/e370d86fea3da13ec0256df30323dd26c0cb9c8c85f0c6ec42ac9df0106b/google_auth-2.55.2.tar.gz", hash = "sha256:97ae7790ff740f2bc9db60eb864a7804f4ac19f5f02c38b3d942f2fea6e9b9ae", size = 361414, upload-time = "2026-07-07T18:43:21.227Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/1e/c6/02eb5a337ac316a4c30c012e747bad5cea36e1a876efecdf80865541f7d8/google_auth-2.55.2-py3-none-any.whl", hash = "sha256:d715f265f2cafc6a5f1bf0dc19870d20e3119f6f6682785a250bce3d03d38a3b", size = 256778, upload-time = "2026-07-07T18:43:19.52Z" }, +] + +[package.optional-dependencies] +requests = [ + { name = "requests" }, +] + +[[package]] +name = "google-genai" +version = "2.12.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "anyio" }, + { name = "distro" }, + { name = "google-auth", extra = ["requests"] }, + { name = "httpx" }, + { name = "pydantic" }, + { name = "requests" }, + { name = "sniffio" }, + { name = "tenacity" }, + { name = "typing-extensions" }, + { name = "websockets" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/96/59/9ea84cbeb8f09694564d3b0ee9dd59003551b308d47b61f251415df93982/google_genai-2.12.1.tar.gz", hash = "sha256:78c25217885d63dc430ca7c4526853512b164a25a93a8a0d0af5b85971aa1db0", size = 636710, upload-time = "2026-07-16T16:15:02.035Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/b6/b4/1369fb413fc2ba7f78acace5590b6e9990c52ab5d1d166aafaa1ae2c28c8/google_genai-2.12.1-py3-none-any.whl", hash = "sha256:686d5ec39bda345151d3ed1bac3915f01f49138b1ea519af2eb98f11cc55ebc4", size = 1023403, upload-time = "2026-07-16T16:14:59.79Z" }, +] + [[package]] name = "googleapis-common-protos" version = "1.75.0" @@ -1017,6 +1062,27 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/c4/72/02445137af02769918a93807b2b7890047c32bfb9f90371cbc12688819eb/protobuf-6.33.6-py3-none-any.whl", hash = "sha256:77179e006c476e69bf8e8ce866640091ec42e1beb80b213c3900006ecfba6901", size = 170656, upload-time = "2026-03-18T19:04:59.826Z" }, ] +[[package]] +name = "pyasn1" +version = "0.6.4" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/a4/9a/23310166d960def5897e91fe20e5b724601b02a22e84ba1f94232c0b7f67/pyasn1-0.6.4.tar.gz", hash = "sha256:9c447d8431c947fe4c8febc4ed9e760bc29011a5b01e5c74b67025bd9fb8ce81", size = 151262, upload-time = "2026-07-09T01:12:33.988Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/9a/3b/6163796d69c3977d1e4287bea4a6979161cbbdd170ebb430511e8e1999ce/pyasn1-0.6.4-py3-none-any.whl", hash = "sha256:deda9277cfd454080ec40b207fb6df82206a3a2688735233cdcd8d3d565f088b", size = 84410, upload-time = "2026-07-09T01:12:32.92Z" }, +] + +[[package]] +name = "pyasn1-modules" +version = "0.4.2" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pyasn1" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/e9/e6/78ebbb10a8c8e4b61a59249394a4a594c1a7af95593dc933a349c8d00964/pyasn1_modules-0.4.2.tar.gz", hash = "sha256:677091de870a80aae844b1ca6134f54652fa2c8c5a52aa396440ac3106e941e6", size = 307892, upload-time = "2025-03-28T02:41:22.17Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/47/8d/d529b5d697919ba8c11ad626e835d4039be708a35b0d22de83a269a6682c/pyasn1_modules-0.4.2-py3-none-any.whl", hash = "sha256:29253a9207ce32b64c3ac6600edc75368f98473906e8fd1043bd6b5b1de2c14a", size = 181259, upload-time = "2025-03-28T02:41:19.028Z" }, +] + [[package]] name = "pycparser" version = "2.23" @@ -1418,6 +1484,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/67/81/c9b08609e2a92ecf62c97c59cabfa0608337c8d5cc9941eed5d9a7778840/temporalio-1.27.2-cp310-abi3-win_amd64.whl", hash = "sha256:62a84ae9a60c17932971e4ca3b0f3cd6f32f173b8183e759989376503fb95af6", size = 14981897, upload-time = "2026-05-14T02:17:27.333Z" }, ] +[[package]] +name = "tenacity" +version = "9.1.4" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/47/c6/ee486fd809e357697ee8a44d3d69222b344920433d3b6666ccd9b374630c/tenacity-9.1.4.tar.gz", hash = "sha256:adb31d4c263f2bd041081ab33b498309a57c77f9acf2db65aadf0898179cf93a", size = 49413, upload-time = "2026-02-07T10:45:33.841Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/d7/c1/eb8f9debc45d3b7918a32ab756658a0904732f75e555402972246b0b8e71/tenacity-9.1.4-py3-none-any.whl", hash = "sha256:6095a360c919085f28c6527de529e76a06ad89b23659fa881ae0649b867a9d55", size = 28926, upload-time = "2026-02-07T10:45:32.24Z" }, +] + [[package]] name = "textual" version = "8.2.6"