feat: support Gemini via Vertex AI as LLM provider (#4030)

Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
Aayush Soni
2026-06-16 16:59:06 +05:30
committed by GitHub
parent 7c841a2bce
commit d772f9a961
4 changed files with 160 additions and 5 deletions
+64
View File
@@ -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")
+23 -2
View File
@@ -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,16 +8,37 @@ 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"
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)
+2 -1
View File
@@ -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),
+69
View File
@@ -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")