mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
fix: Remove 'per_agent' from kwargs when using remote session service URIs
Co-authored-by: George Weale <gweale@google.com> PiperOrigin-RevId: 837169299
This commit is contained in:
committed by
Copybara-Service
parent
ec4ccd718f
commit
73e5687b9a
@@ -106,7 +106,6 @@ def get_fast_api_app(
|
|||||||
base_dir=agents_dir,
|
base_dir=agents_dir,
|
||||||
session_service_uri=session_service_uri,
|
session_service_uri=session_service_uri,
|
||||||
session_db_kwargs=session_db_kwargs,
|
session_db_kwargs=session_db_kwargs,
|
||||||
per_agent=True, # Multi-agent mode
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Build the Artifact service
|
# Build the Artifact service
|
||||||
@@ -114,7 +113,6 @@ def get_fast_api_app(
|
|||||||
artifact_service = create_artifact_service_from_options(
|
artifact_service = create_artifact_service_from_options(
|
||||||
base_dir=agents_dir,
|
base_dir=agents_dir,
|
||||||
artifact_service_uri=artifact_service_uri,
|
artifact_service_uri=artifact_service_uri,
|
||||||
per_agent=True, # Multi-agent mode
|
|
||||||
)
|
)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
raise click.ClickException(str(exc)) from exc
|
raise click.ClickException(str(exc)) from exc
|
||||||
|
|||||||
@@ -271,18 +271,14 @@ def _register_builtin_services(registry: ServiceRegistry) -> None:
|
|||||||
|
|
||||||
kwargs_copy = kwargs.copy()
|
kwargs_copy = kwargs.copy()
|
||||||
kwargs_copy.pop("agents_dir", None)
|
kwargs_copy.pop("agents_dir", None)
|
||||||
|
kwargs_copy.pop("per_agent", None)
|
||||||
parsed_uri = urlparse(uri)
|
parsed_uri = urlparse(uri)
|
||||||
bucket_name = parsed_uri.netloc
|
bucket_name = parsed_uri.netloc
|
||||||
return GcsArtifactService(bucket_name=bucket_name, **kwargs_copy)
|
return GcsArtifactService(bucket_name=bucket_name, **kwargs_copy)
|
||||||
|
|
||||||
def file_artifact_factory(uri: str, **kwargs):
|
def file_artifact_factory(uri: str, **_):
|
||||||
from ..artifacts.file_artifact_service import FileArtifactService
|
from ..artifacts.file_artifact_service import FileArtifactService
|
||||||
|
|
||||||
per_agent = kwargs.get("per_agent", False)
|
|
||||||
if per_agent:
|
|
||||||
raise ValueError(
|
|
||||||
"file:// artifact URIs are not supported in multi-agent mode."
|
|
||||||
)
|
|
||||||
parsed_uri = urlparse(uri)
|
parsed_uri = urlparse(uri)
|
||||||
if parsed_uri.netloc not in ("", "localhost"):
|
if parsed_uri.netloc not in ("", "localhost"):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
@@ -58,13 +58,12 @@ def create_local_database_session_service(
|
|||||||
|
|
||||||
|
|
||||||
def create_local_artifact_service(
|
def create_local_artifact_service(
|
||||||
*, base_dir: Path | str, per_agent: bool = False
|
*, base_dir: Path | str
|
||||||
) -> BaseArtifactService:
|
) -> BaseArtifactService:
|
||||||
"""Creates a file-backed artifact service rooted in `.adk/artifacts`.
|
"""Creates a file-backed artifact service rooted in `.adk/artifacts`.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
base_dir: Directory whose `.adk` folder will store artifacts.
|
base_dir: Directory whose `.adk` folder will store artifacts.
|
||||||
per_agent: Indicates whether the service is being used in multi-agent mode.
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A `FileArtifactService` scoped to the derived root directory.
|
A `FileArtifactService` scoped to the derived root directory.
|
||||||
@@ -72,12 +71,6 @@ def create_local_artifact_service(
|
|||||||
manager = DotAdkFolder(base_dir)
|
manager = DotAdkFolder(base_dir)
|
||||||
artifact_root = manager.artifacts_dir
|
artifact_root = manager.artifacts_dir
|
||||||
artifact_root.mkdir(parents=True, exist_ok=True)
|
artifact_root.mkdir(parents=True, exist_ok=True)
|
||||||
if per_agent:
|
|
||||||
logger.info(
|
|
||||||
"Using shared file artifact service rooted at %s for multi-agent mode",
|
|
||||||
artifact_root,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
logger.info("Using file artifact service at %s", artifact_root)
|
logger.info("Using file artifact service at %s", artifact_root)
|
||||||
return FileArtifactService(root_dir=artifact_root)
|
return FileArtifactService(root_dir=artifact_root)
|
||||||
|
|
||||||
|
|||||||
@@ -32,7 +32,6 @@ def create_session_service_from_options(
|
|||||||
base_dir: Path | str,
|
base_dir: Path | str,
|
||||||
session_service_uri: Optional[str] = None,
|
session_service_uri: Optional[str] = None,
|
||||||
session_db_kwargs: Optional[dict[str, Any]] = None,
|
session_db_kwargs: Optional[dict[str, Any]] = None,
|
||||||
per_agent: bool = False,
|
|
||||||
) -> BaseSessionService:
|
) -> BaseSessionService:
|
||||||
"""Creates a session service based on CLI/web options."""
|
"""Creates a session service based on CLI/web options."""
|
||||||
base_path = Path(base_dir)
|
base_path = Path(base_dir)
|
||||||
@@ -40,17 +39,11 @@ def create_session_service_from_options(
|
|||||||
|
|
||||||
kwargs: dict[str, Any] = {
|
kwargs: dict[str, Any] = {
|
||||||
"agents_dir": str(base_path),
|
"agents_dir": str(base_path),
|
||||||
"per_agent": per_agent,
|
|
||||||
}
|
}
|
||||||
if session_db_kwargs:
|
if session_db_kwargs:
|
||||||
kwargs.update(session_db_kwargs)
|
kwargs.update(session_db_kwargs)
|
||||||
|
|
||||||
if session_service_uri:
|
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)
|
logger.info("Using session service URI: %s", session_service_uri)
|
||||||
service = registry.create_session_service(session_service_uri, **kwargs)
|
service = registry.create_session_service(session_service_uri, **kwargs)
|
||||||
if service is not None:
|
if service is not None:
|
||||||
@@ -63,7 +56,6 @@ def create_session_service_from_options(
|
|||||||
|
|
||||||
fallback_kwargs = dict(kwargs)
|
fallback_kwargs = dict(kwargs)
|
||||||
fallback_kwargs.pop("agents_dir", None)
|
fallback_kwargs.pop("agents_dir", None)
|
||||||
fallback_kwargs.pop("per_agent", None)
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"Falling back to DatabaseSessionService for URI: %s",
|
"Falling back to DatabaseSessionService for URI: %s",
|
||||||
session_service_uri,
|
session_service_uri,
|
||||||
@@ -105,23 +97,16 @@ def create_artifact_service_from_options(
|
|||||||
*,
|
*,
|
||||||
base_dir: Path | str,
|
base_dir: Path | str,
|
||||||
artifact_service_uri: Optional[str] = None,
|
artifact_service_uri: Optional[str] = None,
|
||||||
per_agent: bool = False,
|
|
||||||
) -> BaseArtifactService:
|
) -> BaseArtifactService:
|
||||||
"""Creates an artifact service based on CLI/web options."""
|
"""Creates an artifact service based on CLI/web options."""
|
||||||
base_path = Path(base_dir)
|
base_path = Path(base_dir)
|
||||||
registry = get_service_registry()
|
registry = get_service_registry()
|
||||||
|
|
||||||
if artifact_service_uri:
|
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)
|
logger.info("Using artifact service URI: %s", artifact_service_uri)
|
||||||
service = registry.create_artifact_service(
|
service = registry.create_artifact_service(
|
||||||
artifact_service_uri,
|
artifact_service_uri,
|
||||||
agents_dir=str(base_path),
|
agents_dir=str(base_path),
|
||||||
per_agent=per_agent,
|
|
||||||
)
|
)
|
||||||
if service is None:
|
if service is None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@@ -133,6 +118,4 @@ def create_artifact_service_from_options(
|
|||||||
return InMemoryArtifactService()
|
return InMemoryArtifactService()
|
||||||
return service
|
return service
|
||||||
|
|
||||||
if per_agent:
|
return create_local_artifact_service(base_dir=base_path)
|
||||||
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)
|
|
||||||
|
|||||||
@@ -300,13 +300,11 @@ def test_create_artifact_service_defaults_to_file(tmp_path: Path) -> None:
|
|||||||
assert expected_root.exists()
|
assert expected_root.exists()
|
||||||
|
|
||||||
|
|
||||||
def test_create_artifact_service_per_agent_uses_shared_root(
|
def test_create_artifact_service_uses_shared_root(
|
||||||
tmp_path: Path,
|
tmp_path: Path,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Multi-agent mode should still use a single file artifact service."""
|
"""Artifact service should use a single file artifact service."""
|
||||||
service = create_artifact_service_from_options(
|
service = create_artifact_service_from_options(base_dir=tmp_path)
|
||||||
base_dir=tmp_path, per_agent=True
|
|
||||||
)
|
|
||||||
assert isinstance(service, FileArtifactService)
|
assert isinstance(service, FileArtifactService)
|
||||||
expected_root = Path(tmp_path) / ".adk" / "artifacts"
|
expected_root = Path(tmp_path) / ".adk" / "artifacts"
|
||||||
assert service.root_dir == expected_root
|
assert service.root_dir == expected_root
|
||||||
@@ -332,17 +330,6 @@ def test_create_artifact_service_accepts_file_uri(tmp_path: Path) -> None:
|
|||||||
assert custom_root.exists()
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_run_cli_accepts_memory_scheme(
|
async def test_run_cli_accepts_memory_scheme(
|
||||||
fake_agent, tmp_path: Path
|
fake_agent, tmp_path: Path
|
||||||
|
|||||||
@@ -41,35 +41,12 @@ def test_create_session_service_uses_registry(tmp_path: Path, monkeypatch):
|
|||||||
registry.create_session_service.assert_called_once_with(
|
registry.create_session_service.assert_called_once_with(
|
||||||
"sqlite:///test.db",
|
"sqlite:///test.db",
|
||||||
agents_dir=str(tmp_path),
|
agents_dir=str(tmp_path),
|
||||||
per_agent=False,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_create_session_service_per_agent_uri(tmp_path: Path, monkeypatch):
|
def test_create_session_service_defaults_to_memory(tmp_path: Path):
|
||||||
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(
|
service = service_factory.create_session_service_from_options(
|
||||||
base_dir=tmp_path,
|
base_dir=tmp_path,
|
||||||
per_agent=per_agent,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert isinstance(service, InMemorySessionService)
|
assert isinstance(service, InMemorySessionService)
|
||||||
@@ -94,15 +71,11 @@ def test_create_session_service_fallbacks_to_database(
|
|||||||
registry.create_session_service.assert_called_once_with(
|
registry.create_session_service.assert_called_once_with(
|
||||||
"sqlite+aiosqlite:///:memory:",
|
"sqlite+aiosqlite:///:memory:",
|
||||||
agents_dir=str(tmp_path),
|
agents_dir=str(tmp_path),
|
||||||
per_agent=False,
|
|
||||||
echo=True,
|
echo=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("per_agent", [True, False])
|
def test_create_artifact_service_uses_registry(tmp_path: Path, monkeypatch):
|
||||||
def test_create_artifact_service_uses_registry(
|
|
||||||
tmp_path: Path, monkeypatch, per_agent: bool
|
|
||||||
):
|
|
||||||
registry = Mock()
|
registry = Mock()
|
||||||
expected = object()
|
expected = object()
|
||||||
registry.create_artifact_service.return_value = expected
|
registry.create_artifact_service.return_value = expected
|
||||||
@@ -111,14 +84,12 @@ def test_create_artifact_service_uses_registry(
|
|||||||
result = service_factory.create_artifact_service_from_options(
|
result = service_factory.create_artifact_service_from_options(
|
||||||
base_dir=tmp_path,
|
base_dir=tmp_path,
|
||||||
artifact_service_uri="gs://bucket/path",
|
artifact_service_uri="gs://bucket/path",
|
||||||
per_agent=per_agent,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result is expected
|
assert result is expected
|
||||||
registry.create_artifact_service.assert_called_once_with(
|
registry.create_artifact_service.assert_called_once_with(
|
||||||
"gs://bucket/path",
|
"gs://bucket/path",
|
||||||
agents_dir=str(tmp_path),
|
agents_dir=str(tmp_path),
|
||||||
per_agent=per_agent,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user