feat: support Gemini via Vertex AI as LLM provider (#4030)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -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")
|
||||
+25
-4
@@ -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):
|
||||
"""
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user