diff --git a/mem0/configs/llms/gemini.py b/mem0/configs/llms/gemini.py new file mode 100644 index 000000000..4b62fe0f8 --- /dev/null +++ b/mem0/configs/llms/gemini.py @@ -0,0 +1,64 @@ +import os +from typing import Optional + +from mem0.configs.llms.base import BaseLlmConfig + + +class GeminiConfig(BaseLlmConfig): + """ + Configuration class for Google Gemini LLM. + + Supports both the Gemini Developer API (via API key) and Vertex AI (via GCP credentials). + """ + + def __init__( + self, + # Base parameters + model: Optional[str] = None, + temperature: float = 0.1, + api_key: Optional[str] = None, + max_tokens: int = 2000, + top_p: float = 0.1, + top_k: int = 1, + enable_vision: bool = False, + vision_details: Optional[str] = "auto", + http_client_proxies: Optional[dict] = None, + # Gemini-specific parameters + vertexai: Optional[bool] = None, + project: Optional[str] = None, + location: Optional[str] = None, + ): + """ + Initialize Gemini configuration. + + Args: + model: Gemini model to use (e.g., "gemini-2.0-flash"), defaults to None + temperature: Controls randomness, defaults to 0.1 + api_key: Google API key for the Gemini Developer API, defaults to None + max_tokens: Maximum tokens to generate, defaults to 2000 + top_p: Nucleus sampling parameter, defaults to 0.1 + top_k: Top-k sampling parameter, defaults to 1 + enable_vision: Enable vision capabilities, defaults to False + vision_details: Vision detail level, defaults to "auto" + http_client_proxies: HTTP client proxy settings, defaults to None + vertexai: Whether to use Vertex AI backend. If None, checks GOOGLE_GENAI_USE_VERTEXAI env var. + project: GCP project ID for Vertex AI. If None, checks GOOGLE_CLOUD_PROJECT env var. + location: GCP location for Vertex AI. If None, checks GOOGLE_CLOUD_LOCATION env var. + """ + super().__init__( + model=model, + temperature=temperature, + api_key=api_key, + max_tokens=max_tokens, + top_p=top_p, + top_k=top_k, + enable_vision=enable_vision, + vision_details=vision_details, + http_client_proxies=http_client_proxies, + ) + + if vertexai is None: + vertexai = os.getenv("GOOGLE_GENAI_USE_VERTEXAI", "").lower() in ("true", "1", "yes") + self.vertexai = vertexai + self.project = project or os.getenv("GOOGLE_CLOUD_PROJECT") + self.location = location or os.getenv("GOOGLE_CLOUD_LOCATION", "us-central1") diff --git a/mem0/llms/gemini.py b/mem0/llms/gemini.py index 0160d4ad7..c46bcb894 100644 --- a/mem0/llms/gemini.py +++ b/mem0/llms/gemini.py @@ -1,5 +1,5 @@ import os -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Union try: from google import genai @@ -8,18 +8,39 @@ except ImportError: raise ImportError("The 'google-genai' library is required. Please install it using 'pip install google-genai'.") from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.gemini import GeminiConfig from mem0.llms.base import LLMBase class GeminiLLM(LLMBase): - def __init__(self, config: Optional[BaseLlmConfig] = None): + def __init__(self, config: Optional[Union[BaseLlmConfig, GeminiConfig, Dict]] = None): + # Convert to GeminiConfig if needed + if config is None: + config = GeminiConfig() + elif isinstance(config, dict): + config = GeminiConfig(**config) + elif isinstance(config, BaseLlmConfig) and not isinstance(config, GeminiConfig): + config = GeminiConfig( + model=config.model, + temperature=config.temperature, + api_key=config.api_key, + max_tokens=config.max_tokens, + top_p=config.top_p, + top_k=config.top_k, + enable_vision=config.enable_vision, + vision_details=config.vision_details, + ) + super().__init__(config) if not self.config.model: self.config.model = "gemini-2.0-flash" - api_key = self.config.api_key or os.getenv("GOOGLE_API_KEY") - self.client = genai.Client(api_key=api_key) + if self.config.vertexai: + self.client = genai.Client(vertexai=True, project=self.config.project, location=self.config.location) + else: + api_key = self.config.api_key or os.getenv("GOOGLE_API_KEY") + self.client = genai.Client(api_key=api_key) def _parse_response(self, response, tools): """ diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py index 487c10850..f6397ea90 100644 --- a/mem0/utils/factory.py +++ b/mem0/utils/factory.py @@ -7,6 +7,7 @@ from mem0.configs.llms.aws_bedrock import AWSBedrockConfig from mem0.configs.llms.azure import AzureOpenAIConfig from mem0.configs.llms.base import BaseLlmConfig from mem0.configs.llms.deepseek import DeepSeekConfig +from mem0.configs.llms.gemini import GeminiConfig from mem0.configs.llms.lmstudio import LMStudioConfig from mem0.configs.llms.minimax import MinimaxConfig from mem0.configs.llms.ollama import OllamaConfig @@ -46,7 +47,7 @@ class LlmFactory: "openai_structured": ("mem0.llms.openai_structured.OpenAIStructuredLLM", OpenAIConfig), "anthropic": ("mem0.llms.anthropic.AnthropicLLM", AnthropicConfig), "azure_openai_structured": ("mem0.llms.azure_openai_structured.AzureOpenAIStructuredLLM", AzureOpenAIConfig), - "gemini": ("mem0.llms.gemini.GeminiLLM", BaseLlmConfig), + "gemini": ("mem0.llms.gemini.GeminiLLM", GeminiConfig), "deepseek": ("mem0.llms.deepseek.DeepSeekLLM", DeepSeekConfig), "minimax": ("mem0.llms.minimax.MiniMaxLLM", MinimaxConfig), "xai": ("mem0.llms.xai.XAILLM", XAIConfig), diff --git a/tests/llms/test_gemini.py b/tests/llms/test_gemini.py index 4f0c8a9bd..4c67ffc39 100644 --- a/tests/llms/test_gemini.py +++ b/tests/llms/test_gemini.py @@ -4,6 +4,7 @@ import pytest from google.genai import types from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.gemini import GeminiConfig from mem0.llms.gemini import GeminiLLM @@ -229,3 +230,71 @@ def test_explicit_config_values_passed_to_generation_config(mock_gemini_client: assert config_arg.max_output_tokens == 200 assert "top_p" in config_arg.model_fields_set assert config_arg.top_p == 0.9 + + +# --- Vertex AI backend initialization (issue #3990, PR #4030) --- + + +def test_init_default_path_uses_api_key(monkeypatch): + """Default (non-Vertex) path stays backward compatible: client built with api_key, never vertexai.""" + monkeypatch.delenv("GOOGLE_GENAI_USE_VERTEXAI", raising=False) + monkeypatch.setenv("GOOGLE_API_KEY", "test-key") + with patch("mem0.llms.gemini.genai.Client") as mock_client_class: + GeminiLLM(GeminiConfig(model="gemini-2.0-flash")) + mock_client_class.assert_called_once_with(api_key="test-key") + + +def test_init_backward_compat_with_base_config(monkeypatch): + """A legacy BaseLlmConfig still works and uses the API-key path (no missing-attr error).""" + monkeypatch.delenv("GOOGLE_GENAI_USE_VERTEXAI", raising=False) + monkeypatch.setenv("GOOGLE_API_KEY", "legacy-key") + with patch("mem0.llms.gemini.genai.Client") as mock_client_class: + GeminiLLM(BaseLlmConfig(model="gemini-2.0-flash")) + mock_client_class.assert_called_once_with(api_key="legacy-key") + + +def test_init_vertexai_via_explicit_config(monkeypatch): + """vertexai=True in GeminiConfig routes to the Vertex AI client with project/location, no api_key.""" + monkeypatch.delenv("GOOGLE_GENAI_USE_VERTEXAI", raising=False) + with patch("mem0.llms.gemini.genai.Client") as mock_client_class: + GeminiLLM(GeminiConfig(vertexai=True, project="my-project", location="europe-west1")) + mock_client_class.assert_called_once_with(vertexai=True, project="my-project", location="europe-west1") + + +def test_init_vertexai_via_dict_config(monkeypatch): + """The factory hands GeminiLLM a dict; the vertexai key still routes to the Vertex client.""" + monkeypatch.delenv("GOOGLE_GENAI_USE_VERTEXAI", raising=False) + with patch("mem0.llms.gemini.genai.Client") as mock_client_class: + GeminiLLM({"vertexai": True, "project": "p", "location": "us-central1"}) + mock_client_class.assert_called_once_with(vertexai=True, project="p", location="us-central1") + + +def test_init_vertexai_via_env_vars(monkeypatch): + """GOOGLE_GENAI_USE_VERTEXAI + project/location env vars enable Vertex (the exact ask in issue #3990).""" + monkeypatch.setenv("GOOGLE_GENAI_USE_VERTEXAI", "true") + monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "env-project") + monkeypatch.setenv("GOOGLE_CLOUD_LOCATION", "asia-south1") + with patch("mem0.llms.gemini.genai.Client") as mock_client_class: + GeminiLLM(GeminiConfig()) + mock_client_class.assert_called_once_with(vertexai=True, project="env-project", location="asia-south1") + + +def test_init_vertexai_location_defaults_to_us_central1(monkeypatch): + """When Vertex is on but no location is supplied, it defaults to us-central1.""" + monkeypatch.delenv("GOOGLE_CLOUD_LOCATION", raising=False) + monkeypatch.setenv("GOOGLE_GENAI_USE_VERTEXAI", "true") + with patch("mem0.llms.gemini.genai.Client") as mock_client_class: + GeminiLLM(GeminiConfig(project="p")) + mock_client_class.assert_called_once_with(vertexai=True, project="p", location="us-central1") + + +def test_init_base_config_respects_vertexai_env(monkeypatch): + """GOOGLE_GENAI_USE_VERTEXAI is the authoritative global switch (Google's own + convention): a legacy BaseLlmConfig is routed to Vertex when the env var is set, + even if it carries an api_key. This pins the precedence as intentional.""" + monkeypatch.setenv("GOOGLE_GENAI_USE_VERTEXAI", "true") + monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "env-project") + monkeypatch.setenv("GOOGLE_CLOUD_LOCATION", "us-west1") + with patch("mem0.llms.gemini.genai.Client") as mock_client_class: + GeminiLLM(BaseLlmConfig(model="gemini-2.0-flash", api_key="ignored-key")) + mock_client_class.assert_called_once_with(vertexai=True, project="env-project", location="us-west1")