feat: Change service creation and add app name mapping for sessions

This change refactors how session, memory, and artifact services are created in the fast_api server, using the shared service_factory.

Co-authored-by: George Weale <gweale@google.com>
PiperOrigin-RevId: 839997110
This commit is contained in:
George Weale
2025-12-03 18:27:16 -08:00
committed by Copybara-Service
parent e9182e5eb4
commit 0a07a667e9
7 changed files with 121 additions and 60 deletions
+6 -6
View File
@@ -416,15 +416,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(
@@ -556,15 +556,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(
@@ -17,6 +17,7 @@ from __future__ import annotations
from pathlib import Path
from google.adk.cli.utils.local_storage import create_local_database_session_service
from google.adk.cli.utils.local_storage import create_local_session_service
from google.adk.cli.utils.local_storage import PerAgentDatabaseSessionService
from google.adk.sessions.sqlite_session_service import SqliteSessionService
import pytest
@@ -48,6 +49,29 @@ async def test_per_agent_session_service_creates_scoped_dot_adk(
assert agent_b_sessions.sessions[0].app_name == "agent_b"
@pytest.mark.asyncio
async def test_per_agent_session_service_respects_app_name_alias(
tmp_path: Path,
) -> None:
folder_name = "agent_folder"
logical_name = "custom_app"
(tmp_path / folder_name).mkdir()
service = create_local_session_service(
base_dir=tmp_path,
per_agent=True,
app_name_to_dir={logical_name: folder_name},
)
session = await service.create_session(
app_name=logical_name,
user_id="user",
)
assert session.app_name == logical_name
assert (tmp_path / folder_name / ".adk" / "session.db").exists()
def test_create_local_database_session_service_returns_sqlite(
tmp_path: Path,
) -> None:
@@ -60,6 +60,25 @@ async def test_create_session_service_defaults_to_per_agent_sqlite(
assert (agent_dir / ".adk" / "session.db").exists()
@pytest.mark.asyncio
async def test_create_session_service_respects_app_name_mapping(
tmp_path: Path,
) -> None:
agent_dir = tmp_path / "agent_folder"
logical_name = "custom_app"
agent_dir.mkdir()
service = service_factory.create_session_service_from_options(
base_dir=tmp_path,
app_name_to_dir={logical_name: "agent_folder"},
)
assert isinstance(service, PerAgentDatabaseSessionService)
session = await service.create_session(app_name=logical_name, user_id="user")
assert session.app_name == logical_name
assert (agent_dir / ".adk" / "session.db").exists()
def test_create_session_service_fallbacks_to_database(
tmp_path: Path, monkeypatch
):
@@ -101,6 +120,21 @@ def test_create_artifact_service_uses_registry(tmp_path: Path, monkeypatch):
)
def test_create_artifact_service_raises_on_unknown_scheme_when_strict(
tmp_path: Path, monkeypatch
):
registry = Mock()
registry.create_artifact_service.return_value = None
monkeypatch.setattr(service_factory, "get_service_registry", lambda: registry)
with pytest.raises(ValueError):
service_factory.create_artifact_service_from_options(
base_dir=tmp_path,
artifact_service_uri="unknown://foo",
strict_uri=True,
)
def test_create_memory_service_uses_registry(tmp_path: Path, monkeypatch):
registry = Mock()
expected = object()