1212namespace Mcp \Server \Transport ;
1313
1414use Http \Discovery \Psr17FactoryDiscovery ;
15+ use Mcp \Exception \InvalidArgumentException ;
1516use Mcp \Schema \JsonRpc \Error ;
17+ use Mcp \Server \Transport \Http \MiddlewareRequestHandler ;
1618use Psr \Http \Message \ResponseFactoryInterface ;
1719use Psr \Http \Message \ResponseInterface ;
1820use Psr \Http \Message \ServerRequestInterface ;
1921use Psr \Http \Message \StreamFactoryInterface ;
22+ use Psr \Http \Server \MiddlewareInterface ;
2023use Psr \Log \LoggerInterface ;
2124use Symfony \Component \Uid \Uuid ;
2225
@@ -36,19 +39,22 @@ class StreamableHttpTransport extends BaseTransport
3639 /** @var array<string, string> */
3740 private array $ corsHeaders ;
3841
42+ /** @var list<MiddlewareInterface> */
43+ private array $ middleware = [];
44+
3945 /**
40- * @param array<string, string> $corsHeaders
46+ * @param array<string, string> $corsHeaders
47+ * @param iterable<MiddlewareInterface> $middleware
4148 */
4249 public function __construct (
43- private readonly ServerRequestInterface $ request ,
50+ private ServerRequestInterface $ request ,
4451 ?ResponseFactoryInterface $ responseFactory = null ,
4552 ?StreamFactoryInterface $ streamFactory = null ,
4653 array $ corsHeaders = [],
4754 ?LoggerInterface $ logger = null ,
55+ iterable $ middleware = [],
4856 ) {
4957 parent ::__construct ($ logger );
50- $ sessionIdString = $ this ->request ->getHeaderLine ('Mcp-Session-Id ' );
51- $ this ->sessionId = $ sessionIdString ? Uuid::fromString ($ sessionIdString ) : null ;
5258
5359 $ this ->responseFactory = $ responseFactory ?? Psr17FactoryDiscovery::findResponseFactory ();
5460 $ this ->streamFactory = $ streamFactory ?? Psr17FactoryDiscovery::findStreamFactory ();
@@ -59,6 +65,13 @@ public function __construct(
5965 'Access-Control-Allow-Headers ' => 'Content-Type, Mcp-Session-Id, Mcp-Protocol-Version, Last-Event-ID, Authorization, Accept ' ,
6066 'Access-Control-Expose-Headers ' => 'Mcp-Session-Id ' ,
6167 ], $ corsHeaders );
68+
69+ foreach ($ middleware as $ m ) {
70+ if (!$ m instanceof MiddlewareInterface) {
71+ throw new InvalidArgumentException ('Streamable HTTP middleware must implement Psr \\Http \\Server \\MiddlewareInterface. ' );
72+ }
73+ $ this ->middleware [] = $ m ;
74+ }
6275 }
6376
6477 public function send (string $ data , array $ context ): void
@@ -69,17 +82,17 @@ public function send(string $data, array $context): void
6982
7083 public function listen (): ResponseInterface
7184 {
72- return match ( $ this -> request -> getMethod ()) {
73- ' OPTIONS ' => $ this ->handleOptionsRequest () ,
74- ' POST ' => $ this -> handlePostRequest ( ),
75- ' DELETE ' => $ this -> handleDeleteRequest (),
76- default => $ this -> createErrorResponse (Error:: forInvalidRequest ( ' Method Not Allowed ' ), 405 ),
77- } ;
85+ $ handler = new MiddlewareRequestHandler (
86+ $ this ->middleware ,
87+ \Closure:: fromCallable ([ $ this , ' handleRequest ' ] ),
88+ );
89+
90+ return $ this -> withCorsHeaders ( $ handler -> handle ( $ this -> request )) ;
7891 }
7992
8093 protected function handleOptionsRequest (): ResponseInterface
8194 {
82- return $ this ->withCorsHeaders ( $ this -> responseFactory ->createResponse (204 ) );
95+ return $ this ->responseFactory ->createResponse (204 );
8396 }
8497
8598 protected function handlePostRequest (): ResponseInterface
@@ -92,7 +105,7 @@ protected function handlePostRequest(): ResponseInterface
92105 ->withHeader ('Content-Type ' , 'application/json ' )
93106 ->withBody ($ this ->streamFactory ->createStream ($ this ->immediateResponse ));
94107
95- return $ this -> withCorsHeaders ( $ response) ;
108+ return $ response ;
96109 }
97110
98111 if (null !== $ this ->sessionFiber ) {
@@ -112,15 +125,15 @@ protected function handleDeleteRequest(): ResponseInterface
112125
113126 $ this ->handleSessionEnd ($ this ->sessionId );
114127
115- return $ this ->withCorsHeaders ( $ this -> responseFactory ->createResponse (200 ) );
128+ return $ this ->responseFactory ->createResponse (200 );
116129 }
117130
118131 protected function createJsonResponse (): ResponseInterface
119132 {
120133 $ outgoingMessages = $ this ->getOutgoingMessages ($ this ->sessionId );
121134
122135 if (empty ($ outgoingMessages )) {
123- return $ this ->withCorsHeaders ( $ this -> responseFactory ->createResponse (202 ) );
136+ return $ this ->responseFactory ->createResponse (202 );
124137 }
125138
126139 $ messages = array_column ($ outgoingMessages , 'message ' );
@@ -134,7 +147,7 @@ protected function createJsonResponse(): ResponseInterface
134147 $ response = $ response ->withHeader ('Mcp-Session-Id ' , $ this ->sessionId ->toRfc4122 ());
135148 }
136149
137- return $ this -> withCorsHeaders ( $ response) ;
150+ return $ response ;
138151 }
139152
140153 protected function createStreamedResponse (): ResponseInterface
@@ -201,7 +214,7 @@ protected function createStreamedResponse(): ResponseInterface
201214 $ response = $ response ->withHeader ('Mcp-Session-Id ' , $ this ->sessionId ->toRfc4122 ());
202215 }
203216
204- return $ this -> withCorsHeaders ( $ response) ;
217+ return $ response ;
205218 }
206219
207220 protected function handleFiberTermination (): void
@@ -246,15 +259,31 @@ protected function createErrorResponse(Error $jsonRpcError, int $statusCode): Re
246259 $ response = $ response ->withHeader ('Allow ' , 'POST, DELETE, OPTIONS ' );
247260 }
248261
249- return $ this -> withCorsHeaders ( $ response) ;
262+ return $ response ;
250263 }
251264
252265 protected function withCorsHeaders (ResponseInterface $ response ): ResponseInterface
253266 {
254267 foreach ($ this ->corsHeaders as $ name => $ value ) {
255- $ response = $ response ->withHeader ($ name , $ value );
268+ if (!$ response ->hasHeader ($ name )) {
269+ $ response = $ response ->withHeader ($ name , $ value );
270+ }
256271 }
257272
258273 return $ response ;
259274 }
275+
276+ private function handleRequest (ServerRequestInterface $ request ): ResponseInterface
277+ {
278+ $ this ->request = $ request ;
279+ $ sessionIdString = $ request ->getHeaderLine ('Mcp-Session-Id ' );
280+ $ this ->sessionId = $ sessionIdString ? Uuid::fromString ($ sessionIdString ) : null ;
281+
282+ return match ($ request ->getMethod ()) {
283+ 'OPTIONS ' => $ this ->handleOptionsRequest (),
284+ 'POST ' => $ this ->handlePostRequest (),
285+ 'DELETE ' => $ this ->handleDeleteRequest (),
286+ default => $ this ->createErrorResponse (Error::forInvalidRequest ('Method Not Allowed ' ), 405 ),
287+ };
288+ }
260289}
0 commit comments