lowlevel Server: widen on_* return types for InputRequiredResult; add subscriptions/listen slot (#2967)

This commit is contained in:
Max
2026-06-25 14:20:01 +02:00
committed by GitHub
parent a527142312
commit ae13ede143
3 changed files with 29 additions and 6 deletions
+9 -3
View File
@@ -148,7 +148,7 @@ class Server(Generic[LifespanResultT]):
| None = None,
on_call_tool: Callable[
[ServerRequestContext[LifespanResultT], types.CallToolRequestParams],
Awaitable[types.CallToolResult],
Awaitable[types.CallToolResult | types.InputRequiredResult],
]
| None = None,
on_list_resources: Callable[
@@ -163,7 +163,7 @@ class Server(Generic[LifespanResultT]):
| None = None,
on_read_resource: Callable[
[ServerRequestContext[LifespanResultT], types.ReadResourceRequestParams],
Awaitable[types.ReadResourceResult],
Awaitable[types.ReadResourceResult | types.InputRequiredResult],
]
| None = None,
on_subscribe_resource: Callable[
@@ -176,6 +176,11 @@ class Server(Generic[LifespanResultT]):
Awaitable[types.EmptyResult],
]
| None = None,
on_subscriptions_listen: Callable[
[ServerRequestContext[LifespanResultT], types.SubscriptionsListenRequestParams],
Awaitable[types.EmptyResult],
]
| None = None,
on_list_prompts: Callable[
[ServerRequestContext[LifespanResultT], types.PaginatedRequestParams | None],
Awaitable[types.ListPromptsResult],
@@ -183,7 +188,7 @@ class Server(Generic[LifespanResultT]):
| None = None,
on_get_prompt: Callable[
[ServerRequestContext[LifespanResultT], types.GetPromptRequestParams],
Awaitable[types.GetPromptResult],
Awaitable[types.GetPromptResult | types.InputRequiredResult],
]
| None = None,
on_completion: Callable[
@@ -242,6 +247,7 @@ class Server(Generic[LifespanResultT]):
("resources/read", types.ReadResourceRequestParams, on_read_resource),
("resources/subscribe", types.SubscribeRequestParams, on_subscribe_resource),
("resources/unsubscribe", types.UnsubscribeRequestParams, on_unsubscribe_resource),
("subscriptions/listen", types.SubscriptionsListenRequestParams, on_subscriptions_listen),
("tools/list", types.PaginatedRequestParams, on_list_tools),
("tools/call", types.CallToolRequestParams, on_call_tool),
("logging/setLevel", types.SetLevelRequestParams, on_set_logging_level),
+9 -2
View File
@@ -16,9 +16,10 @@ from pydantic import (
Field,
FileUrl,
TypeAdapter,
model_validator,
)
from pydantic.alias_generators import to_camel
from typing_extensions import NotRequired, TypedDict
from typing_extensions import NotRequired, Self, TypedDict
from mcp.types.jsonrpc import RequestId
@@ -2052,7 +2053,7 @@ class InputRequiredResult(Result):
(`tools/call`, `prompts/get`, `resources/read`). The client fulfills
`input_requests` and retries the original request, carrying the responses
and the echoed `request_state`. At least one of those two fields is
present on the wire (spec MUST; not enforced by the model).
present on the wire (spec MUST).
"""
result_type: Literal["input_required"] = "input_required"
@@ -2064,6 +2065,12 @@ class InputRequiredResult(Result):
request_state: str | None = None
"""Opaque state to pass back verbatim when the client retries the original request."""
@model_validator(mode="after")
def _require_one_field(self) -> Self:
if self.input_requests is None and self.request_state is None:
raise ValueError("InputRequiredResult requires at least one of input_requests or request_state")
return self
# Forward refs to InputResponses; rebuild at import time rather than first use.
InputResponseRequestParams.model_rebuild()
+11 -1
View File
@@ -2,6 +2,7 @@ from typing import Any
import pytest
from inline_snapshot import snapshot
from pydantic import ValidationError
from mcp.types import (
LATEST_PROTOCOL_VERSION,
@@ -436,4 +437,13 @@ def test_empty_result_dumps_result_type_only_when_explicitly_tagged():
def test_input_required_result_dumps_its_discriminating_tag():
assert _wire_dump(InputRequiredResult()) == snapshot({"resultType": "input_required"})
assert _wire_dump(InputRequiredResult(request_state="s")) == snapshot(
{"resultType": "input_required", "requestState": "s"}
)
def test_input_required_result_requires_at_least_one_of_input_requests_or_request_state():
with pytest.raises(ValidationError):
InputRequiredResult()
assert InputRequiredResult(input_requests={}).request_state is None
assert InputRequiredResult(request_state="s").input_requests is None