mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat: Add a service registry to provide a generic way to register custom service implementations to be used in FastAPI server
To register a custom service: - Create a factory function that takes a URI and returns an instance of your custom service. This function will parse any details it needs from the URI. - Register your factory with the global service registry. You need to define a unique URI scheme for your service (e.g., custom). PiperOrigin-RevId: 822310466
This commit is contained in:
committed by
Copybara-Service
parent
409df1378f
commit
391628fcdc
@@ -33,18 +33,16 @@ from opentelemetry.sdk.trace import TracerProvider
|
|||||||
from starlette.types import Lifespan
|
from starlette.types import Lifespan
|
||||||
from watchdog.observers import Observer
|
from watchdog.observers import Observer
|
||||||
|
|
||||||
from ..artifacts.gcs_artifact_service import GcsArtifactService
|
|
||||||
from ..artifacts.in_memory_artifact_service import InMemoryArtifactService
|
from ..artifacts.in_memory_artifact_service import InMemoryArtifactService
|
||||||
from ..auth.credential_service.in_memory_credential_service import InMemoryCredentialService
|
from ..auth.credential_service.in_memory_credential_service import InMemoryCredentialService
|
||||||
from ..evaluation.local_eval_set_results_manager import LocalEvalSetResultsManager
|
from ..evaluation.local_eval_set_results_manager import LocalEvalSetResultsManager
|
||||||
from ..evaluation.local_eval_sets_manager import LocalEvalSetsManager
|
from ..evaluation.local_eval_sets_manager import LocalEvalSetsManager
|
||||||
from ..memory.in_memory_memory_service import InMemoryMemoryService
|
from ..memory.in_memory_memory_service import InMemoryMemoryService
|
||||||
from ..memory.vertex_ai_memory_bank_service import VertexAiMemoryBankService
|
|
||||||
from ..runners import Runner
|
from ..runners import Runner
|
||||||
from ..sessions.in_memory_session_service import InMemorySessionService
|
from ..sessions.in_memory_session_service import InMemorySessionService
|
||||||
from ..sessions.vertex_ai_session_service import VertexAiSessionService
|
|
||||||
from ..utils.feature_decorator import working_in_progress
|
from ..utils.feature_decorator import working_in_progress
|
||||||
from .adk_web_server import AdkWebServer
|
from .adk_web_server import AdkWebServer
|
||||||
|
from .service_registry import get_service_registry
|
||||||
from .utils import envs
|
from .utils import envs
|
||||||
from .utils import evals
|
from .utils import evals
|
||||||
from .utils.agent_change_handler import AgentChangeEventHandler
|
from .utils.agent_change_handler import AgentChangeEventHandler
|
||||||
@@ -85,54 +83,14 @@ def get_fast_api_app(
|
|||||||
eval_sets_manager = LocalEvalSetsManager(agents_dir=agents_dir)
|
eval_sets_manager = LocalEvalSetsManager(agents_dir=agents_dir)
|
||||||
eval_set_results_manager = LocalEvalSetResultsManager(agents_dir=agents_dir)
|
eval_set_results_manager = LocalEvalSetResultsManager(agents_dir=agents_dir)
|
||||||
|
|
||||||
def _parse_agent_engine_resource_name(agent_engine_id_or_resource_name):
|
service_registry = get_service_registry()
|
||||||
if not agent_engine_id_or_resource_name:
|
|
||||||
raise click.ClickException(
|
|
||||||
"Agent engine resource name or resource id can not be empty."
|
|
||||||
)
|
|
||||||
|
|
||||||
# "projects/my-project/locations/us-central1/reasoningEngines/1234567890",
|
|
||||||
if "/" in agent_engine_id_or_resource_name:
|
|
||||||
# Validate resource name.
|
|
||||||
if len(agent_engine_id_or_resource_name.split("/")) != 6:
|
|
||||||
raise click.ClickException(
|
|
||||||
"Agent engine resource name is mal-formatted. It should be of"
|
|
||||||
" format :"
|
|
||||||
" projects/{project_id}/locations/{location}/reasoningEngines/{resource_id}"
|
|
||||||
)
|
|
||||||
project = agent_engine_id_or_resource_name.split("/")[1]
|
|
||||||
location = agent_engine_id_or_resource_name.split("/")[3]
|
|
||||||
agent_engine_id = agent_engine_id_or_resource_name.split("/")[-1]
|
|
||||||
else:
|
|
||||||
envs.load_dotenv_for_agent("", agents_dir)
|
|
||||||
project = os.environ.get("GOOGLE_CLOUD_PROJECT", None)
|
|
||||||
location = os.environ.get("GOOGLE_CLOUD_LOCATION", None)
|
|
||||||
agent_engine_id = agent_engine_id_or_resource_name
|
|
||||||
return project, location, agent_engine_id
|
|
||||||
|
|
||||||
# Build the Memory service
|
# Build the Memory service
|
||||||
if memory_service_uri:
|
if memory_service_uri:
|
||||||
if memory_service_uri.startswith("rag://"):
|
memory_service = service_registry.create_memory_service(
|
||||||
from ..memory.vertex_ai_rag_memory_service import VertexAiRagMemoryService
|
memory_service_uri, agents_dir=agents_dir
|
||||||
|
)
|
||||||
rag_corpus = memory_service_uri.split("://")[1]
|
if not memory_service:
|
||||||
if not rag_corpus:
|
|
||||||
raise click.ClickException("Rag corpus can not be empty.")
|
|
||||||
envs.load_dotenv_for_agent("", agents_dir)
|
|
||||||
memory_service = VertexAiRagMemoryService(
|
|
||||||
rag_corpus=f'projects/{os.environ["GOOGLE_CLOUD_PROJECT"]}/locations/{os.environ["GOOGLE_CLOUD_LOCATION"]}/ragCorpora/{rag_corpus}'
|
|
||||||
)
|
|
||||||
elif memory_service_uri.startswith("agentengine://"):
|
|
||||||
agent_engine_id_or_resource_name = memory_service_uri.split("://")[1]
|
|
||||||
project, location, agent_engine_id = _parse_agent_engine_resource_name(
|
|
||||||
agent_engine_id_or_resource_name
|
|
||||||
)
|
|
||||||
memory_service = VertexAiMemoryBankService(
|
|
||||||
project=project,
|
|
||||||
location=location,
|
|
||||||
agent_engine_id=agent_engine_id,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
raise click.ClickException(
|
raise click.ClickException(
|
||||||
"Unsupported memory service URI: %s" % memory_service_uri
|
"Unsupported memory service URI: %s" % memory_service_uri
|
||||||
)
|
)
|
||||||
@@ -141,34 +99,27 @@ def get_fast_api_app(
|
|||||||
|
|
||||||
# Build the Session service
|
# Build the Session service
|
||||||
if session_service_uri:
|
if session_service_uri:
|
||||||
if session_service_uri.startswith("agentengine://"):
|
session_kwargs = session_db_kwargs or {}
|
||||||
agent_engine_id_or_resource_name = session_service_uri.split("://")[1]
|
session_service = service_registry.create_session_service(
|
||||||
project, location, agent_engine_id = _parse_agent_engine_resource_name(
|
session_service_uri, agents_dir=agents_dir, **session_kwargs
|
||||||
agent_engine_id_or_resource_name
|
)
|
||||||
)
|
if not session_service:
|
||||||
session_service = VertexAiSessionService(
|
# Fallback to DatabaseSessionService if the service registry doesn't
|
||||||
project=project,
|
# support the session service URI scheme.
|
||||||
location=location,
|
|
||||||
agent_engine_id=agent_engine_id,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
from ..sessions.database_session_service import DatabaseSessionService
|
from ..sessions.database_session_service import DatabaseSessionService
|
||||||
|
|
||||||
# Database session additional settings
|
|
||||||
if session_db_kwargs is None:
|
|
||||||
session_db_kwargs = {}
|
|
||||||
session_service = DatabaseSessionService(
|
session_service = DatabaseSessionService(
|
||||||
db_url=session_service_uri, **session_db_kwargs
|
db_url=session_service_uri, **session_kwargs
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
session_service = InMemorySessionService()
|
session_service = InMemorySessionService()
|
||||||
|
|
||||||
# Build the Artifact service
|
# Build the Artifact service
|
||||||
if artifact_service_uri:
|
if artifact_service_uri:
|
||||||
if artifact_service_uri.startswith("gs://"):
|
artifact_service = service_registry.create_artifact_service(
|
||||||
gcs_bucket = artifact_service_uri.split("://")[1]
|
artifact_service_uri, agents_dir=agents_dir
|
||||||
artifact_service = GcsArtifactService(bucket_name=gcs_bucket)
|
)
|
||||||
else:
|
if not artifact_service:
|
||||||
raise click.ClickException(
|
raise click.ClickException(
|
||||||
"Unsupported artifact service URI: %s" % artifact_service_uri
|
"Unsupported artifact service URI: %s" % artifact_service_uri
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -0,0 +1,224 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
from typing import Any
|
||||||
|
from typing import Dict
|
||||||
|
from typing import Protocol
|
||||||
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
|
from ..artifacts.base_artifact_service import BaseArtifactService
|
||||||
|
from ..memory.base_memory_service import BaseMemoryService
|
||||||
|
from ..sessions.base_session_service import BaseSessionService
|
||||||
|
|
||||||
|
|
||||||
|
def _load_gcp_config(
|
||||||
|
agents_dir: str | None, service_name: str
|
||||||
|
) -> tuple[str, str]:
|
||||||
|
"""Loads GCP project and location from environment."""
|
||||||
|
if not agents_dir:
|
||||||
|
raise ValueError(f"agents_dir must be provided for {service_name}")
|
||||||
|
|
||||||
|
from .utils import envs
|
||||||
|
|
||||||
|
envs.load_dotenv_for_agent("", agents_dir)
|
||||||
|
|
||||||
|
project = os.environ.get("GOOGLE_CLOUD_PROJECT")
|
||||||
|
location = os.environ.get("GOOGLE_CLOUD_LOCATION")
|
||||||
|
|
||||||
|
if not project or not location:
|
||||||
|
raise ValueError("GOOGLE_CLOUD_PROJECT or GOOGLE_CLOUD_LOCATION not set.")
|
||||||
|
|
||||||
|
return project, location
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_agent_engine_kwargs(
|
||||||
|
uri_part: str, agents_dir: str | None
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Helper to parse agent engine resource name."""
|
||||||
|
if not uri_part:
|
||||||
|
raise ValueError(
|
||||||
|
"Agent engine resource name or resource id can not be empty."
|
||||||
|
)
|
||||||
|
if "/" in uri_part:
|
||||||
|
parts = uri_part.split("/")
|
||||||
|
if not (
|
||||||
|
len(parts) == 6
|
||||||
|
and parts[0] == "projects"
|
||||||
|
and parts[2] == "locations"
|
||||||
|
and parts[4] == "reasoningEngines"
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
"Agent engine resource name is mal-formatted. It should be of"
|
||||||
|
" format :"
|
||||||
|
" projects/{project_id}/locations/{location}/reasoningEngines/{resource_id}"
|
||||||
|
)
|
||||||
|
project = parts[1]
|
||||||
|
location = parts[3]
|
||||||
|
agent_engine_id = parts[5]
|
||||||
|
else:
|
||||||
|
project, location = _load_gcp_config(
|
||||||
|
agents_dir, "short-form agent engine IDs"
|
||||||
|
)
|
||||||
|
agent_engine_id = uri_part
|
||||||
|
return {
|
||||||
|
"project": project,
|
||||||
|
"location": location,
|
||||||
|
"agent_engine_id": agent_engine_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class ServiceFactory(Protocol):
|
||||||
|
"""Protocol for service factory functions."""
|
||||||
|
|
||||||
|
def __call__(
|
||||||
|
self, uri: str, **kwargs
|
||||||
|
) -> BaseSessionService | BaseArtifactService | BaseMemoryService:
|
||||||
|
...
|
||||||
|
|
||||||
|
|
||||||
|
class ServiceRegistry:
|
||||||
|
"""Registry for custom service URI schemes."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._session_factories: Dict[str, ServiceFactory] = {}
|
||||||
|
self._artifact_factories: Dict[str, ServiceFactory] = {}
|
||||||
|
self._memory_factories: Dict[str, ServiceFactory] = {}
|
||||||
|
|
||||||
|
def register_session_service(
|
||||||
|
self, scheme: str, factory: ServiceFactory
|
||||||
|
) -> None:
|
||||||
|
"""Register a factory for a custom session service URI scheme.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
scheme: URI scheme (e.g., 'custom')
|
||||||
|
factory: Callable that takes (uri, **kwargs) and returns
|
||||||
|
BaseSessionService
|
||||||
|
"""
|
||||||
|
self._session_factories[scheme] = factory
|
||||||
|
|
||||||
|
def register_artifact_service(
|
||||||
|
self, scheme: str, factory: ServiceFactory
|
||||||
|
) -> None:
|
||||||
|
"""Register a factory for a custom artifact service URI scheme."""
|
||||||
|
self._artifact_factories[scheme] = factory
|
||||||
|
|
||||||
|
def register_memory_service(
|
||||||
|
self, scheme: str, factory: ServiceFactory
|
||||||
|
) -> None:
|
||||||
|
"""Register a factory for a custom memory service URI scheme."""
|
||||||
|
self._memory_factories[scheme] = factory
|
||||||
|
|
||||||
|
def create_session_service(
|
||||||
|
self, uri: str, **kwargs
|
||||||
|
) -> BaseSessionService | None:
|
||||||
|
"""Create session service from URI using registered factories."""
|
||||||
|
scheme = urlparse(uri).scheme
|
||||||
|
if scheme and scheme in self._session_factories:
|
||||||
|
return self._session_factories[scheme](uri, **kwargs)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def create_artifact_service(
|
||||||
|
self, uri: str, **kwargs
|
||||||
|
) -> BaseArtifactService | None:
|
||||||
|
"""Create artifact service from URI using registered factories."""
|
||||||
|
scheme = urlparse(uri).scheme
|
||||||
|
if scheme and scheme in self._artifact_factories:
|
||||||
|
return self._artifact_factories[scheme](uri, **kwargs)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def create_memory_service(
|
||||||
|
self, uri: str, **kwargs
|
||||||
|
) -> BaseMemoryService | None:
|
||||||
|
"""Create memory service from URI using registered factories."""
|
||||||
|
scheme = urlparse(uri).scheme
|
||||||
|
if scheme and scheme in self._memory_factories:
|
||||||
|
return self._memory_factories[scheme](uri, **kwargs)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _register_builtin_services(registry: ServiceRegistry) -> None:
|
||||||
|
"""Register built-in service implementations."""
|
||||||
|
|
||||||
|
# -- Session Services --
|
||||||
|
def agentengine_session_factory(uri: str, **kwargs):
|
||||||
|
from ..sessions.vertex_ai_session_service import VertexAiSessionService
|
||||||
|
|
||||||
|
parsed = urlparse(uri)
|
||||||
|
params = _parse_agent_engine_kwargs(
|
||||||
|
parsed.netloc + parsed.path, kwargs.get("agents_dir")
|
||||||
|
)
|
||||||
|
return VertexAiSessionService(**params)
|
||||||
|
|
||||||
|
def database_session_factory(uri: str, **kwargs):
|
||||||
|
from ..sessions.database_session_service import DatabaseSessionService
|
||||||
|
|
||||||
|
kwargs_copy = kwargs.copy()
|
||||||
|
kwargs_copy.pop("agents_dir", None)
|
||||||
|
return DatabaseSessionService(db_url=uri, **kwargs_copy)
|
||||||
|
|
||||||
|
registry.register_session_service("agentengine", agentengine_session_factory)
|
||||||
|
for scheme in ["sqlite", "postgresql", "mysql"]:
|
||||||
|
registry.register_session_service(scheme, database_session_factory)
|
||||||
|
|
||||||
|
# -- Artifact Services --
|
||||||
|
def gcs_artifact_factory(uri: str, **kwargs):
|
||||||
|
from ..artifacts.gcs_artifact_service import GcsArtifactService
|
||||||
|
|
||||||
|
kwargs_copy = kwargs.copy()
|
||||||
|
kwargs_copy.pop("agents_dir", None)
|
||||||
|
parsed_uri = urlparse(uri)
|
||||||
|
bucket_name = parsed_uri.netloc
|
||||||
|
return GcsArtifactService(bucket_name=bucket_name, **kwargs_copy)
|
||||||
|
|
||||||
|
registry.register_artifact_service("gs", gcs_artifact_factory)
|
||||||
|
|
||||||
|
# -- Memory Services --
|
||||||
|
def rag_memory_factory(uri: str, **kwargs):
|
||||||
|
from ..memory.vertex_ai_rag_memory_service import VertexAiRagMemoryService
|
||||||
|
|
||||||
|
rag_corpus = urlparse(uri).netloc
|
||||||
|
if not rag_corpus:
|
||||||
|
raise ValueError("Rag corpus can not be empty.")
|
||||||
|
agents_dir = kwargs.get("agents_dir")
|
||||||
|
project, location = _load_gcp_config(agents_dir, "RAG memory service")
|
||||||
|
return VertexAiRagMemoryService(
|
||||||
|
rag_corpus=(
|
||||||
|
f"projects/{project}/locations/{location}/ragCorpora/{rag_corpus}"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def agentengine_memory_factory(uri: str, **kwargs):
|
||||||
|
from ..memory.vertex_ai_memory_bank_service import VertexAiMemoryBankService
|
||||||
|
|
||||||
|
parsed = urlparse(uri)
|
||||||
|
params = _parse_agent_engine_kwargs(
|
||||||
|
parsed.netloc + parsed.path, kwargs.get("agents_dir")
|
||||||
|
)
|
||||||
|
return VertexAiMemoryBankService(**params)
|
||||||
|
|
||||||
|
registry.register_memory_service("rag", rag_memory_factory)
|
||||||
|
registry.register_memory_service("agentengine", agentengine_memory_factory)
|
||||||
|
|
||||||
|
|
||||||
|
# Global registry instance
|
||||||
|
_global_registry = ServiceRegistry()
|
||||||
|
_register_builtin_services(_global_registry)
|
||||||
|
|
||||||
|
|
||||||
|
def get_service_registry() -> ServiceRegistry:
|
||||||
|
"""Get the global service registry instance."""
|
||||||
|
return _global_registry
|
||||||
@@ -0,0 +1,169 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def mock_services():
|
||||||
|
"""Mock all service implementation classes to avoid real instantiation."""
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"google.adk.sessions.vertex_ai_session_service.VertexAiSessionService"
|
||||||
|
) as mock_vertex_session,
|
||||||
|
patch(
|
||||||
|
"google.adk.sessions.database_session_service.DatabaseSessionService"
|
||||||
|
) as mock_db_session,
|
||||||
|
patch(
|
||||||
|
"google.adk.artifacts.gcs_artifact_service.GcsArtifactService"
|
||||||
|
) as mock_gcs_artifact,
|
||||||
|
patch(
|
||||||
|
"google.adk.memory.vertex_ai_rag_memory_service.VertexAiRagMemoryService"
|
||||||
|
) as mock_rag_memory,
|
||||||
|
patch(
|
||||||
|
"google.adk.memory.vertex_ai_memory_bank_service.VertexAiMemoryBankService"
|
||||||
|
) as mock_agentengine_memory,
|
||||||
|
):
|
||||||
|
yield {
|
||||||
|
"vertex_session": mock_vertex_session,
|
||||||
|
"db_session": mock_db_session,
|
||||||
|
"gcs_artifact": mock_gcs_artifact,
|
||||||
|
"rag_memory": mock_rag_memory,
|
||||||
|
"agentengine_memory": mock_agentengine_memory,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def registry():
|
||||||
|
from google.adk.cli.service_registry import get_service_registry
|
||||||
|
|
||||||
|
return get_service_registry()
|
||||||
|
|
||||||
|
|
||||||
|
# Session Service Tests
|
||||||
|
def test_create_session_service_sqlite(registry, mock_services):
|
||||||
|
registry.create_session_service("sqlite:///test.db")
|
||||||
|
mock_services["db_session"].assert_called_once_with(
|
||||||
|
db_url="sqlite:///test.db"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_create_session_service_sqlite_with_kwargs(registry, mock_services):
|
||||||
|
registry.create_session_service(
|
||||||
|
"sqlite:///test.db", pool_size=10, agents_dir="foo"
|
||||||
|
)
|
||||||
|
mock_services["db_session"].assert_called_once_with(
|
||||||
|
db_url="sqlite:///test.db", pool_size=10
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_create_session_service_postgresql(registry, mock_services):
|
||||||
|
registry.create_session_service("postgresql://user:pass@host/db")
|
||||||
|
mock_services["db_session"].assert_called_once_with(
|
||||||
|
db_url="postgresql://user:pass@host/db"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@patch("google.adk.cli.utils.envs.load_dotenv_for_agent")
|
||||||
|
def test_create_session_service_agentengine_short(
|
||||||
|
mock_load_dotenv, registry, mock_services, monkeypatch
|
||||||
|
):
|
||||||
|
monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "test-project")
|
||||||
|
monkeypatch.setenv("GOOGLE_CLOUD_LOCATION", "us-central1")
|
||||||
|
registry.create_session_service(
|
||||||
|
"agentengine://123", agents_dir="/path/to/agents"
|
||||||
|
)
|
||||||
|
mock_services["vertex_session"].assert_called_once_with(
|
||||||
|
project="test-project", location="us-central1", agent_engine_id="123"
|
||||||
|
)
|
||||||
|
mock_load_dotenv.assert_called_once_with("", "/path/to/agents")
|
||||||
|
|
||||||
|
|
||||||
|
def test_create_session_service_agentengine_full(registry, mock_services):
|
||||||
|
uri = "agentengine://projects/p/locations/l/reasoningEngines/123"
|
||||||
|
registry.create_session_service(uri, agents_dir="/path/to/agents")
|
||||||
|
mock_services["vertex_session"].assert_called_once_with(
|
||||||
|
project="p", location="l", agent_engine_id="123"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Artifact Service Tests
|
||||||
|
def test_create_artifact_service_gcs(registry, mock_services):
|
||||||
|
registry.create_artifact_service(
|
||||||
|
"gs://my-bucket/path/prefix", agents_dir="foo", other_kwarg="bar"
|
||||||
|
)
|
||||||
|
mock_services["gcs_artifact"].assert_called_once_with(
|
||||||
|
bucket_name="my-bucket", other_kwarg="bar"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Memory Service Tests
|
||||||
|
@patch("google.adk.cli.utils.envs.load_dotenv_for_agent")
|
||||||
|
def test_create_memory_service_rag(
|
||||||
|
mock_load_dotenv, registry, mock_services, monkeypatch
|
||||||
|
):
|
||||||
|
monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "test-project")
|
||||||
|
monkeypatch.setenv("GOOGLE_CLOUD_LOCATION", "us-central1")
|
||||||
|
registry.create_memory_service(
|
||||||
|
"rag://corpus-123", agents_dir="/path/to/agents"
|
||||||
|
)
|
||||||
|
mock_services["rag_memory"].assert_called_once_with(
|
||||||
|
rag_corpus=(
|
||||||
|
"projects/test-project/locations/us-central1/ragCorpora/corpus-123"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
mock_load_dotenv.assert_called_once_with("", "/path/to/agents")
|
||||||
|
|
||||||
|
|
||||||
|
@patch("google.adk.cli.utils.envs.load_dotenv_for_agent")
|
||||||
|
def test_create_memory_service_agentengine_short(
|
||||||
|
mock_load_dotenv, registry, mock_services, monkeypatch
|
||||||
|
):
|
||||||
|
monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "test-project")
|
||||||
|
monkeypatch.setenv("GOOGLE_CLOUD_LOCATION", "us-central1")
|
||||||
|
registry.create_memory_service(
|
||||||
|
"agentengine://456", agents_dir="/path/to/agents"
|
||||||
|
)
|
||||||
|
mock_services["agentengine_memory"].assert_called_once_with(
|
||||||
|
project="test-project", location="us-central1", agent_engine_id="456"
|
||||||
|
)
|
||||||
|
mock_load_dotenv.assert_called_once_with("", "/path/to/agents")
|
||||||
|
|
||||||
|
|
||||||
|
def test_create_memory_service_agentengine_full(registry, mock_services):
|
||||||
|
uri = "agentengine://projects/p/locations/l/reasoningEngines/456"
|
||||||
|
registry.create_memory_service(uri, agents_dir="/path/to/agents")
|
||||||
|
mock_services["agentengine_memory"].assert_called_once_with(
|
||||||
|
project="p", location="l", agent_engine_id="456"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# General Tests
|
||||||
|
def test_unsupported_scheme(registry, mock_services):
|
||||||
|
session_service = registry.create_session_service("unsupported://foo")
|
||||||
|
artifact_service = registry.create_artifact_service("unsupported://foo")
|
||||||
|
memory_service = registry.create_memory_service("unsupported://foo")
|
||||||
|
assert session_service is None
|
||||||
|
assert artifact_service is None
|
||||||
|
assert memory_service is None
|
||||||
|
for service in [
|
||||||
|
"vertex_session",
|
||||||
|
"db_session",
|
||||||
|
"gcs_artifact",
|
||||||
|
"rag_memory",
|
||||||
|
"agentengine_memory",
|
||||||
|
]:
|
||||||
|
mock_services[service].assert_not_called()
|
||||||
Reference in New Issue
Block a user