From 0a36671a18819f79aed96745e7581a692f9ab9c5 Mon Sep 17 00:00:00 2001 From: Kazuki Nakashima <65545348+lynnswap@users.noreply.github.com> Date: Sat, 22 Aug 2026 17:37:04 +0900 Subject: [PATCH 1/3] Own MCP response work shutdown --- .../CodexReviewMCPHTTPServer.swift | 477 ++++++++++---- .../MCPHTTPNetworkResourceOwner.swift | 610 ++++++++++++++---- .../CodexReviewMCPHTTPServerTests.swift | 399 ++++++++++++ 3 files changed, 1249 insertions(+), 237 deletions(-) diff --git a/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift b/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift index cadfbdf..5ed7d38 100644 --- a/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift +++ b/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift @@ -11,6 +11,7 @@ private let logger = Logger(subsystem: "CodexReviewKit", category: "mcp-http") private struct TrackedHTTPResponse { var response: HTTPResponse var streamCompletion: ActiveRequestCompletion? = nil + var responseSourceKind: MCPHTTPNetworkResourceOwner.ResponseSourceKind? = nil } package extension CodexReviewMCPHTTPServer { @@ -285,6 +286,9 @@ package actor CodexReviewMCPHTTPServer { private let startCompletionGate = MCPHTTPLifecycleCompletionGate() private let joinedStartCompletionGate = MCPHTTPLifecycleCompletionGate() private let stopCompletionGate = MCPHTTPLifecycleCompletionGate() + private let finiteSourceCompletionGate = MCPHTTPLifecycleCompletionGate() + private let writerCompletionGate = MCPHTTPLifecycleCompletionGate() + private let responseEndAcknowledgementGate = MCPHTTPLifecycleCompletionGate() private var eventLoopGroupShutdownCount = 0 private var nextListenerCloseFailureForTesting: LifecycleError.Failure? private var nextEventLoopGroupShutdownFailureForTesting: LifecycleError.Failure? @@ -407,7 +411,9 @@ package actor CodexReviewMCPHTTPServer { guard let connection = networkResources.admitConnection(channel) else { return channel.close(mode: .all) } - return channel.pipeline.configureHTTPServerPipeline().flatMap { + return channel.pipeline.configureHTTPServerPipeline( + withPipeliningAssistance: false + ).flatMap { channel.pipeline.addHandler(CodexReviewMCPHTTPHandler( server: self, connection: connection @@ -868,6 +874,90 @@ package actor CodexReviewMCPHTTPServer { eventLoopGroupShutdownCount } + package func networkSnapshotForTesting() -> MCPHTTPNetworkResourceOwner.Snapshot { + if let networkResources = currentNetworkResources() { + return networkResources.snapshot() + } + return .init( + revision: 0, + generationID: nextGenerationID, + phase: .closed, + connections: [] + ) + } + + package func nextNetworkSnapshotForTesting( + after revision: UInt64 + ) async -> MCPHTTPNetworkResourceOwner.Snapshot { + guard let networkResources = currentNetworkResources() else { + return networkSnapshotForTesting() + } + return await networkResources.nextSnapshot(after: revision) + } + + package func holdNextFiniteSourceCompletionForTesting() async { + await finiteSourceCompletionGate.holdNextCompletion() + } + + package func waitUntilFiniteSourceCompletionIsHeldForTesting() async { + await finiteSourceCompletionGate.waitUntilHolding() + } + + package func releaseFiniteSourceCompletionForTesting() async { + await finiteSourceCompletionGate.release() + } + + package func holdNextWriterCompletionForTesting() async { + await writerCompletionGate.holdNextCompletion() + } + + package func waitUntilWriterCompletionIsHeldForTesting() async { + await writerCompletionGate.waitUntilHolding() + } + + package func releaseWriterCompletionForTesting() async { + await writerCompletionGate.release() + } + + package func holdNextResponseEndAcknowledgementForTesting() async { + await responseEndAcknowledgementGate.holdNextCompletion() + } + + package func waitUntilResponseEndAcknowledgementIsHeldForTesting() async { + await responseEndAcknowledgementGate.waitUntilHolding() + } + + package func releaseResponseEndAcknowledgementForTesting() async { + await responseEndAcknowledgementGate.release() + } + + fileprivate var responseHeartbeatInterval: Duration? { + configuration.streamHeartbeatInterval + } + + fileprivate func waitAfterFiniteSourceCompletionForTesting() async { + await finiteSourceCompletionGate.waitIfNeeded() + } + + fileprivate func waitAfterWriterCompletionForTesting() async { + await writerCompletionGate.waitIfNeeded() + } + + fileprivate func waitAfterResponseEndAcknowledgementForTesting() async { + await responseEndAcknowledgementGate.waitIfNeeded() + } + + private func currentNetworkResources() -> MCPHTTPNetworkResourceOwner? { + switch lifecycleState { + case .starting(let operation): + operation.networkResources + case .running(let resources), .stopping(_, let resources?, _): + resources.networkResources + case .stopped, .stopping: + nil + } + } + package func handleHTTPRequest(_ request: HTTPRequest) async -> HTTPResponse { await handleTrackedHTTPRequest(request).response } @@ -880,7 +970,11 @@ package actor CodexReviewMCPHTTPServer { session.activeRequestCount += 1 sessions[sessionID] = session let response = await session.transport.handleRequest(request) - let (trackedResponse, didFinishRequest) = trackActiveRequest(response, sessionID: sessionID) + let (trackedResponse, didFinishRequest) = trackActiveRequest( + response, + sessionID: sessionID, + responseSourceKind: Self.responseSourceKind(for: request) + ) if didFinishRequest, request.method.uppercased() == "DELETE", trackedResponse.response.statusCode == 200 { await closeSession(sessionID) } @@ -933,7 +1027,11 @@ package actor CodexReviewMCPHTTPServer { ) let response = await transport.handleRequest(request) - let (trackedResponse, didFinishRequest) = trackActiveRequest(response, sessionID: sessionID) + let (trackedResponse, didFinishRequest) = trackActiveRequest( + response, + sessionID: sessionID, + responseSourceKind: Self.responseSourceKind(for: request) + ) if didFinishRequest, case .error = trackedResponse.response { sessions.removeValue(forKey: sessionID) await transport.disconnect() @@ -961,39 +1059,22 @@ package actor CodexReviewMCPHTTPServer { private func trackActiveRequest( _ response: HTTPResponse, - sessionID: String + sessionID: String, + responseSourceKind: MCPHTTPNetworkResourceOwner.ResponseSourceKind ) -> (response: TrackedHTTPResponse, didFinishRequest: Bool) { switch response { - case .stream(let stream, let headers): + case .stream: let completion = ActiveRequestCompletion { Task { await self.finishActiveRequest(sessionID: sessionID) } } - let trackedStream = AsyncThrowingStream(bufferingPolicy: .unbounded) { continuation in - let heartbeatTask = makeStreamHeartbeatTask(continuation: continuation) - let task = Task { - defer { - heartbeatTask?.cancel() - completion.finish() - } - do { - for try await chunk in stream { - continuation.yield(chunk) - } - continuation.finish() - } catch { - continuation.finish(throwing: error) - } - } - continuation.onTermination = { _ in - heartbeatTask?.cancel() - task.cancel() - completion.finish() - } - } return ( - .init(response: .stream(trackedStream, headers: headers), streamCompletion: completion), + .init( + response: response, + streamCompletion: completion, + responseSourceKind: responseSourceKind + ), false ) @@ -1011,27 +1092,6 @@ package actor CodexReviewMCPHTTPServer { } } - private func makeStreamHeartbeatTask( - continuation: AsyncThrowingStream.Continuation - ) -> Task? { - guard let interval = configuration.streamHeartbeatInterval else { - return nil - } - return Task { - while Task.isCancelled == false { - do { - try await Task.sleep(for: interval) - } catch { - return - } - guard Task.isCancelled == false else { - return - } - continuation.yield(Data(": keep-alive\n\n".utf8)) - } - } - } - private func closeAllSessions() async { for sessionID in sessions.keys { await closeSession(sessionID) @@ -1091,6 +1151,19 @@ package actor CodexReviewMCPHTTPServer { return json["method"] as? String == "initialize" && json["id"] != nil } + private static func responseSourceKind( + for request: HTTPRequest + ) -> MCPHTTPNetworkResourceOwner.ResponseSourceKind { + guard let body = request.body, + let json = try? JSONSerialization.jsonObject(with: body) as? [String: Any], + json["method"] is String, + json["id"] != nil, + (json["id"] is NSNull) == false else { + return .open + } + return .finite + } + private func makeValidationPipeline() -> any HTTPRequestValidationPipeline { let resolvedPort = url.port ?? configuration.port let portPattern = resolvedPort > 0 ? String(resolvedPort) : "*" @@ -1211,12 +1284,16 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked var bodyBuffer: ByteBuffer } + private enum WriterEvent: Sendable { + case body(Data) + case heartbeat + case sourceFinished + case sourceFailed(String) + } + private let server: CodexReviewMCPHTTPServer private let connection: MCPHTTPNetworkResourceOwner.Connection private var requestState: RequestState? - private var activeStreamTask: Task? - private var activeStreamID: UUID? - private var activeStreamCompletion: ActiveRequestCompletion? init( server: CodexReviewMCPHTTPServer, @@ -1253,7 +1330,11 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked guard await admittedRequest.lease.waitUntilStartIsAllowed() else { return } - await handleRequest(state: state, context: context) + await handleRequest( + state: state, + operation: admittedRequest.operation, + context: context + ) } admittedRequest.lease.install(task) } @@ -1266,14 +1347,12 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked func channelInactive(context: ChannelHandlerContext) { connection.peerClosed() - finishActiveStream() context.fireChannelInactive() } func userInboundEventTriggered(context: ChannelHandlerContext, event: Any) { if case ChannelEvent.inputClosed = event { connection.peerClosed() - finishActiveStream() context.close(promise: nil) return } @@ -1282,28 +1361,21 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked func errorCaught(context: ChannelHandlerContext, error: any Error) { connection.transportFailed(error.localizedDescription) - finishActiveStream() context.close(promise: nil) } - private func finishActiveStream() { - activeStreamTask?.cancel() - activeStreamCompletion?.finish() - activeStreamTask = nil - activeStreamID = nil - activeStreamCompletion = nil - } - private func handleRequest( state: RequestState, + operation: MCPHTTPNetworkResourceOwner.RequestOperation, context: ChannelHandlerContext ) async { let head = state.head let path = head.uri.split(separator: "?").first.map(String.init) ?? head.uri let endpoint = await server.endpoint guard path == endpoint else { - await writeResponse( + await prepareAndQueueResponse( .init(response: .error(statusCode: 404, .invalidRequest("Not Found"))), + operation: operation, version: head.version, context: context ) @@ -1312,7 +1384,12 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked let request = makeHTTPRequest(from: state) let response = await server.handleTrackedHTTPRequest(request) - await writeResponse(response, version: head.version, context: context) + await prepareAndQueueResponse( + response, + operation: operation, + version: head.version, + context: context + ) } private func makeHTTPRequest(from state: RequestState) -> HTTPRequest { @@ -1343,90 +1420,232 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked ) } - private func writeResponse( + private func prepareAndQueueResponse( _ trackedResponse: TrackedHTTPResponse, + operation: MCPHTTPNetworkResourceOwner.RequestOperation, version: HTTPVersion, context: ChannelHandlerContext ) async { + let response = trackedResponse.response + let preparedResponse: HTTPResponse + + switch response { + case .stream(let source, let headers): + guard let kind = trackedResponse.responseSourceKind else { + trackedResponse.streamCompletion?.finish() + connection.transportFailed("Streaming response has no source lifetime contract.") + return + } + guard let sourceLease = operation.reserveResponseSource(kind) else { + trackedResponse.streamCompletion?.finish() + return + } + let bridge = AsyncThrowingStream.makeStream( + bufferingPolicy: .unbounded + ) + let sourceTask = Task { + let started = await sourceLease.waitUntilStartIsAllowed() + if started { + do { + for try await chunk in source { + try Task.checkCancellation() + bridge.continuation.yield(chunk) + } + if kind == .finite { + await self.server.waitAfterFiniteSourceCompletionForTesting() + } + bridge.continuation.finish() + } catch is CancellationError { + bridge.continuation.finish() + } catch { + bridge.continuation.finish(throwing: error) + self.connection.transportFailed(error.localizedDescription) + } + } else { + bridge.continuation.finish() + } + trackedResponse.streamCompletion?.finish() + sourceLease.acknowledgeCompletion() + } + sourceLease.install(sourceTask) + preparedResponse = .stream(bridge.stream, headers: headers) + + default: + guard operation.markResponseSourceNotRequired() else { + return + } + preparedResponse = response + } + + guard await connection.supplyResponse(for: operation), + let writerLease = operation.reserveWriter() else { + return + } nonisolated(unsafe) let context = context + let heartbeatInterval = await server.responseHeartbeatInterval + let writerTask = Task { + guard await writerLease.waitUntilStartIsAllowed() else { + writerLease.acknowledgeCompletion() + return + } + let result = await self.writeResponse( + preparedResponse, + version: version, + context: context, + heartbeatInterval: heartbeatInterval + ) + switch result { + case .responded: + operation.acknowledgeResponseEnd() + await self.server.waitAfterResponseEndAcknowledgementForTesting() + case .cancelled: + break + case .sourceFailed(let message): + logger.error("MCP SSE source failed: \(message, privacy: .public)") + self.connection.transportFailed(message) + case .transportFailed(let message): + logger.error("MCP HTTP response failed: \(message, privacy: .public)") + self.connection.transportFailed(message) + } + await self.server.waitAfterWriterCompletionForTesting() + writerLease.acknowledgeCompletion() + } + writerLease.install(writerTask) + } + + private enum WriterCompletion { + case responded + case cancelled + case sourceFailed(String) + case transportFailed(String) + } + + private func writeResponse( + _ response: HTTPResponse, + version: HTTPVersion, + context: ChannelHandlerContext, + heartbeatInterval: Duration? + ) async -> WriterCompletion { let eventLoop = context.eventLoop - let response = trackedResponse.response let status = HTTPResponseStatus(statusCode: response.statusCode) let headers = response.headers + var head = HTTPResponseHead(version: version, status: status) + for (name, value) in headers { + head.headers.add(name: name, value: value) + } switch response { case .stream(let stream, _): - let streamID = UUID() - let streamTask = Task { - var head = HTTPResponseHead(version: version, status: status) - for (name, value) in headers { - head.headers.add(name: name, value: value) - } - - var iterator = stream.makeAsyncIterator() - do { - try Task.checkCancellation() - try await writeResponsePart(.head(head), context: context, eventLoop: eventLoop) - while let chunk = try await iterator.next() { - try Task.checkCancellation() - try await writeResponseBody(chunk, context: context, eventLoop: eventLoop) + do { + try Task.checkCancellation() + try await writeResponsePart(.head(head), context: context, eventLoop: eventLoop) + let events = AsyncStream.makeStream(bufferingPolicy: .unbounded) + let sourceResult = await withTaskGroup( + of: Void.self, + returning: WriterCompletion.self + ) { group in + group.addTask { + do { + for try await chunk in stream { + try Task.checkCancellation() + events.continuation.yield(.body(chunk)) + } + events.continuation.yield(.sourceFinished) + } catch is CancellationError { + events.continuation.yield(.sourceFinished) + } catch { + events.continuation.yield(.sourceFailed(error.localizedDescription)) + } + } + if let heartbeatInterval { + group.addTask { + while Task.isCancelled == false { + do { + try await Task.sleep(for: heartbeatInterval) + } catch { + return + } + guard Task.isCancelled == false else { + return + } + events.continuation.yield(.heartbeat) + } + } } - } catch is CancellationError { - trackedResponse.streamCompletion?.finish() - return - } catch { - trackedResponse.streamCompletion?.finish() - logger.error("MCP SSE stream failed: \(error.localizedDescription, privacy: .public)") - } - guard Task.isCancelled == false else { - return - } - try? await writeResponsePart(.end(nil), context: context, eventLoop: eventLoop) - } - eventLoop.execute { - context.channel.closeFuture.whenComplete { _ in - trackedResponse.streamCompletion?.finish() - streamTask.cancel() - } - guard context.channel.isActive else { - trackedResponse.streamCompletion?.finish() - streamTask.cancel() - return + var result: WriterCompletion = .cancelled + eventLoopLoop: for await event in events.stream { + if Task.isCancelled { + result = .cancelled + break eventLoopLoop + } + switch event { + case .body(let data): + do { + try await writeResponseBody( + data, + context: context, + eventLoop: eventLoop + ) + } catch { + result = .transportFailed(error.localizedDescription) + break eventLoopLoop + } + case .heartbeat: + do { + try await writeResponseBody( + Data(": keep-alive\n\n".utf8), + context: context, + eventLoop: eventLoop + ) + } catch { + result = .transportFailed(error.localizedDescription) + break eventLoopLoop + } + case .sourceFinished: + result = .responded + break eventLoopLoop + case .sourceFailed(let message): + result = .sourceFailed(message) + break eventLoopLoop + } + } + group.cancelAll() + events.continuation.finish() + return result } - self.activeStreamTask?.cancel() - self.activeStreamCompletion?.finish() - self.activeStreamTask = streamTask - self.activeStreamID = streamID - self.activeStreamCompletion = trackedResponse.streamCompletion - context.read() - } - await streamTask.value - eventLoop.execute { - if self.activeStreamID == streamID { - self.activeStreamTask = nil - self.activeStreamID = nil - self.activeStreamCompletion = nil + switch sourceResult { + case .responded: + try Task.checkCancellation() + try await writeResponsePart(.end(nil), context: context, eventLoop: eventLoop) + return .responded + case .cancelled, .sourceFailed(_), .transportFailed(_): + return sourceResult } + } catch is CancellationError { + return .cancelled + } catch { + return .transportFailed(error.localizedDescription) } default: let body = response.bodyData - eventLoop.execute { - var head = HTTPResponseHead(version: version, status: status) - for (name, value) in headers { - head.headers.add(name: name, value: value) - } - if let body { - head.headers.add(name: "Content-Length", value: "\(body.count)") - } - context.write(self.wrapOutboundOut(.head(head)), promise: nil) + if let body { + head.headers.add(name: "Content-Length", value: "\(body.count)") + } + do { + try Task.checkCancellation() + try await writeResponsePart(.head(head), context: context, eventLoop: eventLoop) if let body { - var buffer = context.channel.allocator.buffer(capacity: body.count) - buffer.writeBytes(body) - context.write(self.wrapOutboundOut(.body(.byteBuffer(buffer))), promise: nil) + try await writeResponseBody(body, context: context, eventLoop: eventLoop) } - context.writeAndFlush(self.wrapOutboundOut(.end(nil)), promise: nil) + try Task.checkCancellation() + try await writeResponsePart(.end(nil), context: context, eventLoop: eventLoop) + return .responded + } catch is CancellationError { + return .cancelled + } catch { + return .transportFailed(error.localizedDescription) } } } diff --git a/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift b/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift index 25a9e92..be4a51a 100644 --- a/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift +++ b/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift @@ -49,14 +49,53 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { case reserved case installed case running + case responding case closing(TerminalCause) case closed(TerminalCause?) } + package enum ResponseSourceKind: Equatable, Sendable { + case finite + case open + } + + package enum ResponseWorkPhase: Equatable, Sendable { + case notReserved + case reserved + case installed + case running + case completed + } + + package enum ResponseEndPhase: Equatable, Sendable { + case notExpected + case pending + case acknowledged + case closed + } + package struct RequestSnapshot: Equatable, Sendable { package let id: UUID package let admissionOrdinal: UInt64 package let phase: RequestWorkPhase + package let responseSourceKind: ResponseSourceKind? + package let responseSource: ResponseWorkPhase + package let writer: ResponseWorkPhase + package let responseEnd: ResponseEndPhase + package let responseIsReady: Bool + + package var writerIsRunning: Bool { + writer == .running + } + + package var terminalCause: TerminalCause? { + switch phase { + case .closing(let cause), .closed(let cause?): + cause + case .reserved, .installed, .running, .responding, .closed: + nil + } + } } package struct ConnectionSnapshot: Equatable, Sendable { @@ -68,6 +107,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { } package struct Snapshot: Equatable, Sendable { + package let revision: UInt64 package let generationID: UInt64 package let phase: GenerationPhase package let connections: [ConnectionSnapshot] @@ -77,74 +117,42 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { } } - package final class WorkLease: @unchecked Sendable { - fileprivate let id: UUID - private weak var operation: RequestOperation? - - fileprivate init(operation: RequestOperation, id: UUID) { - self.id = id - self.operation = operation - } - - package func install(_ task: Task) { - operation?.install(task: task, leaseID: id) - } - - package func waitUntilStartIsAllowed() async -> Bool { - guard let operation else { - return false - } - return await operation.waitUntilStartIsAllowed(leaseID: id) - } - - package func acknowledgeCompletion() { - operation?.acknowledgeCompletion(leaseID: id) - } - } - - package final class RequestOperation: @unchecked Sendable { - private enum WorkState { + fileprivate final class WorkSlot: @unchecked Sendable { + private enum State { case reserved case installed(@Sendable () -> Void) case running(@Sendable () -> Void) - case acknowledged + case completed } - package let id = UUID() - package let admissionOrdinal: UInt64 - private weak var connection: Connection? + let id = UUID() + private weak var operation: RequestOperation? private let lock = NSLock() - private let leaseID: UUID - private var workState: WorkState = .reserved - private var terminalCause: TerminalCause? + private var state: State = .reserved + private var cancellationWasRequested = false private var startWasRequested = false private var startWaiter: CheckedContinuation? - private var closeWaiters: [CheckedContinuation] = [] - fileprivate init(admissionOrdinal: UInt64, connection: Connection) { - self.admissionOrdinal = admissionOrdinal - self.connection = connection - let leaseID = UUID() - self.leaseID = leaseID + func attach(to operation: RequestOperation) { + precondition(self.operation == nil) + self.operation = operation } - fileprivate func makeLease() -> WorkLease { - WorkLease(operation: self, id: leaseID) + func makeLease() -> WorkLease { + WorkLease(slot: self, id: id) } - fileprivate func beginClosing(_ cause: TerminalCause) { + func requestCancellation() { var cancellation: (@Sendable () -> Void)? var waiter: CheckedContinuation? lock.lock() - if terminalCause == nil { - terminalCause = cause - } - switch workState { + cancellationWasRequested = true + switch state { case .installed(let cancel), .running(let cancel): cancellation = cancel waiter = startWaiter startWaiter = nil - case .reserved, .acknowledged: + case .reserved, .completed: break } lock.unlock() @@ -152,42 +160,28 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { waiter?.resume(returning: false) } - package func waitUntilClosed() async -> TerminalCause? { - await withCheckedContinuation { continuation in - lock.lock() - if case .acknowledged = workState { - let cause = terminalCause - lock.unlock() - continuation.resume(returning: cause) - } else { - closeWaiters.append(continuation) - lock.unlock() - } - } - } - - fileprivate func install(task: Task, leaseID: UUID) { + func install(task: Task, leaseID: UUID) { var shouldCancel = false var waiter: CheckedContinuation? lock.lock() - precondition(self.leaseID == leaseID, "A request WorkLease belongs to exactly one admitted operation.") - guard case .reserved = workState else { + precondition(id == leaseID, "A WorkLease belongs to exactly one response operation slot.") + guard case .reserved = state else { lock.unlock() - preconditionFailure("A request WorkLease can be installed exactly once.") + preconditionFailure("A WorkLease can be installed exactly once.") } let cancellation: @Sendable () -> Void = { task.cancel() } if startWasRequested { waiter = startWaiter startWaiter = nil - if terminalCause == nil { - workState = .running(cancellation) + if cancellationWasRequested == false { + state = .running(cancellation) } else { - workState = .installed(cancellation) + state = .installed(cancellation) shouldCancel = true } } else { - workState = .installed(cancellation) - shouldCancel = terminalCause != nil + state = .installed(cancellation) + shouldCancel = cancellationWasRequested } lock.unlock() if shouldCancel { @@ -196,20 +190,20 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { waiter?.resume(returning: shouldCancel == false) } - fileprivate func waitUntilStartIsAllowed(leaseID: UUID) async -> Bool { + func waitUntilStartIsAllowed(leaseID: UUID) async -> Bool { await withCheckedContinuation { continuation in var cancellation: (@Sendable () -> Void)? lock.lock() - precondition(self.leaseID == leaseID, "A request WorkLease belongs to exactly one admitted operation.") - precondition(startWasRequested == false, "A request WorkLease can start exactly once.") + precondition(id == leaseID, "A WorkLease belongs to exactly one response operation slot.") + precondition(startWasRequested == false, "A WorkLease can start exactly once.") startWasRequested = true - switch workState { + switch state { case .reserved: startWaiter = continuation lock.unlock() case .installed(let cancel): - if terminalCause == nil { - workState = .running(cancel) + if cancellationWasRequested == false { + state = .running(cancel) lock.unlock() continuation.resume(returning: true) } else { @@ -218,48 +212,358 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { cancellation?() continuation.resume(returning: false) } - case .running, .acknowledged: + case .running, .completed: lock.unlock() - preconditionFailure("A request WorkLease can start exactly once.") + preconditionFailure("A WorkLease can start exactly once.") } } } - fileprivate func acknowledgeCompletion(leaseID: UUID) { - let waiters: [CheckedContinuation] - let cause: TerminalCause? + func acknowledgeCompletion(leaseID: UUID) { lock.lock() - precondition(self.leaseID == leaseID, "A request WorkLease belongs to exactly one admitted operation.") - guard case .acknowledged = workState else { - workState = .acknowledged - cause = terminalCause - waiters = closeWaiters - closeWaiters.removeAll(keepingCapacity: false) + precondition(id == leaseID, "A WorkLease belongs to exactly one response operation slot.") + guard case .completed = state else { + state = .completed + lock.unlock() + operation?.workSlotDidComplete(self) + return + } + lock.unlock() + } + + var isCompleted: Bool { + lock.lock() + let result = if case .completed = state { true } else { false } + lock.unlock() + return result + } + + var snapshot: ResponseWorkPhase { + lock.lock() + let snapshot: ResponseWorkPhase = switch state { + case .reserved: .reserved + case .installed: .installed + case .running: .running + case .completed: .completed + } + lock.unlock() + return snapshot + } + } + + package final class WorkLease: @unchecked Sendable { + fileprivate let id: UUID + private let slot: WorkSlot + + fileprivate init(slot: WorkSlot, id: UUID) { + self.id = id + self.slot = slot + } + + package func install(_ task: Task) { + slot.install(task: task, leaseID: id) + } + + package func waitUntilStartIsAllowed() async -> Bool { + await slot.waitUntilStartIsAllowed(leaseID: id) + } + + package func acknowledgeCompletion() { + slot.acknowledgeCompletion(leaseID: id) + } + } + + package final class RequestOperation: @unchecked Sendable { + private enum Outcome { + case open + case responded + case cancelled(TerminalCause) + } + + private enum ResponseQueueState { + case handling + case ready(CheckedContinuation?) + case turnGranted + } + + package let id = UUID() + package let admissionOrdinal: UInt64 + private weak var connection: Connection? + private let lock = NSLock() + private let handlerSlot = WorkSlot() + private var sourceSlot: (kind: ResponseSourceKind, slot: WorkSlot)? + private var writerSlot: WorkSlot? + private var responseQueueState: ResponseQueueState = .handling + private var responseEnd: ResponseEndPhase = .notExpected + private var outcome: Outcome = .open + private var didClose = false + private var closeWaiters: [CheckedContinuation] = [] + + fileprivate init(admissionOrdinal: UInt64, connection: Connection) { + self.admissionOrdinal = admissionOrdinal + self.connection = connection + handlerSlot.attach(to: self) + } + + fileprivate func makeLease() -> WorkLease { + handlerSlot.makeLease() + } + + func reserveResponseSource(_ kind: ResponseSourceKind) -> WorkLease? { + lock.lock() + guard case .open = outcome, + responseEnd == .notExpected, + sourceSlot == nil else { + lock.unlock() + return nil + } + let slot = WorkSlot() + slot.attach(to: self) + sourceSlot = (kind, slot) + responseEnd = .pending + lock.unlock() + notifyChanged() + return slot.makeLease() + } + + func markResponseSourceNotRequired() -> Bool { + lock.lock() + guard case .open = outcome, + responseEnd == .notExpected, + sourceSlot == nil else { + lock.unlock() + return false + } + responseEnd = .pending + lock.unlock() + notifyChanged() + return true + } + + fileprivate func markResponseReady() -> Bool { + lock.lock() + guard case .open = outcome, + responseEnd == .pending, + case .handling = responseQueueState else { + lock.unlock() + return false + } + responseQueueState = .ready(nil) + lock.unlock() + notifyChanged() + return true + } + + fileprivate var isResponseReady: Bool { + lock.lock() + let result = if case .ready = responseQueueState { true } else { false } + lock.unlock() + return result + } + + fileprivate func grantWriterTurn() -> Bool { + let waiter: CheckedContinuation? + lock.lock() + guard case .open = outcome, + case .ready(let continuation) = responseQueueState else { lock.unlock() - for waiter in waiters { - waiter.resume(returning: cause) + return false + } + waiter = continuation + responseQueueState = .turnGranted + lock.unlock() + waiter?.resume(returning: true) + notifyChanged() + return true + } + + fileprivate func waitForWriterTurn() async -> Bool { + await withCheckedContinuation { continuation in + lock.lock() + guard case .open = outcome else { + lock.unlock() + continuation.resume(returning: false) + return } - connection?.requestDidClose(self) + switch responseQueueState { + case .turnGranted: + lock.unlock() + continuation.resume(returning: true) + case .ready(nil): + responseQueueState = .ready(continuation) + lock.unlock() + case .handling, .ready: + lock.unlock() + preconditionFailure("A response can wait for its FIFO writer turn exactly once.") + } + } + } + + func reserveWriter() -> WorkLease? { + lock.lock() + guard case .open = outcome, + case .turnGranted = responseQueueState, + writerSlot == nil else { + lock.unlock() + return nil + } + let slot = WorkSlot() + slot.attach(to: self) + writerSlot = slot + lock.unlock() + notifyChanged() + return slot.makeLease() + } + + func acknowledgeResponseEnd() { + lock.lock() + guard case .open = outcome, responseEnd == .pending else { + lock.unlock() return } + responseEnd = .acknowledged + outcome = .responded lock.unlock() + notifyChanged() + finishIfPossible() + } + + fileprivate func beginClosing(_ cause: TerminalCause) { + transitionToTerminal(.cancelled(cause)) + } + + package func waitUntilClosed() async -> TerminalCause? { + await withCheckedContinuation { continuation in + lock.lock() + if didClose { + let cause = terminalCauseLocked + lock.unlock() + continuation.resume(returning: cause) + } else { + closeWaiters.append(continuation) + lock.unlock() + } + } + } + + fileprivate func workSlotDidComplete(_ slot: WorkSlot) { + lock.lock() + if slot === handlerSlot, + case .open = outcome, + responseEnd == .notExpected { + responseEnd = .acknowledged + outcome = .responded + } + lock.unlock() + notifyChanged() + finishIfPossible() } fileprivate func snapshot() -> RequestSnapshot { lock.lock() + let terminalCause = terminalCauseLocked + let handlerPhase = handlerSlot.snapshot let phase: RequestWorkPhase - switch workState { - case .reserved: - phase = terminalCause.map(RequestWorkPhase.closing) ?? .reserved - case .installed: - phase = terminalCause.map(RequestWorkPhase.closing) ?? .installed - case .running: - phase = terminalCause.map(RequestWorkPhase.closing) ?? .running - case .acknowledged: + if didClose { phase = .closed(terminalCause) + } else if let terminalCause { + phase = .closing(terminalCause) + } else { + phase = switch handlerPhase { + case .reserved: .reserved + case .installed: .installed + case .running: .running + case .completed, .notReserved: .responding + } + } + let sourceKind = sourceSlot?.kind + let sourcePhase = sourceSlot?.slot.snapshot ?? .notReserved + let writerPhase = writerSlot?.snapshot ?? .notReserved + let responseEnd = responseEnd + let responseIsReady = switch responseQueueState { + case .handling: + false + case .ready, .turnGranted: + true + } + lock.unlock() + return .init( + id: id, + admissionOrdinal: admissionOrdinal, + phase: phase, + responseSourceKind: sourceKind, + responseSource: sourcePhase, + writer: writerPhase, + responseEnd: responseEnd, + responseIsReady: responseIsReady + ) + } + + private var terminalCauseLocked: TerminalCause? { + if case .cancelled(let cause) = outcome { + return cause + } + return nil + } + + private func transitionToTerminal(_ requestedOutcome: Outcome) { + let slots: [WorkSlot] + let waiter: CheckedContinuation? + lock.lock() + if case .open = outcome { + outcome = requestedOutcome + if responseEnd == .pending || responseEnd == .notExpected { + responseEnd = .closed + } + } + if case .ready(let continuation) = responseQueueState { + waiter = continuation + responseQueueState = .turnGranted + } else { + waiter = nil + } + slots = [handlerSlot, sourceSlot?.slot, writerSlot].compactMap { $0 } + lock.unlock() + waiter?.resume(returning: false) + for slot in slots { + slot.requestCancellation() + } + notifyChanged() + finishIfPossible() + } + + private func finishIfPossible() { + let waiters: [CheckedContinuation] + let cause: TerminalCause? + lock.lock() + guard didClose == false, + isTerminalLocked, + handlerSlot.isCompleted, + sourceSlot?.slot.isCompleted != false, + writerSlot?.isCompleted != false else { + lock.unlock() + return } + didClose = true + cause = terminalCauseLocked + waiters = closeWaiters + closeWaiters.removeAll(keepingCapacity: false) lock.unlock() - return .init(id: id, admissionOrdinal: admissionOrdinal, phase: phase) + for waiter in waiters { + waiter.resume(returning: cause) + } + connection?.requestDidClose(self) + } + + private var isTerminalLocked: Bool { + if case .open = outcome { + return false + } + return true + } + + private func notifyChanged() { + connection?.requestDidChange() } } @@ -276,7 +580,8 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { private let lock = NSLock() private var phase: ConnectionPhase = .accepting private var nextRequestOrdinal: UInt64 = 0 - private var requests: [UUID: RequestOperation] = [:] + private var requests: [RequestOperation] = [] + private var activeWriterRequestID: UUID? private var closeAcknowledged = false private var closeWaiters: [CheckedContinuation] = [] @@ -308,11 +613,20 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { connection: self ) let lease = operation.makeLease() - requests[operation.id] = operation + requests.append(operation) lock.unlock() + owner?.changed() return .init(operation: operation, lease: lease) } + package func supplyResponse(for operation: RequestOperation) async -> Bool { + guard operation.markResponseReady() else { + return false + } + pumpWriterQueue() + return await operation.waitForWriterTurn() + } + package func peerClosed() { beginClosing(.peerClosed, signalResourceClose: false) } @@ -340,6 +654,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { phase = .admissionClosed } lock.unlock() + owner?.changed() } fileprivate func beginClosing(_ cause: TerminalCause) { @@ -350,7 +665,10 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { var didClose = false var waiters: [CheckedContinuation] = [] lock.lock() - requests.removeValue(forKey: operation.id) + requests.removeAll { $0 === operation } + if activeWriterRequestID == operation.id { + activeWriterRequestID = nil + } if case .closing = phase, closeAcknowledged, requests.isEmpty { phase = .closed waiters = closeWaiters @@ -363,16 +681,22 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { } if didClose { owner?.connectionDidClose(self) + } else { + pumpWriterQueue() + owner?.changed() } } + fileprivate func requestDidChange() { + pumpWriterQueue() + owner?.changed() + } + fileprivate func snapshot() -> ConnectionSnapshot { lock.lock() let phase = phase let closeAcknowledged = closeAcknowledged - let requests = requests.values.sorted { - $0.admissionOrdinal < $1.admissionOrdinal - } + let requests = requests lock.unlock() return .init( id: id, @@ -393,7 +717,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { switch phase { case .accepting, .admissionClosed: phase = .closing(cause) - requests = Array(self.requests.values) + requests = self.requests shouldSignalClose = signalResourceClose case .closing, .closed: requests = [] @@ -415,7 +739,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { switch phase { case .accepting, .admissionClosed: phase = .closing(.peerClosed) - requests = Array(self.requests.values) + requests = self.requests case .closing, .closed: requests = [] } @@ -444,6 +768,30 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { owner?.connectionDidClose(self) } } + + private func pumpWriterQueue() { + let operation: RequestOperation? + lock.lock() + if activeWriterRequestID == nil, + let head = requests.first, + head.isResponseReady { + activeWriterRequestID = head.id + operation = head + } else { + operation = nil + } + lock.unlock() + guard let operation else { + return + } + if operation.grantWriterTurn() == false { + lock.lock() + if activeWriterRequestID == operation.id { + activeWriterRequestID = nil + } + lock.unlock() + } + } } package final class ClosingGeneration: @unchecked Sendable { @@ -469,11 +817,18 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { case closed } + private struct SnapshotWaiter { + let revision: UInt64 + let continuation: CheckedContinuation + } + package let generationID: UInt64 private let lock = NSLock() private var state: State = .accepting(.init(connections: [:])) private var nextConnectionOrdinal: UInt64 = 0 + private var revision: UInt64 = 0 private var closeWaiters: [CheckedContinuation] = [] + private var snapshotWaiters: [SnapshotWaiter] = [] package init(generationID: UInt64) { self.generationID = generationID @@ -500,6 +855,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { state = .accepting(accepting) lock.unlock() connection.installCloseAcknowledgement() + changed() return connection } @@ -519,6 +875,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { for connection in connections { connection.closeAdmission() } + changed() } package func beginClosing(_ cause: TerminalCause) -> ClosingGeneration { @@ -556,11 +913,13 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { for waiter in waiters { waiter.resume() } + changed() return ClosingGeneration(owner: self) } package func snapshot() -> Snapshot { lock.lock() + let revision = revision let phase: GenerationPhase let connections: [Connection] switch state { @@ -579,6 +938,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { } lock.unlock() return .init( + revision: revision, generationID: generationID, phase: phase, connections: connections.sorted { @@ -587,6 +947,22 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { ) } + package func nextSnapshot(after priorRevision: UInt64) async -> Snapshot { + await withCheckedContinuation { continuation in + lock.lock() + if revision > priorRevision { + lock.unlock() + continuation.resume(returning: snapshot()) + } else { + snapshotWaiters.append(.init( + revision: priorRevision, + continuation: continuation + )) + lock.unlock() + } + } + } + fileprivate func connectionDidClose(_ connection: Connection) { var waiters: [CheckedContinuation] = [] lock.lock() @@ -613,6 +989,24 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { for waiter in waiters { waiter.resume() } + changed() + } + + fileprivate func changed() { + let waiters: [CheckedContinuation] + lock.lock() + revision &+= 1 + let ready = snapshotWaiters.filter { revision > $0.revision } + snapshotWaiters.removeAll { revision > $0.revision } + waiters = ready.map(\.continuation) + lock.unlock() + guard waiters.isEmpty == false else { + return + } + let current = snapshot() + for waiter in waiters { + waiter.resume(returning: current) + } } private func waitUntilClosed() async { diff --git a/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift b/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift index d949238..184acf5 100644 --- a/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift +++ b/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift @@ -389,6 +389,286 @@ struct CodexReviewMCPHTTPServerTests { try await server.stop() } + @Test func stopRacingReturnedFiniteReviewListResponseJoinsOneOperation() async throws { + let backend = FakeCodexReviewBackend() + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: backend) + ) + let server = CodexReviewMCPHTTPServer( + adapter: CodexReviewMCPServer(store: store), + configuration: .init(host: "127.0.0.1", port: 0) + ) + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + await server.holdNextFiniteSourceCompletionForTesting() + let responseTask = Task { + try await postJSONRPCData( + endpoint: endpoint, + sessionID: sessionID, + bodyData: makeJSONBody([ + "jsonrpc": "2.0", + "id": 2, + "method": "tools/call", + "params": [ + "name": "review_list", + "arguments": ["limit": 20], + ], + ]) + ) + } + await server.waitUntilFiniteSourceCompletionIsHeldForTesting() + let held = await server.networkSnapshotForTesting() + .connections.flatMap(\.requests) + #expect(held.count == 1) + #expect(held.first?.responseSourceKind == .finite) + #expect(held.first?.responseSource == .running) + + let stopTask = Task { try await server.stop() } + _ = await waitForNetworkSnapshot(on: server) { + $0.phase == .closing(.serverStop) + } + #expect(await server.eventLoopGroupShutdownCountForTesting() == 0) + await server.releaseFiniteSourceCompletionForTesting() + _ = try? await responseTask.value + try await stopTask.value + + #expect(await server.eventLoopGroupShutdownCountForTesting() == 1) + #expect((await server.networkSnapshotForTesting()).isClosed) + } + + @Test func oneConnectionEmitsPipelinedPOSTResponsesInAdmissionOrder() async throws { + let backend = FakeCodexReviewBackend() + let firstResponseGate = AsyncGate() + await backend.holdStartReview(with: firstResponseGate) + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: backend), + idGenerator: .init(next: { "job-1" }) + ) + let server = CodexReviewMCPHTTPServer( + adapter: CodexReviewMCPServer(store: store), + configuration: .init(host: "127.0.0.1", port: 0) + ) + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + let descriptor = try await openRawTCPConnection(endpoint: endpoint) + defer { Darwin.close(descriptor) } + let firstBody = try makeJSONBody([ + "jsonrpc": "2.0", + "id": 2, + "method": "tools/call", + "params": [ + "name": "review_start", + "arguments": [ + "cwd": "/tmp/project", + "target": ["type": "uncommittedChanges"], + ], + ], + ]) + let secondBody = try makeJSONBody([ + "jsonrpc": "2.0", + "id": 3, + "method": "tools/list", + ]) + try await sendRawPipelinedPOSTs( + descriptor: descriptor, + endpoint: endpoint, + sessionID: sessionID, + bodies: [firstBody, secondBody] + ) + await backend.waitForStartReview() + let secondReady = await waitForNetworkSnapshot(on: server) { snapshot in + guard let requests = snapshot.connections.first(where: { + $0.requests.count == 2 + })?.requests else { return false } + return requests[1].responseIsReady + } + let queued = try #require(secondReady.connections.first(where: { + $0.requests.count == 2 + })?.requests) + #expect(queued[0].admissionOrdinal < queued[1].admissionOrdinal) + #expect(queued[1].writer == .notReserved) + + await firstResponseGate.open() + await backend.yield(.completed(summary: "Done", result: "review text")) + let responseText = String( + decoding: try await readTwoRawHTTPResponses(descriptor: descriptor), + as: UTF8.self + ) + let firstID = try #require(responseText.range(of: "\"id\":2")) + let secondID = try #require(responseText.range(of: "\"id\":3")) + #expect(firstID.lowerBound < secondID.lowerBound) + try await server.stop() + } + + @Test func stopAwaitsSSEWriterCompletionBeforeEventLoopShutdown() async throws { + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) + ) + let server = CodexReviewMCPHTTPServer( + adapter: CodexReviewMCPServer(store: store), + configuration: .init(port: 0) + ) + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + await server.holdNextWriterCompletionForTesting() + var request = URLRequest(url: endpoint) + request.httpMethod = "GET" + request.setValue("text/event-stream, application/json", forHTTPHeaderField: "Accept") + request.setValue(sessionID, forHTTPHeaderField: "MCP-Session-Id") + let (bytes, response) = try await URLSession.shared.bytes(for: request) + #expect((response as? HTTPURLResponse)?.statusCode == 200) + + let stopTask = Task { try await server.stop() } + await server.waitUntilWriterCompletionIsHeldForTesting() + #expect(await server.eventLoopGroupShutdownCountForTesting() == 0) + await server.releaseWriterCompletionForTesting() + try await stopTask.value + _ = bytes + #expect(await server.eventLoopGroupShutdownCountForTesting() == 1) + } + + @Test func acceptedChildCloseAcknowledgementPrecedesEventLoopShutdown() async throws { + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) + ) + let server = CodexReviewMCPHTTPServer( + adapter: CodexReviewMCPServer(store: store), + configuration: .init(host: "127.0.0.1", port: 0) + ) + try await server.start() + let descriptor = try await openRawTCPConnection(endpoint: await server.url) + defer { Darwin.close(descriptor) } + _ = await waitForNetworkSnapshot(on: server) { $0.connections.isEmpty == false } + + try await server.stop() + + #expect(await rawConnectionReachedEOF(descriptor: descriptor)) + #expect(await server.eventLoopGroupShutdownCountForTesting() == 1) + } + + @Test func channelCloseOwnsSSETerminationWithoutLateEventLoopCleanup() async throws { + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) + ) + let server = CodexReviewMCPHTTPServer( + adapter: CodexReviewMCPServer(store: store), + configuration: .init( + host: "127.0.0.1", + port: 0, + streamHeartbeatInterval: .milliseconds(50) + ) + ) + try await server.start() + let sessionID = try await initializeSession(endpoint: await server.url) + + try await openAndCloseRawEventStream(endpoint: await server.url, sessionID: sessionID) + _ = await waitForNetworkSnapshot(on: server) { snapshot in + snapshot.connections.flatMap(\.requests).contains { + $0.responseSourceKind == .open + } == false + } + try await server.stop() + + #expect(await server.eventLoopGroupShutdownCountForTesting() == 1) + } + + @Test func acknowledgedResponseEndWinsConcurrentStop() async throws { + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) + ) + let server = CodexReviewMCPHTTPServer( + adapter: CodexReviewMCPServer(store: store), + configuration: .init(host: "127.0.0.1", port: 0) + ) + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + await server.holdNextResponseEndAcknowledgementForTesting() + let body = try makeJSONBody([ + "jsonrpc": "2.0", + "id": 2, + "method": "tools/list", + ]) + let responseTask = Task { + try await postJSONRPCData( + endpoint: endpoint, + sessionID: sessionID, + bodyData: body + ) + } + await server.waitUntilResponseEndAcknowledgementIsHeldForTesting() + let beforeStop = try #require( + await server.networkSnapshotForTesting().connections + .flatMap(\.requests).first + ) + #expect(beforeStop.responseEnd == .acknowledged) + #expect(beforeStop.terminalCause == nil) + + let stopTask = Task { try await server.stop() } + let stopping = await waitForNetworkSnapshot(on: server) { + $0.phase == .closing(.serverStop) + } + let afterStop = try #require( + stopping.connections + .flatMap(\.requests).first + ) + #expect(afterStop.responseEnd == .acknowledged) + #expect(afterStop.terminalCause == nil) + #expect(await server.eventLoopGroupShutdownCountForTesting() == 0) + await server.releaseResponseEndAcknowledgementForTesting() + _ = try decodeSSEJSON(from: try await responseTask.value) + try await stopTask.value + } + + @Test func clientDisconnectDrainsOwnedFiniteResponseSource() async throws { + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) + ) + let server = CodexReviewMCPHTTPServer( + adapter: CodexReviewMCPServer(store: store), + configuration: .init(host: "127.0.0.1", port: 0) + ) + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + await server.holdNextFiniteSourceCompletionForTesting() + let descriptor = try await openRawTCPConnection(endpoint: endpoint) + try await sendRawPipelinedPOSTs( + descriptor: descriptor, + endpoint: endpoint, + sessionID: sessionID, + bodies: [makeJSONBody([ + "jsonrpc": "2.0", + "id": 2, + "method": "tools/list", + ])] + ) + await server.waitUntilFiniteSourceCompletionIsHeldForTesting() + try await readRawHTTPHeaders(descriptor: descriptor) + Darwin.shutdown(descriptor, SHUT_RDWR) + Darwin.close(descriptor) + let closing = await waitForNetworkSnapshot(on: server) { snapshot in + snapshot.connections.flatMap(\.requests).contains { + guard $0.responseSourceKind == .finite else { return false } + switch $0.terminalCause { + case .peerClosed, .transportFailure(_): + return true + case .serverStop, nil: + return false + } + } + } + #expect(closing.connections.flatMap(\.requests).count == 1) + await server.releaseFiniteSourceCompletionForTesting() + _ = await waitForNetworkSnapshot(on: server) { snapshot in + snapshot.connections.flatMap(\.requests).isEmpty + } + try await server.stop() + } + @Test func streamableHTTPCallsReviewStartWithCustomTarget() async throws { let backend = FakeCodexReviewBackend() let store = CodexReviewStore.makeTestingStore( @@ -1585,6 +1865,114 @@ struct CodexReviewMCPHTTPServerTests { return data } + private nonisolated func openRawTCPConnection(endpoint: URL) async throws -> Int32 { + try await Task.detached { + let components = try #require(URLComponents(url: endpoint, resolvingAgainstBaseURL: false)) + let host = try #require(components.host) + let port = try #require(components.port) + let descriptor = Darwin.socket(AF_INET, SOCK_STREAM, 0) + guard descriptor >= 0 else { throw currentPOSIXError() } + var address = sockaddr_in() + address.sin_len = UInt8(MemoryLayout.size) + address.sin_family = sa_family_t(AF_INET) + address.sin_port = in_port_t(port).bigEndian + let ipv4Host = host == "localhost" ? "127.0.0.1" : host + guard inet_pton(AF_INET, ipv4Host, &address.sin_addr) == 1 else { + Darwin.close(descriptor) + throw testError("Unable to resolve IPv4 loopback host \(host)") + } + let connected = withUnsafePointer(to: &address) { pointer in + pointer.withMemoryRebound(to: sockaddr.self, capacity: 1) { + Darwin.connect(descriptor, $0, socklen_t(MemoryLayout.size)) + } + } + guard connected == 0 else { + let error = currentPOSIXError() + Darwin.close(descriptor) + throw error + } + return descriptor + }.value + } + + private nonisolated func sendRawPipelinedPOSTs( + descriptor: Int32, + endpoint: URL, + sessionID: String, + bodies: [Data] + ) async throws { + try await Task.detached { + let components = try #require(URLComponents(url: endpoint, resolvingAgainstBaseURL: false)) + let host = try #require(components.host) + let port = try #require(components.port) + var request = Data() + for body in bodies { + request.append(Data([ + "POST \(endpoint.path) HTTP/1.1", + "Host: \(host):\(port)", + "Accept: text/event-stream, application/json", + "Content-Type: application/json", + "MCP-Session-Id: \(sessionID)", + "MCP-Protocol-Version: 2025-11-25", + "Content-Length: \(body.count)", + "Connection: keep-alive", + "", + "", + ].joined(separator: "\r\n").utf8)) + request.append(body) + } + try request.withUnsafeBytes { buffer in + guard let base = buffer.baseAddress else { throw testError("Empty HTTP request") } + var sent = 0 + while sent < buffer.count { + let count = Darwin.send(descriptor, base.advanced(by: sent), buffer.count - sent, 0) + guard count > 0 else { throw currentPOSIXError() } + sent += count + } + } + }.value + } + + private nonisolated func readTwoRawHTTPResponses(descriptor: Int32) async throws -> Data { + try await Task.detached { + var response = Data() + var buffer = [UInt8](repeating: 0, count: 4096) + while true { + let text = String(decoding: response, as: UTF8.self) + if text.components(separatedBy: "HTTP/1.1 200").count - 1 >= 2, + text.components(separatedBy: "\r\n0\r\n\r\n").count - 1 >= 2 { + return response + } + let count = Darwin.recv(descriptor, &buffer, buffer.count, 0) + guard count > 0 else { throw testError("Connection closed before both responses ended") } + response.append(contentsOf: buffer.prefix(count)) + guard response.count <= 2 * 1024 * 1024 else { + throw testError("Pipelined responses exceeded the test bound") + } + } + }.value + } + + private nonisolated func readRawHTTPHeaders(descriptor: Int32) async throws { + try await Task.detached { + var response = Data() + var buffer = [UInt8](repeating: 0, count: 1024) + while String(decoding: response, as: UTF8.self).contains("\r\n\r\n") == false { + let count = Darwin.recv(descriptor, &buffer, buffer.count, 0) + guard count > 0 else { throw testError("Connection closed before response headers") } + response.append(contentsOf: buffer.prefix(count)) + guard response.count < 8192 else { throw testError("Response headers exceeded test bound") } + } + }.value + } + + private nonisolated func rawConnectionReachedEOF(descriptor: Int32) async -> Bool { + await Task.detached { + var byte: UInt8 = 0 + return Darwin.recv(descriptor, &byte, 1, 0) == 0 + }.value + } + private nonisolated func openAndCloseRawEventStream(endpoint: URL, sessionID: String) async throws { try await Task.detached { let components = try #require(URLComponents(url: endpoint, resolvingAgainstBaseURL: false)) @@ -1721,6 +2109,17 @@ struct CodexReviewMCPHTTPServerTests { } return true } + + private func waitForNetworkSnapshot( + on server: CodexReviewMCPHTTPServer, + satisfying condition: (MCPHTTPNetworkResourceOwner.Snapshot) -> Bool + ) async -> MCPHTTPNetworkResourceOwner.Snapshot { + var snapshot = await server.networkSnapshotForTesting() + while condition(snapshot) == false { + snapshot = await server.nextNetworkSnapshotForTesting(after: snapshot.revision) + } + return snapshot + } } private nonisolated func currentPOSIXError() -> NSError { From a2ea85ed53e653ac0e52a52507439080acc99259 Mon Sep 17 00:00:00 2001 From: Kazuki Nakashima <65545348+lynnswap@users.noreply.github.com> Date: Sat, 22 Aug 2026 18:09:44 +0900 Subject: [PATCH 2/3] fix(mcp): preserve response transport contracts --- .../CodexReviewMCPHTTPServer.swift | 138 ++++++-- .../MCPHTTPNetworkResourceOwner.swift | 39 ++- .../CodexReviewMCPHTTPServerTests.swift | 299 +++++++++++++++++- 3 files changed, 445 insertions(+), 31 deletions(-) diff --git a/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift b/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift index 5ed7d38..30b2c33 100644 --- a/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift +++ b/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift @@ -959,7 +959,59 @@ package actor CodexReviewMCPHTTPServer { } package func handleHTTPRequest(_ request: HTTPRequest) async -> HTTPResponse { - await handleTrackedHTTPRequest(request).response + directResponse(from: await handleTrackedHTTPRequest(request)) + } + + private func directResponse(from trackedResponse: TrackedHTTPResponse) -> HTTPResponse { + guard case .stream(let source, let headers) = trackedResponse.response, + let completion = trackedResponse.streamCompletion else { + return trackedResponse.response + } + let heartbeatInterval = configuration.streamHeartbeatInterval + let stream = AsyncThrowingStream(bufferingPolicy: .unbounded) { continuation in + let task = Task { + await withTaskGroup(of: Void.self) { group in + group.addTask { + do { + for try await chunk in source { + try Task.checkCancellation() + continuation.yield(chunk) + } + await completion.finishAndWait() + continuation.finish() + } catch is CancellationError { + await completion.finishAndWait() + continuation.finish() + } catch { + await completion.finishAndWait() + continuation.finish(throwing: error) + } + } + if let heartbeatInterval { + group.addTask { + while Task.isCancelled == false { + do { + try await Task.sleep(for: heartbeatInterval) + } catch { + return + } + guard Task.isCancelled == false else { + return + } + continuation.yield(Data(": keep-alive\n\n".utf8)) + } + } + } + _ = await group.next() + group.cancelAll() + } + } + continuation.onTermination = { _ in + task.cancel() + completion.finish() + } + } + return .stream(stream, headers: headers) } fileprivate func handleTrackedHTTPRequest(_ request: HTTPRequest) async -> TrackedHTTPResponse { @@ -1065,9 +1117,7 @@ package actor CodexReviewMCPHTTPServer { switch response { case .stream: let completion = ActiveRequestCompletion { - Task { - await self.finishActiveRequest(sessionID: sessionID) - } + await self.finishActiveRequest(sessionID: sessionID) } return ( .init( @@ -1235,22 +1285,38 @@ package actor CodexReviewMCPHTTPServer { private final class ActiveRequestCompletion: @unchecked Sendable { private let lock = NSLock() - private let onFinish: @Sendable () -> Void + private let onFinish: @Sendable () async -> Void private var didFinish = false - init(onFinish: @escaping @Sendable () -> Void) { + init(onFinish: @escaping @Sendable () async -> Void) { self.onFinish = onFinish } func finish() { + guard claim() else { + return + } + Task { + await onFinish() + } + } + + func finishAndWait() async { + guard claim() else { + return + } + await onFinish() + } + + private func claim() -> Bool { lock.lock() if didFinish { lock.unlock() - return + return false } didFinish = true lock.unlock() - onFinish() + return true } } @@ -1377,6 +1443,7 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked .init(response: .error(statusCode: 404, .invalidRequest("Not Found"))), operation: operation, version: head.version, + closeAfterResponse: head.isKeepAlive == false, context: context ) return @@ -1388,6 +1455,7 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked response, operation: operation, version: head.version, + closeAfterResponse: head.isKeepAlive == false, context: context ) } @@ -1424,20 +1492,46 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked _ trackedResponse: TrackedHTTPResponse, operation: MCPHTTPNetworkResourceOwner.RequestOperation, version: HTTPVersion, + closeAfterResponse: Bool, context: ChannelHandlerContext ) async { let response = trackedResponse.response - let preparedResponse: HTTPResponse switch response { - case .stream(let source, let headers): + case .stream: guard let kind = trackedResponse.responseSourceKind else { - trackedResponse.streamCompletion?.finish() + await trackedResponse.streamCompletion?.finishAndWait() connection.transportFailed("Streaming response has no source lifetime contract.") return } - guard let sourceLease = operation.reserveResponseSource(kind) else { - trackedResponse.streamCompletion?.finish() + guard operation.declareResponseSource(kind) else { + await trackedResponse.streamCompletion?.finishAndWait() + return + } + + default: + guard operation.markResponseSourceNotRequired() else { + return + } + } + + guard await connection.supplyResponse(for: operation) else { + await trackedResponse.streamCompletion?.finishAndWait() + return + } + guard let writerLease = operation.reserveWriter() else { + await trackedResponse.streamCompletion?.finishAndWait() + return + } + + let preparedResponse: HTTPResponse + var admittedSource: (lease: MCPHTTPNetworkResourceOwner.WorkLease, task: Task)? + switch response { + case .stream(let source, let headers): + guard let kind = trackedResponse.responseSourceKind, + let sourceLease = operation.reserveResponseSource() else { + await trackedResponse.streamCompletion?.finishAndWait() + writerLease.acknowledgeCompletion() return } let bridge = AsyncThrowingStream.makeStream( @@ -1464,23 +1558,17 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked } else { bridge.continuation.finish() } - trackedResponse.streamCompletion?.finish() + await trackedResponse.streamCompletion?.finishAndWait() sourceLease.acknowledgeCompletion() } - sourceLease.install(sourceTask) + admittedSource = (sourceLease, sourceTask) preparedResponse = .stream(bridge.stream, headers: headers) default: - guard operation.markResponseSourceNotRequired() else { - return - } + admittedSource = nil preparedResponse = response } - guard await connection.supplyResponse(for: operation), - let writerLease = operation.reserveWriter() else { - return - } nonisolated(unsafe) let context = context let heartbeatInterval = await server.responseHeartbeatInterval let writerTask = Task { @@ -1497,6 +1585,9 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked switch result { case .responded: operation.acknowledgeResponseEnd() + if closeAfterResponse { + self.connection.closeAfterResponse() + } await self.server.waitAfterResponseEndAcknowledgementForTesting() case .cancelled: break @@ -1511,6 +1602,9 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked writerLease.acknowledgeCompletion() } writerLease.install(writerTask) + if let admittedSource { + admittedSource.lease.install(admittedSource.task) + } } private enum WriterCompletion { diff --git a/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift b/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift index be4a51a..3fee056 100644 --- a/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift +++ b/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift @@ -28,6 +28,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { package enum TerminalCause: Equatable, Sendable { case serverStop case peerClosed + case responseComplete case transportFailure(String) } @@ -291,7 +292,8 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { private weak var connection: Connection? private let lock = NSLock() private let handlerSlot = WorkSlot() - private var sourceSlot: (kind: ResponseSourceKind, slot: WorkSlot)? + private var responseSourceKind: ResponseSourceKind? + private var sourceSlot: WorkSlot? private var writerSlot: WorkSlot? private var responseQueueState: ResponseQueueState = .handling private var responseEnd: ResponseEndPhase = .notExpected @@ -309,18 +311,34 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { handlerSlot.makeLease() } - func reserveResponseSource(_ kind: ResponseSourceKind) -> WorkLease? { + func declareResponseSource(_ kind: ResponseSourceKind) -> Bool { lock.lock() guard case .open = outcome, responseEnd == .notExpected, + responseSourceKind == nil, + sourceSlot == nil else { + lock.unlock() + return false + } + responseSourceKind = kind + responseEnd = .pending + lock.unlock() + notifyChanged() + return true + } + + func reserveResponseSource() -> WorkLease? { + lock.lock() + guard case .open = outcome, + case .turnGranted = responseQueueState, + responseSourceKind != nil, sourceSlot == nil else { lock.unlock() return nil } let slot = WorkSlot() slot.attach(to: self) - sourceSlot = (kind, slot) - responseEnd = .pending + sourceSlot = slot lock.unlock() notifyChanged() return slot.makeLease() @@ -330,6 +348,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { lock.lock() guard case .open = outcome, responseEnd == .notExpected, + responseSourceKind == nil, sourceSlot == nil else { lock.unlock() return false @@ -476,8 +495,8 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { case .completed, .notReserved: .responding } } - let sourceKind = sourceSlot?.kind - let sourcePhase = sourceSlot?.slot.snapshot ?? .notReserved + let sourceKind = responseSourceKind + let sourcePhase = sourceSlot?.snapshot ?? .notReserved let writerPhase = writerSlot?.snapshot ?? .notReserved let responseEnd = responseEnd let responseIsReady = switch responseQueueState { @@ -522,7 +541,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { } else { waiter = nil } - slots = [handlerSlot, sourceSlot?.slot, writerSlot].compactMap { $0 } + slots = [handlerSlot, sourceSlot, writerSlot].compactMap { $0 } lock.unlock() waiter?.resume(returning: false) for slot in slots { @@ -539,7 +558,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { guard didClose == false, isTerminalLocked, handlerSlot.isCompleted, - sourceSlot?.slot.isCompleted != false, + sourceSlot?.isCompleted != false, writerSlot?.isCompleted != false else { lock.unlock() return @@ -635,6 +654,10 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { beginClosing(.transportFailure(message), signalResourceClose: true) } + package func closeAfterResponse() { + beginClosing(.responseComplete, signalResourceClose: true) + } + package func waitUntilClosed() async { await withCheckedContinuation { continuation in lock.lock() diff --git a/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift b/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift index 184acf5..5dea1f6 100644 --- a/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift +++ b/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift @@ -144,6 +144,80 @@ struct CodexReviewMCPHTTPServerTests { try await server.stop() } + @Test func directHTTPStreamsFinishFiniteAndCancelledOpenRequests() async throws { + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) + ) + let server = CodexReviewMCPHTTPServer( + adapter: CodexReviewMCPServer(store: store), + configuration: .init(host: "review.local", port: 9417) + ) + let initializeBody = try makeJSONBody([ + "jsonrpc": "2.0", + "id": 1, + "method": "initialize", + "params": [ + "protocolVersion": "2025-11-25", + "capabilities": [:], + "clientInfo": ["name": "DirectTests", "version": "0.0.0"], + ], + ]) + let finite = await server.handleHTTPRequest(HTTPRequest( + method: "POST", + headers: [ + HTTPHeaderName.host: "review.local:9417", + HTTPHeaderName.accept: "text/event-stream, application/json", + HTTPHeaderName.contentType: "application/json", + ], + body: initializeBody, + path: "/mcp" + )) + let sessionID = try #require(finite.headers[HTTPHeaderName.sessionID]) + guard case .stream(let finiteStream, _) = finite else { + Issue.record("Initialize must return a finite direct stream.") + return + } + for try await _ in finiteStream {} + #expect(await server.sessionActiveRequestCountForTesting(sessionID: sessionID) == 0) + + let open = await server.handleHTTPRequest(HTTPRequest( + method: "GET", + headers: [ + HTTPHeaderName.host: "review.local:9417", + HTTPHeaderName.accept: "text/event-stream, application/json", + HTTPHeaderName.protocolVersion: "2025-11-25", + HTTPHeaderName.sessionID: sessionID, + ], + path: "/mcp" + )) + guard case .stream(let openStream, _) = open else { + Issue.record("GET must return an open direct stream.") + return + } + let firstElement = AsyncGate() + let consumer = Task { + do { + for try await _ in openStream { + await firstElement.open() + try Task.checkCancellation() + } + } catch is CancellationError { + return + } catch { + Issue.record("Unexpected direct stream failure: \(error)") + } + } + await firstElement.wait() + #expect(await server.sessionActiveRequestCountForTesting(sessionID: sessionID) == 1) + consumer.cancel() + await consumer.value + #expect(await waitUntil(timeout: .seconds(2)) { + await server.sessionActiveRequestCountForTesting(sessionID: sessionID) == 0 + }) + await server.runSessionCleanupForTesting(now: .distantFuture) + #expect(await server.sessionActiveRequestCountForTesting(sessionID: sessionID) == nil) + } + @Test func streamableHTTPClassifiesAddressInUseBindError() { let configuration = CodexReviewMCPHTTPServer.Configuration( host: "127.0.0.1", @@ -502,6 +576,84 @@ struct CodexReviewMCPHTTPServerTests { try await server.stop() } + @Test func pipelinedOpenSourceStartsOnlyAfterItsFIFOResponseTurn() async throws { + let backend = FakeCodexReviewBackend() + let firstResponseGate = AsyncGate() + await backend.holdStartReview(with: firstResponseGate) + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: backend), + idGenerator: .init(next: { "job-1" }) + ) + let server = CodexReviewMCPHTTPServer( + adapter: CodexReviewMCPServer(store: store), + configuration: .init( + host: "127.0.0.1", + port: 0, + streamHeartbeatInterval: .milliseconds(1) + ) + ) + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + let descriptor = try await openRawTCPConnection(endpoint: endpoint) + var descriptorIsOpen = true + defer { + if descriptorIsOpen { Darwin.close(descriptor) } + } + try await sendRawPipelinedPOSTs( + descriptor: descriptor, + endpoint: endpoint, + sessionID: sessionID, + bodies: [makeJSONBody([ + "jsonrpc": "2.0", + "id": 2, + "method": "tools/call", + "params": [ + "name": "review_start", + "arguments": [ + "cwd": "/tmp/project", + "target": ["type": "uncommittedChanges"], + ], + ], + ])] + ) + await backend.waitForStartReview() + try await sendRawEventStreamRequest( + descriptor: descriptor, + endpoint: endpoint, + sessionID: sessionID + ) + let queued = await waitForNetworkSnapshot(on: server) { snapshot in + guard let requests = snapshot.connections.first(where: { + $0.requests.count == 2 + })?.requests else { return false } + return requests[1].responseSourceKind == .open && requests[1].responseIsReady + } + let requests = try #require(queued.connections.first(where: { + $0.requests.count == 2 + })?.requests) + #expect(requests[1].responseSource == .notReserved) + #expect(requests[1].writer == .notReserved) + + await firstResponseGate.open() + await backend.yield(.completed(summary: "Done", result: "review text")) + _ = await waitForNetworkSnapshot(on: server) { snapshot in + guard let openRequest = snapshot.connections + .flatMap(\.requests) + .first(where: { $0.responseSourceKind == .open }) else { + return false + } + return openRequest.responseSource == .running && openRequest.writer == .running + } + Darwin.shutdown(descriptor, SHUT_RDWR) + Darwin.close(descriptor) + descriptorIsOpen = false + _ = await waitForNetworkSnapshot(on: server) { + $0.connections.flatMap(\.requests).isEmpty + } + try await server.stop() + } + @Test func stopAwaitsSSEWriterCompletionBeforeEventLoopShutdown() async throws { let store = CodexReviewStore.makeTestingStore( backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) @@ -549,6 +701,67 @@ struct CodexReviewMCPHTTPServerTests { #expect(await server.eventLoopGroupShutdownCountForTesting() == 1) } + @Test func nonPersistentRequestsCloseAfterAcknowledgedResponseEnd() async throws { + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) + ) + let server = CodexReviewMCPHTTPServer( + adapter: CodexReviewMCPServer(store: store), + configuration: .init(host: "127.0.0.1", port: 0) + ) + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + let cases: [(version: String, connection: String?, id: Int)] = [ + ("HTTP/1.0", nil, 10), + ("HTTP/1.1", "close", 11), + ] + for testCase in cases { + await server.holdNextResponseEndAcknowledgementForTesting() + let descriptor = try await openRawTCPConnection(endpoint: endpoint) + try await sendRawPOST( + descriptor: descriptor, + endpoint: endpoint, + sessionID: sessionID, + version: testCase.version, + connection: testCase.connection, + body: makeJSONBody([ + "jsonrpc": "2.0", + "id": testCase.id, + "method": "tools/list", + ]) + ) + await server.waitUntilResponseEndAcknowledgementIsHeldForTesting() + let closing = try #require( + await server.networkSnapshotForTesting().connections.first(where: { + $0.phase == .closing(.responseComplete) + && $0.requests.first?.responseEnd == .acknowledged + }) + ) + #expect(closing.requests.first?.responseEnd == .acknowledged) + let response = String( + decoding: try await readRawResponseUntilEOF(descriptor: descriptor), + as: UTF8.self + ) + Darwin.close(descriptor) + #expect(response.contains("\(testCase.version) 200")) + #expect(response.contains("\"id\":\(testCase.id)")) + let acknowledged = try #require( + await waitForNetworkSnapshot(on: server) { snapshot in + snapshot.connections.contains { + $0.id == closing.id && $0.closeAcknowledged + } + }.connections.first(where: { $0.id == closing.id }) + ) + #expect(acknowledged.closeAcknowledged) + await server.releaseResponseEndAcknowledgementForTesting() + _ = await waitForNetworkSnapshot(on: server) { snapshot in + snapshot.connections.contains { $0.id == closing.id } == false + } + } + try await server.stop() + } + @Test func channelCloseOwnsSSETerminationWithoutLateEventLoopCleanup() async throws { let store = CodexReviewStore.makeTestingStore( backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) @@ -656,7 +869,7 @@ struct CodexReviewMCPHTTPServerTests { switch $0.terminalCause { case .peerClosed, .transportFailure(_): return true - case .serverStop, nil: + case .serverStop, .responseComplete, nil: return false } } @@ -1933,6 +2146,74 @@ struct CodexReviewMCPHTTPServerTests { }.value } + private nonisolated func sendRawEventStreamRequest( + descriptor: Int32, + endpoint: URL, + sessionID: String + ) async throws { + try await Task.detached { + let components = try #require(URLComponents(url: endpoint, resolvingAgainstBaseURL: false)) + let host = try #require(components.host) + let port = try #require(components.port) + let request = Data([ + "GET \(endpoint.path) HTTP/1.1", + "Host: \(host):\(port)", + "Accept: text/event-stream, application/json", + "MCP-Session-Id: \(sessionID)", + "MCP-Protocol-Version: 2025-11-25", + "Connection: keep-alive", + "", + "", + ].joined(separator: "\r\n").utf8) + try request.withUnsafeBytes { buffer in + guard let base = buffer.baseAddress else { throw testError("Empty GET request") } + var sent = 0 + while sent < buffer.count { + let count = Darwin.send(descriptor, base.advanced(by: sent), buffer.count - sent, 0) + guard count > 0 else { throw currentPOSIXError() } + sent += count + } + } + }.value + } + + private nonisolated func sendRawPOST( + descriptor: Int32, + endpoint: URL, + sessionID: String, + version: String, + connection: String?, + body: Data + ) async throws { + try await Task.detached { + let components = try #require(URLComponents(url: endpoint, resolvingAgainstBaseURL: false)) + let host = try #require(components.host) + let port = try #require(components.port) + var headers = [ + "POST \(endpoint.path) \(version)", + "Host: \(host):\(port)", + "Accept: text/event-stream, application/json", + "Content-Type: application/json", + "MCP-Session-Id: \(sessionID)", + "MCP-Protocol-Version: 2025-11-25", + "Content-Length: \(body.count)", + ] + if let connection { headers.append("Connection: \(connection)") } + headers.append(contentsOf: ["", ""]) + var request = Data(headers.joined(separator: "\r\n").utf8) + request.append(body) + try request.withUnsafeBytes { buffer in + guard let base = buffer.baseAddress else { throw testError("Empty POST request") } + var sent = 0 + while sent < buffer.count { + let count = Darwin.send(descriptor, base.advanced(by: sent), buffer.count - sent, 0) + guard count > 0 else { throw currentPOSIXError() } + sent += count + } + } + }.value + } + private nonisolated func readTwoRawHTTPResponses(descriptor: Int32) async throws -> Data { try await Task.detached { var response = Data() @@ -1966,6 +2247,22 @@ struct CodexReviewMCPHTTPServerTests { }.value } + private nonisolated func readRawResponseUntilEOF(descriptor: Int32) async throws -> Data { + try await Task.detached { + var response = Data() + var buffer = [UInt8](repeating: 0, count: 4096) + while true { + let count = Darwin.recv(descriptor, &buffer, buffer.count, 0) + if count == 0 { return response } + guard count > 0 else { throw currentPOSIXError() } + response.append(contentsOf: buffer.prefix(count)) + guard response.count <= 2 * 1024 * 1024 else { + throw testError("HTTP response exceeded the test bound") + } + } + }.value + } + private nonisolated func rawConnectionReachedEOF(descriptor: Int32) async -> Bool { await Task.detached { var byte: UInt8 = 0 From 84e3f6c020d962337379db862715527c8e5d4d7b Mon Sep 17 00:00:00 2001 From: Kazuki Nakashima <65545348+lynnswap@users.noreply.github.com> Date: Sat, 22 Aug 2026 19:22:04 +0900 Subject: [PATCH 3/3] fix(mcp): bound response writes and expectations --- .../CodexReviewMCPHTTPServer.swift | 569 ++++++++++++++++-- .../MCPHTTPNetworkResourceOwner.swift | 31 +- .../CodexReviewMCPHTTPServerTests.swift | 182 ++++++ 3 files changed, 722 insertions(+), 60 deletions(-) diff --git a/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift b/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift index 30b2c33..6b79f5f 100644 --- a/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift +++ b/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift @@ -289,6 +289,7 @@ package actor CodexReviewMCPHTTPServer { private let finiteSourceCompletionGate = MCPHTTPLifecycleCompletionGate() private let writerCompletionGate = MCPHTTPLifecycleCompletionGate() private let responseEndAcknowledgementGate = MCPHTTPLifecycleCompletionGate() + private let responseBackpressureProbe = MCPHTTPResponseBackpressureProbe() private var eventLoopGroupShutdownCount = 0 private var nextListenerCloseFailureForTesting: LifecycleError.Failure? private var nextEventLoopGroupShutdownFailureForTesting: LifecycleError.Failure? @@ -403,6 +404,7 @@ package actor CodexReviewMCPHTTPServer { id: UInt64, networkResources: MCPHTTPNetworkResourceOwner ) async -> StartingGenerationResult { + let responseBackpressureProbe = responseBackpressureProbe let group = MultiThreadedEventLoopGroup(numberOfThreads: 1) let bootstrap = ServerBootstrap(group: group) .serverChannelOption(ChannelOptions.backlog, value: 128) @@ -416,6 +418,7 @@ package actor CodexReviewMCPHTTPServer { ).flatMap { channel.pipeline.addHandler(CodexReviewMCPHTTPHandler( server: self, + responseBackpressureProbe: responseBackpressureProbe, connection: connection )) } @@ -931,6 +934,22 @@ package actor CodexReviewMCPHTTPServer { await responseEndAcknowledgementGate.release() } + package func holdNextResponseBodyWriteForTesting() { + responseBackpressureProbe.holdNextBodyWriteForTesting() + } + + package func waitUntilResponseBodyWriteIsHeldForTesting() async { + await responseBackpressureProbe.waitUntilBodyWriteIsHeldForTesting() + } + + package func releaseResponseBodyWriteForTesting() { + responseBackpressureProbe.releaseBodyWriteForTesting() + } + + package func responseSourceReadCountForTesting() -> Int { + responseBackpressureProbe.sourceReadCountForTesting() + } + fileprivate var responseHeartbeatInterval: Duration? { configuration.streamHeartbeatInterval } @@ -1320,6 +1339,314 @@ private final class ActiveRequestCompletion: @unchecked Sendable { } } +private final class MCPHTTPResponseEventChannel: @unchecked Sendable { + enum Event: Sendable { + case body(id: UUID, data: Data) + case heartbeat + case sourceFinished + case sourceFailed(String) + case cancelled + } + + private struct PendingBody { + let id: UUID + let data: Data + let acknowledgement: CheckedContinuation + } + + private let lock = NSLock() + private var pendingBody: PendingBody? + private var inFlightBody: PendingBody? + private var heartbeatPending = false + private var terminal: Event? + private var receiver: CheckedContinuation? + private var isClosed = false + + func sendBody(_ data: Data) async -> Bool { + let id = UUID() + return await withCheckedContinuation { acknowledgement in + var receiver: CheckedContinuation? + lock.lock() + if isClosed || terminal != nil { + lock.unlock() + acknowledgement.resume(returning: false) + return + } + precondition( + pendingBody == nil && inFlightBody == nil, + "The response source waits for each physical body-write acknowledgement." + ) + let body = PendingBody( + id: id, + data: data, + acknowledgement: acknowledgement + ) + if let waitingReceiver = self.receiver { + self.receiver = nil + inFlightBody = body + receiver = waitingReceiver + } else { + pendingBody = body + } + lock.unlock() + receiver?.resume(returning: .body(id: id, data: data)) + } + } + + func receive() async -> Event { + await withCheckedContinuation { continuation in + let immediate: Event? + lock.lock() + if isClosed { + immediate = .cancelled + } else if let terminal { + immediate = terminal + } else if heartbeatPending { + heartbeatPending = false + immediate = .heartbeat + } else if let pendingBody { + self.pendingBody = nil + inFlightBody = pendingBody + immediate = .body(id: pendingBody.id, data: pendingBody.data) + } else { + precondition(receiver == nil, "One response writer owns the channel receiver.") + receiver = continuation + immediate = nil + } + lock.unlock() + if let immediate { + continuation.resume(returning: immediate) + } + } + } + + func acknowledgeBody(id: UUID, wasWritten: Bool) { + let acknowledgement: CheckedContinuation? + lock.lock() + if let inFlightBody, inFlightBody.id == id { + self.inFlightBody = nil + acknowledgement = inFlightBody.acknowledgement + } else { + acknowledgement = nil + } + lock.unlock() + acknowledgement?.resume(returning: wasWritten) + } + + func offerHeartbeat() { + var receiver: CheckedContinuation? + lock.lock() + guard isClosed == false, terminal == nil else { + lock.unlock() + return + } + if pendingBody == nil, inFlightBody == nil, let waitingReceiver = self.receiver { + self.receiver = nil + receiver = waitingReceiver + } else { + heartbeatPending = true + } + lock.unlock() + receiver?.resume(returning: .heartbeat) + } + + func finishSource(_ result: Event) { + precondition({ + switch result { + case .sourceFinished, .sourceFailed, .cancelled: + true + case .body, .heartbeat: + false + } + }()) + var receiver: CheckedContinuation? + lock.lock() + guard isClosed == false, terminal == nil else { + lock.unlock() + return + } + terminal = result + heartbeatPending = false + if pendingBody == nil, inFlightBody == nil { + receiver = self.receiver + self.receiver = nil + } + lock.unlock() + receiver?.resume(returning: result) + } + + func close() { + let acknowledgements: [CheckedContinuation] + let receiver: CheckedContinuation? + lock.lock() + guard isClosed == false else { + lock.unlock() + return + } + isClosed = true + acknowledgements = [pendingBody?.acknowledgement, inFlightBody?.acknowledgement] + .compactMap { $0 } + pendingBody = nil + inFlightBody = nil + receiver = self.receiver + self.receiver = nil + lock.unlock() + for acknowledgement in acknowledgements { + acknowledgement.resume(returning: false) + } + receiver?.resume(returning: .cancelled) + } +} + +private final class MCPHTTPResponseBackpressureProbe: @unchecked Sendable { + private let lock = NSLock() + private var sourceReadCount = 0 + private var holdNextBodyWrite = false + private var releaseWasRequested = false + private var bodyWriteContinuation: CheckedContinuation? + private var bodyWriteWaiters: [CheckedContinuation] = [] + + func holdNextBodyWriteForTesting() { + lock.lock() + precondition(holdNextBodyWrite == false && bodyWriteContinuation == nil) + sourceReadCount = 0 + holdNextBodyWrite = true + releaseWasRequested = false + lock.unlock() + } + + func recordSourceRead() { + lock.lock() + sourceReadCount += 1 + lock.unlock() + } + + func waitBeforeBodyWriteIfNeeded() async { + await withCheckedContinuation { continuation in + let waiters: [CheckedContinuation] + lock.lock() + guard holdNextBodyWrite else { + lock.unlock() + continuation.resume() + return + } + waiters = bodyWriteWaiters + bodyWriteWaiters.removeAll(keepingCapacity: false) + if releaseWasRequested { + resetLocked() + lock.unlock() + for waiter in waiters { waiter.resume() } + continuation.resume() + return + } + bodyWriteContinuation = continuation + lock.unlock() + for waiter in waiters { waiter.resume() } + } + } + + func waitUntilBodyWriteIsHeldForTesting() async { + await withCheckedContinuation { continuation in + lock.lock() + if bodyWriteContinuation != nil { + lock.unlock() + continuation.resume() + } else { + bodyWriteWaiters.append(continuation) + lock.unlock() + } + } + } + + func releaseBodyWriteForTesting() { + let continuation: CheckedContinuation? + lock.lock() + guard holdNextBodyWrite else { + lock.unlock() + return + } + if let held = bodyWriteContinuation { + bodyWriteContinuation = nil + continuation = held + resetLocked() + } else { + releaseWasRequested = true + continuation = nil + } + lock.unlock() + continuation?.resume() + } + + func sourceReadCountForTesting() -> Int { + lock.lock() + let count = sourceReadCount + lock.unlock() + return count + } + + private func resetLocked() { + holdNextBodyWrite = false + releaseWasRequested = false + } +} + +private final class MCPHTTPRequestBodyReceipt: @unchecked Sendable { + private enum State { + case waiting + case suspended(CheckedContinuation) + case completed(Value?) + } + + private let lock = NSLock() + private var state: State = .waiting + + func wait() async -> Value? { + await withTaskCancellationHandler { + await withCheckedContinuation { continuation in + lock.lock() + switch state { + case .waiting: + state = .suspended(continuation) + lock.unlock() + case .completed(let value): + lock.unlock() + continuation.resume(returning: value) + case .suspended: + lock.unlock() + preconditionFailure("One request handler owns the body receipt.") + } + } + } onCancel: { + self.cancel() + } + } + + func complete(_ value: Value) { + finish(value) + } + + func cancel() { + finish(nil) + } + + private func finish(_ value: Value?) { + let continuation: CheckedContinuation? + lock.lock() + switch state { + case .waiting: + state = .completed(value) + continuation = nil + case .suspended(let suspended): + state = .completed(value) + continuation = suspended + case .completed: + continuation = nil + } + lock.unlock() + continuation?.resume(returning: value) + } +} + private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked Sendable { typealias InboundIn = HTTPServerRequestPart typealias OutboundOut = HTTPServerResponsePart @@ -1345,27 +1672,29 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked } } - private struct RequestState { + private struct CompletedRequestState: @unchecked Sendable { var head: HTTPRequestHead var bodyBuffer: ByteBuffer } - private enum WriterEvent: Sendable { - case body(Data) - case heartbeat - case sourceFinished - case sourceFailed(String) + private struct RequestState { + var head: HTTPRequestHead + var bodyBuffer: ByteBuffer + let bodyReceipt: MCPHTTPRequestBodyReceipt } private let server: CodexReviewMCPHTTPServer + private let responseBackpressureProbe: MCPHTTPResponseBackpressureProbe private let connection: MCPHTTPNetworkResourceOwner.Connection private var requestState: RequestState? init( server: CodexReviewMCPHTTPServer, + responseBackpressureProbe: MCPHTTPResponseBackpressureProbe, connection: MCPHTTPNetworkResourceOwner.Connection ) { self.server = server + self.responseBackpressureProbe = responseBackpressureProbe self.connection = connection } @@ -1373,9 +1702,50 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked let part = unwrapInboundIn(data) switch part { case .head(let head): + precondition(requestState == nil, "HTTP decoding serializes request bodies on one connection.") + let sendsContinue: Bool + if let expectation = head.headers.first(name: "Expect") { + let isContinue = expectation.trimmingCharacters(in: .whitespacesAndNewlines) + .caseInsensitiveCompare("100-continue") == .orderedSame + guard isContinue else { + admitUnsupportedExpectation( + expectation, + version: head.version, + context: context + ) + return + } + if head.version.major == 1, head.version.minor == 0 { + sendsContinue = false + } else if head.version.major == 1, head.version.minor >= 1 { + sendsContinue = true + } else { + admitUnsupportedExpectation( + expectation, + version: head.version, + context: context + ) + return + } + } else { + sendsContinue = false + } + guard let admittedRequest = connection.admitRequest() else { + context.close(promise: nil) + return + } + let bodyReceipt = MCPHTTPRequestBodyReceipt() requestState = RequestState( head: head, - bodyBuffer: context.channel.allocator.buffer(capacity: 0) + bodyBuffer: context.channel.allocator.buffer(capacity: 0), + bodyReceipt: bodyReceipt + ) + startRequestHandler( + admittedRequest, + bodyReceipt: bodyReceipt, + sendsContinue: sendsContinue, + version: head.version, + context: context ) case .body(var buffer): requestState?.bodyBuffer.writeBuffer(&buffer) @@ -1384,40 +1754,106 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked return } requestState = nil - guard let admittedRequest = connection.admitRequest() else { - context.close(promise: nil) + state.bodyReceipt.complete(.init( + head: state.head, + bodyBuffer: state.bodyBuffer + )) + } + } + + private func startRequestHandler( + _ admittedRequest: MCPHTTPNetworkResourceOwner.Connection.AdmittedRequest, + bodyReceipt: MCPHTTPRequestBodyReceipt, + sendsContinue: Bool, + version: HTTPVersion, + context: ChannelHandlerContext + ) { + nonisolated(unsafe) let context = context + let task = Task { [self] in + defer { admittedRequest.lease.acknowledgeCompletion() } + guard await admittedRequest.lease.waitUntilStartIsAllowed() else { + bodyReceipt.cancel() return } - nonisolated(unsafe) let context = context - let task = Task { [self] in - defer { - admittedRequest.lease.acknowledgeCompletion() - } - guard await admittedRequest.lease.waitUntilStartIsAllowed() else { + if sendsContinue { + guard await connection.supplyExpectation(for: admittedRequest.operation), + await writeContinue(version: version, context: context) else { + bodyReceipt.cancel() return } - await handleRequest( - state: state, - operation: admittedRequest.operation, - context: context - ) } - admittedRequest.lease.install(task) + guard let state = await bodyReceipt.wait(), Task.isCancelled == false else { + return + } + await handleRequest( + state: state, + operation: admittedRequest.operation, + context: context + ) + } + admittedRequest.lease.install(task) + } + + private func writeContinue( + version: HTTPVersion, + context: ChannelHandlerContext + ) async -> Bool { + do { + try await writeResponsePart( + .head(.init(version: version, status: .continue)), + context: context, + eventLoop: context.eventLoop + ) + return true + } catch { + connection.transportFailed(error.localizedDescription) + return false } } + private func admitUnsupportedExpectation( + _ expectation: String, + version: HTTPVersion, + context: ChannelHandlerContext + ) { + guard let admittedRequest = connection.admitRequest() else { + context.close(promise: nil) + return + } + nonisolated(unsafe) let context = context + let task = Task { [self] in + defer { admittedRequest.lease.acknowledgeCompletion() } + guard await admittedRequest.lease.waitUntilStartIsAllowed() else { return } + await prepareAndQueueResponse( + .init(response: .error( + statusCode: Int(HTTPResponseStatus.expectationFailed.code), + .invalidRequest("Unsupported HTTP expectation: \(expectation)") + )), + operation: admittedRequest.operation, + version: version, + closeAfterResponse: true, + context: context + ) + } + admittedRequest.lease.install(task) + } + func channelReadComplete(context: ChannelHandlerContext) { context.flush() context.read() } func channelInactive(context: ChannelHandlerContext) { + requestState?.bodyReceipt.cancel() + requestState = nil connection.peerClosed() context.fireChannelInactive() } func userInboundEventTriggered(context: ChannelHandlerContext, event: Any) { if case ChannelEvent.inputClosed = event { + requestState?.bodyReceipt.cancel() + requestState = nil connection.peerClosed() context.close(promise: nil) return @@ -1426,12 +1862,14 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked } func errorCaught(context: ChannelHandlerContext, error: any Error) { + requestState?.bodyReceipt.cancel() + requestState = nil connection.transportFailed(error.localizedDescription) context.close(promise: nil) } private func handleRequest( - state: RequestState, + state: CompletedRequestState, operation: MCPHTTPNetworkResourceOwner.RequestOperation, context: ChannelHandlerContext ) async { @@ -1460,7 +1898,7 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked ) } - private func makeHTTPRequest(from state: RequestState) -> HTTPRequest { + private func makeHTTPRequest(from state: CompletedRequestState) -> HTTPRequest { var headers: [String: String] = [:] for (name, value) in state.head.headers { if let existing = headers[name] { @@ -1525,63 +1963,77 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked } let preparedResponse: HTTPResponse + let responseEvents: MCPHTTPResponseEventChannel? var admittedSource: (lease: MCPHTTPNetworkResourceOwner.WorkLease, task: Task)? switch response { - case .stream(let source, let headers): + case .stream(let source, _): guard let kind = trackedResponse.responseSourceKind, let sourceLease = operation.reserveResponseSource() else { await trackedResponse.streamCompletion?.finishAndWait() writerLease.acknowledgeCompletion() return } - let bridge = AsyncThrowingStream.makeStream( - bufferingPolicy: .unbounded - ) + let events = MCPHTTPResponseEventChannel() let sourceTask = Task { let started = await sourceLease.waitUntilStartIsAllowed() if started { do { for try await chunk in source { try Task.checkCancellation() - bridge.continuation.yield(chunk) +#if DEBUG + self.responseBackpressureProbe.recordSourceRead() +#endif + guard await events.sendBody(chunk) else { + break + } } if kind == .finite { await self.server.waitAfterFiniteSourceCompletionForTesting() } - bridge.continuation.finish() + events.finishSource(.sourceFinished) } catch is CancellationError { - bridge.continuation.finish() + events.finishSource(.cancelled) } catch { - bridge.continuation.finish(throwing: error) + events.finishSource(.sourceFailed(error.localizedDescription)) self.connection.transportFailed(error.localizedDescription) } } else { - bridge.continuation.finish() + events.finishSource(.cancelled) } await trackedResponse.streamCompletion?.finishAndWait() sourceLease.acknowledgeCompletion() } admittedSource = (sourceLease, sourceTask) - preparedResponse = .stream(bridge.stream, headers: headers) + preparedResponse = response + responseEvents = events default: admittedSource = nil preparedResponse = response + responseEvents = nil } nonisolated(unsafe) let context = context let heartbeatInterval = await server.responseHeartbeatInterval let writerTask = Task { - guard await writerLease.waitUntilStartIsAllowed() else { + let result: WriterCompletion? = await withTaskCancellationHandler { + guard await writerLease.waitUntilStartIsAllowed() else { + return nil + } + return await self.writeResponse( + preparedResponse, + responseEvents: responseEvents, + version: version, + context: context, + heartbeatInterval: heartbeatInterval + ) + } onCancel: { + responseEvents?.close() + } + guard let result else { writerLease.acknowledgeCompletion() return } - let result = await self.writeResponse( - preparedResponse, - version: version, - context: context, - heartbeatInterval: heartbeatInterval - ) switch result { case .responded: operation.acknowledgeResponseEnd() @@ -1616,6 +2068,7 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked private func writeResponse( _ response: HTTPResponse, + responseEvents: MCPHTTPResponseEventChannel?, version: HTTPVersion, context: ChannelHandlerContext, heartbeatInterval: Duration? @@ -1629,28 +2082,18 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked } switch response { - case .stream(let stream, _): + case .stream: + guard let responseEvents else { + return .transportFailed("Streaming response has no bounded event channel.") + } + defer { responseEvents.close() } do { try Task.checkCancellation() try await writeResponsePart(.head(head), context: context, eventLoop: eventLoop) - let events = AsyncStream.makeStream(bufferingPolicy: .unbounded) let sourceResult = await withTaskGroup( of: Void.self, returning: WriterCompletion.self ) { group in - group.addTask { - do { - for try await chunk in stream { - try Task.checkCancellation() - events.continuation.yield(.body(chunk)) - } - events.continuation.yield(.sourceFinished) - } catch is CancellationError { - events.continuation.yield(.sourceFinished) - } catch { - events.continuation.yield(.sourceFailed(error.localizedDescription)) - } - } if let heartbeatInterval { group.addTask { while Task.isCancelled == false { @@ -1662,26 +2105,29 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked guard Task.isCancelled == false else { return } - events.continuation.yield(.heartbeat) + responseEvents.offerHeartbeat() } } } var result: WriterCompletion = .cancelled - eventLoopLoop: for await event in events.stream { + eventLoopLoop: while true { + let event = await responseEvents.receive() if Task.isCancelled { result = .cancelled break eventLoopLoop } switch event { - case .body(let data): + case .body(let id, let data): do { try await writeResponseBody( data, context: context, eventLoop: eventLoop ) + responseEvents.acknowledgeBody(id: id, wasWritten: true) } catch { + responseEvents.acknowledgeBody(id: id, wasWritten: false) result = .transportFailed(error.localizedDescription) break eventLoopLoop } @@ -1702,10 +2148,12 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked case .sourceFailed(let message): result = .sourceFailed(message) break eventLoopLoop + case .cancelled: + result = .cancelled + break eventLoopLoop } } group.cancelAll() - events.continuation.finish() return result } switch sourceResult { @@ -1762,6 +2210,9 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked context: ChannelHandlerContext, eventLoop: any EventLoop ) async throws { +#if DEBUG + await responseBackpressureProbe.waitBeforeBodyWriteIfNeeded() +#endif let writer = ResponsePartWriter(handler: self, context: context) let promise = eventLoop.makePromise(of: Void.self) eventLoop.execute { diff --git a/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift b/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift index 3fee056..4c58a47 100644 --- a/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift +++ b/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift @@ -362,7 +362,28 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { fileprivate func markResponseReady() -> Bool { lock.lock() guard case .open = outcome, - responseEnd == .pending, + responseEnd == .pending else { + lock.unlock() + return false + } + switch responseQueueState { + case .handling: + responseQueueState = .ready(nil) + case .turnGranted: + break + case .ready: + lock.unlock() + return false + } + lock.unlock() + notifyChanged() + return true + } + + fileprivate func markExpectationReady() -> Bool { + lock.lock() + guard case .open = outcome, + responseEnd == .notExpected, case .handling = responseQueueState else { lock.unlock() return false @@ -646,6 +667,14 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { return await operation.waitForWriterTurn() } + package func supplyExpectation(for operation: RequestOperation) async -> Bool { + guard operation.markExpectationReady() else { + return false + } + pumpWriterQueue() + return await operation.waitForWriterTurn() + } + package func peerClosed() { beginClosing(.peerClosed, signalResourceClose: false) } diff --git a/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift b/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift index 5dea1f6..8dd1540 100644 --- a/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift +++ b/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift @@ -654,6 +654,38 @@ struct CodexReviewMCPHTTPServerTests { try await server.stop() } + @Test func activeWriterAcknowledgesEachBodyBeforeReadingTheNextSourceChunk() async throws { + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) + ) + let server = CodexReviewMCPHTTPServer( + adapter: CodexReviewMCPServer(store: store), + configuration: .init(host: "127.0.0.1", port: 0) + ) + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + await server.holdNextResponseBodyWriteForTesting() + let responseTask = Task { + try await postJSONRPCData( + endpoint: endpoint, + sessionID: sessionID, + bodyData: makeJSONBody([ + "jsonrpc": "2.0", + "id": 2, + "method": "tools/list", + ]) + ) + } + await server.waitUntilResponseBodyWriteIsHeldForTesting() + + #expect(await server.responseSourceReadCountForTesting() == 1) + await server.releaseResponseBodyWriteForTesting() + _ = try decodeSSEJSON(from: try await responseTask.value) + #expect(await server.responseSourceReadCountForTesting() >= 2) + try await server.stop() + } + @Test func stopAwaitsSSEWriterCompletionBeforeEventLoopShutdown() async throws { let store = CodexReviewStore.makeTestingStore( backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) @@ -762,6 +794,87 @@ struct CodexReviewMCPHTTPServerTests { try await server.stop() } + @Test func expectContinueIsAcknowledgedBeforeTheRequestBody() async throws { + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) + ) + let server = CodexReviewMCPHTTPServer( + adapter: CodexReviewMCPServer(store: store), + configuration: .init(host: "127.0.0.1", port: 0) + ) + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + let body = try makeJSONBody([ + "jsonrpc": "2.0", + "id": 20, + "method": "tools/list", + ]) + let descriptor = try await openRawTCPConnection(endpoint: endpoint) + defer { Darwin.close(descriptor) } + try await sendRawExpectHeaders( + descriptor: descriptor, + endpoint: endpoint, + sessionID: sessionID, + expectation: "100-continue", + contentLength: body.count + ) + let interim = String( + decoding: try await readRawHTTPHeadersData(descriptor: descriptor), + as: UTF8.self + ) + #expect(interim.hasPrefix("HTTP/1.1 100 Continue")) + + try await sendRawBytes(descriptor: descriptor, bytes: body) + let final = String( + decoding: try await readOneRawHTTPResponse(descriptor: descriptor), + as: UTF8.self + ) + #expect(final.contains("HTTP/1.1 200")) + #expect(final.contains("\"id\":20")) + + let legacyBody = try makeJSONBody([ + "jsonrpc": "2.0", + "id": 21, + "method": "tools/list", + ]) + let legacy = try await openRawTCPConnection(endpoint: endpoint) + try await sendRawExpectHeaders( + descriptor: legacy, + endpoint: endpoint, + sessionID: sessionID, + version: "HTTP/1.0", + connection: nil, + expectation: "100-continue", + contentLength: legacyBody.count + ) + try await sendRawBytes(descriptor: legacy, bytes: legacyBody) + let legacyResponse = String( + decoding: try await readRawResponseUntilEOF(descriptor: legacy), + as: UTF8.self + ) + Darwin.close(legacy) + #expect(legacyResponse.contains("100 Continue") == false) + #expect(legacyResponse.contains("HTTP/1.0 200")) + #expect(legacyResponse.contains("\"id\":21")) + + let unsupported = try await openRawTCPConnection(endpoint: endpoint) + try await sendRawExpectHeaders( + descriptor: unsupported, + endpoint: endpoint, + sessionID: sessionID, + expectation: "unsupported", + contentLength: body.count + ) + let rejected = String( + decoding: try await readRawResponseUntilEOF(descriptor: unsupported), + as: UTF8.self + ) + Darwin.close(unsupported) + #expect(rejected.hasPrefix("HTTP/1.1 417 Expectation Failed")) + try await server.stop() + } + @Test func channelCloseOwnsSSETerminationWithoutLateEventLoopCleanup() async throws { let store = CodexReviewStore.makeTestingStore( backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) @@ -2214,6 +2327,51 @@ struct CodexReviewMCPHTTPServerTests { }.value } + private nonisolated func sendRawExpectHeaders( + descriptor: Int32, + endpoint: URL, + sessionID: String, + version: String = "HTTP/1.1", + connection: String? = "keep-alive", + expectation: String, + contentLength: Int + ) async throws { + let components = try #require(URLComponents(url: endpoint, resolvingAgainstBaseURL: false)) + let host = try #require(components.host) + let port = try #require(components.port) + var headerLines = [ + "POST \(endpoint.path) \(version)", + "Host: \(host):\(port)", + "Accept: text/event-stream, application/json", + "Content-Type: application/json", + "MCP-Session-Id: \(sessionID)", + "MCP-Protocol-Version: 2025-11-25", + "Expect: \(expectation)", + "Content-Length: \(contentLength)", + ] + if let connection { headerLines.append("Connection: \(connection)") } + headerLines.append(contentsOf: ["", ""]) + let headers = Data(headerLines.joined(separator: "\r\n").utf8) + try await sendRawBytes(descriptor: descriptor, bytes: headers) + } + + private nonisolated func sendRawBytes( + descriptor: Int32, + bytes: Data + ) async throws { + try await Task.detached { + try bytes.withUnsafeBytes { buffer in + guard let base = buffer.baseAddress else { throw testError("Empty socket write") } + var sent = 0 + while sent < buffer.count { + let count = Darwin.send(descriptor, base.advanced(by: sent), buffer.count - sent, 0) + guard count > 0 else { throw currentPOSIXError() } + sent += count + } + } + }.value + } + private nonisolated func readTwoRawHTTPResponses(descriptor: Int32) async throws -> Data { try await Task.detached { var response = Data() @@ -2234,7 +2392,30 @@ struct CodexReviewMCPHTTPServerTests { }.value } + private nonisolated func readOneRawHTTPResponse(descriptor: Int32) async throws -> Data { + try await Task.detached { + var response = Data() + var buffer = [UInt8](repeating: 0, count: 4096) + while true { + let text = String(decoding: response, as: UTF8.self) + if text.contains("HTTP/1.1 200"), text.contains("\r\n0\r\n\r\n") { + return response + } + let count = Darwin.recv(descriptor, &buffer, buffer.count, 0) + guard count > 0 else { throw testError("Connection closed before the response ended") } + response.append(contentsOf: buffer.prefix(count)) + guard response.count <= 2 * 1024 * 1024 else { + throw testError("HTTP response exceeded the test bound") + } + } + }.value + } + private nonisolated func readRawHTTPHeaders(descriptor: Int32) async throws { + _ = try await readRawHTTPHeadersData(descriptor: descriptor) + } + + private nonisolated func readRawHTTPHeadersData(descriptor: Int32) async throws -> Data { try await Task.detached { var response = Data() var buffer = [UInt8](repeating: 0, count: 1024) @@ -2244,6 +2425,7 @@ struct CodexReviewMCPHTTPServerTests { response.append(contentsOf: buffer.prefix(count)) guard response.count < 8192 else { throw testError("Response headers exceeded test bound") } } + return response }.value }