diff --git a/CHANGELOG.md b/CHANGELOG.md index eb1ea471..458177c8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -34,6 +34,7 @@ All notable changes to `mcp/sdk` will be documented in this file. * [BC Break] `ProtectedResourceMetadata` requires `$resource`, serves at the path derived from it (RFC 9728 ยง3.1) and requires https except for loopback hosts; drops localized, policy, ToS, extra fields and `$metadataPaths`. * [BC Break] Add `ScopePolicy` as third argument of `AuthorizationMiddleware`, answering `403 insufficient_scope` per method and tool, with scope hierarchies; the `resource_metadata` challenge URL comes from the configured resource instead of the `Host` header. * Expose `WWW-Authenticate` in the default `CorsMiddleware`. +* [BC Break] Fix concurrent Streamable HTTP streams on one session resuming each other's fibers: each stream now polls only the client request its own fiber sent, so an elicitation answer reaches the tool call that asked for it. `Protocol::handleFiberYield()` returns the ID of the request it sent. 0.8.0 ----- diff --git a/src/Server/Protocol.php b/src/Server/Protocol.php index 71fc9b1b..e31afcc0 100644 --- a/src/Server/Protocol.php +++ b/src/Server/Protocol.php @@ -72,6 +72,14 @@ class Protocol */ private const INTERNAL_ERROR_MESSAGE = 'Internal server error.'; + /** + * The client request each transport's fiber is suspended on. Pending requests are stored in the + * session, which concurrent streams share, so a stream must only poll the one its fiber sent. + * + * @var \WeakMap, int> + */ + private \WeakMap $awaitedRequestIds; + /** * @param array>> $requestHandlers * @param array $notificationHandlers @@ -86,6 +94,7 @@ public function __construct( private readonly ?InputRequiredShim $inputRequiredShim = null, private readonly ?RequestStateCodec $requestStateCodec = null, ) { + $this->awaitedRequestIds = new \WeakMap(); } /** @@ -103,11 +112,20 @@ public function connect(TransportInterface $transport): void $transport->setOutgoingMessagesProvider($this->consumeOutgoingMessages(...)); - $transport->setPendingRequestsProvider($this->getPendingRequests(...)); + // The transport keeps these callbacks, so they reference it weakly to not keep it alive. + $transportRef = \WeakReference::create($transport); + + $transport->setPendingRequestsProvider(fn (Uuid $sessionId): array => $this->getAwaitedPendingRequests($transportRef->get(), $sessionId)); $transport->setResponseFinder($this->checkResponse(...)); - $transport->setFiberYieldHandler($this->handleFiberYield(...)); + $transport->setFiberYieldHandler(function (mixed $yieldedValue, ?Uuid $sessionId) use ($transportRef): void { + $requestId = $this->handleFiberYield($yieldedValue, $sessionId); + + if (null !== $transport = $transportRef->get()) { + $this->trackAwaitedRequest($transport, $requestId); + } + }); $this->logger->info('Protocol connected to transport', ['transport' => $transport::class]); } @@ -324,12 +342,14 @@ private function handleRequest(TransportInterface $transport, Request $request, $result = $fiber->start(); if ($fiber->isSuspended()) { + $awaitedRequestId = null; if ($result instanceof NotificationSuspension) { $this->sendNotification($result->notification, $session); } elseif ($result instanceof RequestSuspension) { - $this->sendRequest($result->request, $result->timeout, $session); + $awaitedRequestId = $this->sendRequest($result->request, $result->timeout, $session); } + $this->trackAwaitedRequest($transport, $awaitedRequestId); $transport->attachFiberToSession($fiber, $session->getId()); return; @@ -604,13 +624,15 @@ public function getPendingRequests(Uuid $sessionId): array * Handle values yielded by Fibers during transport-managed resumes. * * @param FiberSuspend|null $yieldedValue + * + * @return int|null the ID of the request sent to the client, which the fiber now waits on */ - public function handleFiberYield(mixed $yieldedValue, ?Uuid $sessionId): void + public function handleFiberYield(mixed $yieldedValue, ?Uuid $sessionId): ?int { if (!$sessionId) { $this->logger->warning('Fiber yielded value without associated session context.'); - return; + return null; } if (!$yieldedValue instanceof NotificationSuspension && !$yieldedValue instanceof RequestSuspension) { @@ -619,7 +641,7 @@ public function handleFiberYield(mixed $yieldedValue, ?Uuid $sessionId): void 'session_id' => $sessionId->toRfc4122(), ]); - return; + return null; } $session = $this->sessionManager->createWithId($sessionId); @@ -632,13 +654,45 @@ public function handleFiberYield(mixed $yieldedValue, ?Uuid $sessionId): void } try { - match (true) { - $yieldedValue instanceof NotificationSuspension => $this->sendNotification($yieldedValue->notification, $session), - $yieldedValue instanceof RequestSuspension => $this->sendRequest($yieldedValue->request, $yieldedValue->timeout, $session), - }; + if ($yieldedValue instanceof RequestSuspension) { + return $this->sendRequest($yieldedValue->request, $yieldedValue->timeout, $session); + } + + $this->sendNotification($yieldedValue->notification, $session); } finally { $session->save(); } + + return null; + } + + /** + * @param TransportInterface $transport + */ + private function trackAwaitedRequest(TransportInterface $transport, ?int $requestId): void + { + if (null === $requestId) { + unset($this->awaitedRequestIds[$transport]); + + return; + } + + $this->awaitedRequestIds[$transport] = $requestId; + } + + /** + * @param TransportInterface|null $transport + * + * @return array + */ + private function getAwaitedPendingRequests(?TransportInterface $transport, Uuid $sessionId): array + { + $requestId = null !== $transport ? $this->awaitedRequestIds[$transport] ?? null : null; + if (null === $requestId) { + return []; + } + + return array_intersect_key($this->getPendingRequests($sessionId), [$requestId => true]); } /** diff --git a/tests/Unit/Fixtures/PollingLoopTransport.php b/tests/Unit/Fixtures/PollingLoopTransport.php new file mode 100644 index 00000000..34317a18 --- /dev/null +++ b/tests/Unit/Fixtures/PollingLoopTransport.php @@ -0,0 +1,42 @@ + + */ + public function getPendingRequestIds(): array + { + return array_keys($this->getPendingRequests($this->sessionId)); + } + + /** + * @param FiberSuspend $yielded + */ + public function yieldFromFiber(NotificationSuspension|RequestSuspension $yielded): void + { + $this->handleFiberYield($yielded, $this->sessionId); + } +} diff --git a/tests/Unit/Server/ProtocolTest.php b/tests/Unit/Server/ProtocolTest.php index 93773cd9..54dd81ad 100644 --- a/tests/Unit/Server/ProtocolTest.php +++ b/tests/Unit/Server/ProtocolTest.php @@ -35,6 +35,7 @@ use Mcp\Server\Suspension\NotificationSuspension; use Mcp\Server\Suspension\RequestSuspension; use Mcp\Server\Transport\TransportInterface; +use Mcp\Tests\Unit\Fixtures\PollingLoopTransport; use Mcp\Tests\Unit\Fixtures\ThrowingRequest; use PHPUnit\Framework\Attributes\DataProvider; use PHPUnit\Framework\Attributes\TestDox; @@ -859,6 +860,94 @@ public function testOutboundRequestFailureIsAnsweredUnderInboundRequestId(): voi $this->assertSame(Error::INTERNAL_ERROR, $errors[0]['error']['code']); } + #[TestDox('Concurrent streams on one session each poll only the client request their own fiber sent')] + public function testConcurrentStreamsPollOnlyTheirOwnPendingRequest(): void + { + [$protocol, $sessionId, $firstStream, $secondStream] = $this->startTwoStreamsWaitingOnClient(); + + $protocol->processInput($secondStream, '{"jsonrpc": "2.0", "id": 1001, "result": {}}', $sessionId); + + $this->assertSame([1000], $firstStream->getPendingRequestIds()); + $this->assertSame([1001], $secondStream->getPendingRequestIds()); + } + + #[TestDox('A client request a fiber sends after resuming is polled only by its own stream')] + public function testRequestYieldedOnResumeIsPolledOnlyByItsOwnStream(): void + { + [, $sessionId, $firstStream, $secondStream] = $this->startTwoStreamsWaitingOnClient(); + + $firstStream->yieldFromFiber(new RequestSuspension(new PingRequest(), $sessionId->toRfc4122(), 5)); + + $this->assertSame([1002], $firstStream->getPendingRequestIds()); + $this->assertSame([1001], $secondStream->getPendingRequestIds()); + } + + #[TestDox('A stream whose fiber resumes and sends a notification no longer polls the request it was waiting on')] + public function testNotificationYieldedOnResumeClearsTheAwaitedRequest(): void + { + [, $sessionId, $firstStream, $secondStream] = $this->startTwoStreamsWaitingOnClient(); + + $firstStream->yieldFromFiber(new NotificationSuspension(new LoggingMessageNotification(LoggingLevel::Info, 'hello'), $sessionId->toRfc4122())); + + $this->assertSame([], $firstStream->getPendingRequestIds()); + $this->assertSame([1001], $secondStream->getPendingRequestIds()); + } + + #[TestDox('A stream whose fiber first suspends on a notification polls none of the session\'s pending requests')] + public function testStreamSuspendedOnNotificationPollsNoPendingRequest(): void + { + [$protocol, $sessionId, , $secondStream] = $this->startTwoStreamsWaitingOnClient(); + + $thirdStream = new PollingLoopTransport(); + $protocol->connect($thirdStream); + $protocol->processInput($thirdStream, '{"jsonrpc": "2.0", "id": 3, "method": "ping"}', $sessionId); + + $this->assertSame([], $thirdStream->getPendingRequestIds()); + $this->assertSame([1001], $secondStream->getPendingRequestIds()); + } + + /** + * Two tool calls on one session, each suspended on a request to the client, as with elicitation. + * A further call with ID 3 suspends on a notification instead. + * + * @return array{Protocol, Uuid, PollingLoopTransport, PollingLoopTransport} + */ + private function startTwoStreamsWaitingOnClient(): array + { + $handler = $this->createMock(RequestHandlerInterface::class); + $handler->method('supports')->willReturn(true); + $handler->method('handle')->willReturnCallback(static function (Request $request, SessionInterface $session): Response { + $sessionId = $session->getId()->toRfc4122(); + \Fiber::suspend(3 === $request->getId() + ? new NotificationSuspension(new LoggingMessageNotification(LoggingLevel::Info, 'hello'), $sessionId) + : new RequestSuspension(new PingRequest(), $sessionId, 5)); + + return new Response(1, []); + }); + + $sessionManager = new SessionManager(new InMemorySessionStore()); + $protocol = new Protocol( + requestHandlers: [$handler], + notificationHandlers: [], + messageFactory: MessageFactory::make(), + sessionManager: $sessionManager, + ); + + $session = $sessionManager->create(); + $session->save(); + $sessionId = $session->getId(); + + $firstStream = new PollingLoopTransport(); + $protocol->connect($firstStream); + $protocol->processInput($firstStream, '{"jsonrpc": "2.0", "id": 1, "method": "ping"}', $sessionId); + + $secondStream = new PollingLoopTransport(); + $protocol->connect($secondStream); + $protocol->processInput($secondStream, '{"jsonrpc": "2.0", "id": 2, "method": "ping"}', $sessionId); + + return [$protocol, $sessionId, $firstStream, $secondStream]; + } + #[TestDox('Notification handler exceptions are caught and logged')] public function testNotificationHandlerExceptionsAreCaught(): void {