mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat: wire runtime entrypoints to service factory defaults
This change routes adk run and the FastAPI server through the new session/artifact service factory, keeps the default experience backed by per-agent .adk storage Co-authored-by: George Weale <gweale@google.com> PiperOrigin-RevId: 836733234
This commit is contained in:
committed by
Copybara-Service
parent
5453b5bfde
commit
06e6fc9132
+58
-23
@@ -15,6 +15,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
from typing import Union
|
||||
|
||||
@@ -22,7 +23,6 @@ import click
|
||||
from google.genai import types
|
||||
from pydantic import BaseModel
|
||||
|
||||
from ..agents.base_agent import BaseAgent
|
||||
from ..agents.llm_agent import LlmAgent
|
||||
from ..apps.app import App
|
||||
from ..artifacts.base_artifact_service import BaseArtifactService
|
||||
@@ -35,8 +35,11 @@ from ..sessions.in_memory_session_service import InMemorySessionService
|
||||
from ..sessions.session import Session
|
||||
from ..utils.context_utils import Aclosing
|
||||
from ..utils.env_utils import is_env_enabled
|
||||
from .service_registry import load_services_module
|
||||
from .utils import envs
|
||||
from .utils.agent_loader import AgentLoader
|
||||
from .utils.service_factory import create_artifact_service_from_options
|
||||
from .utils.service_factory import create_session_service_from_options
|
||||
|
||||
|
||||
class InputFile(BaseModel):
|
||||
@@ -66,7 +69,7 @@ async def run_input_file(
|
||||
)
|
||||
with open(input_path, 'r', encoding='utf-8') as f:
|
||||
input_file = InputFile.model_validate_json(f.read())
|
||||
input_file.state['_time'] = datetime.now()
|
||||
input_file.state['_time'] = datetime.now().isoformat()
|
||||
|
||||
session = await session_service.create_session(
|
||||
app_name=app_name, user_id=user_id, state=input_file.state
|
||||
@@ -134,6 +137,8 @@ async def run_cli(
|
||||
saved_session_file: Optional[str] = None,
|
||||
save_session: bool,
|
||||
session_id: Optional[str] = None,
|
||||
session_service_uri: Optional[str] = None,
|
||||
artifact_service_uri: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Runs an interactive CLI for a certain agent.
|
||||
|
||||
@@ -148,24 +153,47 @@ async def run_cli(
|
||||
contains a previously saved session, exclusive with input_file.
|
||||
save_session: bool, whether to save the session on exit.
|
||||
session_id: Optional[str], the session ID to save the session to on exit.
|
||||
session_service_uri: Optional[str], custom session service URI.
|
||||
artifact_service_uri: Optional[str], custom artifact service URI.
|
||||
"""
|
||||
|
||||
artifact_service = InMemoryArtifactService()
|
||||
session_service = InMemorySessionService()
|
||||
credential_service = InMemoryCredentialService()
|
||||
|
||||
agent_parent_path = Path(agent_parent_dir).resolve()
|
||||
agent_root = agent_parent_path / agent_folder_name
|
||||
load_services_module(str(agent_root))
|
||||
user_id = 'test_user'
|
||||
agent_or_app = AgentLoader(agents_dir=agent_parent_dir).load_agent(
|
||||
|
||||
# Create session and artifact services using factory functions
|
||||
session_service = create_session_service_from_options(
|
||||
base_dir=agent_root,
|
||||
session_service_uri=session_service_uri,
|
||||
)
|
||||
|
||||
artifact_service = create_artifact_service_from_options(
|
||||
base_dir=agent_root,
|
||||
artifact_service_uri=artifact_service_uri,
|
||||
)
|
||||
|
||||
credential_service = InMemoryCredentialService()
|
||||
agents_dir = str(agent_parent_path)
|
||||
agent_or_app = AgentLoader(agents_dir=agents_dir).load_agent(
|
||||
agent_folder_name
|
||||
)
|
||||
session_app_name = (
|
||||
agent_or_app.name if isinstance(agent_or_app, App) else agent_folder_name
|
||||
)
|
||||
session = await session_service.create_session(
|
||||
app_name=session_app_name, user_id=user_id
|
||||
)
|
||||
if not is_env_enabled('ADK_DISABLE_LOAD_DOTENV'):
|
||||
envs.load_dotenv_for_agent(agent_folder_name, agent_parent_dir)
|
||||
envs.load_dotenv_for_agent(agent_folder_name, agents_dir)
|
||||
|
||||
# Helper function for printing events
|
||||
def _print_event(event) -> None:
|
||||
content = event.content
|
||||
if not content or not content.parts:
|
||||
return
|
||||
text_parts = [part.text for part in content.parts if part.text]
|
||||
if not text_parts:
|
||||
return
|
||||
author = event.author or 'system'
|
||||
click.echo(f'[{author}]: {"".join(text_parts)}')
|
||||
|
||||
if input_file:
|
||||
session = await run_input_file(
|
||||
app_name=session_app_name,
|
||||
@@ -177,16 +205,22 @@ async def run_cli(
|
||||
input_path=input_file,
|
||||
)
|
||||
elif saved_session_file:
|
||||
# Load the saved session from file
|
||||
with open(saved_session_file, 'r', encoding='utf-8') as f:
|
||||
loaded_session = Session.model_validate_json(f.read())
|
||||
|
||||
# Create a new session in the service, copying state from the file
|
||||
session = await session_service.create_session(
|
||||
app_name=session_app_name,
|
||||
user_id=user_id,
|
||||
state=loaded_session.state if loaded_session else None,
|
||||
)
|
||||
|
||||
# Append events from the file to the new session and display them
|
||||
if loaded_session:
|
||||
for event in loaded_session.events:
|
||||
await session_service.append_event(session, event)
|
||||
content = event.content
|
||||
if not content or not content.parts or not content.parts[0].text:
|
||||
continue
|
||||
click.echo(f'[{event.author}]: {content.parts[0].text}')
|
||||
_print_event(event)
|
||||
|
||||
await run_interactively(
|
||||
agent_or_app,
|
||||
@@ -196,6 +230,9 @@ async def run_cli(
|
||||
credential_service,
|
||||
)
|
||||
else:
|
||||
session = await session_service.create_session(
|
||||
app_name=session_app_name, user_id=user_id
|
||||
)
|
||||
click.echo(f'Running agent {agent_or_app.name}, type exit to exit.')
|
||||
await run_interactively(
|
||||
agent_or_app,
|
||||
@@ -207,9 +244,7 @@ async def run_cli(
|
||||
|
||||
if save_session:
|
||||
session_id = session_id or input('Session ID to save: ')
|
||||
session_path = (
|
||||
f'{agent_parent_dir}/{agent_folder_name}/{session_id}.session.json'
|
||||
)
|
||||
session_path = agent_root / f'{session_id}.session.json'
|
||||
|
||||
# Fetch the session again to get all the details.
|
||||
session = await session_service.get_session(
|
||||
@@ -217,9 +252,9 @@ async def run_cli(
|
||||
user_id=session.user_id,
|
||||
session_id=session.id,
|
||||
)
|
||||
with open(session_path, 'w', encoding='utf-8') as f:
|
||||
f.write(
|
||||
session.model_dump_json(indent=2, exclude_none=True, by_alias=True)
|
||||
)
|
||||
session_path.write_text(
|
||||
session.model_dump_json(indent=2, exclude_none=True, by_alias=True),
|
||||
encoding='utf-8',
|
||||
)
|
||||
|
||||
print('Session saved to', session_path)
|
||||
|
||||
@@ -34,20 +34,19 @@ from opentelemetry.sdk.trace import TracerProvider
|
||||
from starlette.types import Lifespan
|
||||
from watchdog.observers import Observer
|
||||
|
||||
from ..artifacts.in_memory_artifact_service import InMemoryArtifactService
|
||||
from ..auth.credential_service.in_memory_credential_service import InMemoryCredentialService
|
||||
from ..evaluation.local_eval_set_results_manager import LocalEvalSetResultsManager
|
||||
from ..evaluation.local_eval_sets_manager import LocalEvalSetsManager
|
||||
from ..memory.in_memory_memory_service import InMemoryMemoryService
|
||||
from ..runners import Runner
|
||||
from ..sessions.in_memory_session_service import InMemorySessionService
|
||||
from .adk_web_server import AdkWebServer
|
||||
from .service_registry import get_service_registry
|
||||
from .service_registry import load_services_module
|
||||
from .utils import envs
|
||||
from .utils import evals
|
||||
from .utils.agent_change_handler import AgentChangeEventHandler
|
||||
from .utils.agent_loader import AgentLoader
|
||||
from .utils.service_factory import create_artifact_service_from_options
|
||||
from .utils.service_factory import create_memory_service_from_options
|
||||
from .utils.service_factory import create_session_service_from_options
|
||||
|
||||
logger = logging.getLogger("google_adk." + __name__)
|
||||
|
||||
@@ -74,6 +73,8 @@ def get_fast_api_app(
|
||||
logo_text: Optional[str] = None,
|
||||
logo_image_url: Optional[str] = None,
|
||||
) -> FastAPI:
|
||||
# Convert to absolute path for consistency
|
||||
agents_dir = str(Path(agents_dir).resolve())
|
||||
|
||||
# Set up eval managers.
|
||||
if eval_storage_uri:
|
||||
@@ -91,48 +92,32 @@ def get_fast_api_app(
|
||||
# Load services.py from agents_dir for custom service registration.
|
||||
load_services_module(agents_dir)
|
||||
|
||||
service_registry = get_service_registry()
|
||||
|
||||
# Build the Memory service
|
||||
if memory_service_uri:
|
||||
memory_service = service_registry.create_memory_service(
|
||||
memory_service_uri, agents_dir=agents_dir
|
||||
try:
|
||||
memory_service = create_memory_service_from_options(
|
||||
base_dir=agents_dir,
|
||||
memory_service_uri=memory_service_uri,
|
||||
)
|
||||
if not memory_service:
|
||||
raise click.ClickException(
|
||||
"Unsupported memory service URI: %s" % memory_service_uri
|
||||
)
|
||||
else:
|
||||
memory_service = InMemoryMemoryService()
|
||||
except ValueError as exc:
|
||||
raise click.ClickException(str(exc)) from exc
|
||||
|
||||
# Build the Session service
|
||||
if session_service_uri:
|
||||
session_kwargs = session_db_kwargs or {}
|
||||
session_service = service_registry.create_session_service(
|
||||
session_service_uri, agents_dir=agents_dir, **session_kwargs
|
||||
)
|
||||
if not session_service:
|
||||
# Fallback to DatabaseSessionService if the service registry doesn't
|
||||
# support the session service URI scheme.
|
||||
from ..sessions.database_session_service import DatabaseSessionService
|
||||
|
||||
session_service = DatabaseSessionService(
|
||||
db_url=session_service_uri, **session_kwargs
|
||||
)
|
||||
else:
|
||||
session_service = InMemorySessionService()
|
||||
session_service = create_session_service_from_options(
|
||||
base_dir=agents_dir,
|
||||
session_service_uri=session_service_uri,
|
||||
session_db_kwargs=session_db_kwargs,
|
||||
per_agent=True, # Multi-agent mode
|
||||
)
|
||||
|
||||
# Build the Artifact service
|
||||
if artifact_service_uri:
|
||||
artifact_service = service_registry.create_artifact_service(
|
||||
artifact_service_uri, agents_dir=agents_dir
|
||||
try:
|
||||
artifact_service = create_artifact_service_from_options(
|
||||
base_dir=agents_dir,
|
||||
artifact_service_uri=artifact_service_uri,
|
||||
per_agent=True, # Multi-agent mode
|
||||
)
|
||||
if not artifact_service:
|
||||
raise click.ClickException(
|
||||
"Unsupported artifact service URI: %s" % artifact_service_uri
|
||||
)
|
||||
else:
|
||||
artifact_service = InMemoryArtifactService()
|
||||
except ValueError as exc:
|
||||
raise click.ClickException(str(exc)) from exc
|
||||
|
||||
# Build the Credential service
|
||||
credential_service = InMemoryCredentialService()
|
||||
|
||||
@@ -0,0 +1,138 @@
|
||||
# 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 logging
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Optional
|
||||
|
||||
from ...artifacts.base_artifact_service import BaseArtifactService
|
||||
from ...memory.base_memory_service import BaseMemoryService
|
||||
from ...sessions.base_session_service import BaseSessionService
|
||||
from ..service_registry import get_service_registry
|
||||
from .local_storage import create_local_artifact_service
|
||||
|
||||
logger = logging.getLogger("google_adk." + __name__)
|
||||
|
||||
|
||||
def create_session_service_from_options(
|
||||
*,
|
||||
base_dir: Path | str,
|
||||
session_service_uri: Optional[str] = None,
|
||||
session_db_kwargs: Optional[dict[str, Any]] = None,
|
||||
per_agent: bool = False,
|
||||
) -> BaseSessionService:
|
||||
"""Creates a session service based on CLI/web options."""
|
||||
base_path = Path(base_dir)
|
||||
registry = get_service_registry()
|
||||
|
||||
kwargs: dict[str, Any] = {
|
||||
"agents_dir": str(base_path),
|
||||
"per_agent": per_agent,
|
||||
}
|
||||
if session_db_kwargs:
|
||||
kwargs.update(session_db_kwargs)
|
||||
|
||||
if session_service_uri:
|
||||
if per_agent:
|
||||
logger.warning(
|
||||
"per_agent is not supported with remote session service URIs,"
|
||||
" ignoring"
|
||||
)
|
||||
logger.info("Using session service URI: %s", session_service_uri)
|
||||
service = registry.create_session_service(session_service_uri, **kwargs)
|
||||
if service is not None:
|
||||
return service
|
||||
|
||||
# Fallback to DatabaseSessionService if the registry doesn't support the
|
||||
# session service URI scheme. This keeps support for SQLAlchemy-compatible
|
||||
# databases like AlloyDB or Cloud Spanner without explicit registration.
|
||||
from ...sessions.database_session_service import DatabaseSessionService
|
||||
|
||||
fallback_kwargs = dict(kwargs)
|
||||
fallback_kwargs.pop("agents_dir", None)
|
||||
fallback_kwargs.pop("per_agent", None)
|
||||
logger.info(
|
||||
"Falling back to DatabaseSessionService for URI: %s",
|
||||
session_service_uri,
|
||||
)
|
||||
return DatabaseSessionService(db_url=session_service_uri, **fallback_kwargs)
|
||||
|
||||
logger.info("Using in-memory session service")
|
||||
from ...sessions.in_memory_session_service import InMemorySessionService
|
||||
|
||||
return InMemorySessionService()
|
||||
|
||||
|
||||
def create_memory_service_from_options(
|
||||
*,
|
||||
base_dir: Path | str,
|
||||
memory_service_uri: Optional[str] = None,
|
||||
) -> BaseMemoryService:
|
||||
"""Creates a memory service based on CLI/web options."""
|
||||
base_path = Path(base_dir)
|
||||
registry = get_service_registry()
|
||||
|
||||
if memory_service_uri:
|
||||
logger.info("Using memory service URI: %s", memory_service_uri)
|
||||
service = registry.create_memory_service(
|
||||
memory_service_uri,
|
||||
agents_dir=str(base_path),
|
||||
)
|
||||
if service is None:
|
||||
raise ValueError(f"Unsupported memory service URI: {memory_service_uri}")
|
||||
return service
|
||||
|
||||
logger.info("Using in-memory memory service")
|
||||
from ...memory.in_memory_memory_service import InMemoryMemoryService
|
||||
|
||||
return InMemoryMemoryService()
|
||||
|
||||
|
||||
def create_artifact_service_from_options(
|
||||
*,
|
||||
base_dir: Path | str,
|
||||
artifact_service_uri: Optional[str] = None,
|
||||
per_agent: bool = False,
|
||||
) -> BaseArtifactService:
|
||||
"""Creates an artifact service based on CLI/web options."""
|
||||
base_path = Path(base_dir)
|
||||
registry = get_service_registry()
|
||||
|
||||
if artifact_service_uri:
|
||||
if per_agent:
|
||||
logger.warning(
|
||||
"per_agent is not supported with remote artifact service URIs,"
|
||||
" ignoring"
|
||||
)
|
||||
logger.info("Using artifact service URI: %s", artifact_service_uri)
|
||||
service = registry.create_artifact_service(
|
||||
artifact_service_uri,
|
||||
agents_dir=str(base_path),
|
||||
per_agent=per_agent,
|
||||
)
|
||||
if service is None:
|
||||
logger.warning(
|
||||
"Unsupported artifact service URI: %s, falling back to in-memory",
|
||||
artifact_service_uri,
|
||||
)
|
||||
from ...artifacts.in_memory_artifact_service import InMemoryArtifactService
|
||||
|
||||
return InMemoryArtifactService()
|
||||
return service
|
||||
|
||||
if per_agent:
|
||||
logger.info("Using shared file artifact service rooted at %s", base_dir)
|
||||
return create_local_artifact_service(base_dir=base_path, per_agent=per_agent)
|
||||
@@ -327,15 +327,15 @@ def test_app(
|
||||
with (
|
||||
patch("signal.signal", return_value=None),
|
||||
patch(
|
||||
"google.adk.cli.fast_api.InMemorySessionService",
|
||||
"google.adk.cli.fast_api.create_session_service_from_options",
|
||||
return_value=mock_session_service,
|
||||
),
|
||||
patch(
|
||||
"google.adk.cli.fast_api.InMemoryArtifactService",
|
||||
"google.adk.cli.fast_api.create_artifact_service_from_options",
|
||||
return_value=mock_artifact_service,
|
||||
),
|
||||
patch(
|
||||
"google.adk.cli.fast_api.InMemoryMemoryService",
|
||||
"google.adk.cli.fast_api.create_memory_service_from_options",
|
||||
return_value=mock_memory_service,
|
||||
),
|
||||
patch(
|
||||
@@ -472,15 +472,15 @@ def test_app_with_a2a(
|
||||
with (
|
||||
patch("signal.signal", return_value=None),
|
||||
patch(
|
||||
"google.adk.cli.fast_api.InMemorySessionService",
|
||||
"google.adk.cli.fast_api.create_session_service_from_options",
|
||||
return_value=mock_session_service,
|
||||
),
|
||||
patch(
|
||||
"google.adk.cli.fast_api.InMemoryArtifactService",
|
||||
"google.adk.cli.fast_api.create_artifact_service_from_options",
|
||||
return_value=mock_artifact_service,
|
||||
),
|
||||
patch(
|
||||
"google.adk.cli.fast_api.InMemoryMemoryService",
|
||||
"google.adk.cli.fast_api.create_memory_service_from_options",
|
||||
return_value=mock_memory_service,
|
||||
),
|
||||
patch(
|
||||
|
||||
@@ -28,7 +28,12 @@ from typing import Tuple
|
||||
import click
|
||||
from google.adk.agents.base_agent import BaseAgent
|
||||
from google.adk.apps.app import App
|
||||
from google.adk.artifacts.file_artifact_service import FileArtifactService
|
||||
from google.adk.artifacts.in_memory_artifact_service import InMemoryArtifactService
|
||||
from google.adk.auth.credential_service.in_memory_credential_service import InMemoryCredentialService
|
||||
import google.adk.cli.cli as cli
|
||||
from google.adk.cli.utils.service_factory import create_artifact_service_from_options
|
||||
from google.adk.sessions.in_memory_session_service import InMemorySessionService
|
||||
import pytest
|
||||
|
||||
|
||||
@@ -151,9 +156,9 @@ async def test_run_input_file_outputs(
|
||||
input_path = tmp_path / "input.json"
|
||||
input_path.write_text(json.dumps(input_json))
|
||||
|
||||
artifact_service = cli.InMemoryArtifactService()
|
||||
session_service = cli.InMemorySessionService()
|
||||
credential_service = cli.InMemoryCredentialService()
|
||||
artifact_service = InMemoryArtifactService()
|
||||
session_service = InMemorySessionService()
|
||||
credential_service = InMemoryCredentialService()
|
||||
dummy_root = BaseAgent(name="root")
|
||||
|
||||
session = await cli.run_input_file(
|
||||
@@ -189,6 +194,34 @@ async def test_run_cli_with_input_file(fake_agent, tmp_path: Path) -> None:
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_cli_loads_services_module(
|
||||
fake_agent, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""run_cli should load custom services from the agents directory."""
|
||||
parent_dir, folder_name = fake_agent
|
||||
input_json = {"state": {}, "queries": ["ping"]}
|
||||
input_path = tmp_path / "input.json"
|
||||
input_path.write_text(json.dumps(input_json))
|
||||
|
||||
loaded_dirs: list[str] = []
|
||||
monkeypatch.setattr(
|
||||
cli, "load_services_module", lambda path: loaded_dirs.append(path)
|
||||
)
|
||||
|
||||
agent_root = parent_dir / folder_name
|
||||
|
||||
await cli.run_cli(
|
||||
agent_parent_dir=str(parent_dir),
|
||||
agent_folder_name=folder_name,
|
||||
input_file=str(input_path),
|
||||
saved_session_file=None,
|
||||
save_session=False,
|
||||
)
|
||||
|
||||
assert loaded_dirs == [str(agent_root.resolve())]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_cli_app_uses_app_name_for_sessions(
|
||||
fake_app_agent, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
@@ -197,15 +230,20 @@ async def test_run_cli_app_uses_app_name_for_sessions(
|
||||
parent_dir, folder_name, app_name = fake_app_agent
|
||||
created_app_names: List[str] = []
|
||||
|
||||
original_session_cls = cli.InMemorySessionService
|
||||
|
||||
class _SpySessionService(original_session_cls):
|
||||
class _SpySessionService(InMemorySessionService):
|
||||
|
||||
async def create_session(self, *, app_name: str, **kwargs: Any) -> Any:
|
||||
created_app_names.append(app_name)
|
||||
return await super().create_session(app_name=app_name, **kwargs)
|
||||
|
||||
monkeypatch.setattr(cli, "InMemorySessionService", _SpySessionService)
|
||||
spy_session_service = _SpySessionService()
|
||||
|
||||
def _session_factory(**_: Any) -> InMemorySessionService:
|
||||
return spy_session_service
|
||||
|
||||
monkeypatch.setattr(
|
||||
cli, "create_session_service_from_options", _session_factory
|
||||
)
|
||||
|
||||
input_json = {"state": {}, "queries": ["ping"]}
|
||||
input_path = tmp_path / "input_app.json"
|
||||
@@ -253,16 +291,89 @@ async def test_run_cli_save_session(
|
||||
assert "id" in data and "events" in data
|
||||
|
||||
|
||||
def test_create_artifact_service_defaults_to_file(tmp_path: Path) -> None:
|
||||
"""Service factory should default to FileArtifactService when URI is unset."""
|
||||
service = create_artifact_service_from_options(base_dir=tmp_path)
|
||||
assert isinstance(service, FileArtifactService)
|
||||
expected_root = Path(tmp_path) / ".adk" / "artifacts"
|
||||
assert service.root_dir == expected_root
|
||||
assert expected_root.exists()
|
||||
|
||||
|
||||
def test_create_artifact_service_per_agent_uses_shared_root(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""Multi-agent mode should still use a single file artifact service."""
|
||||
service = create_artifact_service_from_options(
|
||||
base_dir=tmp_path, per_agent=True
|
||||
)
|
||||
assert isinstance(service, FileArtifactService)
|
||||
expected_root = Path(tmp_path) / ".adk" / "artifacts"
|
||||
assert service.root_dir == expected_root
|
||||
assert expected_root.exists()
|
||||
|
||||
|
||||
def test_create_artifact_service_respects_memory_uri(tmp_path: Path) -> None:
|
||||
"""Service factory should honor memory:// URIs."""
|
||||
service = create_artifact_service_from_options(
|
||||
base_dir=tmp_path, artifact_service_uri="memory://"
|
||||
)
|
||||
assert isinstance(service, InMemoryArtifactService)
|
||||
|
||||
|
||||
def test_create_artifact_service_accepts_file_uri(tmp_path: Path) -> None:
|
||||
"""Service factory should allow custom local roots via file:// URIs."""
|
||||
custom_root = tmp_path / "custom_artifacts"
|
||||
service = create_artifact_service_from_options(
|
||||
base_dir=tmp_path, artifact_service_uri=custom_root.as_uri()
|
||||
)
|
||||
assert isinstance(service, FileArtifactService)
|
||||
assert service.root_dir == custom_root
|
||||
assert custom_root.exists()
|
||||
|
||||
|
||||
def test_create_artifact_service_file_uri_rejects_per_agent(tmp_path: Path):
|
||||
"""file:// URIs are incompatible with per-agent mode."""
|
||||
custom_root = tmp_path / "custom"
|
||||
with pytest.raises(ValueError, match="multi-agent"):
|
||||
create_artifact_service_from_options(
|
||||
base_dir=tmp_path,
|
||||
artifact_service_uri=custom_root.as_uri(),
|
||||
per_agent=True,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_cli_accepts_memory_scheme(
|
||||
fake_agent, tmp_path: Path
|
||||
) -> None:
|
||||
"""run_cli should allow configuring in-memory services via memory:// URIs."""
|
||||
parent_dir, folder_name = fake_agent
|
||||
input_json = {"state": {}, "queries": []}
|
||||
input_path = tmp_path / "noop.json"
|
||||
input_path.write_text(json.dumps(input_json))
|
||||
|
||||
await cli.run_cli(
|
||||
agent_parent_dir=str(parent_dir),
|
||||
agent_folder_name=folder_name,
|
||||
input_file=str(input_path),
|
||||
saved_session_file=None,
|
||||
save_session=False,
|
||||
session_service_uri="memory://",
|
||||
artifact_service_uri="memory://",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_interactively_whitespace_and_exit(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""run_interactively should skip blank input, echo once, then exit."""
|
||||
# make a session that belongs to dummy agent
|
||||
session_service = cli.InMemorySessionService()
|
||||
session_service = InMemorySessionService()
|
||||
sess = await session_service.create_session(app_name="dummy", user_id="u")
|
||||
artifact_service = cli.InMemoryArtifactService()
|
||||
credential_service = cli.InMemoryCredentialService()
|
||||
artifact_service = InMemoryArtifactService()
|
||||
credential_service = InMemoryCredentialService()
|
||||
root_agent = BaseAgent(name="root")
|
||||
|
||||
# fake user input: blank -> 'hello' -> 'exit'
|
||||
|
||||
@@ -0,0 +1,162 @@
|
||||
# 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.
|
||||
|
||||
"""Tests for service factory helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import Mock
|
||||
|
||||
import google.adk.cli.utils.service_factory as service_factory
|
||||
from google.adk.memory.in_memory_memory_service import InMemoryMemoryService
|
||||
from google.adk.sessions.database_session_service import DatabaseSessionService
|
||||
from google.adk.sessions.in_memory_session_service import InMemorySessionService
|
||||
import pytest
|
||||
|
||||
|
||||
def test_create_session_service_uses_registry(tmp_path: Path, monkeypatch):
|
||||
registry = Mock()
|
||||
expected = object()
|
||||
registry.create_session_service.return_value = expected
|
||||
monkeypatch.setattr(service_factory, "get_service_registry", lambda: registry)
|
||||
|
||||
result = service_factory.create_session_service_from_options(
|
||||
base_dir=tmp_path,
|
||||
session_service_uri="sqlite:///test.db",
|
||||
)
|
||||
|
||||
assert result is expected
|
||||
registry.create_session_service.assert_called_once_with(
|
||||
"sqlite:///test.db",
|
||||
agents_dir=str(tmp_path),
|
||||
per_agent=False,
|
||||
)
|
||||
|
||||
|
||||
def test_create_session_service_per_agent_uri(tmp_path: Path, monkeypatch):
|
||||
registry = Mock()
|
||||
expected = object()
|
||||
registry.create_session_service.return_value = expected
|
||||
monkeypatch.setattr(service_factory, "get_service_registry", lambda: registry)
|
||||
|
||||
result = service_factory.create_session_service_from_options(
|
||||
base_dir=tmp_path,
|
||||
session_service_uri="memory://",
|
||||
per_agent=True,
|
||||
)
|
||||
|
||||
assert result is expected
|
||||
registry.create_session_service.assert_called_once_with(
|
||||
"memory://", agents_dir=str(tmp_path), per_agent=True
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("per_agent", [True, False])
|
||||
def test_create_session_service_defaults_to_memory(
|
||||
tmp_path: Path, per_agent: bool
|
||||
):
|
||||
service = service_factory.create_session_service_from_options(
|
||||
base_dir=tmp_path,
|
||||
per_agent=per_agent,
|
||||
)
|
||||
|
||||
assert isinstance(service, InMemorySessionService)
|
||||
|
||||
|
||||
def test_create_session_service_fallbacks_to_database(
|
||||
tmp_path: Path, monkeypatch
|
||||
):
|
||||
registry = Mock()
|
||||
registry.create_session_service.return_value = None
|
||||
monkeypatch.setattr(service_factory, "get_service_registry", lambda: registry)
|
||||
|
||||
service = service_factory.create_session_service_from_options(
|
||||
base_dir=tmp_path,
|
||||
session_service_uri="sqlite+aiosqlite:///:memory:",
|
||||
session_db_kwargs={"echo": True},
|
||||
)
|
||||
|
||||
assert isinstance(service, DatabaseSessionService)
|
||||
assert service.db_engine.url.drivername == "sqlite+aiosqlite"
|
||||
assert service.db_engine.echo is True
|
||||
registry.create_session_service.assert_called_once_with(
|
||||
"sqlite+aiosqlite:///:memory:",
|
||||
agents_dir=str(tmp_path),
|
||||
per_agent=False,
|
||||
echo=True,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("per_agent", [True, False])
|
||||
def test_create_artifact_service_uses_registry(
|
||||
tmp_path: Path, monkeypatch, per_agent: bool
|
||||
):
|
||||
registry = Mock()
|
||||
expected = object()
|
||||
registry.create_artifact_service.return_value = expected
|
||||
monkeypatch.setattr(service_factory, "get_service_registry", lambda: registry)
|
||||
|
||||
result = service_factory.create_artifact_service_from_options(
|
||||
base_dir=tmp_path,
|
||||
artifact_service_uri="gs://bucket/path",
|
||||
per_agent=per_agent,
|
||||
)
|
||||
|
||||
assert result is expected
|
||||
registry.create_artifact_service.assert_called_once_with(
|
||||
"gs://bucket/path",
|
||||
agents_dir=str(tmp_path),
|
||||
per_agent=per_agent,
|
||||
)
|
||||
|
||||
|
||||
def test_create_memory_service_uses_registry(tmp_path: Path, monkeypatch):
|
||||
registry = Mock()
|
||||
expected = object()
|
||||
registry.create_memory_service.return_value = expected
|
||||
monkeypatch.setattr(service_factory, "get_service_registry", lambda: registry)
|
||||
|
||||
result = service_factory.create_memory_service_from_options(
|
||||
base_dir=tmp_path,
|
||||
memory_service_uri="rag://my-corpus",
|
||||
)
|
||||
|
||||
assert result is expected
|
||||
registry.create_memory_service.assert_called_once_with(
|
||||
"rag://my-corpus",
|
||||
agents_dir=str(tmp_path),
|
||||
)
|
||||
|
||||
|
||||
def test_create_memory_service_defaults_to_in_memory(tmp_path: Path):
|
||||
service = service_factory.create_memory_service_from_options(
|
||||
base_dir=tmp_path
|
||||
)
|
||||
|
||||
assert isinstance(service, InMemoryMemoryService)
|
||||
|
||||
|
||||
def test_create_memory_service_raises_on_unknown_scheme(
|
||||
tmp_path: Path, monkeypatch
|
||||
):
|
||||
registry = Mock()
|
||||
registry.create_memory_service.return_value = None
|
||||
monkeypatch.setattr(service_factory, "get_service_registry", lambda: registry)
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
service_factory.create_memory_service_from_options(
|
||||
base_dir=tmp_path,
|
||||
memory_service_uri="unknown://foo",
|
||||
)
|
||||
Reference in New Issue
Block a user