mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
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:
committed by
Copybara-Service
parent
e9182e5eb4
commit
0a07a667e9
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user