From eccc77898fb92bade2119358f7549bb7c9749cc4 Mon Sep 17 00:00:00 2001 From: Nathan Brake Date: Fri, 5 Sep 2025 10:22:47 -0400 Subject: [PATCH] split google into gemini and vertexai --- pyproject.toml | 8 ++- src/any_llm/provider.py | 3 +- src/any_llm/providers/gemini/__init__.py | 3 + .../{google/google.py => gemini/base.py} | 64 ++++--------------- src/any_llm/providers/gemini/gemini.py | 26 ++++++++ .../providers/{google => gemini}/utils.py | 11 +--- src/any_llm/providers/google/__init__.py | 3 - src/any_llm/providers/vertexai/__init__.py | 3 + src/any_llm/providers/vertexai/vertexai.py | 35 ++++++++++ tests/conftest.py | 9 ++- tests/integration/test_embedding.py | 2 +- tests/integration/test_reasoning.py | 4 +- tests/unit/providers/test_google_provider.py | 54 ++++++++++------ tests/unit/providers/test_google_utils.py | 2 +- tests/unit/test_provider.py | 12 ++-- 15 files changed, 141 insertions(+), 98 deletions(-) create mode 100644 src/any_llm/providers/gemini/__init__.py rename src/any_llm/providers/{google/google.py => gemini/base.py} (76%) create mode 100644 src/any_llm/providers/gemini/gemini.py rename src/any_llm/providers/{google => gemini}/utils.py (93%) delete mode 100644 src/any_llm/providers/google/__init__.py create mode 100644 src/any_llm/providers/vertexai/__init__.py create mode 100644 src/any_llm/providers/vertexai/vertexai.py diff --git a/pyproject.toml b/pyproject.toml index 08681f148..07c565ca6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -19,7 +19,7 @@ dependencies = [ [project.optional-dependencies] all = [ - "any-llm-sdk[mistral,anthropic,huggingface,google,cohere,cerebras,fireworks,groq,aws,azure,azureopenai,watsonx,together,sambanova,ollama,moonshot,nebius,xai,databricks,deepseek,inception,openai,openrouter,portkey,lmstudio,llama,voyage,perplexity,llamafile,llamacpp]" + "any-llm-sdk[mistral,anthropic,huggingface,gemini,vertexai,cohere,cerebras,fireworks,groq,aws,azure,azureopenai,watsonx,together,sambanova,ollama,moonshot,nebius,xai,databricks,deepseek,inception,openai,openrouter,portkey,lmstudio,llama,voyage,perplexity,llamafile,llamacpp]" ] perplexity = [] @@ -32,7 +32,11 @@ anthropic = [ "anthropic", ] -google = [ +gemini = [ + "google-genai", +] + +vertexai = [ "google-genai", ] diff --git a/src/any_llm/provider.py b/src/any_llm/provider.py index 34858a1b9..f5c139498 100644 --- a/src/any_llm/provider.py +++ b/src/any_llm/provider.py @@ -43,7 +43,7 @@ class ProviderName(StrEnum): DATABRICKS = "databricks" DEEPSEEK = "deepseek" FIREWORKS = "fireworks" - GOOGLE = "google" + GEMINI = "gemini" GROQ = "groq" HUGGINGFACE = "huggingface" INCEPTION = "inception" @@ -60,6 +60,7 @@ class ProviderName(StrEnum): PORTKEY = "portkey" SAMBANOVA = "sambanova" TOGETHER = "together" + VERTEXAI = "vertexai" VOYAGE = "voyage" WATSONX = "watsonx" XAI = "xai" diff --git a/src/any_llm/providers/gemini/__init__.py b/src/any_llm/providers/gemini/__init__.py new file mode 100644 index 000000000..8a9116b4c --- /dev/null +++ b/src/any_llm/providers/gemini/__init__.py @@ -0,0 +1,3 @@ +from .gemini import GeminiProvider + +__all__ = ["GeminiProvider"] diff --git a/src/any_llm/providers/google/google.py b/src/any_llm/providers/gemini/base.py similarity index 76% rename from src/any_llm/providers/google/google.py rename to src/any_llm/providers/gemini/base.py index b4e30b361..bc4f00ad0 100644 --- a/src/any_llm/providers/google/google.py +++ b/src/any_llm/providers/gemini/base.py @@ -1,10 +1,10 @@ -import os +from abc import abstractmethod from collections.abc import AsyncIterator, Sequence -from typing import Any +from typing import TYPE_CHECKING, Any, Literal, cast from pydantic import BaseModel -from any_llm.exceptions import MissingApiKeyError, UnsupportedParameterError +from any_llm.exceptions import UnsupportedParameterError from any_llm.provider import ClientConfig, Provider from any_llm.types.completion import ( ChatCompletion, @@ -38,16 +38,14 @@ except ImportError as e: MISSING_PACKAGES_ERROR = e -# From https://ai.google.dev/gemini-api/docs/openai#thinking +if TYPE_CHECKING: + from google import genai + REASONING_EFFORT_TO_THINKING_BUDGETS = {"minimal": 256, "low": 1024, "medium": 8192, "high": 24576} class GoogleProvider(Provider): - """Google Provider using the new response conversion utilities.""" - - PROVIDER_NAME = "google" - PROVIDER_DOCUMENTATION_URL = "https://cloud.google.com/vertex-ai/docs" - ENV_API_KEY_NAME = "GOOGLE_API_KEY/GEMINI_API_KEY" + """Base Google Provider class with common functionality for Gemini and Vertex AI.""" SUPPORTS_COMPLETION_STREAMING = True SUPPORTS_COMPLETION = True @@ -58,35 +56,9 @@ class GoogleProvider(Provider): MISSING_PACKAGES_ERROR = MISSING_PACKAGES_ERROR - def __init__(self, config: ClientConfig) -> None: - """Initialize Google GenAI provider.""" - self._verify_no_missing_packages() - self.config = config - self.use_vertex_ai = os.getenv("GOOGLE_USE_VERTEX_AI", "false").lower() == "true" - - def _get_client(self, use_vertex_ai: bool, config: ClientConfig) -> "genai.Client": - if use_vertex_ai: - project_id = os.getenv("GOOGLE_PROJECT_ID") - location = os.getenv("GOOGLE_REGION", "us-central1") - - if not project_id: - msg = "Google Vertex AI" - raise MissingApiKeyError(msg, "GOOGLE_PROJECT_ID") - - return genai.Client( - vertexai=True, - project=project_id, - location=location, - **(config.client_args if config.client_args else {}), - ) - - api_key = getattr(config, "api_key", None) or os.getenv("GEMINI_API_KEY") or os.getenv("GOOGLE_API_KEY") - - if not api_key: - msg = "Google Gemini Developer API" - raise MissingApiKeyError(msg, "GEMINI_API_KEY/GOOGLE_API_KEY") - - return genai.Client(api_key=api_key, **(config.client_args if config.client_args else {})) + @abstractmethod + def _get_client(self, config: ClientConfig) -> "genai.Client": + """Get the appropriate client for this provider implementation.""" async def aembedding( self, @@ -94,7 +66,7 @@ async def aembedding( inputs: str | list[str], **kwargs: Any, ) -> CreateEmbeddingResponse: - client = self._get_client(self.use_vertex_ai, self.config) + client = self._get_client(self.config) result = await client.aio.models.embed_content( model=model, contents=inputs, # type: ignore[arg-type] @@ -125,7 +97,6 @@ async def acompletion( if params.reasoning_effort is None: kwargs["thinking_config"] = types.ThinkingConfig(include_thoughts=False) - # in "auto" mode, we just don't pass a `thinking_config` elif params.reasoning_effort != "auto": kwargs["thinking_config"] = types.ThinkingConfig( include_thoughts=True, thinking_budget=REASONING_EFFORT_TO_THINKING_BUDGETS[params.reasoning_effort] @@ -133,7 +104,6 @@ async def acompletion( stream = bool(params.stream) response_format = params.response_format - # Build generation config without duplicating keys (e.g., tools) base_kwargs = params.model_dump( exclude_none=True, exclude={ @@ -148,7 +118,6 @@ async def acompletion( }, ) - # Convert max_tokens to max_output_tokens for Google if params.max_tokens is not None: base_kwargs["max_output_tokens"] = params.max_tokens @@ -162,7 +131,7 @@ async def acompletion( if system_instruction: generation_config.system_instruction = system_instruction - client = self._get_client(self.use_vertex_ai, self.config) + client = self._get_client(self.config) if stream: response_stream = await client.aio.models.generate_content_stream( model=params.model_id, @@ -184,7 +153,6 @@ async def _stream() -> AsyncIterator[ChatCompletionChunk]: response_dict = _convert_response_to_response_dict(response) - # Directly construct ChatCompletion choices_out: list[Choice] = [] for i, choice_item in enumerate(response_dict.get("choices", [])): message_dict: dict[str, Any] = choice_item.get("message", {}) @@ -211,8 +179,6 @@ async def _stream() -> AsyncIterator[ChatCompletionChunk]: tool_calls=tool_calls, reasoning=Reasoning(content=reasoning_content) if reasoning_content else None, ) - from typing import Literal, cast - choices_out.append( Choice( index=i, @@ -241,9 +207,7 @@ async def _stream() -> AsyncIterator[ChatCompletionChunk]: ) def list_models(self, **kwargs: Any) -> Sequence[Model]: - """ - Fetch available models from the /v1/models endpoint. - """ - client = self._get_client(self.use_vertex_ai, self.config) + """Fetch available models from the /v1/models endpoint.""" + client = self._get_client(self.config) models_list = client.models.list(**kwargs) return _convert_models_list(models_list) diff --git a/src/any_llm/providers/gemini/gemini.py b/src/any_llm/providers/gemini/gemini.py new file mode 100644 index 000000000..136a7aea4 --- /dev/null +++ b/src/any_llm/providers/gemini/gemini.py @@ -0,0 +1,26 @@ +import os + +from google import genai + +from any_llm.exceptions import MissingApiKeyError +from any_llm.provider import ClientConfig + +from .base import GoogleProvider + + +class GeminiProvider(GoogleProvider): + """Gemini Provider using the Google GenAI Developer API.""" + + PROVIDER_NAME = "gemini" + PROVIDER_DOCUMENTATION_URL = "https://ai.google.dev/gemini-api/docs" + ENV_API_KEY_NAME = "GEMINI_API_KEY/GOOGLE_API_KEY" + + def _get_client(self, config: ClientConfig) -> "genai.Client": + """Get Gemini API client.""" + api_key = getattr(config, "api_key", None) or os.getenv("GEMINI_API_KEY") or os.getenv("GOOGLE_API_KEY") + + if not api_key: + msg = "Google Gemini Developer API" + raise MissingApiKeyError(msg, "GEMINI_API_KEY/GOOGLE_API_KEY") + + return genai.Client(api_key=api_key, **(config.client_args if config.client_args else {})) diff --git a/src/any_llm/providers/google/utils.py b/src/any_llm/providers/gemini/utils.py similarity index 93% rename from src/any_llm/providers/google/utils.py rename to src/any_llm/providers/gemini/utils.py index 3f2b6f78b..347d09195 100644 --- a/src/any_llm/providers/google/utils.py +++ b/src/any_llm/providers/gemini/utils.py @@ -26,7 +26,6 @@ def _convert_tool_spec(openai_tools: list[dict[str, Any]]) -> list[types.Tool]: continue function = tool["function"] - # Preserve nested schema details such as items/additionalProperties for arrays/objects properties: dict[str, dict[str, Any]] = {} for param_name, param_info in function["parameters"]["properties"].items(): prop: dict[str, Any] = { @@ -35,12 +34,10 @@ def _convert_tool_spec(openai_tools: list[dict[str, Any]]) -> list[types.Tool]: } if "enum" in param_info: prop["enum"] = param_info["enum"] - # Google requires explicit items for arrays if "items" in param_info: prop["items"] = param_info["items"] if prop.get("type") == "array" and "items" not in prop: prop["items"] = {"type": "string"} - # Google tool schema does not accept additionalProperties; drop it properties[param_name] = prop parameters_dict = { @@ -85,7 +82,7 @@ def _convert_messages(messages: list[dict[str, Any]]) -> tuple[list[types.Conten formatted_messages.append(types.Content(role="user", parts=parts)) elif message["role"] == "assistant": if message.get("tool_calls"): - tool_call = message["tool_calls"][0] # Assuming single function call for now + tool_call = message["tool_calls"][0] function_call = tool_call["function"] parts = [ @@ -202,7 +199,6 @@ def _create_openai_embedding_response_from_google( if embedding.values ] - # Google does not provide usage data in the embedding response usage = Usage(prompt_tokens=0, total_tokens=0) return CreateEmbeddingResponse( @@ -228,10 +224,8 @@ def _create_openai_chunk_from_google_chunk( for part in candidate.content.parts: if part.thought: - # This is a thinking/reasoning part reasoning_content += part.text or "" else: - # Regular content part content += part.text or "" delta = ChoiceDelta( @@ -247,7 +241,7 @@ def _create_openai_chunk_from_google_chunk( ) return ChatCompletionChunk( - id=f"chatcmpl-{time()}", # Google doesn't provide an ID in the chunk + id=f"chatcmpl-{time()}", choices=[choice], created=int(time()), model=str(response.model_version), @@ -256,5 +250,4 @@ def _create_openai_chunk_from_google_chunk( def _convert_models_list(models_list: Pager[types.Model]) -> list[Model]: - # Google doesn't provide a creation date for models return [Model(id=model.name or "Unknown", object="model", created=0, owned_by="google") for model in models_list] diff --git a/src/any_llm/providers/google/__init__.py b/src/any_llm/providers/google/__init__.py deleted file mode 100644 index 6f5834801..000000000 --- a/src/any_llm/providers/google/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from .google import GoogleProvider - -__all__ = ["GoogleProvider"] diff --git a/src/any_llm/providers/vertexai/__init__.py b/src/any_llm/providers/vertexai/__init__.py new file mode 100644 index 000000000..46095028f --- /dev/null +++ b/src/any_llm/providers/vertexai/__init__.py @@ -0,0 +1,3 @@ +from .vertexai import VertexaiProvider + +__all__ = ["VertexaiProvider"] diff --git a/src/any_llm/providers/vertexai/vertexai.py b/src/any_llm/providers/vertexai/vertexai.py new file mode 100644 index 000000000..b46e9f82e --- /dev/null +++ b/src/any_llm/providers/vertexai/vertexai.py @@ -0,0 +1,35 @@ +import os +from typing import TYPE_CHECKING + +from any_llm.exceptions import MissingApiKeyError +from any_llm.provider import ClientConfig +from any_llm.providers.gemini.base import GoogleProvider + +if TYPE_CHECKING: + from google import genai + + +class VertexaiProvider(GoogleProvider): + """Vertex AI Provider using Google Cloud Vertex AI.""" + + PROVIDER_NAME = "vertexai" + PROVIDER_DOCUMENTATION_URL = "https://cloud.google.com/vertex-ai/docs" + ENV_API_KEY_NAME = "GOOGLE_PROJECT_ID" + + def _get_client(self, config: ClientConfig) -> "genai.Client": + """Get Vertex AI client.""" + from google import genai + + project_id = os.getenv("GOOGLE_PROJECT_ID") + location = os.getenv("GOOGLE_REGION", "us-central1") + + if not project_id: + msg = "Google Vertex AI" + raise MissingApiKeyError(msg, "GOOGLE_PROJECT_ID") + + return genai.Client( + vertexai=True, + project=project_id, + location=location, + **(config.client_args if config.client_args else {}), + ) diff --git a/tests/conftest.py b/tests/conftest.py index f9d20ee16..359499778 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -10,7 +10,8 @@ def provider_reasoning_model_map() -> dict[ProviderName, str]: return { ProviderName.ANTHROPIC: "claude-sonnet-4-20250514", ProviderName.MISTRAL: "magistral-small-latest", - ProviderName.GOOGLE: "gemini-2.5-flash", + ProviderName.GEMINI: "gemini-2.5-flash", + ProviderName.VERTEXAI: "gemini-2.5-flash", ProviderName.GROQ: "openai/gpt-oss-20b", ProviderName.FIREWORKS: "accounts/fireworks/models/deepseek-r1", ProviderName.OPENAI: "gpt-5-nano", @@ -33,7 +34,8 @@ def provider_model_map() -> dict[ProviderName, str]: ProviderName.DATABRICKS: "databricks-meta-llama-3-1-8b-instruct", ProviderName.DEEPSEEK: "deepseek-chat", ProviderName.OPENAI: "gpt-5-nano", - ProviderName.GOOGLE: "gemini-2.5-flash", + ProviderName.GEMINI: "gemini-2.5-flash", + ProviderName.VERTEXAI: "gemini-2.5-flash", ProviderName.MOONSHOT: "moonshot-v1-8k", ProviderName.SAMBANOVA: "Meta-Llama-3.1-8B-Instruct", ProviderName.TOGETHER: "meta-llama/Llama-3.3-70B-Instruct-Turbo-Free", @@ -73,7 +75,8 @@ def embedding_provider_model_map() -> dict[ProviderName, str]: ProviderName.OLLAMA: "gpt-oss:20b", ProviderName.LLAMAFILE: "N/A", ProviderName.LMSTUDIO: "text-embedding-nomic-embed-text-v1.5", - ProviderName.GOOGLE: "gemini-embedding-001", + ProviderName.GEMINI: "gemini-embedding-001", + ProviderName.VERTEXAI: "gemini-embedding-001", ProviderName.AZURE: "openai/text-embedding-3-small", ProviderName.AZUREOPENAI: "azure/", ProviderName.VOYAGE: "voyage-3.5-lite", diff --git a/tests/integration/test_embedding.py b/tests/integration/test_embedding.py index 56ea220ff..3c22c3b16 100644 --- a/tests/integration/test_embedding.py +++ b/tests/integration/test_embedding.py @@ -38,6 +38,6 @@ async def test_embedding_providers_async( assert len(result.data) > 0 for entry in result.data: assert all(isinstance(v, float) for v in entry.embedding) - if provider not in (ProviderName.GOOGLE, ProviderName.LMSTUDIO): + if provider not in (ProviderName.GEMINI, ProviderName.VERTEXAI, ProviderName.LMSTUDIO): assert result.usage.prompt_tokens > 0 assert result.usage.total_tokens > 0 diff --git a/tests/integration/test_reasoning.py b/tests/integration/test_reasoning.py index adc2143df..9c018e43e 100644 --- a/tests/integration/test_reasoning.py +++ b/tests/integration/test_reasoning.py @@ -25,7 +25,7 @@ async def test_completion_reasoning( model_id = provider_reasoning_model_map[provider] extra_kwargs = provider_extra_kwargs_map.get(provider, {}) - if provider in (ProviderName.ANTHROPIC, ProviderName.GOOGLE, ProviderName.OLLAMA): + if provider in (ProviderName.ANTHROPIC, ProviderName.GEMINI, ProviderName.VERTEXAI, ProviderName.OLLAMA): extra_kwargs["reasoning_effort"] = "low" try: @@ -62,7 +62,7 @@ async def test_completion_reasoning_streaming( model_id = provider_reasoning_model_map[provider] extra_kwargs = provider_extra_kwargs_map.get(provider, {}) - if provider in (ProviderName.ANTHROPIC, ProviderName.GOOGLE, ProviderName.OLLAMA): + if provider in (ProviderName.ANTHROPIC, ProviderName.GEMINI, ProviderName.VERTEXAI, ProviderName.OLLAMA): extra_kwargs["reasoning_effort"] = "low" try: diff --git a/tests/unit/providers/test_google_provider.py b/tests/unit/providers/test_google_provider.py index f2322e26d..03575b942 100644 --- a/tests/unit/providers/test_google_provider.py +++ b/tests/unit/providers/test_google_provider.py @@ -6,20 +6,29 @@ from google.genai import types from any_llm.exceptions import UnsupportedParameterError -from any_llm.provider import ClientConfig -from any_llm.providers.google.google import REASONING_EFFORT_TO_THINKING_BUDGETS, GoogleProvider +from any_llm.provider import ClientConfig, Provider +from any_llm.providers.gemini import GeminiProvider +from any_llm.providers.gemini.base import REASONING_EFFORT_TO_THINKING_BUDGETS +from any_llm.providers.vertexai import VertexaiProvider from any_llm.types.completion import CompletionParams +@pytest.fixture(params=[GeminiProvider, VertexaiProvider]) +def google_provider_class(request: pytest.FixtureRequest) -> type[Provider]: + """Parametrized fixture that provides both GeminiProvider and VertexaiProvider classes.""" + return request.param # type: ignore[no-any-return] + + @contextmanager def mock_google_provider(): # type: ignore[no-untyped-def] with ( - patch("any_llm.providers.google.google.genai.Client") as mock_genai, - patch("any_llm.providers.google.google._convert_response_to_response_dict") as mock_convert_response, + patch("any_llm.providers.gemini.base.genai.Client") as mock_genai, + patch("any_llm.providers.gemini.base._convert_response_to_response_dict") as mock_convert_response, + patch.dict("os.environ", {"GOOGLE_PROJECT_ID": "test-project", "GOOGLE_REGION": "us-central1"}), ): mock_convert_response.return_value = { "id": "google_genai_response", - "model": "google/genai", + "model": "gemini/genai", "created": 0, "choices": [ { @@ -39,14 +48,14 @@ def mock_google_provider(): # type: ignore[no-untyped-def] @pytest.mark.asyncio -async def test_completion_with_system_instruction() -> None: +async def test_completion_with_system_instruction(google_provider_class: type[Provider]) -> None: """Test that completion works correctly with system_instruction.""" api_key = "test-api-key" model = "gemini-pro" messages = [{"role": "system", "content": "You are a helpful assistant"}, {"role": "user", "content": "Hello"}] with mock_google_provider() as mock_genai: - provider = GoogleProvider(ClientConfig(api_key=api_key)) + provider = google_provider_class(ClientConfig(api_key=api_key)) await provider.acompletion(CompletionParams(model_id=model, messages=messages)) _, call_kwargs = mock_genai.return_value.aio.models.generate_content.call_args @@ -65,14 +74,16 @@ async def test_completion_with_system_instruction() -> None: ], ) @pytest.mark.asyncio -async def test_completion_with_tool_choice_auto(tool_choice: str, expected_mode: str) -> None: +async def test_completion_with_tool_choice_auto( + google_provider_class: type[Provider], tool_choice: str, expected_mode: str +) -> None: """Test that completion correctly processes tool_choice='auto'.""" api_key = "test-api-key" model = "gemini-pro" messages = [{"role": "user", "content": "Hello"}] with mock_google_provider() as mock_genai: - provider = GoogleProvider(ClientConfig(api_key=api_key)) + provider = google_provider_class(ClientConfig(api_key=api_key)) await provider.acompletion(CompletionParams(model_id=model, messages=messages, tool_choice=tool_choice)) _, call_kwargs = mock_genai.return_value.aio.models.generate_content.call_args @@ -82,14 +93,14 @@ async def test_completion_with_tool_choice_auto(tool_choice: str, expected_mode: @pytest.mark.asyncio -async def test_completion_without_tool_choice() -> None: +async def test_completion_without_tool_choice(google_provider_class: type[Provider]) -> None: """Test that completion works correctly without tool_choice.""" api_key = "test-api-key" model = "gemini-pro" messages = [{"role": "user", "content": "Hello"}] with mock_google_provider() as mock_genai: - provider = GoogleProvider(ClientConfig(api_key=api_key)) + provider = google_provider_class(ClientConfig(api_key=api_key)) await provider.acompletion(CompletionParams(model_id=model, messages=messages)) _, call_kwargs = mock_genai.return_value.aio.models.generate_content.call_args @@ -99,13 +110,13 @@ async def test_completion_without_tool_choice() -> None: @pytest.mark.asyncio -async def test_completion_with_stream_and_response_format_raises() -> None: +async def test_completion_with_stream_and_response_format_raises(google_provider_class: type[Provider]) -> None: api_key = "test-api-key" model = "gemini-pro" messages = [{"role": "user", "content": "Hello"}] with mock_google_provider(): - provider = GoogleProvider(ClientConfig(api_key=api_key)) + provider = google_provider_class(ClientConfig(api_key=api_key)) with pytest.raises(UnsupportedParameterError): await provider.acompletion( CompletionParams( @@ -118,13 +129,13 @@ async def test_completion_with_stream_and_response_format_raises() -> None: @pytest.mark.asyncio -async def test_completion_with_parallel_tool_calls_raises() -> None: +async def test_completion_with_parallel_tool_calls_raises(google_provider_class: type[Provider]) -> None: api_key = "test-api-key" model = "gemini-pro" messages = [{"role": "user", "content": "Hello"}] with mock_google_provider(): - provider = GoogleProvider(ClientConfig(api_key=api_key)) + provider = google_provider_class(ClientConfig(api_key=api_key)) with pytest.raises(UnsupportedParameterError): await provider.acompletion( CompletionParams( @@ -136,12 +147,14 @@ async def test_completion_with_parallel_tool_calls_raises() -> None: @pytest.mark.asyncio -async def test_completion_inside_agent_loop(agent_loop_messages: list[dict[str, Any]]) -> None: +async def test_completion_inside_agent_loop( + google_provider_class: type[Provider], agent_loop_messages: list[dict[str, Any]] +) -> None: api_key = "test-api-key" model = "gemini-pro" with mock_google_provider() as mock_genai: - provider = GoogleProvider(ClientConfig(api_key=api_key)) + provider = google_provider_class(ClientConfig(api_key=api_key)) await provider.acompletion(CompletionParams(model_id=model, messages=agent_loop_messages)) _, call_kwargs = mock_genai.return_value.aio.models.generate_content.call_args @@ -164,6 +177,7 @@ async def test_completion_inside_agent_loop(agent_loop_messages: list[dict[str, ) @pytest.mark.asyncio async def test_completion_with_custom_reasoning_effort( + google_provider_class: type[Provider], reasoning_effort: Literal["low", "medium", "high"] | None, ) -> None: api_key = "test-api-key" @@ -171,7 +185,7 @@ async def test_completion_with_custom_reasoning_effort( messages = [{"role": "user", "content": "Hello"}] with mock_google_provider() as mock_genai: - provider = GoogleProvider(ClientConfig(api_key=api_key)) + provider = google_provider_class(ClientConfig(api_key=api_key)) await provider.acompletion( CompletionParams(model_id=model, messages=messages, reasoning_effort=reasoning_effort) ) @@ -187,7 +201,7 @@ async def test_completion_with_custom_reasoning_effort( @pytest.mark.asyncio -async def test_completion_with_max_tokens_conversion() -> None: +async def test_completion_with_max_tokens_conversion(google_provider_class: type[Provider]) -> None: """Test that max_tokens parameter gets converted to max_output_tokens.""" api_key = "test-api-key" model = "gemini-pro" @@ -195,7 +209,7 @@ async def test_completion_with_max_tokens_conversion() -> None: max_tokens = 100 with mock_google_provider() as mock_genai: - provider = GoogleProvider(ClientConfig(api_key=api_key)) + provider = google_provider_class(ClientConfig(api_key=api_key)) await provider.acompletion(CompletionParams(model_id=model, messages=messages, max_tokens=max_tokens)) _, call_kwargs = mock_genai.return_value.aio.models.generate_content.call_args diff --git a/tests/unit/providers/test_google_utils.py b/tests/unit/providers/test_google_utils.py index 5fbddad16..fe7e96e70 100644 --- a/tests/unit/providers/test_google_utils.py +++ b/tests/unit/providers/test_google_utils.py @@ -1,4 +1,4 @@ -from any_llm.providers.google.utils import _convert_tool_spec +from any_llm.providers.gemini.utils import _convert_tool_spec def test_convert_tool_spec_basic_mapping() -> None: diff --git a/tests/unit/test_provider.py b/tests/unit/test_provider.py index b7f5bcd40..61ade28d6 100644 --- a/tests/unit/test_provider.py +++ b/tests/unit/test_provider.py @@ -108,7 +108,7 @@ def test_unsupported_provider_error_attributes() -> None: assert "Supported providers:" in str(e) -def test_all_providers_have_required_attributes(provider: str) -> None: +def test_all_providers_have_required_attributes(provider: ProviderName) -> None: """Test that all supported providers can be loaded with sample config parameters. This test verifies that providers can handle common configuration parameters @@ -116,7 +116,7 @@ def test_all_providers_have_required_attributes(provider: str) -> None: """ sample_config = ClientConfig(api_key="test_key", api_base="https://test.example.com") - provider_instance = ProviderFactory.create_provider(provider, sample_config) + provider_instance = ProviderFactory.create_provider(provider.value, sample_config) assert provider_instance.PROVIDER_NAME is not None assert provider_instance.PROVIDER_DOCUMENTATION_URL is not None @@ -128,12 +128,12 @@ def test_all_providers_have_required_attributes(provider: str) -> None: assert provider_instance.SUPPORTS_RESPONSES is not None -def test_providers_raise_MissingApiKeyError(provider: str) -> None: - if provider in ("aws", "google", "ollama", "lmstudio", "llamafile"): +def test_providers_raise_MissingApiKeyError(provider: ProviderName) -> None: + if provider.value in ("aws", "ollama", "lmstudio", "llamafile"): pytest.skip("This provider handles `api_key` differently.") with patch.dict(os.environ, {}, clear=True): with pytest.raises(MissingApiKeyError): - ProviderFactory.create_provider(provider, ClientConfig()) + ProviderFactory.create_provider(provider.value, ClientConfig()) @pytest.mark.parametrize( @@ -144,7 +144,7 @@ def test_providers_raise_MissingApiKeyError(provider: str) -> None: ("azure", "azure"), ("cerebras", "cerebras"), ("cohere", "cohere"), - ("google", "google"), + ("gemini", "google"), ("groq", "groq"), ("huggingface", "huggingface_hub"), ("mistral", "mistralai"),