diff --git a/src/google/adk/cli/api_server.py b/src/google/adk/cli/api_server.py index 12d56bc07a7..00799f68655 100644 --- a/src/google/adk/cli/api_server.py +++ b/src/google/adk/cli/api_server.py @@ -21,6 +21,7 @@ import asyncio from contextlib import asynccontextmanager import importlib +import inspect import json import logging import os @@ -54,6 +55,7 @@ from opentelemetry.sdk.trace import SpanProcessor from opentelemetry.sdk.trace import TracerProvider from pydantic import Field +from pydantic import model_validator from pydantic import ValidationError from starlette.types import Lifespan from typing_extensions import deprecated @@ -282,6 +284,14 @@ def _is_request_origin_allowed( return origin == request_origin +def _accepts_arbitrary_kwargs(func: Callable[..., Any]) -> bool: + """Returns True if func accepts arbitrary keyword arguments (**kwargs).""" + return any( + param.kind == inspect.Parameter.VAR_KEYWORD + for param in inspect.signature(func).parameters.values() + ) + + _SAFE_HTTP_METHODS = frozenset({"GET", "HEAD", "OPTIONS"}) @@ -471,6 +481,32 @@ class CreateSessionRequest(common.BaseModel): default=None, description="A list of events to initialize the session with.", ) + expire_time: Optional[str] = Field( + default=None, + description=( + "The absolute expiration time of the session as an RFC 3339" + " timestamp, e.g. '2026-08-01T00:00:00Z'. Mutually exclusive with" + " ttl. Only supported by session services that accept expiration" + " arguments, such as VertexAiSessionService." + ), + ) + ttl: Optional[str] = Field( + default=None, + description=( + "The time-to-live of the session as a duration, e.g. '7200s'." + " Mutually exclusive with expire_time. Only supported by session" + " services that accept expiration arguments, such as" + " VertexAiSessionService." + ), + ) + + @model_validator(mode="after") + def _validate_expiration_mutually_exclusive(self) -> CreateSessionRequest: + if self.ttl is not None and self.expire_time is not None: + raise ValueError( + "Cannot specify both 'ttl' and 'expire_time' simultaneously." + ) + return self class SaveArtifactRequest(common.BaseModel): @@ -937,13 +973,31 @@ async def _create_session( user_id: str, session_id: Optional[str] = None, state: Optional[dict[str, Any]] = None, + expire_time: Optional[str] = None, + ttl: Optional[str] = None, ) -> Session: + expiration_kwargs: dict[str, Any] = {} + if expire_time is not None: + expiration_kwargs["expire_time"] = expire_time + if ttl is not None: + expiration_kwargs["ttl"] = ttl + if expiration_kwargs and not _accepts_arbitrary_kwargs( + self.session_service.create_session + ): + raise HTTPException( + status_code=400, + detail=( + "Session expiration (ttl/expire_time) is not supported by the" + " configured session service." + ), + ) try: session = await self.session_service.create_session( app_name=app_name, user_id=user_id, state=state, session_id=session_id, + **expiration_kwargs, ) logger.info("New session created: %s", session.id) return session @@ -1251,6 +1305,8 @@ async def create_session( user_id=user_id, state=req.state, session_id=req.session_id, + expire_time=req.expire_time, + ttl=req.ttl, ) if req.events: diff --git a/tests/unittests/cli/test_fast_api.py b/tests/unittests/cli/test_fast_api.py index 42f311f5a9d..52fcca08e7d 100755 --- a/tests/unittests/cli/test_fast_api.py +++ b/tests/unittests/cli/test_fast_api.py @@ -45,6 +45,7 @@ from google.adk.plugins.bigquery_agent_analytics_plugin import BigQueryAgentAnalyticsPlugin from google.adk.runners import Runner from google.adk.sessions.in_memory_session_service import InMemorySessionService +from google.adk.sessions.session import Session from google.genai import types from pydantic import BaseModel import pytest @@ -1370,6 +1371,106 @@ def test_create_session_without_id(test_app, test_session_info): logger.info(f"Created session with generated ID: {data['id']}") +class _ExpirationRecordingSessionService(InMemorySessionService): + """In-memory session service that accepts and records expiration kwargs.""" + + def __init__(self): + super().__init__() + self.recorded_kwargs: dict[str, Any] = {} + + async def create_session( + self, + *, + app_name: str, + user_id: str, + state: Optional[dict[str, Any]] = None, + session_id: Optional[str] = None, + **kwargs: Any, + ) -> Session: + self.recorded_kwargs = dict(kwargs) + return await super().create_session( + app_name=app_name, + user_id=user_id, + state=state, + session_id=session_id, + ) + + +@pytest.fixture +def expiration_session_service(): + """Create a session service whose create_session accepts expiration kwargs.""" + return _ExpirationRecordingSessionService() + + +@pytest.fixture +def expiration_test_app( + expiration_session_service, + mock_artifact_service, + mock_memory_service, + mock_agent_loader, + mock_eval_sets_manager, + mock_eval_set_results_manager, +): + """Create a TestClient backed by an expiration-capable session service.""" + return _create_test_client( + expiration_session_service, + mock_artifact_service, + mock_memory_service, + mock_agent_loader, + mock_eval_sets_manager, + mock_eval_set_results_manager, + ) + + +def test_create_session_forwards_expire_time( + expiration_test_app, expiration_session_service, test_session_info +): + """Test that expire_time is forwarded to a supporting session service.""" + url = f"/apps/{test_session_info['app_name']}/users/{test_session_info['user_id']}/sessions" + response = expiration_test_app.post( + url, json={"expireTime": "2026-08-01T00:00:00Z"} + ) + + assert response.status_code == 200 + assert expiration_session_service.recorded_kwargs == { + "expire_time": "2026-08-01T00:00:00Z" + } + + +def test_create_session_forwards_ttl( + expiration_test_app, expiration_session_service, test_session_info +): + """Test that ttl is forwarded to a supporting session service.""" + url = f"/apps/{test_session_info['app_name']}/users/{test_session_info['user_id']}/sessions" + response = expiration_test_app.post(url, json={"ttl": "7200s"}) + + assert response.status_code == 200 + assert expiration_session_service.recorded_kwargs == {"ttl": "7200s"} + + +def test_create_session_rejects_ttl_with_expire_time( + expiration_test_app, test_session_info +): + """Test that requesting both ttl and expire_time is rejected.""" + url = f"/apps/{test_session_info['app_name']}/users/{test_session_info['user_id']}/sessions" + response = expiration_test_app.post( + url, json={"ttl": "7200s", "expireTime": "2026-08-01T00:00:00Z"} + ) + + assert response.status_code == 422 + + +def test_create_session_expiration_with_unsupported_service( + test_app, test_session_info +): + """Test 400 when expiration is requested but the service cannot honor it.""" + url = f"/apps/{test_session_info['app_name']}/users/{test_session_info['user_id']}/sessions" + response = test_app.post(url, json={"ttl": "7200s"}) + + assert response.status_code == 400 + assert "not supported" in response.json()["detail"] + + def test_get_session(test_app, create_test_session): """Test retrieving a session by ID.""" info = create_test_session