diff --git a/src/acp/agent/acp_requests.py b/src/acp/agent/acp_requests.py index e23e1cfab..45d272c23 100644 --- a/src/acp/agent/acp_requests.py +++ b/src/acp/agent/acp_requests.py @@ -3,7 +3,7 @@ from __future__ import annotations import asyncio -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Literal import structlog @@ -214,6 +214,7 @@ async def elicitation_create( self, message: str, *, + mode: Literal["form", "url"], requested_schema: dict[str, Any], url: str | None = None, elicitation_id: str | None = None, @@ -222,6 +223,7 @@ async def elicitation_create( Args: message: Human-readable message describing what input is being requested + mode: Elicitation mode (``form`` or ``url``) requested_schema: JSON Schema object describing the expected input structure url: Optional URL for URL-based elicitation (e.g., OAuth flows) elicitation_id: Optional unique identifier for this elicitation request @@ -232,6 +234,7 @@ async def elicitation_create( request = ElicitationCreateRequest( session_id=self.id, message=message, + mode=mode, requested_schema=requested_schema, url=url, elicitation_id=elicitation_id, diff --git a/src/acp/schema/agent_responses.py b/src/acp/schema/agent_responses.py index 4c702fd96..5b58f87da 100644 --- a/src/acp/schema/agent_responses.py +++ b/src/acp/schema/agent_responses.py @@ -261,7 +261,7 @@ class InitializeResponse(Response): See protocol docs: [Initialization](https://agentclientprotocol.com/protocol/initialization) """ - agent_capabilities: AgentCapabilities | None = Field(default_factory=AgentCapabilities) + agent_capabilities: AgentCapabilities | None = None """Capabilities supported by the agent.""" agent_info: Implementation | None = None diff --git a/src/acp/schema/capabilities.py b/src/acp/schema/capabilities.py index d0a5918c0..d4900d920 100644 --- a/src/acp/schema/capabilities.py +++ b/src/acp/schema/capabilities.py @@ -4,7 +4,7 @@ from typing import Self -from pydantic import Field +from pydantic import Field, field_validator from acp.schema.base import AnnotatedObject from acp.schema.slash_commands import AvailableCommand # noqa: TC001 @@ -37,14 +37,30 @@ class AuthCapabilities(AnnotatedObject): class ElicitationCapabilities(AnnotatedObject): """Elicitation capabilities supported by the client. - Advertised during initialization to inform the agent whether - the client supports the `elicitation/create` method. + Advertised during initialization to inform the agent which + elicitation modes the client supports. See protocol docs: [Elicitation](https://agentclientprotocol.com/protocol/elicitation) """ - create: bool | None = False - """Whether the Client supports `elicitation/create` requests.""" + form: bool | None = False + """Whether the Client supports form-mode `elicitation/create` requests.""" + + url: bool | None = False + """Whether the Client supports URL-mode `elicitation/create` requests.""" + + @field_validator("form", "url", mode="before") + @classmethod + def convert_empty_object(cls, v): + """Convert empty object {} to True for compatibility. + + Some clients may send empty objects {} instead of booleans for optional fields. + If field is empty object {}, interpret as True. + Otherwise preserve original value (True/False/None). + """ + if isinstance(v, dict) and len(v) == 0: + return True + return v class ClientCapabilities(AnnotatedObject): @@ -59,7 +75,7 @@ class ClientCapabilities(AnnotatedObject): auth: AuthCapabilities | None = None """**UNSTABLE**: Authentication capabilities supported by the client.""" - fs: FileSystemCapability | None = Field(default_factory=FileSystemCapability) + fs: FileSystemCapability | None = None """File system capabilities supported by the client. Determines which file operations the agent can request. @@ -241,13 +257,13 @@ class AgentCapabilities(AnnotatedObject): load_session: bool | None = False """Whether the agent supports `session/load`.""" - mcp_capabilities: McpCapabilities | None = Field(default_factory=McpCapabilities) + mcp_capabilities: McpCapabilities | None = None """MCP capabilities supported by the agent.""" - prompt_capabilities: PromptCapabilities | None = Field(default_factory=PromptCapabilities) + prompt_capabilities: PromptCapabilities | None = None """Prompt capabilities supported by the agent.""" - session_capabilities: SessionCapabilities | None = Field(default_factory=SessionCapabilities) + session_capabilities: SessionCapabilities | None = None """Session capabilities supported by the agent.""" slash_commands: list[AvailableCommand] = Field(default_factory=list) diff --git a/src/acp/schema/client_requests.py b/src/acp/schema/client_requests.py index 102d8cdfd..78bff025d 100644 --- a/src/acp/schema/client_requests.py +++ b/src/acp/schema/client_requests.py @@ -196,7 +196,7 @@ class InitializeRequest(Request): See protocol docs: [Initialization](https://agentclientprotocol.com/protocol/initialization) """ - client_capabilities: ClientCapabilities | None = Field(default_factory=ClientCapabilities) + client_capabilities: ClientCapabilities | None = None """Capabilities supported by the client.""" client_info: Implementation | None = None diff --git a/src/acp/schema/elicitation.py b/src/acp/schema/elicitation.py index 540427ab6..9ffc303b2 100644 --- a/src/acp/schema/elicitation.py +++ b/src/acp/schema/elicitation.py @@ -25,6 +25,9 @@ class ElicitationCreateRequest(Request): message: str """A human-readable message describing what input is being requested.""" + mode: Literal["form", "url"] + """The elicitation mode: ``form`` for schema-based input, ``url`` for external URL flows.""" + requested_schema: dict[str, Any] = Field(alias="requestedSchema") """A JSON Schema object describing the expected input structure.""" diff --git a/src/acp/transports.py b/src/acp/transports.py index 49a9ead21..f1b6faa27 100644 --- a/src/acp/transports.py +++ b/src/acp/transports.py @@ -261,24 +261,21 @@ async def handle_client(websocket: ServerConnection) -> None: connections.append(conn) try: - # Keep connection alive until client disconnects or shutdown - client_done = asyncio.Event() - - async def monitor_websocket() -> None: - try: - async for _ in websocket: - pass # Messages handled by ws_reader - except websockets.exceptions.ConnectionClosed: - pass - finally: - client_done.set() - - monitor_task = asyncio.create_task(monitor_websocket()) - tasks = [asyncio.create_task(client_done.wait()), asyncio.create_task(shutdown.wait())] - _done, _ = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED) - monitor_task.cancel() - with contextlib.suppress(asyncio.CancelledError): - await monitor_task + # Wait for shutdown or for the receive loop to end (client disconnect) + _recv_conn = getattr(conn, "_conn", None) + recv_task = getattr(_recv_conn, "_recv_task", None) if _recv_conn else None + waitables: list[asyncio.Future[Any]] = [asyncio.create_task(shutdown.wait())] + if isinstance(recv_task, asyncio.Task): + waitables.append(recv_task) + + done, _ = await asyncio.wait(waitables, return_when=asyncio.FIRST_COMPLETED) + + # Cancel any remaining tasks + for w in waitables: + if isinstance(w, asyncio.Task) and w not in done: + w.cancel() + with contextlib.suppress(asyncio.CancelledError): + await w except websockets.exceptions.ConnectionClosed: logger.info("WebSocket client disconnected") finally: @@ -334,15 +331,12 @@ async def handle_acp(websocket: Any) -> None: try: # Wait for shutdown or for the receive loop to end (client disconnect) - recv_task = conn._conn._recv_task - waitables: list[asyncio.Future[Any]] = [ - asyncio.create_task(shutdown.wait()) - ] - if isinstance(recv_task, asyncio.Task) and not recv_task.done(): + _recv_conn = getattr(conn, "_conn", None) + recv_task = getattr(_recv_conn, "_recv_task", None) if _recv_conn else None + waitables: list[asyncio.Future[Any]] = [asyncio.create_task(shutdown.wait())] + if isinstance(recv_task, asyncio.Task): waitables.append(recv_task) - done, _pending = await asyncio.wait( - waitables, return_when=asyncio.FIRST_COMPLETED - ) + done, _pending = await asyncio.wait(waitables, return_when=asyncio.FIRST_COMPLETED) # Cancel any remaining tasks for w in waitables: if isinstance(w, asyncio.Task) and w not in done: diff --git a/src/agentpool_server/acp_server/input_provider.py b/src/agentpool_server/acp_server/input_provider.py index 9d17cb764..72282c634 100644 --- a/src/agentpool_server/acp_server/input_provider.py +++ b/src/agentpool_server/acp_server/input_provider.py @@ -196,10 +196,17 @@ def _handle_permission_response(self, option_id: str, tool_name: str) -> Confirm logger.warning("Unknown permission option", option_id=option_id) return "abort_run" - def _client_supports_elicitation(self) -> bool: - """Check if the client supports the elicitation/create method.""" + def _client_supports_elicitation(self, mode: Literal["form", "url"]) -> bool: + """Check if the client supports the given elicitation mode.""" caps = self.session.client_capabilities - return caps.elicitation is not None and bool(caps.elicitation.create) + if caps.elicitation is None: + return False + logger.info("Checking elicitation capability", mode=mode, capabilities=caps.elicitation) + match mode: + case "form": + return bool(caps.elicitation.form) + case "url": + return bool(caps.elicitation.url) @staticmethod def _map_elicitation_create_response( @@ -222,9 +229,9 @@ async def get_elicitation( ) -> types.ElicitResult | types.ErrorData: """Get user response to elicitation request with capability-gated dual path. - When the client declares the ``elicitation.create`` capability, uses the - native ``elicitation/create`` protocol method. Otherwise falls back to - the legacy ``request_permission`` approach for backward compatibility. + When the client declares elicitation capability for the requested mode, + uses the native ``elicitation/create`` protocol method. Otherwise falls + back to the legacy ``request_permission`` approach for backward compatibility. Args: params: MCP elicit request parameters @@ -257,7 +264,7 @@ async def _get_url_elicitation( elicitation_id=elicit_id, ) - if self._client_supports_elicitation(): + if self._client_supports_elicitation("url"): # TODO: URL-mode elicitation currently returns the immediate response # from ``elicitation/create``. For full URL flows where the user # completes an external action (OAuth, payments), the result arrives @@ -266,6 +273,7 @@ async def _get_url_elicitation( # deferred to a future PR. response = await self.session.requests.elicitation_create( message=params.message, + mode="url", requested_schema={"type": "object"}, url=params.url, elicitation_id=elicit_id, @@ -307,13 +315,13 @@ async def _get_form_elicitation( schema = params.requestedSchema logger.info("Elicitation request", message=params.message, schema=schema) - if self._client_supports_elicitation(): + if self._client_supports_elicitation("form"): response = await self.session.requests.elicitation_create( message=params.message, + mode="form", requested_schema=schema, ) return self._map_elicitation_create_response(response) - # Fallback: request_permission with schema-specific handling tool_call_id = f"elicit_{hash(params.message)}" title = f"Elicitation: {params.message}"