From e14ef606392293bae34acf0d72d2a1c35575218f Mon Sep 17 00:00:00 2001 From: King Star Date: Sun, 26 Jul 2026 22:20:19 +0800 Subject: [PATCH] fix: isolate stateless HTTP request exchanges Signed-off-by: King Star --- .../HTTPServer/HTTPServerTypes.swift | 8 + .../StatelessHTTPServerTransport.swift | 167 +++++++++- Sources/MCP/Server/Server.swift | 15 +- Tests/MCPTests/HTTPServerTransportTests.swift | 315 +++++++++++++++++- 4 files changed, 477 insertions(+), 28 deletions(-) diff --git a/Sources/MCP/Base/Transports/HTTPServer/HTTPServerTypes.swift b/Sources/MCP/Base/Transports/HTTPServer/HTTPServerTypes.swift index bd56191c..7f84af90 100644 --- a/Sources/MCP/Base/Transports/HTTPServer/HTTPServerTypes.swift +++ b/Sources/MCP/Base/Transports/HTTPServer/HTTPServerTypes.swift @@ -225,6 +225,14 @@ public protocol HTTPContextProviding: Sendable { func httpRequestContext(for id: ID) async -> HTTPRequest? } +/// A transport-internal mapping between private routed ids and client wire ids. +package protocol RoutedRequestIDProviding: Sendable { + func originalRequestID(for routedID: ID) async -> ID? + + /// Returns `nil` when the wire id has no active exchange or is ambiguous. + func routedRequestID(for originalID: ID) async -> ID? +} + // MARK: - JSON-RPC Message Classification /// Classifies a raw JSON-RPC message for routing purposes. diff --git a/Sources/MCP/Base/Transports/HTTPServer/StatelessHTTPServerTransport.swift b/Sources/MCP/Base/Transports/HTTPServer/StatelessHTTPServerTransport.swift index abe9576c..6457e72e 100644 --- a/Sources/MCP/Base/Transports/HTTPServer/StatelessHTTPServerTransport.swift +++ b/Sources/MCP/Base/Transports/HTTPServer/StatelessHTTPServerTransport.swift @@ -29,7 +29,9 @@ import Logging /// - Session management is handled externally or not needed /// /// For full streaming and session support, use ``StatefulHTTPServerTransport`` instead. -public actor StatelessHTTPServerTransport: Transport, HTTPContextProviding { +public actor StatelessHTTPServerTransport: + Transport, HTTPContextProviding, RoutedRequestIDProviding +{ public nonisolated let logger: Logger // MARK: - Dependencies @@ -48,15 +50,23 @@ public actor StatelessHTTPServerTransport: Transport, HTTPContextProviding { // MARK: - Response waiters - /// Maps request ID → continuation waiting for the server's response. - /// When the server calls `send()` with a response, the matching continuation is resumed. - private var responseWaiters: [String: CheckedContinuation] = [:] + private struct ResponseWaiter { + let originalID: ID + let continuation: CheckedContinuation + } + + /// Maps a transport-private exchange ID to the matching HTTP response waiter. + private var responseWaiters: [String: ResponseWaiter] = [:] /// Maps request ID → originating HTTP request, surfaced to handlers via /// ``Server/currentHTTPContext``. Entries live only while a JSON-RPC request /// is in flight. private var httpRequestContexts: [String: HTTPRequest] = [:] + /// Preserves the existing raw-ID context lookup for callers outside `Server`. + /// Multiple entries are possible because JSON-RPC IDs are scoped to each client. + private var exchangeIDsByRequestID: [ID: [String]] = [:] + // MARK: - Init /// Creates a new stateless HTTP server transport. @@ -117,14 +127,20 @@ public actor StatelessHTTPServerTransport: Transport, HTTPContextProviding { switch kind { case .response(let id): - guard let continuation = responseWaiters.removeValue(forKey: id) else { + guard let waiter = responseWaiters.removeValue(forKey: id) else { logger.debug( "No waiter for response, may have timed out", metadata: ["requestID": "\(id)"] ) return } - continuation.resume(returning: data) + do { + let response = try restoringResponseID(in: data, to: waiter.originalID) + waiter.continuation.resume(returning: response) + } catch { + waiter.continuation.resume(throwing: error) + throw error + } case .notification(let method): logger.debug( @@ -221,32 +237,146 @@ public actor StatelessHTTPServerTransport: Transport, HTTPContextProviding { requestID: String, request: HTTPRequest ) async -> HTTPResponse { - httpRequestContexts[requestID] = request + let exchangeID = makeExchangeID(excluding: requestID) + let routedBody: Data + let originalID: ID + do { + (routedBody, originalID) = try routingRequest(body, as: exchangeID) + } catch { + return .error( + statusCode: 400, + .parseError("Invalid JSON-RPC request id") + ) + } + + registerHTTPContext(request, exchangeID: exchangeID, requestID: originalID) // Yield the incoming message to the server - incomingContinuation.yield(body) + incomingContinuation.yield(routedBody) // Wait for the server to process and send a response let responseData: Data do { responseData = try await withCheckedThrowingContinuation { continuation in - responseWaiters[requestID] = continuation + responseWaiters[exchangeID] = ResponseWaiter( + originalID: originalID, + continuation: continuation + ) } } catch { - httpRequestContexts.removeValue(forKey: requestID) + removeHTTPContext(exchangeID: exchangeID, requestID: originalID) return .error( statusCode: 500, .internalError("Error processing request: \(error.localizedDescription)") ) } - httpRequestContexts.removeValue(forKey: requestID) + removeHTTPContext(exchangeID: exchangeID, requestID: originalID) return .data(responseData, headers: [HTTPHeaderName.contentType: ContentType.json]) } + private func makeExchangeID(excluding requestID: String) -> String { + var exchangeID: String + repeat { + exchangeID = UUID().uuidString + } while exchangeID == requestID + || responseWaiters[exchangeID] != nil + || httpRequestContexts[exchangeID] != nil + || exchangeIDsByRequestID[.string(exchangeID)] != nil + return exchangeID + } + + private func routingRequest(_ body: Data, as exchangeID: String) throws -> (Data, ID) { + guard var json = try JSONSerialization.jsonObject(with: body) as? [String: Any] else { + throw MCPError.parseError("Invalid JSON-RPC request") + } + + guard let originalID = requestID(from: json["id"]) else { + throw MCPError.parseError("Invalid JSON-RPC request id") + } + + json["id"] = exchangeID + let routedBody = try JSONSerialization.data( + withJSONObject: json, + options: [.sortedKeys, .withoutEscapingSlashes] + ) + return (routedBody, originalID) + } + + private func requestID(from value: Any?) -> ID? { + if let stringID = value as? String { + return .string(stringID) + } + if let numberID = value as? Int { + return .number(numberID) + } + return nil + } + + private func restoringResponseID(in data: Data, to originalID: ID) throws -> Data { + guard var json = try JSONSerialization.jsonObject(with: data) as? [String: Any] else { + throw MCPError.parseError("Invalid JSON-RPC response") + } + + switch originalID { + case .string(let value): + json["id"] = value + case .number(let value): + json["id"] = value + } + + return try JSONSerialization.data( + withJSONObject: json, + options: [.sortedKeys, .withoutEscapingSlashes] + ) + } + + private func registerHTTPContext( + _ request: HTTPRequest, + exchangeID: String, + requestID: ID + ) { + httpRequestContexts[exchangeID] = request + exchangeIDsByRequestID[requestID, default: []].append(exchangeID) + } + + private func removeHTTPContext(exchangeID: String, requestID: ID) { + httpRequestContexts.removeValue(forKey: exchangeID) + guard var exchangeIDs = exchangeIDsByRequestID[requestID] else { return } + exchangeIDs.removeAll { $0 == exchangeID } + if exchangeIDs.isEmpty { + exchangeIDsByRequestID.removeValue(forKey: requestID) + } else { + exchangeIDsByRequestID[requestID] = exchangeIDs + } + } + // MARK: - HTTPContextProviding public func httpRequestContext(for id: ID) -> HTTPRequest? { - httpRequestContexts[id.description] + if case .string(let exchangeID) = id, + let request = httpRequestContexts[exchangeID] + { + return request + } + guard let exchangeID = exchangeIDsByRequestID[id]?.last else { + return nil + } + return httpRequestContexts[exchangeID] + } + + package func originalRequestID(for routedID: ID) -> ID? { + guard case .string(let exchangeID) = routedID else { return nil } + return responseWaiters[exchangeID]?.originalID + } + + package func routedRequestID(for originalID: ID) -> ID? { + guard let exchangeIDs = exchangeIDsByRequestID[originalID], + exchangeIDs.count == 1, + let exchangeID = exchangeIDs.first + else { + return nil + } + return .string(exchangeID) } // MARK: - Termination @@ -258,12 +388,19 @@ public actor StatelessHTTPServerTransport: Transport, HTTPContextProviding { logger.debug("Stateless HTTP server transport terminated") // Cancel all waiting continuations - for (id, continuation) in responseWaiters { - continuation.resume(throwing: MCPError.connectionClosed) - logger.debug("Cancelled waiter for request", metadata: ["requestID": "\(id)"]) + for (exchangeID, waiter) in responseWaiters { + waiter.continuation.resume(throwing: MCPError.connectionClosed) + logger.debug( + "Cancelled waiter for request", + metadata: [ + "exchangeID": "\(exchangeID)", + "requestID": "\(waiter.originalID)", + ] + ) } responseWaiters.removeAll() httpRequestContexts.removeAll() + exchangeIDsByRequestID.removeAll() // Close incoming stream incomingContinuation.finish() diff --git a/Sources/MCP/Server/Server.swift b/Sources/MCP/Server/Server.swift index 060b51f8..07057eac 100644 --- a/Sources/MCP/Server/Server.swift +++ b/Sources/MCP/Server/Server.swift @@ -783,7 +783,9 @@ public actor Server { // that don't carry HTTP context (stdio, in-memory) don't conform. let httpContext = await (connection as? any HTTPContextProviding)? .httpRequestContext(for: request.id) - let handlerContext = HandlerContext(id: request.id, httpContext: httpContext) + let handlerRequestID = await (connection as? any RoutedRequestIDProviding)? + .originalRequestID(for: request.id) ?? request.id + let handlerContext = HandlerContext(id: handlerRequestID, httpContext: httpContext) // Create a task to handle the request with cancellation support. // Set currentHandlerContext as a task local so handlers see it. @@ -998,7 +1000,9 @@ public actor Server { } // Cancel the pending request task if it exists and remove from tracking - if let task = await self.removePendingRequest(id: requestId) { + if let pendingRequestID = await self.routedRequestID(for: requestId), + let task = await self.removePendingRequest(id: pendingRequestID) + { task.cancel() await self.logger?.debug( "Cancelled request", @@ -1015,6 +1019,13 @@ public actor Server { } } + private func routedRequestID(for originalID: ID) async -> ID? { + guard let provider = connection as? any RoutedRequestIDProviding else { + return originalID + } + return await provider.routedRequestID(for: originalID) + } + /// Cancel a request by sending a CancelledNotification to the client. /// /// This is used when the server needs to cancel an in-progress request it made to the client diff --git a/Tests/MCPTests/HTTPServerTransportTests.swift b/Tests/MCPTests/HTTPServerTransportTests.swift index 8e1c89b4..b97516c1 100644 --- a/Tests/MCPTests/HTTPServerTransportTests.swift +++ b/Tests/MCPTests/HTTPServerTransportTests.swift @@ -29,6 +29,15 @@ private func makeNotificationBody(method: String = "notifications/initialized") return try! JSONSerialization.data(withJSONObject: json) } +private func makeCancelledNotificationBody(requestID: Any) -> Data { + let json: [String: Any] = [ + "jsonrpc": "2.0", + "method": "notifications/cancelled", + "params": ["requestId": requestID], + ] + return try! JSONSerialization.data(withJSONObject: json) +} + private func makeRequestBody(id: String = "2", method: String = "tools/list") -> Data { let json: [String: Any] = [ "jsonrpc": "2.0", @@ -158,6 +167,25 @@ private func drainSSEStream( return await collector.getChunks() } +private func nextValue( + from stream: AsyncStream, + timeout: Duration = .seconds(1) +) async -> T? { + await withTaskGroup(of: T?.self) { group in + group.addTask { + var iterator = stream.makeAsyncIterator() + return await iterator.next() + } + group.addTask { + try? await Task.sleep(for: timeout) + return nil + } + let value = await group.next() ?? nil + group.cancelAll() + return value + } +} + /// Initializes a stateful transport session and returns the session ID. /// Spawns a background task to consume the receive stream and send the init response. private func initializeSession( @@ -895,7 +923,17 @@ struct StatelessHTTPServerTransportTests { try await transport.connect() let requestBody = makeRequestBody(id: "42", method: "tools/list") - let responseBody = makeResponseBody(id: "42") + + let responseRouter = Task { + let stream = await transport.receive() + var iterator = stream.makeAsyncIterator() + let body = try #require(try await iterator.next()) + let json = try #require( + try JSONSerialization.jsonObject(with: body) as? [String: Any] + ) + let exchangeID = try #require(json["id"] as? String) + try await transport.send(makeResponseBody(id: exchangeID)) + } // handleRequest blocks waiting for response let handleTask = Task { @@ -904,16 +942,183 @@ struct StatelessHTTPServerTransportTests { ) } - // Give handleRequest time to register the waiter - try await Task.sleep(for: .milliseconds(50)) - - // Consume the request from receive and send the response - try await transport.send(responseBody) - + try await responseRouter.value let httpResponse = await handleTask.value #expect(httpResponse.statusCode == 200) - #expect(httpResponse.bodyData == responseBody) #expect(httpResponse.headers[HTTPHeaderName.contentType] == ContentType.json) + let responseData = try #require(httpResponse.bodyData) + let responseJSON = try #require( + try JSONSerialization.jsonObject(with: responseData) as? [String: Any] + ) + #expect(responseJSON["id"] as? String == "42") + let result = try #require(responseJSON["result"] as? [String: Any]) + let tools = try #require(result["tools"] as? [Any]) + #expect(tools.isEmpty) + + await transport.disconnect() + } + + @Test("Concurrent requests sharing a wire ID route to their own HTTP exchange") + func testConcurrentRequestsWithSameWireID() async throws { + let transport = makeStatelessTransport() + try await transport.connect() + + func request(marker: String) -> HTTPRequest { + let body = try! JSONSerialization.data(withJSONObject: [ + "jsonrpc": "2.0", + "id": 7, + "method": "tools/list", + "params": ["marker": marker], + ]) + return HTTPRequest( + method: "POST", + headers: [ + "Content-Type": "application/json", + "Accept": "application/json", + "Authorization": "Bearer \(marker)", + ], + body: body, + path: "/mcp/\(marker)" + ) + } + + let responseRouter = Task { + let stream = await transport.receive() + var iterator = stream.makeAsyncIterator() + var responses: [Data] = [] + + for _ in 0..<2 { + let body = try #require(try await iterator.next()) + let json = try #require( + try JSONSerialization.jsonObject(with: body) as? [String: Any] + ) + let exchangeID = try #require(json["id"] as? String) + let params = try #require(json["params"] as? [String: Any]) + let marker = try #require(params["marker"] as? String) + let context = await transport.httpRequestContext(for: .string(exchangeID)) + #expect(context?.header("Authorization") == "Bearer \(marker)") + #expect(context?.path == "/mcp/\(marker)") + responses.append( + try JSONSerialization.data(withJSONObject: [ + "jsonrpc": "2.0", + "id": exchangeID, + "result": ["marker": marker], + ]) + ) + } + + #expect(await transport.routedRequestID(for: .number(7)) == nil) + + for response in responses { + try await transport.send(response) + } + } + + let firstTask = Task { await transport.handleRequest(request(marker: "first")) } + let secondTask = Task { await transport.handleRequest(request(marker: "second")) } + + try await responseRouter.value + let first = await firstTask.value + let second = await secondTask.value + + for (response, expectedMarker) in [(first, "first"), (second, "second")] { + #expect(response.statusCode == 200) + let body = try #require(response.bodyData) + let json = try #require( + try JSONSerialization.jsonObject(with: body) as? [String: Any] + ) + #expect(json["id"] as? Int == 7) + let result = try #require(json["result"] as? [String: Any]) + #expect(result["marker"] as? String == expectedMarker) + } + + await transport.disconnect() + } + + @Test("String and number wire IDs route to distinct HTTP exchanges") + func testConcurrentStringAndNumberWireIDs() async throws { + let transport = makeStatelessTransport() + try await transport.connect() + + func request(id: Any, marker: String) -> HTTPRequest { + let body = try! JSONSerialization.data(withJSONObject: [ + "jsonrpc": "2.0", + "id": id, + "method": "tools/list", + "params": ["marker": marker], + ]) + return HTTPRequest( + method: "POST", + headers: [ + "Content-Type": "application/json", + "Accept": "application/json", + "Authorization": "Bearer \(marker)", + ], + body: body, + path: "/mcp/\(marker)" + ) + } + + let responseRouter = Task { + let stream = await transport.receive() + var iterator = stream.makeAsyncIterator() + var responses: [Data] = [] + + for _ in 0..<2 { + let body = try #require(try await iterator.next()) + let json = try #require( + try JSONSerialization.jsonObject(with: body) as? [String: Any] + ) + let exchangeID = try #require(json["id"] as? String) + let params = try #require(json["params"] as? [String: Any]) + let marker = try #require(params["marker"] as? String) + responses.append( + try JSONSerialization.data(withJSONObject: [ + "jsonrpc": "2.0", + "id": exchangeID, + "result": ["marker": marker], + ]) + ) + } + + let stringContext = await transport.httpRequestContext(for: .string("1")) + #expect(stringContext?.header("Authorization") == "Bearer string") + #expect(stringContext?.path == "/mcp/string") + let numberContext = await transport.httpRequestContext(for: .number(1)) + #expect(numberContext?.header("Authorization") == "Bearer number") + #expect(numberContext?.path == "/mcp/number") + + for response in responses { + try await transport.send(response) + } + } + + let stringTask = Task { + await transport.handleRequest(request(id: "1", marker: "string")) + } + let numberTask = Task { + await transport.handleRequest(request(id: 1, marker: "number")) + } + + try await responseRouter.value + let stringResponse = await stringTask.value + let numberResponse = await numberTask.value + + let stringData = try #require(stringResponse.bodyData) + let stringJSON = try #require( + try JSONSerialization.jsonObject(with: stringData) as? [String: Any] + ) + #expect(stringJSON["id"] as? String == "1") + let stringResult = try #require(stringJSON["result"] as? [String: Any]) + #expect(stringResult["marker"] as? String == "string") + + let numberData = try #require(numberResponse.bodyData) + let numberJSON = try #require( + try JSONSerialization.jsonObject(with: numberData) as? [String: Any] + ) + #expect(numberJSON["id"] as? Int == 1) + let numberResult = try #require(numberJSON["result"] as? [String: Any]) + #expect(numberResult["marker"] as? String == "number") await transport.disconnect() } @@ -957,6 +1162,45 @@ struct StatelessHTTPServerTransportTests { await transport.disconnect() } + @Test("Cancellation lookup maps the wire ID to the active HTTP exchange") + func testCancellationLookupMapsToExchangeID() async throws { + let transport = makeStatelessTransport() + try await transport.connect() + + let stream = await transport.receive() + var iterator = stream.makeAsyncIterator() + let handleTask = Task { + await transport.handleRequest( + makeStatelessPOSTRequest(body: makeRequestBody(id: "cancel-me")) + ) + } + + let routedRequest = try #require(try await iterator.next()) + let requestJSON = try #require( + try JSONSerialization.jsonObject(with: routedRequest) as? [String: Any] + ) + let exchangeID = try #require(requestJSON["id"] as? String) + + #expect( + await transport.routedRequestID(for: .string("cancel-me")) + == .string(exchangeID) + ) + + let cancellationBody = makeCancelledNotificationBody(requestID: "cancel-me") + let cancelResponse = await transport.handleRequest( + makeStatelessPOSTRequest( + body: cancellationBody + ) + ) + #expect(cancelResponse.statusCode == 202) + + let routedCancellation = try #require(try await iterator.next()) + #expect(routedCancellation == cancellationBody) + + await transport.disconnect() + _ = await handleTask.value + } + // MARK: - Unsupported Methods @Test("GET returns 405 Method Not Allowed") @@ -1146,11 +1390,17 @@ struct ServerHandlerContextTests { path: "/mcp" ) + let stream = await transport.receive() + var iterator = stream.makeAsyncIterator() + // handleRequest blocks until a response is sent — run it concurrently. let handleTask = Task { await transport.handleRequest(httpRequest) } - // Give the handler time to register its waiter (and store the context). - try await Task.sleep(for: .milliseconds(50)) + let routedBody = try #require(try await iterator.next()) + let routedJSON = try #require( + try JSONSerialization.jsonObject(with: routedBody) as? [String: Any] + ) + let exchangeID = try #require(routedJSON["id"] as? String) let inFlight = await transport.httpRequestContext(for: .string("99")) #expect(inFlight != nil) @@ -1158,7 +1408,7 @@ struct ServerHandlerContextTests { #expect(inFlight?.path == "/mcp") // Send the response to unblock the waiter. - try await transport.send(makeResponseBody(id: "99")) + try await transport.send(makeResponseBody(id: exchangeID)) _ = await handleTask.value let afterResponse = await transport.httpRequestContext(for: .string("99")) @@ -1233,6 +1483,49 @@ struct ServerHandlerContextTests { #expect(await captured.id == .string("call-1")) } + @Test("Stateless cancellation finds the handler by its wire ID") + func testStatelessCancellationFindsHandlerByWireID() async throws { + let (entered, enteredContinuation) = AsyncStream.makeStream() + let (cancelled, cancelledContinuation) = AsyncStream.makeStream() + + let transport = makeStatelessTransport() + let server = Server(name: "TestServer", version: "1.0") + await server.withMethodHandler(CallTool.self) { _ in + enteredContinuation.yield(()) + do { + try await Task.sleep(for: .seconds(5)) + } catch is CancellationError { + cancelledContinuation.yield(()) + throw CancellationError() + } + return CallTool.Result(content: [.text(text: "late", annotations: nil, _meta: nil)]) + } + + try await server.start(transport: transport) + let requestBody = try JSONSerialization.data(withJSONObject: [ + "jsonrpc": "2.0", + "id": "cancel-handler", + "method": "tools/call", + "params": ["name": "slow-tool"] as [String: Any], + ]) + let handleTask = Task { + await transport.handleRequest(makeStatelessPOSTRequest(body: requestBody)) + } + + _ = try #require(await nextValue(from: entered)) + let cancelResponse = await transport.handleRequest( + makeStatelessPOSTRequest( + body: makeCancelledNotificationBody(requestID: "cancel-handler") + ) + ) + #expect(cancelResponse.statusCode == 202) + _ = try #require(await nextValue(from: cancelled)) + + await transport.disconnect() + _ = await handleTask.value + await server.stop() + } + @Test("Non-HTTP transport yields nil httpContext") func testNonHTTPTransportNilContext() async throws { actor Captured {