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
+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")