From 5e58e186f95d42d266070d9deec57ba6793d08c6 Mon Sep 17 00:00:00 2001 From: Abhishek Singh Date: Thu, 3 Sep 2026 10:38:29 -0700 Subject: [PATCH] Refactor A2A gateway routing and session lifecycle - simplify gateway request handling and A2UI bootstrap flow - add database-backed task routing and worker protocol types - clean up expired worker child processes - improve GCP gateway architecture and mTLS deployment - expand A2A gateway and worker tests --- gcloud/gateway/README.md | 78 +- gcloud/gateway/deploy.sh | 13 +- src/select_ai/agent/a2a/a2ui.py | 78 +- src/select_ai/agent/a2a/forms.py | 27 +- src/select_ai/agent/a2a/gateway.py | 450 +++++++++--- src/select_ai/agent/a2a/session_runtime.py | 120 ++++ src/select_ai/agent/a2a/task_store.py | 57 +- src/select_ai/agent/a2a/worker.py | 352 ++++++--- src/select_ai/agent/a2a/worker_client.py | 194 ++++- src/select_ai/agent/a2a/worker_protocol.py | 103 +++ src/select_ai/cli/a2a.py | 2 +- tests/a2a/test_worker_runtime.py | 789 +++++++++++++++++++-- 12 files changed, 1898 insertions(+), 365 deletions(-) create mode 100644 src/select_ai/agent/a2a/session_runtime.py create mode 100644 src/select_ai/agent/a2a/worker_protocol.py diff --git a/gcloud/gateway/README.md b/gcloud/gateway/README.md index a01ee79..a117207 100644 --- a/gcloud/gateway/README.md +++ b/gcloud/gateway/README.md @@ -1,4 +1,4 @@ -# Dynamic gateway deployment +# Dynamic gateway Dynamic gateway mode exposes one public A2A endpoint. Each user dynamically selects an Oracle database connection and Select AI team through the A2UI @@ -6,15 +6,62 @@ connection form. A session remains available for 15 minutes by default. Set a different lifetime in seconds with `--session-ttl-seconds`; for example, `--session-ttl-seconds 1800` keeps sessions for 30 minutes. +## Protocol architecture + +```text +┌──────────────┐ A2A JSON-RPC/HTTP ┌──────────────────┐ ┌────────────────────────────┐ +│ A2A client │────────────────────►│ Gateway instances│──── route lookup/update ──────────► │ Service Registry │ +└──────────────┘ │ Public A2A API │ │ Service discovery │ + │ A2UI bootstrap │ │ Session routes │ + └────────┬─────────┘ │ Task routes │ + │ Internal protobuf │ │ + ▼ │ │ + ┌──────────────────┐ │ │ + │ Worker pool │──── registration / heartbeat------->│ │ + │ worker-0, ... │ │ │ + └────────┬─────────┘ └────────────────────────────┘ + │ one child per database session + ▼ + ┌──────────────────┐ + │ Session runtime │ + │ A2A handler │ + │ Task/context │ + │ stores │ + │ Database session │ + └────────┬─────────┘ + │ SQL / Select AI + ▼ + ┌──────────────────┐ + │ Oracle Database │ + └──────────────────┘ +``` + +The gateway is the only public A2A application. It selects a worker through +the Service Registry, opens a session there, and proxies subsequent A2A calls +using the internal protobuf protocol. The selected worker starts one child +runtime for that session. The child owns the database connection, +`DefaultRequestHandler`, `OracleTaskStore`, and `OracleContextStore`. + +The Service Registry stores only service-discovery and non-secret +session/task-to-worker metadata. Task payloads and context mappings remain in +Oracle. Connection-form tasks are response-only bootstrap tasks: they are +created by the gateway before a database session exists and are not persisted +or routed. + +## GCP deployment + +The protocol architecture above is implemented on GCP as follows: + ```text ┌──────────────────────┐ │ A2A / Gemini client │ └──────────┬───────────┘ │ public A2A v - ┌──────────────────────┐ - │ Cloud Run gateway │ - └──────┬───────┬───────┘ + ┌────────────────────────────┐ + │ Cloud Run gateway │ + │ A2A proxy + form bootstrap │ + └──────┬───────────┬─────────┘ │ │ private VPC: mTLS request to worker hostname │ │ │ │ ┌─────────────────────── GKE ───────────────────────┐ @@ -23,19 +70,16 @@ different lifetime in seconds with `--session-ttl-seconds`; for example, │ │ │ │ │ │ v │ │ │ [StatefulSet worker-0 / worker-1 / ...] │ - │ │ session child process → Oracle Database │ + │ │ session child: A2A handler + Oracle stores │ + │ │ │ │ + │ │ v │ + │ │ Oracle Database │ │ │ │ │ │ [Consul] │ └────────────>│ selects healthy worker; returns worker hostname │ └───────────────────────────────────────────────────┘ ``` -The gateway is the only public A2A application. Consul and workers are a GKE -clustered service: Consul selects a worker for each new dynamic session, and -the chosen worker retains that session's process and Oracle conversation. - -## Deploy the complete stack - Run this from the repository root: ```bash @@ -68,10 +112,14 @@ gcloud/gateway/deploy.sh \ ``` The Cloud Run gateway uses direct VPC egress to reach the internal Consul load -balancer and GKE worker pod addresses. The default one-instance gateway limit -is intentional: gateway A2A task and context/session state is currently in -memory. Workers, rather than the gateway, provide the clustered capacity for -dynamic sessions. +balancer and GKE worker pod addresses. The gateway keeps only the connection +form task transiently, before a database session exists. Connected task and +context state is stored in Oracle on the selected worker. Workers, rather +than the gateway, provide the clustered capacity for dynamic sessions. +The deployment currently keeps one gateway instance as an operational default; +the gateway does not cache forms or connected task/context state. Gateway +scaling does not change session affinity because Consul stores the session and +task routes. `cloudbuild.yaml` is the complete build and deployment workflow. It supplies the generated image and Consul endpoint values to the Cloud Run gateway at diff --git a/gcloud/gateway/deploy.sh b/gcloud/gateway/deploy.sh index 55959c8..08a93ab 100755 --- a/gcloud/gateway/deploy.sh +++ b/gcloud/gateway/deploy.sh @@ -166,10 +166,15 @@ if [[ "$enable_worker_mtls" == "true" ]]; then fi done if [[ "$create_mtls_material" != "true" ]]; then - existing_worker_certificate_sans="$(gcloud secrets versions access latest \ - --secret=select-ai-worker-mtls-cert --project="$project_id" 2>/dev/null | \ - openssl x509 -noout -ext subjectAltName 2>/dev/null || true)" - if [[ "$existing_worker_certificate_sans" != *"DNS:$worker_certificate_dns_name"* ]]; then + # macOS ships LibreSSL, which does not support x509's -ext option. + # Read the portable text representation and perform a literal SAN check. + existing_worker_certificate_text="$( + gcloud secrets versions access latest \ + --secret=select-ai-worker-mtls-cert --project="$project_id" 2>/dev/null | \ + openssl x509 -noout -text 2>/dev/null || true + )" + if ! printf '%s\n' "$existing_worker_certificate_text" | \ + grep -F -- "DNS:$worker_certificate_dns_name" >/dev/null; then create_mtls_material="true" echo "Replacing mTLS material because the worker certificate does not" echo "match the GKE DNS domain." diff --git a/src/select_ai/agent/a2a/a2ui.py b/src/select_ai/agent/a2a/a2ui.py index a25b7f9..331e8c5 100644 --- a/src/select_ai/agent/a2a/a2ui.py +++ b/src/select_ai/agent/a2a/a2ui.py @@ -5,24 +5,34 @@ # https://oss.oracle.com/licenses/upl. # ----------------------------------------------------------------------------- -"""Shared A2UI protocol declarations for Select AI A2A agents.""" +"""Shared A2UI protocol helpers for Select AI A2A agents.""" +from __future__ import annotations + +from collections.abc import Iterator + +from a2a.helpers import new_data_part from a2a.types import AgentExtension -from google.protobuf.json_format import ParseDict +from a2a.types.a2a_pb2 import Message, Part +from google.protobuf.json_format import MessageToDict, ParseDict from google.protobuf.struct_pb2 import Struct -from select_ai.agent.a2a.forms import _CATALOG - -A2UI_EXTENSION_URI = "https://a2ui.org/a2a-extension/a2ui/v0.9" +A2UI_VERSION = "v0.9" +A2UI_EXTENSION_URI = f"https://a2ui.org/a2a-extension/a2ui/{A2UI_VERSION}" A2UI_MIME_TYPE = "application/json+a2ui" +A2UI_CATALOG_ID = ( + "https://www.gstatic.com/vertexaisearch/a2ui/" + f"{A2UI_VERSION.replace('.', '_')}/" + "gemini_enterprise_composite_catalog.json" +) def a2ui_extension() -> AgentExtension: - """Return the A2UI v0.9 capability used by Gemini Enterprise.""" + """Return the supported A2UI capability used by Gemini Enterprise.""" params = ParseDict( { "acceptsInlineCatalogs": True, - "supportedCatalogIds": [_CATALOG], + "supportedCatalogIds": [A2UI_CATALOG_ID], }, Struct(), ) @@ -31,3 +41,57 @@ def a2ui_extension() -> AgentExtension: description="Provides agent driven UI using the A2UI JSON format.", params=params, ) + + +def a2ui_part(operation: dict) -> Part: + """Encode one A2UI operation in an A2A data part.""" + part = new_data_part(operation) + ParseDict({"mimeType": A2UI_MIME_TYPE}, part.metadata) + return part + + +def a2ui_operations(message: Message) -> Iterator[dict]: + """Yield operations from data parts marked with the A2UI MIME type.""" + for part in message.parts: + if not _is_a2ui_part(part): + continue + data = MessageToDict(part.data) + operations = data if isinstance(data, list) else [data] + yield from ( + operation + for operation in operations + if isinstance(operation, dict) + ) + + +def find_action(message: Message, name: str) -> dict | None: + """Return a named A2UI action, including Gemini's unmarked input form.""" + for part in message.parts: + if part.WhichOneof("content") != "data": + continue + mime_type = MessageToDict(part.metadata).get("mimeType") + # Gemini Enterprise does not currently echo the A2UI MIME metadata on + # a submitted form action. Accept that legacy input shape, while + # still rejecting data parts explicitly marked as another format. + if mime_type not in (None, A2UI_MIME_TYPE): + continue + data = MessageToDict(part.data) + operations = data if isinstance(data, list) else [data] + for operation in operations: + if ( + not isinstance(operation, dict) + or operation.get("version") != A2UI_VERSION + ): + continue + action = operation.get("action") + if isinstance(action, dict) and action.get("name") == name: + return action + return None + + +def _is_a2ui_part(part: Part) -> bool: + """Identify an A2UI data part by its standard MIME type.""" + return ( + part.WhichOneof("content") == "data" + and MessageToDict(part.metadata).get("mimeType") == A2UI_MIME_TYPE + ) diff --git a/src/select_ai/agent/a2a/forms.py b/src/select_ai/agent/a2a/forms.py index 6823464..9d82bff 100644 --- a/src/select_ai/agent/a2a/forms.py +++ b/src/select_ai/agent/a2a/forms.py @@ -7,27 +7,30 @@ """A2UI connection form emitted by the public gateway.""" +from uuid import uuid4 -_CATALOG = ( - "https://www.gstatic.com/vertexaisearch/a2ui/v0_9/" - "gemini_enterprise_composite_catalog.json" -) +from select_ai.agent.a2a.a2ui import A2UI_CATALOG_ID, A2UI_VERSION -def connection_form() -> list[dict]: +def connection_form(surface_id: str | None = None) -> list[dict]: """Return the non-persistent database connection form.""" + # A2UI surface IDs must be globally unique for the renderer's lifetime. + # Gemini retains surfaces for an A2A conversation after the connection + # form is submitted, so reusing a fixed ID prevents a reconnect form from + # being created in that same conversation. + surface_id = surface_id or f"db-connect-{uuid4().hex}" return [ { - "version": "v0.9", + "version": A2UI_VERSION, "createSurface": { - "surfaceId": "db-connect", - "catalogId": _CATALOG, + "surfaceId": surface_id, + "catalogId": A2UI_CATALOG_ID, }, }, { - "version": "v0.9", + "version": A2UI_VERSION, "updateComponents": { - "surfaceId": "db-connect", + "surfaceId": surface_id, "components": [ {"id": "root", "component": "Card", "child": "column"}, { @@ -102,9 +105,9 @@ def connection_form() -> list[dict]: }, }, { - "version": "v0.9", + "version": A2UI_VERSION, "updateDataModel": { - "surfaceId": "db-connect", + "surfaceId": surface_id, "path": "/", "value": { "dsn": "", diff --git a/src/select_ai/agent/a2a/gateway.py b/src/select_ai/agent/a2a/gateway.py index 622e2c7..1fc9834 100644 --- a/src/select_ai/agent/a2a/gateway.py +++ b/src/select_ai/agent/a2a/gateway.py @@ -10,25 +10,43 @@ from __future__ import annotations import asyncio +from uuid import uuid4 import requests from a2a.compat.v0_3.conversions import to_compat_agent_card from a2a.helpers import ( - new_data_part, - new_task_from_user_message, + new_artifact, new_text_part, ) -from a2a.server.agent_execution import AgentExecutor -from a2a.server.request_handlers import DefaultRequestHandler +from a2a.server.context import ServerCallContext +from a2a.server.request_handlers import RequestHandler from a2a.server.routes import create_jsonrpc_routes -from a2a.server.tasks import InMemoryTaskStore, TaskUpdater from a2a.types import ( AgentCapabilities, AgentCard, AgentInterface, AgentSkill, ) -from google.protobuf.json_format import MessageToDict, ParseDict +from a2a.types.a2a_pb2 import ( + AgentCard as AgentCardMessage, +) +from a2a.types.a2a_pb2 import ( + CancelTaskRequest, + GetExtendedAgentCardRequest, + GetTaskRequest, + ListTasksRequest, + ListTasksResponse, + Message, + SendMessageRequest, + Task, + TaskState, + TaskStatus, +) +from a2a.utils.errors import ( + InvalidParamsError, + TaskNotFoundError, + UnsupportedOperationError, +) from starlette.applications import Starlette from starlette.responses import JSONResponse from starlette.routing import Route @@ -37,63 +55,198 @@ A2UI_EXTENSION_URI, A2UI_MIME_TYPE, a2ui_extension, + a2ui_part, + find_action, ) from select_ai.agent.a2a.forms import connection_form from select_ai.agent.a2a.models import GatewaySettings, SessionInfo -from select_ai.agent.a2a.results import add_team_result from select_ai.agent.a2a.worker_client import ReconnectRequired, WorkerClient from select_ai.version import __version__ +_CONNECTION_ACTION_NAME = "submit_database_connection" +_CONNECTION_SUCCESS_MESSAGE = "Connected. Ask a database question." +_CONNECTION_FAILURE_MESSAGE = ( + "Could not connect. Check the DSN, credentials, and team name." +) +_BOOTSTRAP_TASK_ID_PREFIX = "gateway-bootstrap-" + + +async def _unsupported_operation(*_args, **_kwargs): + """Reject an A2A operation not advertised by the gateway.""" + raise UnsupportedOperationError + + +class _UnsupportedStream: + """Async iterator that reports unsupported streaming when consumed.""" + + def __aiter__(self): + return self -class GatewayExecutor(AgentExecutor): - """Route each A2A context to one in-memory worker session.""" + async def __anext__(self): + raise UnsupportedOperationError - def __init__(self, worker_client: WorkerClient): + +def _unsupported_stream(*_args, **_kwargs): + """Return an async iterator that rejects an unsupported A2A stream.""" + return _UnsupportedStream() + + +class GatewayRequestHandler(RequestHandler): + """Proxy A2A requests to the Oracle-backed handler in a worker child.""" + + def __init__(self, worker_client: WorkerClient, agent_card: AgentCard): self.worker_client = worker_client - self.sessions: dict[str, str] = {} + self.agent_card = agent_card + + # RequestHandler requires every operation even when the Agent Card does + # not advertise streaming or push notifications. + on_message_send_stream = _unsupported_stream + on_create_task_push_notification_config = _unsupported_operation + on_get_task_push_notification_config = _unsupported_operation + on_list_task_push_notification_configs = _unsupported_operation + on_delete_task_push_notification_config = _unsupported_operation + on_subscribe_to_task = _unsupported_stream + + async def on_message_send( + self, + params: SendMessageRequest, + _context: ServerCallContext, + ) -> Task | Message: + """Handle connection bootstrap or forward the request to a worker.""" + message = params.message + if action := find_action(message, _CONNECTION_ACTION_NAME): + return await self._handle_connection_action(message, action) + + response = await self._recover_task_context(message) + if response is not None: + return response + + context_id = self._ensure_context_id(message) + had_task_id = bool(message.task_id) - async def execute(self, context, event_queue): - task = context.current_task or new_task_from_user_message( - context.message + if not await asyncio.to_thread( + self.worker_client.session_exists, + context_id, + ): + return self._build_connection_form_task( + task_id=message.task_id or None, + context_id=context_id, + history=None if had_task_id else [message], + ) + + return await self._forward_message(context_id, params) + + async def _handle_connection_action( + self, + message: Message, + action: dict, + ) -> Task: + """Open a database session from the submitted A2UI form.""" + if not message.context_id: + raise InvalidParamsError( + "A database connection action requires contextId." + ) + + task = _new_bootstrap_task( + task_id=message.task_id or None, + context_id=message.context_id, ) - if context.current_task is None: - await event_queue.enqueue_event(task) - context_id = task.context_id or task.id - updater = TaskUpdater( - event_queue=event_queue, - task_id=task.id, - context_id=task.context_id, + session_id = await self._open_session( + action.get("context") or {}, + message.context_id, ) - await updater.start_work() - action = self._a2ui_action(context) - if action and action.get("name") == "submit_database_connection": - parts = await self._open_session( - action.get("context", {}), - context_id, - ) - artifact_name = "database-session" - extensions = None - elif context_id not in self.sessions: - parts = [_a2ui_part(message) for message in connection_form()] - artifact_name = "database-connection-form" - extensions = [A2UI_EXTENSION_URI] - else: - result = await self._send_prompt( - context_id, - context.get_user_input(), - ) - await add_team_result(updater, result) - await updater.complete() - return - await updater.add_artifact( - parts=parts, - name=artifact_name, - last_chunk=True, - extensions=extensions, + text = ( + _CONNECTION_SUCCESS_MESSAGE + if session_id is not None + else _CONNECTION_FAILURE_MESSAGE ) - await updater.complete() + return self._complete_task( + task, + [new_text_part(text)], + "database-session", + ) + + async def _recover_task_context(self, message: Message) -> Task | None: + """Recover a context when a client sends only a previous task ID.""" + if message.context_id or not message.task_id: + return None + + try: + existing_task = await asyncio.to_thread( + self.worker_client.get_task, + GetTaskRequest(id=message.task_id), + ) + except ReconnectRequired as error: + if error.session_id is None: + raise + return self._build_connection_form_task( + task_id=message.task_id or None, + context_id=error.session_id, + ) + if existing_task is None: + # Gemini Enterprise may resume a conversation with only its + # previous taskId. Task routing metadata is intentionally + # ephemeral, so it can be absent after a deployment or Consul + # restart even though the Gemini conversation still exists. + return self._build_connection_form_task( + task_id=message.task_id, + context_id=str(uuid4()), + ) + message.context_id = existing_task.context_id + return None - async def _open_session(self, action_context: dict, context_id: str): + @staticmethod + def _ensure_context_id(message: Message) -> str: + """Assign a context ID to a new message when one was not provided.""" + if message.context_id: + return message.context_id + context_id = str(uuid4()) + message.context_id = context_id + return context_id + + async def _forward_message( + self, + context_id: str, + params: SendMessageRequest, + ) -> Task | Message: + """Forward a message and turn worker loss into a reconnect task.""" + try: + try: + result = await asyncio.to_thread( + self.worker_client.send_message, + context_id, + params, + ) + except TaskNotFoundError: + # Gemini continues the gateway's transient connection-form + # task after the database session opens. That task was never + # persisted in Oracle, so let the worker create its first + # durable task while retaining the connected context. + if not params.message.task_id.startswith( + _BOOTSTRAP_TASK_ID_PREFIX + ): + raise + params.message.ClearField("task_id") + result = await asyncio.to_thread( + self.worker_client.send_message, + context_id, + params, + ) + except ReconnectRequired: + return self._build_connection_form_task( + task_id=params.message.task_id or None, + context_id=context_id, + history=(None if params.message.task_id else [params.message]), + ) + if isinstance(result, (Task, Message)): + return result + raise RuntimeError("The worker returned no A2A response.") + + async def _open_session( + self, + action_context: dict, + context_id: str, + ) -> str | None: try: session_info = SessionInfo.from_a2ui_event(action_context) session_id = await asyncio.to_thread( @@ -101,71 +254,156 @@ async def _open_session(self, action_context: dict, context_id: str): context_id, session_info, ) + return session_id except (requests.RequestException, ValueError, RuntimeError): - return [ - new_text_part( - "Could not connect. Check the DSN, credentials, and team " - "name." - ) - ] - self.sessions[context_id] = session_id - return [new_text_part("Connected. Ask a database question.")] + return None - async def _send_prompt(self, context_id: str, prompt: str) -> str | None: + async def on_get_task( + self, + params: GetTaskRequest, + _context: ServerCallContext, + ) -> Task | None: try: - result = await asyncio.to_thread( - self.worker_client.send_prompt, - self.sessions[context_id], - prompt, + task = await asyncio.to_thread( + self.worker_client.get_task, + params, ) - except ReconnectRequired: - self.sessions.pop(context_id, None) - return "Your database session ended. Please reconnect." - return result + except ReconnectRequired as error: + if error.session_id is None: + raise + return self._build_connection_form_task( + task_id=params.id, + context_id=error.session_id, + ) + if task is None: + return self._build_connection_form_task( + task_id=params.id, + context_id=str(uuid4()), + ) + return task - async def cancel(self, context, event_queue): - """Close the live worker session when the A2A task is cancelled.""" - task = context.current_task + async def on_list_tasks( + self, + params: ListTasksRequest, + _context: ServerCallContext, + ) -> ListTasksResponse: + if not params.context_id: + raise InvalidParamsError( + "ListTasks requires contextId for this gateway." + ) + try: + return await asyncio.to_thread( + self.worker_client.list_tasks, + params.context_id, + params, + ) + except ReconnectRequired as error: + raise self._session_expired_error(params.context_id) from error + + async def on_cancel_task( + self, + params: CancelTaskRequest, + _context: ServerCallContext, + ) -> Task | None: + try: + task = await asyncio.to_thread( + self.worker_client.cancel_task, + params.id, + ) + except ReconnectRequired as error: + if error.session_id is None: + raise + return self._build_connection_form_task( + task_id=params.id, + context_id=error.session_id, + ) if task is None: - return - context_id = task.context_id or task.id - session_id = self.sessions.pop(context_id, None) - if session_id: - await asyncio.to_thread( - self.worker_client.close_session, - session_id, - ) - updater = TaskUpdater( - event_queue=event_queue, - task_id=task.id, - context_id=task.context_id, + return self._build_connection_form_task( + task_id=params.id, + context_id=str(uuid4()), + ) + return task + + async def on_get_extended_agent_card( + self, + _params: GetExtendedAgentCardRequest, + _context: ServerCallContext, + ) -> AgentCardMessage: + return self.agent_card + + @staticmethod + def _session_expired_error(context_id: str) -> InvalidParamsError: + """Build the reconnect error returned by the task-list operation.""" + return InvalidParamsError( + "Database session expired. Reconnect using contextId.", + data={ + "reason": "SESSION_EXPIRED", + "contextId": context_id, + "connectionForm": connection_form(), + }, ) - await updater.cancel() @staticmethod - def _a2ui_action(context) -> dict | None: - for part in context.message.parts: - if part.WhichOneof("content") != "data": - continue - data = MessageToDict(part.data) - messages = data if isinstance(data, list) else [data] - for message in messages: - if ( - not isinstance(message, dict) - or message.get("version") != "v0.9" - ): - continue - action = message.get("action") - if isinstance(action, dict): - return action - return None + def _add_artifact( + task: Task, + parts: list, + name: str, + extensions: list[str] | None = None, + ) -> None: + artifact = new_artifact(parts, name) + if extensions: + artifact.extensions.extend(extensions) + task.artifacts.add().CopyFrom(artifact) + + def _complete_task( + self, + task: Task, + parts: list, + name: str, + extensions: list[str] | None = None, + ) -> Task: + self._add_artifact(task, parts, name, extensions) + task.status.state = TaskState.TASK_STATE_COMPLETED + task.status.timestamp.GetCurrentTime() + return task + + def _build_connection_form_task( + self, + *, + task_id: str | None, + context_id: str, + history: list[Message] | None = None, + ) -> Task: + """Build a completed, transient task containing the connection form.""" + task = _new_bootstrap_task( + task_id=task_id, + context_id=context_id, + history=history, + ) + return self._complete_task( + task, + [ + a2ui_part(item) + for item in connection_form(f"db-connect-{task.id}") + ], + "database-connection-form", + [A2UI_EXTENSION_URI], + ) -def _a2ui_part(message: dict): - """Encode one A2UI operation in the format used by Gemini Enterprise.""" - part = new_data_part(message) - ParseDict({"mimeType": A2UI_MIME_TYPE}, part.metadata) - return part +def _new_bootstrap_task( + *, + task_id: str | None, + context_id: str, + history: list[Message] | None = None, +) -> Task: + """Build a transient response task for the connection bootstrap.""" + return Task( + id=task_id or f"{_BOOTSTRAP_TASK_ID_PREFIX}{uuid4()}", + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_SUBMITTED), + history=history or [], + ) def create_gateway_app(settings: GatewaySettings) -> Starlette: @@ -206,11 +444,7 @@ def create_gateway_app(settings: GatewaySettings) -> Starlette: ) ], ) - handler = DefaultRequestHandler( - agent_executor=GatewayExecutor(WorkerClient(settings)), - task_store=InMemoryTaskStore(), - agent_card=card, - ) + handler = GatewayRequestHandler(WorkerClient(settings), card) compat_card = to_compat_agent_card(card).model_dump( by_alias=True, exclude_none=True, diff --git a/src/select_ai/agent/a2a/session_runtime.py b/src/select_ai/agent/a2a/session_runtime.py new file mode 100644 index 0000000..f516557 --- /dev/null +++ b/src/select_ai/agent/a2a/session_runtime.py @@ -0,0 +1,120 @@ +# ----------------------------------------------------------------------------- +# Copyright (c) 2026, Oracle and/or its affiliates. +# +# Licensed under the Universal Permissive License v 1.0 as shown at +# https://oss.oracle.com/licenses/upl. +# ----------------------------------------------------------------------------- + +"""A database-bound A2A runtime for one temporary worker session.""" + +from __future__ import annotations + +from a2a.auth.user import User +from a2a.server.context import ServerCallContext +from a2a.server.request_handlers import DefaultRequestHandler +from a2a.types.a2a_pb2 import ( + CancelTaskRequest, + GetTaskRequest, + ListTasksRequest, + SendMessageRequest, +) +from a2a.utils.errors import TaskNotFoundError + +from select_ai.agent import AsyncTeam +from select_ai.agent.a2a.context_store import OracleContextStore +from select_ai.agent.a2a.server import DatabaseTeamExecutor, _build_agent_card +from select_ai.agent.a2a.task_store import OracleTaskStore +from select_ai.agent.a2a.worker_protocol import ( + A2AMethod, + ResultKind, + WorkerResult, + encode_result, + parse_request, +) + +_METHOD_HANDLERS = { + A2AMethod.SEND_MESSAGE: (SendMessageRequest, "on_message_send"), + A2AMethod.GET_TASK: (GetTaskRequest, "on_get_task"), + A2AMethod.LIST_TASKS: (ListTasksRequest, "on_list_tasks"), + A2AMethod.CANCEL_TASK: (CancelTaskRequest, "on_cancel_task"), +} + + +class SessionUser(User): + """Internal A2A user used to scope one worker session's database rows.""" + + def __init__(self, session_id: str) -> None: + self.session_id = session_id + + @property + def is_authenticated(self) -> bool: + return True + + @property + def user_name(self) -> str: + return self.session_id + + +class SessionRuntime: + """Own the A2A handler and Oracle stores for one connected database.""" + + def __init__(self, session_id: str, team_name: str) -> None: + self.session_id = session_id + self.team_name = team_name + self.task_store = OracleTaskStore() + self.context_store = OracleContextStore() + self.handler: DefaultRequestHandler | None = None + + async def initialize(self) -> None: + """Initialize Oracle-backed stores after the worker connects.""" + # Validate the user-supplied team before the worker reports this + # database session as ready. Without this check, invalid team names + # appear to connect successfully and fail only on the first prompt. + await AsyncTeam.fetch(self.team_name) + await self.task_store.initialize() + await self.context_store.initialize() + self.handler = DefaultRequestHandler( + agent_executor=DatabaseTeamExecutor( + self.team_name, + self.context_store, + ), + task_store=self.task_store, + agent_card=_build_agent_card( + self.team_name, + f"http://worker/sessions/{self.session_id}", + None, + ), + ) + + async def handle( + self, + method: A2AMethod | str, + payload: bytes, + ) -> WorkerResult: + """Handle one internal A2A operation using protobuf wire bytes.""" + if self.handler is None: + raise RuntimeError("A2A session runtime is not initialized.") + + operation = A2AMethod(method) + context = ServerCallContext(user=SessionUser(self.session_id)) + if operation == A2AMethod.DELETE_TASK: + request = parse_request(GetTaskRequest, payload) + await self.task_store.delete(request.id, context) + return encode_result(None) + + try: + request_type, handler_name = _METHOD_HANDLERS[operation] + except KeyError as error: + raise ValueError( + f"Unsupported worker A2A method: {operation.value}" + ) from error + + request = parse_request(request_type, payload) + try: + result = await getattr(self.handler, handler_name)( + request, + context, + ) + except TaskNotFoundError: + return WorkerResult(ResultKind.TASK_NOT_FOUND) + return encode_result(result) diff --git a/src/select_ai/agent/a2a/task_store.py b/src/select_ai/agent/a2a/task_store.py index 3962e8e..418c145 100644 --- a/src/select_ai/agent/a2a/task_store.py +++ b/src/select_ai/agent/a2a/task_store.py @@ -126,12 +126,24 @@ async def list( SELECT task_json FROM SELECT_AI_A2A_TASKS WHERE owner = :owner + AND ( + :context_id IS NULL OR context_id = :context_id + ) + AND ( + :status IS NULL OR + JSON_VALUE(task_json, '$.status.state') = :status + ) + AND ( + :status_timestamp_after IS NULL OR + JSON_VALUE(task_json, '$.status.timestamp') + >= :status_timestamp_after + ) ORDER BY updated_at DESC, task_id DESC """, owner=self._owner(context), + **self._list_query_parameters(params), ) tasks = [await self._task_from_json(row[0]) for row in rows] - tasks = self._filter_tasks(tasks, params) total_size = len(tasks) start_index = self._page_start_index(tasks, params.page_token) page_size = params.page_size or DEFAULT_LIST_TASKS_PAGE_SIZE @@ -162,31 +174,28 @@ async def delete(self, task_id: str, context: ServerCallContext) -> None: ) @staticmethod - def _filter_tasks( - tasks: list[Task], + def _list_query_parameters( params: a2a_pb2.ListTasksRequest, - ) -> list[Task]: - if params.context_id: - tasks = [ - task for task in tasks if task.context_id == params.context_id - ] + ) -> dict[str, str | None]: + """Convert protobuf list filters to Oracle query parameters.""" + status = None if params.status: - tasks = [ - task - for task in tasks - if task.HasField("status") - and task.status.state == params.status - ] - if params.HasField("status_timestamp_after"): - timestamp_after = params.status_timestamp_after.ToJsonString() - tasks = [ - task - for task in tasks - if task.HasField("status") - and task.status.HasField("timestamp") - and task.status.timestamp.ToJsonString() >= timestamp_after - ] - return tasks + try: + status = a2a_pb2.TaskState.Name(params.status) + except ValueError: + # An unknown enum value cannot match a stored protobuf JSON + # enum name, which preserves the old empty-result behavior. + status = "__UNKNOWN_TASK_STATE__" + timestamp_after = ( + params.status_timestamp_after.ToJsonString() + if params.HasField("status_timestamp_after") + else None + ) + return { + "context_id": params.context_id or None, + "status": status, + "status_timestamp_after": timestamp_after, + } def _owner(self, context: ServerCallContext) -> str: return self.owner_resolver(context) or "anonymous" diff --git a/src/select_ai/agent/a2a/worker.py b/src/select_ai/agent/a2a/worker.py index d6e7b04..5043955 100644 --- a/src/select_ai/agent/a2a/worker.py +++ b/src/select_ai/agent/a2a/worker.py @@ -21,11 +21,22 @@ from multiprocessing.connection import Connection import httpx -from fastapi import FastAPI, HTTPException, status -from fastapi.responses import PlainTextResponse +from fastapi import FastAPI, HTTPException, Request, status from pydantic import BaseModel, Field, SecretStr +from starlette.responses import Response + +from select_ai.agent.a2a.session_runtime import SessionRuntime +from select_ai.agent.a2a.worker_protocol import ( + PROTOBUF_CONTENT_TYPE, + WORKER_A2A_METHOD_HEADER, + WORKER_RESULT_KIND_HEADER, + PipeMessageType, + ResultKind, + WorkerResult, +) LOGGER = logging.getLogger(__name__) +_SESSION_REAPER_INTERVAL_SECONDS = 1 class OpenSessionRequest(BaseModel): @@ -38,12 +49,6 @@ class OpenSessionRequest(BaseModel): team_name: str = Field(min_length=1, max_length=128) -class PromptRequest(BaseModel): - """One user prompt for a previously opened session.""" - - prompt: str = Field(min_length=1, max_length=32_000) - - @dataclass class ChildSession: """In-memory ownership record for one database-session process.""" @@ -65,7 +70,6 @@ def __init__( self.session_ttl_seconds = session_ttl_seconds self.session_start_timeout_seconds = session_start_timeout_seconds self.sessions: dict[str, ChildSession] = {} - self.lock = asyncio.Lock() async def open(self, request: OpenSessionRequest) -> None: """Start a child runtime and wait until its database pool is ready.""" @@ -98,23 +102,24 @@ async def open(self, request: OpenSessionRequest) -> None: except Exception: await self._terminate(session) raise - async with self.lock: - previous = self.sessions.pop(request.session_id, None) - if previous: - await self._terminate(previous) - self.sessions[request.session_id] = session + previous = self.sessions.pop(request.session_id, None) + self.sessions[request.session_id] = session + if previous: + await self._terminate(previous) async def get(self, session_id: str) -> ChildSession: """Return a live session, closing it if its expiry has elapsed.""" - async with self.lock: - session = self.sessions.get(session_id) - if session and ( - session.expires_at <= time.monotonic() - or not session.process.is_alive() - ): - self.sessions.pop(session_id, None) - await self._terminate(session) - session = None + expired_session = None + session = self.sessions.get(session_id) + if session and ( + session.expires_at <= time.monotonic() + or not session.process.is_alive() + ): + self.sessions.pop(session_id, None) + expired_session = session + session = None + if expired_session: + await self._terminate(expired_session) if session is None: raise HTTPException( status_code=404, @@ -122,32 +127,57 @@ async def get(self, session_id: str) -> ChildSession: ) return session - async def send_prompt(self, session_id: str, prompt: str) -> str | None: - """Run one prompt in the process that owns this database session.""" + async def dispatch( + self, + session_id: str, + method: str, + payload: bytes, + ) -> WorkerResult: + """Run one A2A operation in the process owning this session.""" session = await self.get(session_id) + response = None + failure = None async with session.lock: - if not session.process.is_alive(): - await self._discard(session_id, session) - raise HTTPException( - status_code=404, - detail="Database session expired; reconnect required.", - ) - try: - await asyncio.to_thread( - session.connection.send, - {"type": "run", "prompt": prompt}, - ) - response = await self._receive(session, timeout_seconds=120) - except (EOFError, OSError, TimeoutError) as error: - await self._discard(session_id, session) - raise HTTPException( - status_code=502, - detail=( - "Database session is unavailable; reconnect required." - ), - ) from error - if response.get("type") == "result": - return response.get("result") + if session.process.is_alive(): + try: + await asyncio.to_thread( + session.connection.send, + { + "type": PipeMessageType.A2A.value, + "method": method, + "payload": payload, + }, + ) + response = await self._receive( + session, + timeout_seconds=120, + ) + except (EOFError, OSError, TimeoutError) as error: + failure = error + if failure is not None: + await self._discard(session_id, session) + raise HTTPException( + status_code=502, + detail=( + "Database session is unavailable; reconnect required." + ), + ) from failure + if response is None: + await self._discard(session_id, session) + raise HTTPException( + status_code=404, + detail="Database session expired; reconnect required.", + ) + if response.get("type") == PipeMessageType.RESULT.value: + return response.get( + "result", + WorkerResult(ResultKind.NONE), + ) + if response.get("type") == PipeMessageType.A2A_ERROR.value: + raise HTTPException( + status_code=400, + detail=response.get("detail", "A2A request failed."), + ) LOGGER.error("Select AI session process reported a command failure.") raise HTTPException( status_code=502, @@ -156,8 +186,7 @@ async def send_prompt(self, session_id: str, prompt: str) -> str | None: async def close(self, session_id: str) -> None: """Terminate a session on an explicit gateway DELETE request.""" - async with self.lock: - session = self.sessions.pop(session_id, None) + session = self.sessions.pop(session_id, None) if session is None: raise HTTPException( status_code=404, @@ -165,14 +194,38 @@ async def close(self, session_id: str) -> None: ) await self._terminate(session) + async def reap_expired(self) -> None: + """Continuously terminate expired or dead child sessions.""" + while True: + await asyncio.sleep(_SESSION_REAPER_INTERVAL_SECONDS) + await self._reap_expired_sessions() + async def close_all(self) -> None: """Terminate all child sessions during worker shutdown.""" - async with self.lock: - sessions = list(self.sessions.values()) - self.sessions.clear() + sessions = list(self.sessions.values()) + self.sessions.clear() for session in sessions: await self._terminate(session) + async def _reap_expired_sessions(self) -> None: + now = time.monotonic() + expired = [ + (session_id, session) + for session_id, session in self.sessions.items() + if session.expires_at <= now or not session.process.is_alive() + ] + for session_id, _session in expired: + self.sessions.pop(session_id, None) + + for session_id, session in expired: + try: + await self._terminate(session) + except Exception: + LOGGER.exception( + "Failed to terminate expired session %s.", + session_id, + ) + async def _wait_ready( self, session: ChildSession, @@ -188,7 +241,7 @@ async def _wait_ready( status_code=504, detail="Database session start timed out.", ) from error - if response.get("type") == "ready": + if response.get("type") == PipeMessageType.READY.value: return detail = response.get("detail", "Database login failed.") for secret in ( @@ -214,26 +267,30 @@ async def _receive( return await asyncio.to_thread(session.connection.recv) async def _discard(self, session_id: str, session: ChildSession) -> None: - async with self.lock: - if self.sessions.get(session_id) is session: - self.sessions.pop(session_id, None) + if self.sessions.get(session_id) is session: + self.sessions.pop(session_id, None) await self._terminate(session) @staticmethod async def _terminate(session: ChildSession) -> None: - def stop() -> None: + async with session.lock: try: - if session.process.is_alive(): - with contextlib.suppress(OSError): - session.connection.send({"type": "close"}) - session.process.join(timeout=5) - if session.process.is_alive(): - session.process.terminate() - session.process.join(timeout=5) + with contextlib.suppress(OSError): + await asyncio.to_thread( + session.connection.send, + {"type": PipeMessageType.CLOSE.value}, + ) + await asyncio.to_thread(session.process.join, timeout=5) + for stop in ( + session.process.terminate, + session.process.kill, + ): + if not session.process.is_alive(): + break + await asyncio.to_thread(stop) + await asyncio.to_thread(session.process.join, timeout=5) finally: - session.connection.close() - - await asyncio.to_thread(stop) + await asyncio.to_thread(session.connection.close) def _session_process_main( @@ -262,10 +319,10 @@ async def _run_session_process( session_id: str, team_name: str, ) -> None: - """Open one async connection and execute raw ``AsyncTeam`` prompts.""" + """Open one async connection and execute A2A operations.""" import select_ai - from select_ai.agent import AsyncTeam + runtime: SessionRuntime | None = None ready = False try: await select_ai.async_connect( @@ -275,47 +332,98 @@ async def _run_session_process( ) if not await select_ai.async_is_connected(): raise RuntimeError("Database login failed.") - connection.send({"type": "ready"}) + runtime = SessionRuntime(session_id, team_name) + await runtime.initialize() + connection.send({"type": PipeMessageType.READY.value}) ready = True - conversation_id = None - while True: - try: - command = await asyncio.to_thread(connection.recv) - except EOFError: - return - if command.get("type") == "close": - return - if command.get("type") != "run": - connection.send( - {"type": "error", "detail": "Invalid command."} - ) - continue - try: - if conversation_id is None: - conversation = select_ai.AsyncConversation( - attributes=select_ai.ConversationAttributes( - title=f"A2A {team_name}", - description=f"Temporary session {session_id}", - ) - ) - conversation_id = await conversation.create() - result = await AsyncTeam(team_name=team_name).run( - prompt=command["prompt"], - params={"conversation_id": conversation_id}, - ) - connection.send({"type": "result", "result": result}) - except Exception: - # Keep session-process failures private from gateway callers. - LOGGER.error("Select AI session command failed") - connection.send({"type": "error"}) + await _serve_session_commands(connection, runtime) except Exception as error: - LOGGER.error("Select AI session process startup failed") - if not ready: - with contextlib.suppress(OSError): - connection.send({"type": "error", "detail": str(error)}) + _report_session_process_error(connection, error, not ready) finally: + await _close_session_process(runtime, select_ai) + + +async def _serve_session_commands( + connection: Connection, + runtime: SessionRuntime, +) -> None: + """Serve commands for one initialized database session.""" + while True: + command = await _receive_session_command(connection) + if ( + command is None + or command.get("type") == PipeMessageType.CLOSE.value + ): + return + if command.get("type") != PipeMessageType.A2A.value: + connection.send( + { + "type": PipeMessageType.ERROR.value, + "detail": "Invalid command.", + } + ) + continue + await _handle_a2a_command(connection, runtime, command) + + +async def _receive_session_command(connection: Connection) -> dict | None: + """Read one command without blocking the event loop.""" + try: + return await asyncio.to_thread(connection.recv) + except EOFError: + return None + + +async def _handle_a2a_command( + connection: Connection, + runtime: SessionRuntime, + command: dict, +) -> None: + """Execute one internal A2A command and send its result.""" + try: + result = await runtime.handle( + command["method"], + command.get("payload", b""), + ) + except Exception as error: + LOGGER.exception("Select AI session A2A command failed") + connection.send( + {"type": PipeMessageType.A2A_ERROR.value, "detail": str(error)} + ) + return + connection.send({"type": PipeMessageType.RESULT.value, "result": result}) + + +def _report_session_process_error( + connection: Connection, + error: Exception, + during_startup: bool, +) -> None: + """Report startup failures without sending errors after readiness.""" + message = ( + "Select AI session process startup failed" + if during_startup + else "Select AI session process failed" + ) + LOGGER.error(message) + if during_startup: + with contextlib.suppress(OSError): + connection.send( + {"type": PipeMessageType.ERROR.value, "detail": str(error)} + ) + + +async def _close_session_process( + runtime: SessionRuntime | None, + select_ai, +) -> None: + """Close the A2A handler and database connection owned by the child.""" + handler = getattr(runtime, "handler", None) + if handler is not None: with contextlib.suppress(Exception): - await select_ai.async_disconnect() + await handler.aclose() + with contextlib.suppress(Exception): + await select_ai.async_disconnect() async def _register_with_consul( @@ -393,9 +501,13 @@ async def lifespan(_app): worker_endpoint, ) heartbeat = asyncio.create_task(_heartbeat(consul_url, worker_id)) + reaper = asyncio.create_task(worker.reap_expired()) try: yield finally: + reaper.cancel() + with contextlib.suppress(asyncio.CancelledError): + await reaper heartbeat.cancel() with contextlib.suppress(asyncio.CancelledError): await heartbeat @@ -418,13 +530,29 @@ async def open_session(request: OpenSessionRequest) -> dict[str, str]: await worker.open(request) return {"status": "opened"} - @app.post("/sessions/{session_id}/messages") - async def send_message( + @app.post("/sessions/{session_id}/a2a") + async def handle_a2a( session_id: str, - request: PromptRequest, - ) -> PlainTextResponse: - result = await worker.send_prompt(session_id, request.prompt) - return PlainTextResponse(result or "") + request: Request, + ) -> Response: + method = request.headers.get(WORKER_A2A_METHOD_HEADER) + if not method: + raise HTTPException( + status_code=400, + detail=f"Missing {WORKER_A2A_METHOD_HEADER} header.", + ) + result = await worker.dispatch( + session_id, + method, + await request.body(), + ) + return Response( + content=result.payload, + media_type=PROTOBUF_CONTENT_TYPE, + headers={ + WORKER_RESULT_KIND_HEADER: result.kind.value, + }, + ) @app.delete( "/sessions/{session_id}", diff --git a/src/select_ai/agent/a2a/worker_client.py b/src/select_ai/agent/a2a/worker_client.py index b149453..3c38b82 100644 --- a/src/select_ai/agent/a2a/worker_client.py +++ b/src/select_ai/agent/a2a/worker_client.py @@ -11,17 +11,46 @@ import base64 import json +import logging import time from threading import Lock +from urllib.parse import quote import requests +from a2a.types.a2a_pb2 import ( + CancelTaskRequest, + GetTaskRequest, + ListTasksRequest, + ListTasksResponse, + Message, + SendMessageRequest, + Task, +) from select_ai.agent.a2a import GatewaySettings, SessionInfo, SessionRoute +from select_ai.agent.a2a.worker_protocol import ( + PROTOBUF_CONTENT_TYPE, + WORKER_A2A_METHOD_HEADER, + WORKER_RESULT_KIND_HEADER, + A2AMethod, + ResultKind, + WorkerResult, + decode_result, +) + +LOGGER = logging.getLogger(__name__) + +_SESSION_PREFIX = "select-ai/sessions/" +_TASK_PREFIX = "select-ai/tasks/" class ReconnectRequired(RuntimeError): """The worker no longer owns the requested in-memory session.""" + def __init__(self, message: str, session_id: str | None = None): + super().__init__(message) + self.session_id = session_id + class WorkerClient: """Open, route, and close short-lived Select AI worker sessions.""" @@ -59,22 +88,152 @@ def open_session(self, session_id: str, session_info: SessionInfo) -> str: raise RuntimeError("Could not create the database session.") return session_id - def send_prompt(self, session_id: str, prompt: str) -> str | None: - """Forward a prompt and return the worker's raw team result.""" + def send_message( + self, + session_id: str, + request: SendMessageRequest, + ) -> Task | Message | None: + """Forward one A2A message to the worker session.""" + result = self._dispatch( + session_id, + A2AMethod.SEND_MESSAGE, + request, + ) + value = decode_result(result) + if isinstance(value, Task): + self._save_task_route(value.id, session_id) + return value + + def get_task(self, params: GetTaskRequest) -> Task | None: + """Retrieve a task from its owning worker session.""" + session_id = self.task_session_id(params.id) + if session_id is None: + return None + result = self._dispatch(session_id, A2AMethod.GET_TASK, params) + value = decode_result(result) + if isinstance(value, Task): + return value + return None + + def list_tasks( + self, + session_id: str, + params: ListTasksRequest, + ) -> ListTasksResponse: + """Retrieve one session's database-paginated task response.""" + result = self._dispatch( + session_id, + A2AMethod.LIST_TASKS, + params, + ) + return decode_result(result) or ListTasksResponse() + + def cancel_task(self, task_id: str) -> Task | None: + """Cancel a task in the worker that owns it.""" + session_id = self.task_session_id(task_id) + if session_id is None: + return None + result = self._dispatch( + session_id, + A2AMethod.CANCEL_TASK, + CancelTaskRequest(id=task_id), + ) + value = decode_result(result) + return value if isinstance(value, Task) else None + + def delete_task(self, task_id: str) -> None: + """Delete a task and its routing metadata.""" + session_id = self.task_session_id(task_id) + if session_id: + try: + self._dispatch( + session_id, + A2AMethod.DELETE_TASK, + GetTaskRequest(id=task_id), + ) + except ReconnectRequired: + pass + self._delete_task_route(task_id) + + def session_exists(self, session_id: str) -> bool: + """Return whether Consul still has a live route for a session.""" + try: + self._route_for(session_id) + except ReconnectRequired: + return False + return True + + def task_session_id(self, task_id: str) -> str | None: + """Return the worker session recorded for a task, if any.""" + response = requests.get( + f"{self.settings.consul_url}/v1/kv/{_TASK_PREFIX}" + f"{quote(task_id, safe='')}", + timeout=10, + ) + if response.status_code == 404: + return None + response.raise_for_status() + values = response.json() or [] + if not values or not values[0].get("Value"): + return None + return base64.b64decode(values[0]["Value"]).decode() + + def _dispatch( + self, + session_id: str, + method: A2AMethod | str, + message, + ) -> WorkerResult: + """Send one protobuf-serialized A2A operation to a worker session.""" route = self._route_for(session_id) response = requests.post( - f"{route.endpoint}/sessions/{session_id}/messages", - json={"prompt": prompt}, + f"{route.endpoint}/sessions/{quote(session_id, safe='')}/a2a", + data=message.SerializeToString(), + headers={ + "content-type": PROTOBUF_CONTENT_TYPE, + WORKER_A2A_METHOD_HEADER: A2AMethod(method).value, + }, timeout=130, **getattr(self, "_worker_request_kwargs", {}), ) if response.status_code in (404, 502): self._close_worker_session(route, session_id) raise ReconnectRequired( - "Database session ended; reconnect required." + "Database session ended; reconnect required.", + session_id, ) response.raise_for_status() - return response.text or None + return WorkerResult( + kind=ResultKind( + response.headers.get( + WORKER_RESULT_KIND_HEADER, + ResultKind.NONE.value, + ) + ), + payload=response.content, + ) + + def _save_task_route(self, task_id: str, session_id: str) -> None: + """Save only task-to-session routing metadata in Consul.""" + try: + response = requests.put( + f"{self.settings.consul_url}/v1/kv/{_TASK_PREFIX}" + f"{quote(task_id, safe='')}", + data=session_id, + timeout=10, + ) + response.raise_for_status() + except requests.RequestException: + # The task itself is already durable in Oracle, but task routes + # are required because workers may use different databases. + LOGGER.warning("Could not save route for task %s", task_id) + + def _delete_task_route(self, task_id: str) -> None: + requests.delete( + f"{self.settings.consul_url}/v1/kv/{_TASK_PREFIX}" + f"{quote(task_id, safe='')}", + timeout=10, + ) def close_session(self, session_id: str) -> None: """Close the child process and remove the Consul route.""" @@ -100,7 +259,8 @@ def _select_worker(self) -> str: worker = workers[self._next_worker % len(workers)] self._next_worker += 1 service = worker["Service"] - endpoint = service.get("Meta", {}).get("endpoint") + metadata = service.get("Meta") or {} + endpoint = metadata.get("endpoint") if endpoint: endpoint = endpoint.rstrip("/") if self.settings.worker_mtls_enabled and not endpoint.startswith( @@ -121,13 +281,14 @@ def _select_worker(self) -> str: def _route_for(self, session_id: str) -> SessionRoute: response = requests.get( - f"{self.settings.consul_url}/v1/kv/select-ai/sessions/" - f"{session_id}", + f"{self.settings.consul_url}/v1/kv/{_SESSION_PREFIX}" + f"{quote(session_id, safe='')}", timeout=10, ) if response.status_code == 404: raise ReconnectRequired( - "Database session expired; reconnect required." + "Database session expired; reconnect required.", + session_id, ) response.raise_for_status() value = response.json()[0]["Value"] @@ -135,14 +296,15 @@ def _route_for(self, session_id: str) -> SessionRoute: if route.expires_at <= time.time(): self._close_worker_session(route, session_id) raise ReconnectRequired( - "Database session expired; reconnect required." + "Database session expired; reconnect required.", + session_id, ) return route def _save_route(self, session_id: str, route: SessionRoute) -> bool: response = requests.put( - f"{self.settings.consul_url}/v1/kv/select-ai/sessions/" - f"{session_id}?cas=0", + f"{self.settings.consul_url}/v1/kv/{_SESSION_PREFIX}" + f"{quote(session_id, safe='')}?cas=0", data=json.dumps(route.__dict__), timeout=10, ) @@ -150,8 +312,8 @@ def _save_route(self, session_id: str, route: SessionRoute) -> bool: def _delete_route(self, session_id: str) -> None: requests.delete( - f"{self.settings.consul_url}/v1/kv/select-ai/sessions/" - f"{session_id}", + f"{self.settings.consul_url}/v1/kv/{_SESSION_PREFIX}" + f"{quote(session_id, safe='')}", timeout=10, ) @@ -162,7 +324,7 @@ def _close_worker_session( ) -> None: try: response = requests.delete( - f"{route.endpoint}/sessions/{session_id}", + f"{route.endpoint}/sessions/{quote(session_id, safe='')}", timeout=10, **getattr(self, "_worker_request_kwargs", {}), ) diff --git a/src/select_ai/agent/a2a/worker_protocol.py b/src/select_ai/agent/a2a/worker_protocol.py new file mode 100644 index 0000000..5aa6cd7 --- /dev/null +++ b/src/select_ai/agent/a2a/worker_protocol.py @@ -0,0 +1,103 @@ +# ----------------------------------------------------------------------------- +# Copyright (c) 2026, Oracle and/or its affiliates. +# +# Licensed under the Universal Permissive License v 1.0 as shown at +# https://oss.oracle.com/licenses/upl. +# ----------------------------------------------------------------------------- + +"""Typed values for the private Gateway-to-Worker session protocol.""" + +from __future__ import annotations + +from dataclasses import dataclass +from enum import Enum + +from a2a.types.a2a_pb2 import ListTasksResponse, Message, Task +from a2a.utils.errors import TaskNotFoundError + +WORKER_A2A_METHOD_HEADER = "x-select-ai-a2a-method" +WORKER_RESULT_KIND_HEADER = "x-select-ai-a2a-result-kind" +PROTOBUF_CONTENT_TYPE = "application/x-protobuf" + + +class PipeMessageType(str, Enum): + """Message types exchanged over a session process pipe.""" + + A2A = "a2a" + CLOSE = "close" + READY = "ready" + RESULT = "result" + ERROR = "error" + A2A_ERROR = "a2a_error" + + +class A2AMethod(str, Enum): + """A2A operations supported by the private session runtime.""" + + SEND_MESSAGE = "SendMessage" + GET_TASK = "GetTask" + LIST_TASKS = "ListTasks" + CANCEL_TASK = "CancelTask" + DELETE_TASK = "DeleteTask" + + +class ResultKind(str, Enum): + """Response types returned by a session runtime.""" + + NONE = "none" + TASK_NOT_FOUND = "task_not_found" + TASK = "task" + MESSAGE = "message" + LIST_TASKS = "list_tasks" + + +@dataclass(frozen=True) +class WorkerResult: + """One typed response from the worker session runtime.""" + + kind: ResultKind + payload: bytes = b"" + + +def parse_request(message_type, payload: bytes): + """Parse one protobuf request from the internal wire format.""" + message = message_type() + message.ParseFromString(payload) + return message + + +def encode_result(value) -> WorkerResult: + """Encode one supported A2A response as protobuf wire bytes.""" + if value is None: + return WorkerResult(ResultKind.NONE) + if isinstance(value, Task): + kind = ResultKind.TASK + elif isinstance(value, Message): + kind = ResultKind.MESSAGE + elif isinstance(value, ListTasksResponse): + kind = ResultKind.LIST_TASKS + else: + raise TypeError(f"Unsupported A2A response type: {type(value)!r}") + return WorkerResult(kind, value.SerializeToString()) + + +def decode_result(result: WorkerResult, expected_type=None): + """Decode a session runtime response from protobuf wire bytes.""" + kind = result.kind + if kind == ResultKind.TASK_NOT_FOUND: + raise TaskNotFoundError + if kind == ResultKind.NONE: + return None + payload = result.payload + if not payload: + return None + message_type = expected_type or { + ResultKind.TASK: Task, + ResultKind.MESSAGE: Message, + ResultKind.LIST_TASKS: ListTasksResponse, + }.get(kind) + if message_type is None: + raise ValueError(f"Unsupported A2A result kind: {kind}") + message = message_type() + message.ParseFromString(payload) + return message diff --git a/src/select_ai/cli/a2a.py b/src/select_ai/cli/a2a.py index e43342e..cab69e8 100644 --- a/src/select_ai/cli/a2a.py +++ b/src/select_ai/cli/a2a.py @@ -132,7 +132,7 @@ def worker( tls_key_file, tls_ca_file, ): - """Start the internal, in-memory Select AI session worker.""" + """Start the internal Select AI session worker.""" try: from select_ai.agent.a2a import create_worker_app except ImportError as error: diff --git a/tests/a2a/test_worker_runtime.py b/tests/a2a/test_worker_runtime.py index 8ab1cff..89671f6 100644 --- a/tests/a2a/test_worker_runtime.py +++ b/tests/a2a/test_worker_runtime.py @@ -11,14 +11,46 @@ import pytest import requests +from google.protobuf.json_format import MessageToJson pytest.importorskip("fastapi") import select_ai -from select_ai.agent.a2a import GatewaySettings, worker -from select_ai.agent.a2a.gateway import GatewayExecutor, _a2ui_part +from a2a.helpers import new_data_part, new_text_part +from a2a.server.context import ServerCallContext +from a2a.types.a2a_pb2 import ( + CancelTaskRequest, + GetTaskRequest, + ListTasksRequest, + Message, + Role, + SendMessageRequest, + Task, + TaskState, +) +from a2a.utils.errors import ( + InvalidParamsError, + TaskNotFoundError, + UnsupportedOperationError, +) +from select_ai.agent.a2a import GatewaySettings, session_runtime, worker +from select_ai.agent.a2a.a2ui import ( + A2UI_CATALOG_ID, + A2UI_EXTENSION_URI, + A2UI_VERSION, + a2ui_part, + find_action, +) +from select_ai.agent.a2a.forms import connection_form +from select_ai.agent.a2a.gateway import GatewayRequestHandler from select_ai.agent.a2a.results import message_parts -from select_ai.agent.a2a.worker_client import WorkerClient +from select_ai.agent.a2a.worker_client import ReconnectRequired, WorkerClient +from select_ai.agent.a2a.worker_protocol import ( + A2AMethod, + ResultKind, + WorkerResult, + decode_result, +) class FakeConnection: @@ -64,10 +96,19 @@ def join(self, timeout): def terminate(self): self.alive = False + def kill(self): + self.alive = False + -def test_worker_uses_pipe_process_and_returns_raw_result(monkeypatch): +def test_worker_uses_pipe_process_and_dispatches_a2a(monkeypatch): parent = FakeConnection( - [{"type": "ready"}, {"type": "result", "result": "hi"}] + [ + {"type": "ready"}, + { + "type": "result", + "result": WorkerResult(ResultKind.NONE), + }, + ] ) child = FakeConnection() processes = [] @@ -93,13 +134,26 @@ def process_factory(**kwargs): ) asyncio.run(session_worker.open(request)) - result = asyncio.run(session_worker.send_prompt("session-1", "hello")) + task_request = GetTaskRequest(id="task-1") + result = asyncio.run( + session_worker.dispatch( + "session-1", + "GetTask", + task_request.SerializeToString(), + ) + ) - assert result == "hi" + assert result == WorkerResult(ResultKind.NONE) assert child.closed assert processes[0].target is worker._session_process_main assert processes[0].daemon is True - assert parent.sent == [{"type": "run", "prompt": "hello"}] + assert parent.sent == [ + { + "type": "a2a", + "method": "GetTask", + "payload": task_request.SerializeToString(), + } + ] def test_expired_session_closes_the_process_and_pipe(monkeypatch): @@ -138,32 +192,82 @@ def process_factory(**kwargs): assert parent.sent == [{"type": "close"}] -def test_session_runtime_uses_one_async_connection_and_returns_raw_result( +def test_session_reaper_closes_idle_expired_process(monkeypatch): + parent = FakeConnection([{"type": "ready"}]) + child = FakeConnection() + process = None + + def process_factory(**kwargs): + nonlocal process + process = FakeProcess(**kwargs) + return process + + monkeypatch.setattr( + worker.multiprocessing, + "Pipe", + lambda: (parent, child), + ) + monkeypatch.setattr(worker.multiprocessing, "Process", process_factory) + monkeypatch.setattr(worker, "_SESSION_REAPER_INTERVAL_SECONDS", 0.01) + session_worker = worker.SessionWorker(60, 1) + request = worker.OpenSessionRequest( + session_id="session-1", + username="user", + password="password", + dsn="database", + team_name="TEAM", + ) + + async def exercise_reaper(): + await session_worker.open(request) + session_worker.sessions["session-1"].expires_at = 0 + reaper = asyncio.create_task(session_worker.reap_expired()) + try: + await asyncio.wait_for(_wait_for_process_exit(process), 1) + finally: + reaper.cancel() + with pytest.raises(asyncio.CancelledError): + await reaper + + async def _wait_for_process_exit(session_process): + while session_process is None or session_process.alive: + await asyncio.sleep(0.01) + + asyncio.run(exercise_reaper()) + + assert process is not None + assert parent.closed + assert parent.sent == [{"type": "close"}] + assert session_worker.sessions == {} + + +def test_session_runtime_uses_one_async_connection_and_dispatches_a2a( monkeypatch, ): connection = FakeConnection( - [{"type": "run", "prompt": "hello"}, {"type": "close"}] + [ + { + "type": "a2a", + "method": "GetTask", + "payload": GetTaskRequest(id="t1").SerializeToString(), + }, + {"type": "close"}, + ] ) connection_arguments = {} - class Conversation: - def __init__(self, attributes): - self.attributes = attributes - - async def create(self): - return "conversation-1" - - class Team: - def __init__(self, team_name): + class Runtime: + def __init__(self, session_id, team_name): + assert session_id == "session-1" assert team_name == "TEAM" - async def run(self, prompt, params): - assert prompt == "hello" - assert params == {"conversation_id": "conversation-1"} - return ( - '{"metadata":{"mimeType":"application/json+a2ui"},' - '"data":[]}' - ) + async def initialize(self): + return None + + async def handle(self, method, payload): + assert method == "GetTask" + assert GetTaskRequest.FromString(payload).id == "t1" + return WorkerResult(ResultKind.NONE) async def async_connect(**kwargs): connection_arguments.update(kwargs) @@ -177,8 +281,7 @@ async def disconnect(): monkeypatch.setattr(select_ai, "async_connect", async_connect) monkeypatch.setattr(select_ai, "async_is_connected", connected) monkeypatch.setattr(select_ai, "async_disconnect", disconnect) - monkeypatch.setattr(select_ai, "AsyncConversation", Conversation) - monkeypatch.setattr(select_ai.agent, "AsyncTeam", Team) + monkeypatch.setattr(worker, "SessionRuntime", Runtime) asyncio.run( worker._run_session_process( @@ -201,12 +304,115 @@ async def disconnect(): assert connection.sent[0] == {"type": "ready"} assert connection.sent[1] == { "type": "result", - "result": ( - '{"metadata":{"mimeType":"application/json+a2ui"},' '"data":[]}' - ), + "result": WorkerResult(ResultKind.NONE), } +def test_session_runtime_builds_oracle_backed_default_handler(monkeypatch): + initialized = [] + + class Team: + @staticmethod + async def fetch(team_name): + assert team_name == "TEAM" + initialized.append("team") + + class Store: + async def initialize(self): + initialized.append(self) + + class Handler: + def __init__(self, **kwargs): + self.kwargs = kwargs + + monkeypatch.setattr(session_runtime, "OracleTaskStore", Store) + monkeypatch.setattr(session_runtime, "OracleContextStore", Store) + monkeypatch.setattr(session_runtime, "AsyncTeam", Team) + monkeypatch.setattr(session_runtime, "DefaultRequestHandler", Handler) + monkeypatch.setattr( + session_runtime, + "DatabaseTeamExecutor", + lambda team_name, context_store: (team_name, context_store), + ) + monkeypatch.setattr( + session_runtime, + "_build_agent_card", + lambda *args: "agent-card", + ) + + runtime = session_runtime.SessionRuntime("session-1", "TEAM") + asyncio.run(runtime.initialize()) + + assert initialized == [ + "team", + runtime.task_store, + runtime.context_store, + ] + assert runtime.handler.kwargs["task_store"] is runtime.task_store + assert runtime.handler.kwargs["agent_card"] == "agent-card" + + +def test_session_runtime_rejects_unknown_team_before_initializing_stores( + monkeypatch, +): + initialized = [] + + class Team: + @staticmethod + async def fetch(team_name): + assert team_name == "MISSPELLED_TEAM" + raise RuntimeError("team does not exist") + + class Store: + async def initialize(self): + initialized.append(self) + + monkeypatch.setattr(session_runtime, "AsyncTeam", Team) + monkeypatch.setattr(session_runtime, "OracleTaskStore", Store) + monkeypatch.setattr(session_runtime, "OracleContextStore", Store) + + runtime = session_runtime.SessionRuntime("session-1", "MISSPELLED_TEAM") + with pytest.raises(RuntimeError, match="team does not exist"): + asyncio.run(runtime.initialize()) + + assert initialized == [] + assert runtime.handler is None + + +def test_session_runtime_preserves_database_task_not_found(): + class Handler: + async def on_get_task(self, _request, _context): + raise TaskNotFoundError + + runtime = session_runtime.SessionRuntime("session-1", "TEAM") + runtime.handler = Handler() + + result = asyncio.run( + runtime.handle( + A2AMethod.GET_TASK, + GetTaskRequest(id="missing-task").SerializeToString(), + ) + ) + + assert result == WorkerResult(ResultKind.TASK_NOT_FOUND) + + +def test_decode_result_raises_only_for_explicit_database_task_not_found(): + with pytest.raises(TaskNotFoundError): + decode_result(WorkerResult(ResultKind.TASK_NOT_FOUND)) + + assert decode_result(WorkerResult(ResultKind.NONE)) is None + + +def test_gateway_rejects_unsupported_stream(): + async def consume_stream(): + stream = GatewayRequestHandler.on_message_send_stream(None, None) + with pytest.raises(UnsupportedOperationError): + await anext(stream) + + asyncio.run(consume_stream()) + + def test_message_parts_preserves_serialized_text_and_data_parts(): parts = message_parts( """{ @@ -227,19 +433,30 @@ def test_message_parts_preserves_serialized_text_and_data_parts(): assert parts[1].WhichOneof("content") == "data" -def test_worker_client_returns_the_raw_team_result(monkeypatch): +def test_worker_client_forwards_and_parses_a2a_message(monkeypatch): + sent = {} + class Response: status_code = 200 + headers = {"x-select-ai-a2a-result-kind": "message"} + content = Message( + message_id="m1", + role=Role.ROLE_AGENT, + parts=[new_text_part("hello")], + ).SerializeToString() @staticmethod def raise_for_status(): return None - text = "team result" + def post(*args, **kwargs): + del args + sent.update(kwargs) + return Response() monkeypatch.setattr( "select_ai.agent.a2a.worker_client.requests.post", - lambda *args, **kwargs: Response(), + post, ) client = WorkerClient.__new__(WorkerClient) monkeypatch.setattr( @@ -248,7 +465,23 @@ def raise_for_status(): lambda session_id: type("Route", (), {"endpoint": "http://worker"})(), ) - assert client.send_prompt("context-1", "hello") == "team result" + request = SendMessageRequest( + message=Message( + message_id="m1", + role=Role.ROLE_USER, + parts=[new_text_part("hello")], + ) + ) + result = client.send_message("context-1", request) + + assert result is not None + assert result.message_id == "m1" + assert result.parts[0].text == "hello" + assert sent["headers"] == { + "content-type": "application/x-protobuf", + "x-select-ai-a2a-method": "SendMessage", + } + assert sent["data"] == request.SerializeToString() def test_worker_client_uses_consul_https_endpoint_with_mtls(monkeypatch): @@ -341,7 +574,7 @@ async def put(self, _url, json): def test_a2ui_operation_uses_a_metadata_marked_data_part(): from google.protobuf.json_format import MessageToDict - part = _a2ui_part( + part = a2ui_part( {"version": "v0.9", "createSurface": {"surfaceId": "form"}} ) @@ -354,48 +587,88 @@ def test_a2ui_operation_uses_a_metadata_marked_data_part(): } -def test_a2ui_action_reads_an_operation_list_or_single_operation(): - action = GatewayExecutor._a2ui_action( - type( - "Context", - (), - { - "message": type( - "Message", - (), +def test_each_connection_form_uses_one_new_unique_surface(): + first = connection_form() + second = connection_form() + + first_surface_ids = [ + operation[message_type]["surfaceId"] + for operation, message_type in zip( + first, + ("createSurface", "updateComponents", "updateDataModel"), + ) + ] + second_surface_id = second[0]["createSurface"]["surfaceId"] + + assert len(set(first_surface_ids)) == 1 + assert first_surface_ids[0].startswith("db-connect-") + assert first_surface_ids[0] != second_surface_id + + +def test_connection_form_matches_advertised_a2ui_version(): + assert A2UI_EXTENSION_URI.endswith(f"/{A2UI_VERSION}") + assert f"/{A2UI_VERSION.replace('.', '_')}/" in A2UI_CATALOG_ID + assert all( + operation["version"] == A2UI_VERSION for operation in connection_form() + ) + + +def test_connection_action_reads_an_operation_list_or_single_operation(): + action = find_action( + Message( + role=Role.ROLE_USER, + parts=[ + a2ui_part( { - "parts": [ - _a2ui_part( - { - "version": "v0.9", - "action": { - "name": "submit_database_connection" - }, - } - ) - ] - }, - )() - }, - )() + "version": "v0.9", + "action": {"name": "submit_database_connection"}, + } + ) + ], + ), + "submit_database_connection", ) assert action == {"name": "submit_database_connection"} +def test_connection_action_accepts_gemini_unmarked_data_part(): + action = find_action( + Message( + role=Role.ROLE_USER, + parts=[ + new_data_part( + { + "version": "v0.9", + "action": { + "name": "submit_database_connection", + "context": {"team_name": "TEAM"}, + }, + } + ) + ], + ), + "submit_database_connection", + ) + + assert action == { + "name": "submit_database_connection", + "context": {"team_name": "TEAM"}, + } + + def test_gateway_returns_connection_error_when_worker_rejects_opening(): - executor = GatewayExecutor.__new__(GatewayExecutor) - executor.sessions = {} + handler = GatewayRequestHandler.__new__(GatewayRequestHandler) class Client: @staticmethod def open_session(_context_id, _session_info): raise requests.HTTPError("worker rejected the connection") - executor.worker_client = Client() + handler.worker_client = Client() - parts = asyncio.run( - executor._open_session( + session_id = asyncio.run( + handler._open_session( { "dsn": "database", "username": "user", @@ -406,5 +679,389 @@ def open_session(_context_id, _session_info): ) ) - assert parts[0].text.startswith("Could not connect.") - assert executor.sessions == {} + assert session_id is None + + +def test_gateway_bootstrap_is_transient_until_worker_session_opens(): + class Client: + def __init__(self): + self.opened = False + self.saved = None + + def session_exists(self, _session_id): + return self.opened + + def open_session(self, session_id, session_info): + assert session_id == "context-1" + assert session_info.team_name == "TEAM" + self.opened = True + return session_id + + client = Client() + handler = GatewayRequestHandler.__new__(GatewayRequestHandler) + handler.worker_client = client + + form_request = SendMessageRequest( + message=Message( + message_id="m1", + context_id="context-1", + role=Role.ROLE_USER, + parts=[new_text_part("connect")], + ) + ) + form_task = asyncio.run( + handler.on_message_send(form_request, ServerCallContext()) + ) + + assert isinstance(form_task, Task) + assert form_task.context_id == "context-1" + assert form_task.artifacts[0].name == "database-connection-form" + assert len(form_task.artifacts[0].parts) == 3 + + connect_request = SendMessageRequest( + message=Message( + message_id="m2", + task_id=form_task.id, + context_id=form_task.context_id, + role=Role.ROLE_USER, + parts=[ + a2ui_part( + { + "version": "v0.9", + "action": { + "name": "submit_database_connection", + "context": { + "dsn": "database", + "username": "user", + "password": "password", + "team_name": "TEAM", + }, + }, + } + ) + ], + ) + ) + connected_task = asyncio.run( + handler.on_message_send(connect_request, ServerCallContext()) + ) + + assert connected_task.status.state == TaskState.TASK_STATE_COMPLETED + assert all(item.message_id != "m2" for item in connected_task.history) + assert '"password": "password"' not in MessageToJson(connected_task) + + +def test_gateway_replaces_bootstrap_task_with_worker_owned_task(): + class Client: + calls = 0 + + @staticmethod + def session_exists(_session_id): + return True + + def send_message(self, session_id, request): + assert session_id == "context-1" + self.calls += 1 + if self.calls == 1: + assert request.message.task_id == ( + "gateway-bootstrap-form-task" + ) + raise TaskNotFoundError + assert not request.message.task_id + return Task( + id="worker-task", + context_id=session_id, + status={"state": TaskState.TASK_STATE_COMPLETED}, + ) + + client = Client() + handler = GatewayRequestHandler.__new__(GatewayRequestHandler) + handler.worker_client = client + request = SendMessageRequest( + message=Message( + message_id="m3", + task_id="gateway-bootstrap-form-task", + context_id="context-1", + role=Role.ROLE_USER, + parts=[new_text_part("hello")], + ) + ) + + task = asyncio.run(handler.on_message_send(request, ServerCallContext())) + + assert task.id == "worker-task" + assert task.context_id == "context-1" + assert client.calls == 2 + + +def test_gateway_preserves_missing_worker_task_error(): + class Client: + @staticmethod + def session_exists(_session_id): + return True + + @staticmethod + def send_message(_session_id, _request): + raise TaskNotFoundError + + handler = GatewayRequestHandler.__new__(GatewayRequestHandler) + handler.worker_client = Client() + request = SendMessageRequest( + message=Message( + message_id="m3", + task_id="oracle-task-that-no-longer-exists", + context_id="context-1", + role=Role.ROLE_USER, + parts=[new_text_part("hello")], + ) + ) + + with pytest.raises(TaskNotFoundError): + asyncio.run(handler.on_message_send(request, ServerCallContext())) + + +def test_gateway_bootstrap_without_context_returns_transient_form_task(): + class Client: + @staticmethod + def session_exists(session_id): + assert session_id + return False + + handler = GatewayRequestHandler.__new__(GatewayRequestHandler) + handler.worker_client = Client() + request = SendMessageRequest( + message=Message( + message_id="m1", + role=Role.ROLE_USER, + parts=[new_text_part("hello")], + ) + ) + + form_task = asyncio.run( + handler.on_message_send(request, ServerCallContext()) + ) + + assert isinstance(form_task, Task) + assert form_task.id.startswith("gateway-bootstrap-") + assert form_task.context_id + # The connection form is a response-only gateway task. Its ID must not be + # attached to the user's message because no corresponding task exists in + # the worker's Oracle task store. + assert not form_task.history[0].task_id + assert form_task.status.state == TaskState.TASK_STATE_COMPLETED + assert form_task.artifacts[0].name == "database-connection-form" + assert len(form_task.artifacts[0].parts) == 3 + + +def test_gateway_returns_task_form_when_task_session_expires_before_send(): + class Client: + @staticmethod + def get_task(_request): + raise ReconnectRequired( + "Database session ended; reconnect required.", + "context-1", + ) + + handler = GatewayRequestHandler.__new__(GatewayRequestHandler) + handler.worker_client = Client() + request = SendMessageRequest( + message=Message( + message_id="m1", + task_id="task-1", + role=Role.ROLE_USER, + parts=[new_text_part("hello")], + ) + ) + + task = asyncio.run(handler.on_message_send(request, ServerCallContext())) + + assert isinstance(task, Task) + assert task.id == "task-1" + assert task.context_id == "context-1" + assert task.status.state == TaskState.TASK_STATE_COMPLETED + assert task.artifacts[0].name == "database-connection-form" + assert len(task.artifacts[0].parts) == 3 + + +def test_gateway_returns_task_form_when_task_route_is_missing_before_send(): + class Client: + @staticmethod + def get_task(_request): + return None + + handler = GatewayRequestHandler.__new__(GatewayRequestHandler) + handler.worker_client = Client() + request = SendMessageRequest( + message=Message( + message_id="m1", + task_id="task-from-old-deployment", + role=Role.ROLE_USER, + parts=[new_text_part("hello")], + ) + ) + + task = asyncio.run(handler.on_message_send(request, ServerCallContext())) + + assert isinstance(task, Task) + assert task.id == "task-from-old-deployment" + assert task.context_id + assert task.status.state == TaskState.TASK_STATE_COMPLETED + assert task.artifacts[0].name == "database-connection-form" + assert len(task.artifacts[0].parts) == 3 + + +def test_gateway_uses_task_form_for_expired_context_only_message(): + class Client: + @staticmethod + def session_exists(_session_id): + return True + + @staticmethod + def send_message(_session_id, _request): + raise ReconnectRequired( + "Database session ended; reconnect required.", + "context-1", + ) + + handler = GatewayRequestHandler.__new__(GatewayRequestHandler) + handler.worker_client = Client() + request = SendMessageRequest( + message=Message( + message_id="m1", + context_id="context-1", + role=Role.ROLE_USER, + parts=[new_text_part("hello")], + ) + ) + + task = asyncio.run(handler.on_message_send(request, ServerCallContext())) + + assert isinstance(task, Task) + assert task.id + assert task.context_id == "context-1" + assert task.status.state == TaskState.TASK_STATE_COMPLETED + assert task.artifacts[0].name == "database-connection-form" + assert len(task.artifacts[0].parts) == 3 + + +def test_gateway_uses_connection_form_when_get_task_session_expires(): + class Client: + @staticmethod + def get_task(_request): + raise ReconnectRequired( + "Database session ended; reconnect required.", + "context-1", + ) + + handler = GatewayRequestHandler.__new__(GatewayRequestHandler) + handler.worker_client = Client() + + task = asyncio.run( + handler.on_get_task( + GetTaskRequest(id="task-1"), + ServerCallContext(), + ) + ) + + assert task.id == "task-1" + assert task.context_id == "context-1" + assert task.status.state == TaskState.TASK_STATE_COMPLETED + assert task.artifacts[0].name == "database-connection-form" + assert len(task.artifacts[0].parts) == 3 + + +def test_gateway_uses_connection_form_when_get_task_route_is_missing(): + class Client: + @staticmethod + def get_task(_request): + return None + + handler = GatewayRequestHandler.__new__(GatewayRequestHandler) + handler.worker_client = Client() + + task = asyncio.run( + handler.on_get_task( + GetTaskRequest(id="task-from-old-deployment"), + ServerCallContext(), + ) + ) + + assert task.id == "task-from-old-deployment" + assert task.context_id + assert task.status.state == TaskState.TASK_STATE_COMPLETED + assert task.artifacts[0].name == "database-connection-form" + assert len(task.artifacts[0].parts) == 3 + + +def test_gateway_uses_connection_form_when_cancel_task_session_expires(): + class Client: + @staticmethod + def cancel_task(_task_id): + raise ReconnectRequired( + "Database session ended; reconnect required.", + "context-1", + ) + + handler = GatewayRequestHandler.__new__(GatewayRequestHandler) + handler.worker_client = Client() + + task = asyncio.run( + handler.on_cancel_task( + CancelTaskRequest(id="task-1"), + ServerCallContext(), + ) + ) + + assert task.id == "task-1" + assert task.context_id == "context-1" + assert task.status.state == TaskState.TASK_STATE_COMPLETED + assert task.artifacts[0].name == "database-connection-form" + assert len(task.artifacts[0].parts) == 3 + + +def test_gateway_uses_connection_form_when_cancel_task_route_is_missing(): + class Client: + @staticmethod + def cancel_task(_task_id): + return None + + handler = GatewayRequestHandler.__new__(GatewayRequestHandler) + handler.worker_client = Client() + + task = asyncio.run( + handler.on_cancel_task( + CancelTaskRequest(id="task-from-old-deployment"), + ServerCallContext(), + ) + ) + + assert task.id == "task-from-old-deployment" + assert task.context_id + assert task.status.state == TaskState.TASK_STATE_COMPLETED + assert task.artifacts[0].name == "database-connection-form" + assert len(task.artifacts[0].parts) == 3 + + +def test_gateway_exposes_connection_form_data_when_list_session_expires(): + class Client: + @staticmethod + def list_tasks(_context_id, _request): + raise ReconnectRequired( + "Database session ended; reconnect required.", + "context-1", + ) + + handler = GatewayRequestHandler.__new__(GatewayRequestHandler) + handler.worker_client = Client() + + with pytest.raises(InvalidParamsError) as raised: + asyncio.run( + handler.on_list_tasks( + ListTasksRequest(context_id="context-1"), + ServerCallContext(), + ) + ) + + assert raised.value.data["reason"] == "SESSION_EXPIRED" + assert raised.value.data["contextId"] == "context-1" + assert raised.value.data["connectionForm"]