diff --git a/src/Server/ClientGateway.php b/src/Server/ClientGateway.php index d4452b15..da2a38b7 100644 --- a/src/Server/ClientGateway.php +++ b/src/Server/ClientGateway.php @@ -40,6 +40,8 @@ use Mcp\Schema\Tool; use Mcp\Schema\ToolChoice; use Mcp\Server\Session\SessionInterface; +use Mcp\Server\Suspension\NotificationSuspension; +use Mcp\Server\Suspension\RequestSuspension; /** * @final @@ -91,11 +93,7 @@ public function __construct( */ public function notify(Notification $notification): void { - \Fiber::suspend([ - 'type' => 'notification', - 'notification' => $notification, - 'session_id' => $this->session->getId()->toRfc4122(), - ]); + \Fiber::suspend(new NotificationSuspension($notification, $this->session->getId()->toRfc4122())); } /** @@ -450,13 +448,7 @@ public function request(Request $request, int $timeout = 120): Response|Error */ private function suspend(Request $request, int $timeout, ?string $key = null): Response|Error { - $response = \Fiber::suspend([ - 'type' => 'request', - 'request' => $request, - 'session_id' => $this->session->getId()->toRfc4122(), - 'timeout' => $timeout, - 'input_key' => $key, - ]); + $response = \Fiber::suspend(new RequestSuspension($request, $this->session->getId()->toRfc4122(), $timeout, $key)); if (!$response instanceof Response && !$response instanceof Error) { throw new RuntimeException('Transport returned an unexpected payload; expected a Response or Error message.'); diff --git a/src/Server/Protocol.php b/src/Server/Protocol.php index 19e88852..71fc9b1b 100644 --- a/src/Server/Protocol.php +++ b/src/Server/Protocol.php @@ -30,6 +30,8 @@ use Mcp\Server\Session\SessionManagerInterface; use Mcp\Server\Stateless\InputContext; use Mcp\Server\Stateless\RequestStateCodec; +use Mcp\Server\Suspension\NotificationSuspension; +use Mcp\Server\Suspension\RequestSuspension; use Mcp\Server\Transport\TransportInterface; use Psr\EventDispatcher\EventDispatcherInterface; use Psr\Log\LoggerInterface; @@ -322,17 +324,10 @@ private function handleRequest(TransportInterface $transport, Request $request, $result = $fiber->start(); if ($fiber->isSuspended()) { - if (\is_array($result) && isset($result['type'])) { - if ('notification' === $result['type']) { - $notification = $result['notification']; - $this->sendNotification($notification, $session); - } elseif ('request' === $result['type']) { - // Keep $request untouched: it is the inbound request the catch - // blocks below answer under, not this outbound one. - $outboundRequest = $result['request']; - $timeout = $result['timeout'] ?? 120; - $this->sendRequest($outboundRequest, $timeout, $session); - } + if ($result instanceof NotificationSuspension) { + $this->sendNotification($result->notification, $session); + } elseif ($result instanceof RequestSuspension) { + $this->sendRequest($result->request, $result->timeout, $session); } $transport->attachFiberToSession($fiber, $session->getId()); @@ -618,7 +613,7 @@ public function handleFiberYield(mixed $yieldedValue, ?Uuid $sessionId): void return; } - if (!\is_array($yieldedValue) || !isset($yieldedValue['type'])) { + if (!$yieldedValue instanceof NotificationSuspension && !$yieldedValue instanceof RequestSuspension) { $this->logger->warning('Fiber yielded unexpected payload.', [ 'payload' => $yieldedValue, 'session_id' => $sessionId->toRfc4122(), @@ -629,43 +624,18 @@ public function handleFiberYield(mixed $yieldedValue, ?Uuid $sessionId): void $session = $this->sessionManager->createWithId($sessionId); - $payloadSessionId = $yieldedValue['session_id'] ?? null; - if (\is_string($payloadSessionId) && $payloadSessionId !== $sessionId->toRfc4122()) { + if ($yieldedValue->sessionId !== $sessionId->toRfc4122()) { $this->logger->warning('Fiber yielded payload with mismatched session ID.', [ - 'payload_session_id' => $payloadSessionId, + 'payload_session_id' => $yieldedValue->sessionId, 'expected_session_id' => $sessionId->toRfc4122(), ]); } try { - if ('notification' === $yieldedValue['type']) { - $notification = $yieldedValue['notification'] ?? null; - if (!$notification instanceof Notification) { - $this->logger->warning('Fiber yielded notification without Notification instance.', [ - 'payload' => $yieldedValue, - ]); - - return; - } - - $this->sendNotification($notification, $session); - } elseif ('request' === $yieldedValue['type']) { - $request = $yieldedValue['request'] ?? null; - if (!$request instanceof Request) { - $this->logger->warning('Fiber yielded request without Request instance.', [ - 'payload' => $yieldedValue, - ]); - - return; - } - - $timeout = isset($yieldedValue['timeout']) ? (int) $yieldedValue['timeout'] : 120; - $this->sendRequest($request, $timeout, $session); - } else { - $this->logger->warning('Fiber yielded unknown operation type.', [ - 'type' => $yieldedValue['type'], - ]); - } + match (true) { + $yieldedValue instanceof NotificationSuspension => $this->sendNotification($yieldedValue->notification, $session), + $yieldedValue instanceof RequestSuspension => $this->sendRequest($yieldedValue->request, $yieldedValue->timeout, $session), + }; } finally { $session->save(); } diff --git a/src/Server/Stateless/StatelessProtocol.php b/src/Server/Stateless/StatelessProtocol.php index 4d6a8917..fcee2fd9 100644 --- a/src/Server/Stateless/StatelessProtocol.php +++ b/src/Server/Stateless/StatelessProtocol.php @@ -37,6 +37,8 @@ use Mcp\Server\Session\InMemorySessionStore; use Mcp\Server\Session\Session; use Mcp\Server\Subscription\NotificationBusInterface; +use Mcp\Server\Suspension\NotificationSuspension; +use Mcp\Server\Suspension\RequestSuspension; use Mcp\Server\Wire\CachePolicy; use Mcp\Server\Wire\InboundClassifier; use Mcp\Server\Wire\Rev2026Codec; @@ -675,8 +677,8 @@ private function run(RequestHandlerInterface $handler, Request $request, Session */ private function readNotification(mixed $suspended, RequestMeta $meta): ?Notification { - if (!\is_array($suspended) || 'notification' !== ($suspended['type'] ?? null)) { - if (\is_array($suspended) && 'request' === ($suspended['type'] ?? null)) { + if (!$suspended instanceof NotificationSuspension) { + if ($suspended instanceof RequestSuspension) { // Elicitation never reaches here — it is answered in run(). What // is left are the kinds this revision removed outright, and no // multi round-trip shape brings them back. @@ -686,11 +688,7 @@ private function readNotification(mixed $suspended, RequestMeta $meta): ?Notific return null; } - $notification = $suspended['notification'] ?? null; - - if (!$notification instanceof Notification) { - return null; - } + $notification = $suspended->notification; // The client opts into logs per request; with no level named the server // MUST NOT send any, which is why an absent level drops rather than @@ -713,19 +711,11 @@ private function readNotification(mixed $suspended, RequestMeta $meta): ?Notific */ private static function readElicitation(mixed $suspended): ?array { - if (!\is_array($suspended) || 'request' !== ($suspended['type'] ?? null)) { - return null; - } - - $request = $suspended['request'] ?? null; - - if (!$request instanceof ElicitRequest) { + if (!$suspended instanceof RequestSuspension || !$suspended->request instanceof ElicitRequest) { return null; } - $key = $suspended['input_key'] ?? null; - - return [\is_string($key) ? $key : null, $request]; + return [$suspended->inputKey, $suspended->request]; } /** diff --git a/src/Server/Suspension/NotificationSuspension.php b/src/Server/Suspension/NotificationSuspension.php new file mode 100644 index 00000000..b4f5b5b0 --- /dev/null +++ b/src/Server/Suspension/NotificationSuspension.php @@ -0,0 +1,28 @@ + + */ +final class NotificationSuspension +{ + public function __construct( + public readonly Notification $notification, + public readonly string $sessionId, + ) { + } +} diff --git a/src/Server/Suspension/RequestSuspension.php b/src/Server/Suspension/RequestSuspension.php new file mode 100644 index 00000000..fcd5f029 --- /dev/null +++ b/src/Server/Suspension/RequestSuspension.php @@ -0,0 +1,38 @@ + + */ +final class RequestSuspension +{ + /** + * @param int $timeout maximum time to wait for the response (seconds) + * @param string|null $inputKey the name an elicitation's answer is filed under when + * the revision serving the call answers by asking + * ({@see \Mcp\Server\Stateless\ElicitationReplay}); + * ignored by every leg that has a live client to ask + */ + public function __construct( + public readonly Request $request, + public readonly string $sessionId, + public readonly int $timeout = 120, + public readonly ?string $inputKey = null, + ) { + } +} diff --git a/src/Server/Transport/TransportInterface.php b/src/Server/Transport/TransportInterface.php index 6e9a5730..d35fd3e2 100644 --- a/src/Server/Transport/TransportInterface.php +++ b/src/Server/Transport/TransportInterface.php @@ -21,10 +21,7 @@ * * @phpstan-type FiberReturn (Response|Error) * @phpstan-type FiberResume (FiberReturn|null) - * @phpstan-type FiberSuspend ( - * array{type: 'notification', notification: \Mcp\Schema\JsonRpc\Notification}| - * array{type: 'request', request: \Mcp\Schema\JsonRpc\Request, timeout?: int} - * ) + * @phpstan-type FiberSuspend (\Mcp\Server\Suspension\NotificationSuspension|\Mcp\Server\Suspension\RequestSuspension) * @phpstan-type McpFiber \Fiber * * @author Christopher Hertel diff --git a/tests/Unit/Server/ClientGatewayTest.php b/tests/Unit/Server/ClientGatewayTest.php index 2b70fd6e..21073116 100644 --- a/tests/Unit/Server/ClientGatewayTest.php +++ b/tests/Unit/Server/ClientGatewayTest.php @@ -28,6 +28,7 @@ use Mcp\Schema\Result\ListRootsResult; use Mcp\Server\ClientGateway; use Mcp\Server\Session\SessionInterface; +use Mcp\Server\Suspension\RequestSuspension; use PHPUnit\Framework\TestCase; use Symfony\Component\Uid\Uuid; @@ -217,11 +218,10 @@ private function runInFiber(\Closure $call, Response|Error $response, string $ex $fiber = new \Fiber($call); $suspend = $fiber->start(); - $this->assertIsArray($suspend); - $this->assertSame('request', $suspend['type']); - $this->assertInstanceOf($expectedRequest, $suspend['request']); + $this->assertInstanceOf(RequestSuspension::class, $suspend); + $this->assertInstanceOf($expectedRequest, $suspend->request); - $request = $suspend['request']; + $request = $suspend->request; $fiber->resume($response); diff --git a/tests/Unit/Server/InputRequiredShimTest.php b/tests/Unit/Server/InputRequiredShimTest.php index 137270f5..f1c4a736 100644 --- a/tests/Unit/Server/InputRequiredShimTest.php +++ b/tests/Unit/Server/InputRequiredShimTest.php @@ -27,6 +27,7 @@ use Mcp\Server\Session\Session; use Mcp\Server\Session\SessionInterface; use Mcp\Server\Stateless\InputContext; +use Mcp\Server\Suspension\RequestSuspension; use PHPUnit\Framework\Attributes\TestDox; use PHPUnit\Framework\TestCase; @@ -149,8 +150,7 @@ private function drive( $suspended = $fiber->start(); while ($fiber->isSuspended()) { - $this->assertIsArray($suspended); - $this->assertSame('request', $suspended['type'], 'the shim only ever suspends to send a client request'); + $this->assertInstanceOf(RequestSuspension::class, $suspended, 'the shim only ever suspends to send a client request'); $this->assertNotEmpty($answers, 'the shim sent more requests than the test queued answers for'); $suspended = $fiber->resume(array_shift($answers)); diff --git a/tests/Unit/Server/ProtocolTest.php b/tests/Unit/Server/ProtocolTest.php index 44422992..93773cd9 100644 --- a/tests/Unit/Server/ProtocolTest.php +++ b/tests/Unit/Server/ProtocolTest.php @@ -18,17 +18,22 @@ use Mcp\JsonRpc\MessageFactory; use Mcp\Schema\Enum\LoggingLevel; use Mcp\Schema\JsonRpc\Error; +use Mcp\Schema\JsonRpc\Request; use Mcp\Schema\JsonRpc\Response; use Mcp\Schema\Notification\LoggingMessageNotification; use Mcp\Schema\Request\CallToolRequest; use Mcp\Schema\Request\PingRequest; +use Mcp\Server\ClientGateway; use Mcp\Server\Handler\Notification\NotificationHandlerInterface; use Mcp\Server\Handler\Request\RequestHandlerInterface; use Mcp\Server\Protocol; use Mcp\Server\Session\InMemorySessionStore; +use Mcp\Server\Session\Session; use Mcp\Server\Session\SessionInterface; use Mcp\Server\Session\SessionManager; use Mcp\Server\Session\SessionManagerInterface; +use Mcp\Server\Suspension\NotificationSuspension; +use Mcp\Server\Suspension\RequestSuspension; use Mcp\Server\Transport\TransportInterface; use Mcp\Tests\Unit\Fixtures\ThrowingRequest; use PHPUnit\Framework\Attributes\DataProvider; @@ -802,9 +807,9 @@ public function testOutboundRequestFailureIsAnsweredUnderInboundRequestId(): voi { $handler = $this->createMock(RequestHandlerInterface::class); $handler->method('supports')->willReturn(true); - $handler->method('handle')->willReturnCallback(static function (): Response { + $handler->method('handle')->willReturnCallback(static function (Request $request, SessionInterface $session): Response { // Suspend with an outbound, id-less request, as sampling/elicitation handlers do. - \Fiber::suspend(['type' => 'request', 'request' => new PingRequest(), 'timeout' => 5]); + \Fiber::suspend(new RequestSuspension(new PingRequest(), $session->getId()->toRfc4122(), 5)); return new Response(1, []); }); @@ -844,11 +849,11 @@ public function testOutboundRequestFailureIsAnsweredUnderInboundRequestId(): voi $this->assertSame($exception, $errorEvents[0]->getThrowable()); $this->assertSame(1, $errorEvents[0]->getError()->getId()); - $outgoing = $protocol->consumeOutgoingMessages($sessionId); - $errors = array_values(array_filter( - array_map(static fn (array $outgoingMessage): array => json_decode($outgoingMessage['message'], true), $outgoing), - static fn (array $message): bool => isset($message['error']), - )); + $outgoing = array_map(static fn (array $outgoingMessage): array => json_decode($outgoingMessage['message'], true), $protocol->consumeOutgoingMessages($sessionId)); + $outbound = array_values(array_filter($outgoing, static fn (array $message): bool => 'ping' === ($message['method'] ?? null))); + $this->assertCount(1, $outbound); + + $errors = array_values(array_filter($outgoing, static fn (array $message): bool => isset($message['error']))); $this->assertCount(1, $errors); $this->assertSame(1, $errors[0]['id']); $this->assertSame(Error::INTERNAL_ERROR, $errors[0]['error']['code']); @@ -1700,6 +1705,95 @@ public function testMessagePayloadsAreOnlyLoggedAtDebugLevel(string $input, stri $this->assertStringContainsString($identifier, $infoAndAbove); $this->assertStringContainsString('s3cr3t-payload', $debug); } + + #[TestDox('A notification suspension from the gateway round-trips into the outgoing queue')] + public function testFiberYieldedNotificationSuspensionIsQueued(): void + { + $sessionId = Uuid::v4(); + $session = new Session(new InMemorySessionStore(), $sessionId); + + $this->sessionManager->method('createWithId')->willReturn($session); + + $protocol = new Protocol( + requestHandlers: [], + notificationHandlers: [], + messageFactory: MessageFactory::make(), + sessionManager: $this->sessionManager, + ); + + $gateway = new ClientGateway($session); + $notification = new LoggingMessageNotification(LoggingLevel::Info, 'hello'); + + $fiber = new \Fiber(static fn () => $gateway->notify($notification)); + $suspension = $fiber->start(); + + $this->assertInstanceOf(NotificationSuspension::class, $suspension); + $this->assertSame($sessionId->toRfc4122(), $suspension->sessionId); + + $protocol->handleFiberYield($suspension, $sessionId); + + $outgoing = $protocol->consumeOutgoingMessages($sessionId); + $this->assertCount(1, $outgoing); + $this->assertSame(['type' => 'notification'], $outgoing[0]['context']); + $this->assertSame(json_encode($notification), $outgoing[0]['message']); + } + + #[TestDox('A request suspension from the gateway round-trips into the outgoing queue and pending requests')] + public function testFiberYieldedRequestSuspensionIsQueued(): void + { + $sessionId = Uuid::v4(); + $session = new Session(new InMemorySessionStore(), $sessionId); + + $this->sessionManager->method('createWithId')->willReturn($session); + + $protocol = new Protocol( + requestHandlers: [], + notificationHandlers: [], + messageFactory: MessageFactory::make(), + sessionManager: $this->sessionManager, + ); + + $gateway = new ClientGateway($session); + + $fiber = new \Fiber(static fn () => $gateway->listRoots(timeout: 45)); + $suspension = $fiber->start(); + + $this->assertInstanceOf(RequestSuspension::class, $suspension); + $this->assertSame($sessionId->toRfc4122(), $suspension->sessionId); + $this->assertSame(45, $suspension->timeout); + $this->assertNull($suspension->inputKey); + + $protocol->handleFiberYield($suspension, $sessionId); + + $pending = $protocol->getPendingRequests($sessionId); + $this->assertCount(1, $pending); + $this->assertSame(45, $pending[1000]['timeout']); + + $outgoing = $protocol->consumeOutgoingMessages($sessionId); + $this->assertCount(1, $outgoing); + $this->assertSame(['type' => 'request'], $outgoing[0]['context']); + + $message = json_decode($outgoing[0]['message'], true); + $this->assertSame('roots/list', $message['method']); + $this->assertSame(1000, $message['id']); + } + + #[TestDox('A fiber yield that is not a suspension object is dropped without touching the session')] + public function testFiberYieldedUnexpectedPayloadIsIgnored(): void + { + $this->sessionManager->expects($this->never())->method('createWithId'); + + $protocol = new Protocol( + requestHandlers: [], + notificationHandlers: [], + messageFactory: MessageFactory::make(), + sessionManager: $this->sessionManager, + ); + + // The pre-VO array shape is deliberately no longer accepted. + // @phpstan-ignore argument.type + $protocol->handleFiberYield(['type' => 'notification'], Uuid::v4()); + } } /**