From a7b509763c1732f0363e90952bb4c2672572d542 Mon Sep 17 00:00:00 2001 From: George Weale Date: Fri, 20 Feb 2026 11:28:29 -0800 Subject: [PATCH] feat: Use --memory_service_uri in ADK CLI run command Co-authored-by: George Weale PiperOrigin-RevId: 873000092 --- src/google/adk/cli/cli.py | 20 ++- src/google/adk/cli/cli_tools_click.py | 9 +- src/google/adk/cli/service_registry.py | 6 + tests/unittests/cli/test_service_registry.py | 7 + tests/unittests/cli/utils/test_cli.py | 136 ++++++++++++++++++ .../cli/utils/test_cli_tools_click.py | 13 +- .../cli/utils/test_service_factory.py | 9 ++ 7 files changed, 187 insertions(+), 13 deletions(-) diff --git a/src/google/adk/cli/cli.py b/src/google/adk/cli/cli.py index 16eba88b..1d49f50d 100644 --- a/src/google/adk/cli/cli.py +++ b/src/google/adk/cli/cli.py @@ -28,6 +28,7 @@ from ..apps.app import App from ..artifacts.base_artifact_service import BaseArtifactService from ..auth.credential_service.base_credential_service import BaseCredentialService from ..auth.credential_service.in_memory_credential_service import InMemoryCredentialService +from ..memory.base_memory_service import BaseMemoryService from ..runners import Runner from ..sessions.base_session_service import BaseSessionService from ..sessions.session import Session @@ -37,6 +38,7 @@ 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_memory_service_from_options from .utils.service_factory import create_session_service_from_options @@ -53,6 +55,7 @@ async def run_input_file( session_service: BaseSessionService, credential_service: BaseCredentialService, input_path: str, + memory_service: Optional[BaseMemoryService] = None, ) -> Session: app = ( agent_or_app @@ -63,6 +66,7 @@ async def run_input_file( app=app, artifact_service=artifact_service, session_service=session_service, + memory_service=memory_service, credential_service=credential_service, ) with open(input_path, 'r', encoding='utf-8') as f: @@ -93,6 +97,7 @@ async def run_interactively( session: Session, session_service: BaseSessionService, credential_service: BaseCredentialService, + memory_service: Optional[BaseMemoryService] = None, ) -> None: app = ( root_agent_or_app @@ -103,6 +108,7 @@ async def run_interactively( app=app, artifact_service=artifact_service, session_service=session_service, + memory_service=memory_service, credential_service=credential_service, ) while True: @@ -137,6 +143,7 @@ async def run_cli( session_id: Optional[str] = None, session_service_uri: Optional[str] = None, artifact_service_uri: Optional[str] = None, + memory_service_uri: Optional[str] = None, use_local_storage: bool = True, ) -> None: """Runs an interactive CLI for a certain agent. @@ -154,6 +161,7 @@ async def run_cli( 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. + memory_service_uri: Optional[str], custom memory service URI. use_local_storage: bool, whether to use local .adk storage by default. """ agent_parent_path = Path(agent_parent_dir).resolve() @@ -171,6 +179,9 @@ async def run_cli( if isinstance(agent_or_app, App) and agent_or_app.name != agent_folder_name: app_name_to_dir = {agent_or_app.name: agent_folder_name} + if not is_env_enabled('ADK_DISABLE_LOAD_DOTENV'): + envs.load_dotenv_for_agent(agent_folder_name, agents_dir) + # Create session and artifact services using factory functions. # Sessions persist under //.adk/session.db when enabled. session_service = create_session_service_from_options( @@ -185,10 +196,12 @@ async def run_cli( artifact_service_uri=artifact_service_uri, use_local_storage=use_local_storage, ) + memory_service = create_memory_service_from_options( + base_dir=agent_parent_path, + memory_service_uri=memory_service_uri, + ) credential_service = InMemoryCredentialService() - if not is_env_enabled('ADK_DISABLE_LOAD_DOTENV'): - envs.load_dotenv_for_agent(agent_folder_name, agents_dir) # Helper function for printing events def _print_event(event) -> None: @@ -208,6 +221,7 @@ async def run_cli( agent_or_app=agent_or_app, artifact_service=artifact_service, session_service=session_service, + memory_service=memory_service, credential_service=credential_service, input_path=input_file, ) @@ -235,6 +249,7 @@ async def run_cli( session, session_service, credential_service, + memory_service=memory_service, ) else: session = await session_service.create_session( @@ -247,6 +262,7 @@ async def run_cli( session, session_service, credential_service, + memory_service=memory_service, ) if save_session: diff --git a/src/google/adk/cli/cli_tools_click.py b/src/google/adk/cli/cli_tools_click.py index 5b5d3e5c..f55a8f10 100644 --- a/src/google/adk/cli/cli_tools_click.py +++ b/src/google/adk/cli/cli_tools_click.py @@ -645,14 +645,6 @@ def cli_run( """ logs.log_to_tmp_folder() - # Validation warning for memory_service_uri (not supported for adk run) - if memory_service_uri: - click.secho( - "WARNING: --memory_service_uri is not supported for adk run.", - fg="yellow", - err=True, - ) - agent_parent_folder = os.path.dirname(agent) agent_folder_name = os.path.basename(agent) @@ -666,6 +658,7 @@ def cli_run( session_id=session_id, session_service_uri=session_service_uri, artifact_service_uri=artifact_service_uri, + memory_service_uri=memory_service_uri, use_local_storage=use_local_storage, ) ) diff --git a/src/google/adk/cli/service_registry.py b/src/google/adk/cli/service_registry.py index 2ea286ef..b1328958 100644 --- a/src/google/adk/cli/service_registry.py +++ b/src/google/adk/cli/service_registry.py @@ -301,6 +301,11 @@ def _register_builtin_services(registry: ServiceRegistry) -> None: registry.register_artifact_service("file", file_artifact_factory) # -- Memory Services -- + def memory_memory_factory(_uri: str, **_): + from ..memory.in_memory_memory_service import InMemoryMemoryService + + return InMemoryMemoryService() + def rag_memory_factory(uri: str, **kwargs): from ..memory.vertex_ai_rag_memory_service import VertexAiRagMemoryService @@ -324,6 +329,7 @@ def _register_builtin_services(registry: ServiceRegistry) -> None: ) return VertexAiMemoryBankService(**params) + registry.register_memory_service("memory", memory_memory_factory) registry.register_memory_service("rag", rag_memory_factory) registry.register_memory_service("agentengine", agentengine_memory_factory) diff --git a/tests/unittests/cli/test_service_registry.py b/tests/unittests/cli/test_service_registry.py index 37c6e7c2..dd33e006 100644 --- a/tests/unittests/cli/test_service_registry.py +++ b/tests/unittests/cli/test_service_registry.py @@ -165,6 +165,13 @@ def test_create_memory_service_agentengine_full(registry, mock_services): ) +def test_create_memory_service_memory(registry): + from google.adk.memory.in_memory_memory_service import InMemoryMemoryService + + memory_service = registry.create_memory_service("memory://") + assert isinstance(memory_service, InMemoryMemoryService) + + # General Tests def test_unsupported_scheme(registry, mock_services): session_service = registry.create_session_service("unsupported://foo") diff --git a/tests/unittests/cli/utils/test_cli.py b/tests/unittests/cli/utils/test_cli.py index 6814ef97..f7df1bf1 100644 --- a/tests/unittests/cli/utils/test_cli.py +++ b/tests/unittests/cli/utils/test_cli.py @@ -354,9 +354,145 @@ async def test_run_cli_accepts_memory_scheme( save_session=False, session_service_uri="memory://", artifact_service_uri="memory://", + memory_service_uri="memory://", ) +@pytest.mark.asyncio +async def test_run_cli_invalid_memory_uri_surfaces_value_error( + fake_agent, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """run_cli should let ValueError propagate for invalid memory service URIs.""" + parent_dir, folder_name = fake_agent + input_json = {"state": {}, "queries": []} + input_path = tmp_path / "invalid_memory_uri.json" + input_path.write_text(json.dumps(input_json)) + + def _raise_invalid_memory_uri( + *, + base_dir: Path | str, + memory_service_uri: str | None = None, + ) -> object: + del base_dir, memory_service_uri + raise ValueError("Unsupported memory service URI: unknown://x") + + monkeypatch.setattr( + cli, "create_memory_service_from_options", _raise_invalid_memory_uri + ) + + with pytest.raises(ValueError, match="Unsupported memory service URI"): + 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, + memory_service_uri="unknown://x", + ) + + +@pytest.mark.asyncio +async def test_run_cli_passes_memory_service_to_input_file( + fake_agent, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """run_cli should construct and pass the configured memory service.""" + parent_dir, folder_name = fake_agent + input_json = {"state": {}, "queries": []} + input_path = tmp_path / "memory_input.json" + input_path.write_text(json.dumps(input_json)) + + memory_service_sentinel = object() + captured_factory_args: dict[str, Any] = {} + captured_memory_service: dict[str, Any] = {} + + def _memory_factory( + *, + base_dir: Path | str, + memory_service_uri: str | None = None, + ) -> object: + captured_factory_args["base_dir"] = base_dir + captured_factory_args["memory_service_uri"] = memory_service_uri + return memory_service_sentinel + + async def _run_input_file( + app_name: str, + user_id: str, + agent_or_app: BaseAgent | App, + artifact_service: Any, + session_service: Any, + credential_service: InMemoryCredentialService, + input_path: str, + memory_service: Any = None, + ) -> object: + del app_name, user_id, agent_or_app, artifact_service + del session_service, credential_service, input_path + captured_memory_service["value"] = memory_service + return object() + + monkeypatch.setattr( + cli, "create_memory_service_from_options", _memory_factory + ) + monkeypatch.setattr(cli, "run_input_file", _run_input_file) + + 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, + memory_service_uri="memory://", + ) + + assert Path(captured_factory_args["base_dir"]) == parent_dir.resolve() + assert captured_factory_args["memory_service_uri"] == "memory://" + assert captured_memory_service["value"] is memory_service_sentinel + + +@pytest.mark.asyncio +async def test_run_cli_loads_dotenv_before_memory_service_creation( + fake_agent, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """run_cli should load agent .env values before creating memory service.""" + parent_dir, folder_name = fake_agent + input_json = {"state": {}, "queries": []} + input_path = tmp_path / "dotenv_order_input.json" + input_path.write_text(json.dumps(input_json)) + + call_order: list[str] = [] + + def _load_dotenv_for_agent(agent_name: str, agents_dir: str) -> None: + del agent_name, agents_dir + call_order.append("load_dotenv") + + def _memory_factory( + *, + base_dir: Path | str, + memory_service_uri: str | None = None, + ) -> object: + del base_dir, memory_service_uri + call_order.append("create_memory") + return object() + + monkeypatch.setenv("ADK_DISABLE_LOAD_DOTENV", "0") + monkeypatch.setattr(cli.envs, "load_dotenv_for_agent", _load_dotenv_for_agent) + monkeypatch.setattr( + cli, "create_memory_service_from_options", _memory_factory + ) + + 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, + memory_service_uri="memory://", + ) + + assert "create_memory" in call_order + assert "load_dotenv" in call_order + assert call_order.index("load_dotenv") < call_order.index("create_memory") + + @pytest.mark.asyncio async def test_run_interactively_whitespace_and_exit( tmp_path: Path, monkeypatch: pytest.MonkeyPatch diff --git a/tests/unittests/cli/utils/test_cli_tools_click.py b/tests/unittests/cli/utils/test_cli_tools_click.py index 61b1468c..7c642dbb 100644 --- a/tests/unittests/cli/utils/test_cli_tools_click.py +++ b/tests/unittests/cli/utils/test_cli_tools_click.py @@ -23,6 +23,7 @@ from types import SimpleNamespace from typing import Any from typing import Dict from typing import List +from typing import Optional from typing import Tuple from unittest import mock @@ -129,7 +130,7 @@ def test_cli_create_cmd_invokes_run_cmd( # cli run @pytest.mark.parametrize( - "cli_args,expected_session_uri,expected_artifact_uri", + "cli_args,expected_session_uri,expected_artifact_uri,expected_memory_uri", [ pytest.param( [ @@ -137,15 +138,19 @@ def test_cli_create_cmd_invokes_run_cmd( "memory://", "--artifact_service_uri", "memory://", + "--memory_service_uri", + "memory://", ], "memory://", "memory://", + "memory://", id="memory_scheme_uris", ), pytest.param( [], None, None, + None, id="default_uris_none", ), ], @@ -154,8 +159,9 @@ def test_cli_run_service_uris( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, cli_args: list, - expected_session_uri: str, - expected_artifact_uri: str, + expected_session_uri: Optional[str], + expected_artifact_uri: Optional[str], + expected_memory_uri: Optional[str], ) -> None: """`adk run` should forward service URIs correctly to run_cli.""" agent_dir = tmp_path / "agent" @@ -186,6 +192,7 @@ def test_cli_run_service_uris( coro_locals = captured_locals[0] assert coro_locals.get("session_service_uri") == expected_session_uri assert coro_locals.get("artifact_service_uri") == expected_artifact_uri + assert coro_locals.get("memory_service_uri") == expected_memory_uri assert coro_locals["agent_folder_name"] == "agent" diff --git a/tests/unittests/cli/utils/test_service_factory.py b/tests/unittests/cli/utils/test_service_factory.py index 910bf906..d6f1426a 100644 --- a/tests/unittests/cli/utils/test_service_factory.py +++ b/tests/unittests/cli/utils/test_service_factory.py @@ -252,6 +252,15 @@ def test_create_memory_service_defaults_to_in_memory(tmp_path: Path): assert isinstance(service, InMemoryMemoryService) +def test_create_memory_service_supports_memory_uri(tmp_path: Path): + service = service_factory.create_memory_service_from_options( + base_dir=tmp_path, + memory_service_uri="memory://", + ) + + assert isinstance(service, InMemoryMemoryService) + + def test_create_memory_service_raises_on_unknown_scheme( tmp_path: Path, monkeypatch ):