From 7791ee9b66c0a469fcc4f1934cee7401bf134fc0 Mon Sep 17 00:00:00 2001 From: Kazuki Nakashima <65545348+lynnswap@users.noreply.github.com> Date: Sat, 22 Aug 2026 21:34:45 +0900 Subject: [PATCH] fix(mcp): own HTTP response lifecycle --- .../CodexReviewMCPHTTPServer.swift | 811 ++++++++++++++---- .../MCPHTTPNetworkResourceOwner.swift | 146 +++- .../CodexReviewMCPHTTPServerTests.swift | 313 ++++++- 3 files changed, 1091 insertions(+), 179 deletions(-) diff --git a/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift b/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift index f611868..7aa0259 100644 --- a/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift +++ b/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift @@ -8,9 +8,20 @@ import OSLog private let logger = Logger(subsystem: "CodexReviewKit", category: "mcp-http") +private enum MCPHTTPResponseSourceKind: Sendable { + case finite + case open +} + private struct TrackedHTTPResponse { + struct StreamLifecycle: Sendable { + let source: AsyncThrowingStream + let completion: ActiveRequestCompletion + let kind: MCPHTTPResponseSourceKind + } + var response: HTTPResponse - var streamCompletion: ActiveRequestCompletion? = nil + var streamLifecycle: StreamLifecycle? = nil } package extension CodexReviewMCPHTTPServer { @@ -294,6 +305,11 @@ 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 let responseEndWriteGate = MCPHTTPLifecycleCompletionGate() + private let responseBackpressureProbe = MCPHTTPResponseBackpressureProbe() private var eventLoopGroupShutdownCount = 0 private var nextListenerCloseFailureForTesting: LifecycleError.Failure? private var nextEventLoopGroupShutdownFailureForTesting: LifecycleError.Failure? @@ -410,6 +426,7 @@ package actor CodexReviewMCPHTTPServer { ) async -> StartingGenerationResult { let group = MultiThreadedEventLoopGroup(numberOfThreads: 1) let maximumRequestBodyBytes = configuration.maximumRequestBodyBytes + let responseBackpressureProbe = responseBackpressureProbe let bootstrap = ServerBootstrap(group: group) .serverChannelOption(ChannelOptions.backlog, value: 128) .serverChannelOption(ChannelOptions.socketOption(.so_reuseaddr), value: 1) @@ -421,7 +438,8 @@ package actor CodexReviewMCPHTTPServer { channel.pipeline.addHandler(CodexReviewMCPHTTPHandler( server: self, connection: connection, - maximumRequestBodyBytes: maximumRequestBodyBytes + maximumRequestBodyBytes: maximumRequestBodyBytes, + responseBackpressureProbe: responseBackpressureProbe )) } } @@ -888,12 +906,128 @@ package actor CodexReviewMCPHTTPServer { } } + 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() + } + + package func holdNextResponseEndWriteForTesting() async { + await responseEndWriteGate.holdNextCompletion() + } + + package func waitUntilResponseEndWriteIsHeldForTesting() async { + await responseEndWriteGate.waitUntilHolding() + } + + package func releaseResponseEndWriteForTesting() async { + await responseEndWriteGate.release() + } + + package func holdNextResponseBodyWriteForTesting() { + responseBackpressureProbe.holdNextBodyWriteForTesting() + } + + package func waitUntilResponseBodyWriteIsHeldForTesting() async { + await responseBackpressureProbe.waitUntilBodyWriteIsHeldForTesting() + } + + package func releaseResponseBodyWriteForTesting() { + responseBackpressureProbe.releaseBodyWriteForTesting() + } + + package func responseSourceReadCountForTesting() -> Int { + responseBackpressureProbe.sourceReadCountForTesting() + } + + package static func responseRendezvousPrioritizesBodyForTesting() async -> Bool { + let events = MCPHTTPResponseEventChannel() + let body = Data("body".utf8) + let sender = Task { await events.sendBody(body) } + await events.waitUntilBodyIsPendingForTesting() + for _ in 0..<3 { events.offerHeartbeat() } + guard case .body(let id, let received) = await events.receive() else { + events.close() + return false + } + events.acknowledgeBody(id: id, wasWritten: true) + let wasAcknowledged = await sender.value + guard case .heartbeat = await events.receive() else { + events.close() + return false + } + events.finishSource(.sourceFinished) + guard case .sourceFinished = await events.receive() else { + events.close() + return false + } + events.close() + return received == body && wasAcknowledged + } + + 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() + } + + fileprivate func waitBeforeResponseEndWriteForTesting() async { + await responseEndWriteGate.waitIfNeeded() + } + package func eventLoopGroupShutdownCountForTesting() -> Int { eventLoopGroupShutdownCount } - package func handleHTTPRequest(_ request: HTTPRequest) async -> HTTPResponse { - await handleTrackedHTTPRequest(request).response + package func validationResponseForTesting(_ request: HTTPRequest) -> HTTPResponse? { + makeValidationPipeline().validate( + request, + context: .init( + httpMethod: request.method, + sessionID: request.header(HTTPHeaderName.sessionID), + isInitializationRequest: request.body.map(Self.isInitializeRequest) ?? false + ) + ) } fileprivate func handleTrackedHTTPRequest(_ request: HTTPRequest) async -> TrackedHTTPResponse { @@ -904,7 +1038,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) } @@ -957,7 +1095,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() @@ -985,39 +1127,23 @@ package actor CodexReviewMCPHTTPServer { private func trackActiveRequest( _ response: HTTPResponse, - sessionID: String + sessionID: String, + responseSourceKind: MCPHTTPResponseSourceKind ) -> (response: TrackedHTTPResponse, didFinishRequest: Bool) { switch response { - case .stream(let stream, let headers): - 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() - } + case .stream(let source, _): + let completion = ActiveRequestCompletion { [weak self] in + await self?.finishActiveRequest(sessionID: sessionID) } return ( - .init(response: .stream(trackedStream, headers: headers), streamCompletion: completion), + .init( + response: response, + streamLifecycle: .init( + source: source, + completion: completion, + kind: responseSourceKind + ) + ), false ) @@ -1035,27 +1161,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) @@ -1115,6 +1220,17 @@ package actor CodexReviewMCPHTTPServer { return json["method"] as? String == "initialize" && json["id"] != nil } + private static func responseSourceKind(for request: HTTPRequest) -> MCPHTTPResponseSourceKind { + 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) : "*" @@ -1186,22 +1302,282 @@ 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() { + func finishAndWait() async { + let shouldFinish = lock.withLock { + guard didFinish == false else { + return false + } + didFinish = true + return true + } + guard shouldFinish else { return } + await onFinish() + } +} + +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 pendingBodyWaiters: [CheckedContinuation] = [] + private var isClosed = false + + func sendBody(_ data: Data) async -> Bool { + let id = UUID() + return await withCheckedContinuation { acknowledgement in + let receiver: CheckedContinuation? + let waiters: [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 + receiver = nil + } + waiters = pendingBodyWaiters + pendingBodyWaiters.removeAll(keepingCapacity: false) + lock.unlock() + for waiter in waiters { waiter.resume() } + 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 pendingBody { + self.pendingBody = nil + inFlightBody = pendingBody + immediate = .body(id: pendingBody.id, data: pendingBody.data) + } else if let terminal { + immediate = terminal + } else if heartbeatPending { + heartbeatPending = false + immediate = .heartbeat + } 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() { + let 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 + receiver = nil + } + lock.unlock() + receiver?.resume(returning: .heartbeat) + } + + func finishSource(_ result: Event) { + let 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 + } else { + 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) + } + + func waitUntilBodyIsPendingForTesting() async { + await withCheckedContinuation { continuation in + lock.lock() + if pendingBody != nil || inFlightBody != nil { + lock.unlock() + continuation.resume() + } else { + pendingBodyWaiters.append(continuation) + lock.unlock() + } + } + } +} + +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() - if didFinish { + 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 } - didFinish = true + 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() - onFinish() + return count + } + + private func resetLocked() { + holdNextBodyWrite = false + releaseWasRequested = false } } @@ -1334,19 +1710,19 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked private let server: CodexReviewMCPHTTPServer private let connection: MCPHTTPNetworkResourceOwner.Connection private let maximumRequestBodyBytes: Int + private let responseBackpressureProbe: MCPHTTPResponseBackpressureProbe private var requestState: RequestState? - private var activeStreamTask: Task? - private var activeStreamID: UUID? - private var activeStreamCompletion: ActiveRequestCompletion? init( server: CodexReviewMCPHTTPServer, connection: MCPHTTPNetworkResourceOwner.Connection, - maximumRequestBodyBytes: Int + maximumRequestBodyBytes: Int, + responseBackpressureProbe: MCPHTTPResponseBackpressureProbe ) { self.server = server self.connection = connection self.maximumRequestBodyBytes = maximumRequestBodyBytes + self.responseBackpressureProbe = responseBackpressureProbe } func channelRead(context: ChannelHandlerContext, data: NIOAny) { @@ -1400,6 +1776,9 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked guard await admittedRequest.lease.waitUntilStartIsAllowed() else { return } + guard admittedRequest.operation.beginResponse() else { + return + } if shouldSendContinue { do { let responseHead = HTTPResponseHead(version: head.version, status: .continue) @@ -1419,17 +1798,24 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked } switch bodyResult { case .body(let body): - await handleRequest(head: head, body: body, context: context) + await handleRequest( + head: head, + body: body, + operation: admittedRequest.operation, + context: context + ) case .payloadTooLarge: await writeRequestRejection( status: .payloadTooLarge, version: head.version, + operation: admittedRequest.operation, context: context ) case .expectationFailed: await writeRequestRejection( status: .expectationFailed, version: head.version, + operation: admittedRequest.operation, context: context ) case .cancelled: @@ -1478,7 +1864,6 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked requestState?.bodyReceipt.reject(.cancelled) requestState = nil connection.peerClosed() - finishActiveStream() context.fireChannelInactive() } @@ -1487,7 +1872,6 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked requestState?.bodyReceipt.reject(.cancelled) requestState = nil connection.peerClosed() - finishActiveStream() context.close(promise: nil) return } @@ -1498,21 +1882,13 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked requestState?.bodyReceipt.reject(.cancelled) requestState = nil 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( head: HTTPRequestHead, body: Data?, + operation: MCPHTTPNetworkResourceOwner.RequestOperation, context: ChannelHandlerContext ) async { let path = head.uri.split(separator: "?").first.map(String.init) ?? head.uri @@ -1521,6 +1897,8 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked await writeResponse( .init(response: .error(statusCode: 404, .invalidRequest("Not Found"))), version: head.version, + operation: operation, + closeAfterResponse: head.isKeepAlive == false, context: context ) return @@ -1528,7 +1906,13 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked let request = makeHTTPRequest(head: head, body: body) let response = await server.handleTrackedHTTPRequest(request) - await writeResponse(response, version: head.version, context: context) + await writeResponse( + response, + version: head.version, + operation: operation, + closeAfterResponse: head.isKeepAlive == false, + context: context + ) } private func makeHTTPRequest(head: HTTPRequestHead, body: Data?) -> HTTPRequest { @@ -1553,6 +1937,7 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked private func writeRequestRejection( status: HTTPResponseStatus, version: HTTPVersion, + operation: MCPHTTPNetworkResourceOwner.RequestOperation, context: ChannelHandlerContext ) async { nonisolated(unsafe) let context = context @@ -1563,17 +1948,19 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked do { try await writeResponsePart(.head(head), context: context, eventLoop: eventLoop) try await writeResponsePart(.end(nil), context: context, eventLoop: eventLoop) + operation.acknowledgeResponseEnd() + await connection.closeAfterResponse() } catch { logger.error("MCP request rejection write failed: \(error.localizedDescription, privacy: .public)") - } - eventLoop.execute { - context.close(promise: nil) + connection.transportFailed(error.localizedDescription) } } private func writeResponse( _ trackedResponse: TrackedHTTPResponse, version: HTTPVersion, + operation: MCPHTTPNetworkResourceOwner.RequestOperation, + closeAfterResponse: Bool, context: ChannelHandlerContext ) async { nonisolated(unsafe) let context = context @@ -1581,81 +1968,215 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked 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) + } + if closeAfterResponse { + head.headers.replaceOrAdd(name: "Connection", value: "close") + } + let result: WriterCompletion 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) - } + case .stream: + if let lifecycle = trackedResponse.streamLifecycle { + result = await writeStreamingResponse( + lifecycle, + head: head, + context: context, + eventLoop: eventLoop + ) + } else { + result = .transportFailed("Streaming response has no source lifecycle.") + } - 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) - } - } catch is CancellationError { - trackedResponse.streamCompletion?.finish() - return - } catch { - trackedResponse.streamCompletion?.finish() - logger.error("MCP SSE stream failed: \(error.localizedDescription, privacy: .public)") - } + default: + result = await writeNonStreamingResponse( + response, + head: head, + context: context, + eventLoop: eventLoop + ) + } - guard Task.isCancelled == false else { - return - } - try? await writeResponsePart(.end(nil), context: context, eventLoop: eventLoop) + switch result { + case .responded: + operation.acknowledgeResponseEnd() + if closeAfterResponse { + await connection.closeAfterResponse() } - eventLoop.execute { - context.channel.closeFuture.whenComplete { _ in - trackedResponse.streamCompletion?.finish() - streamTask.cancel() + await server.waitAfterResponseEndAcknowledgementForTesting() + case .cancelled: + break + case .sourceFailed(let message): + logger.error("MCP SSE source failed: \(message, privacy: .public)") + connection.transportFailed(message) + case .transportFailed(let message): + logger.error("MCP HTTP response failed: \(message, privacy: .public)") + connection.transportFailed(message) + } + await server.waitAfterWriterCompletionForTesting() + } + + private enum WriterCompletion: Sendable { + case responded + case cancelled + case sourceFailed(String) + case transportFailed(String) + } + + private enum ResponseChildCompletion: Sendable { + case source + case heartbeat + case writer(WriterCompletion) + } + + private func writeStreamingResponse( + _ lifecycle: TrackedHTTPResponse.StreamLifecycle, + head: HTTPResponseHead, + context: ChannelHandlerContext, + eventLoop: any EventLoop + ) async -> WriterCompletion { + nonisolated(unsafe) let context = context + let events = MCPHTTPResponseEventChannel() + let heartbeatInterval = await server.responseHeartbeatInterval + return await withTaskCancellationHandler { + await withTaskGroup(of: ResponseChildCompletion.self) { group in + group.addTask { [self] in + do { + for try await chunk in lifecycle.source { + try Task.checkCancellation() +#if DEBUG + responseBackpressureProbe.recordSourceRead() +#endif + guard await events.sendBody(chunk) else { break } + } + if lifecycle.kind == .finite { + await server.waitAfterFiniteSourceCompletionForTesting() + } + events.finishSource(Task.isCancelled ? .cancelled : .sourceFinished) + } catch let error as CancellationError { + if Task.isCancelled { + events.finishSource(.cancelled) + } else { + events.finishSource(.sourceFailed(error.localizedDescription)) + } + } catch { + events.finishSource(.sourceFailed(error.localizedDescription)) + } + await lifecycle.completion.finishAndWait() + return .source } - guard context.channel.isActive else { - trackedResponse.streamCompletion?.finish() - streamTask.cancel() - return + if let heartbeatInterval { + group.addTask { + while Task.isCancelled == false { + do { + try await Task.sleep(for: heartbeatInterval) + } catch { + return .heartbeat + } + events.offerHeartbeat() + } + return .heartbeat + } } - 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 + group.addTask { [self] in + .writer(await consumeResponseEvents( + events, + head: head, + context: context, + eventLoop: eventLoop + )) } - } - 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)") + var result = WriterCompletion.cancelled + while let completion = await group.next() { + if case .writer(let writerResult) = completion { + result = writerResult + events.close() + group.cancelAll() + break + } } - context.write(self.wrapOutboundOut(.head(head)), promise: nil) - if let body { - var buffer = context.channel.allocator.buffer(capacity: body.count) - buffer.writeBytes(body) - context.write(self.wrapOutboundOut(.body(.byteBuffer(buffer))), promise: nil) + await group.waitForAll() + return result + } + } onCancel: { + events.close() + } + } + + private func consumeResponseEvents( + _ events: MCPHTTPResponseEventChannel, + head: HTTPResponseHead, + context: ChannelHandlerContext, + eventLoop: any EventLoop + ) async -> WriterCompletion { + do { + try Task.checkCancellation() + try await writeResponsePart(.head(head), context: context, eventLoop: eventLoop) + while true { + switch await events.receive() { + case .body(let id, let data): + do { +#if DEBUG + await responseBackpressureProbe.waitBeforeBodyWriteIfNeeded() +#endif + try await writeResponseBody(data, context: context, eventLoop: eventLoop) + events.acknowledgeBody(id: id, wasWritten: true) + } catch { + events.acknowledgeBody(id: id, wasWritten: false) + return .transportFailed(error.localizedDescription) + } + case .heartbeat: + try await writeResponseBody( + Data(": keep-alive\n\n".utf8), + context: context, + eventLoop: eventLoop + ) + case .sourceFinished: + await server.waitBeforeResponseEndWriteForTesting() + try Task.checkCancellation() + try await writeResponsePart(.end(nil), context: context, eventLoop: eventLoop) + return .responded + case .sourceFailed(let message): + return .sourceFailed(message) + case .cancelled: + return .cancelled } - context.writeAndFlush(self.wrapOutboundOut(.end(nil)), promise: nil) } + } catch is CancellationError { + return .cancelled + } catch { + return .transportFailed(error.localizedDescription) + } + } + + private func writeNonStreamingResponse( + _ response: HTTPResponse, + head: HTTPResponseHead, + context: ChannelHandlerContext, + eventLoop: any EventLoop + ) async -> WriterCompletion { + var head = head + let body = response.bodyData + 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 { + try await writeResponseBody(body, context: context, eventLoop: eventLoop) + } + await server.waitBeforeResponseEndWriteForTesting() + 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 85351a1..80bebe6 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) } @@ -53,10 +54,27 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { case closed(TerminalCause?) } + 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 responseEnd: ResponseEndPhase + + package var terminalCause: TerminalCause? { + switch phase { + case .closing(let cause), .closed(let cause?): + cause + case .reserved, .installed, .running, .closed: + nil + } + } } package struct ConnectionSnapshot: Equatable, Sendable { @@ -117,6 +135,8 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { private let leaseID: UUID private var workState: WorkState = .reserved private var terminalCause: TerminalCause? + private var responseEnd: ResponseEndPhase = .notExpected + private var didClose = false private var startWasRequested = false private var startWaiter: CheckedContinuation? private var closeWaiters: [CheckedContinuation] = [] @@ -136,26 +156,52 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { var cancellation: (@Sendable () -> Void)? var waiter: CheckedContinuation? lock.lock() - if terminalCause == nil { + if terminalCause == nil, responseEnd != .acknowledged { terminalCause = cause - } - switch workState { - case .installed(let cancel), .running(let cancel): - cancellation = cancel - waiter = startWaiter - startWaiter = nil - case .reserved, .acknowledged: - break + responseEnd = .closed + switch workState { + case .installed(let cancel), .running(let cancel): + cancellation = cancel + waiter = startWaiter + startWaiter = nil + case .reserved, .acknowledged: + break + } } lock.unlock() cancellation?() waiter?.resume(returning: false) + finishIfPossible() + } + + package func beginResponse() -> Bool { + lock.lock() + guard terminalCause == nil, + didClose == false, + responseEnd == .notExpected else { + lock.unlock() + return false + } + responseEnd = .pending + lock.unlock() + return true + } + + package func acknowledgeResponseEnd() { + lock.lock() + guard terminalCause == nil, responseEnd == .pending else { + lock.unlock() + return + } + responseEnd = .acknowledged + lock.unlock() + finishIfPossible() } package func waitUntilClosed() async -> TerminalCause? { await withCheckedContinuation { continuation in lock.lock() - if case .acknowledged = workState { + if didClose { let cause = terminalCause lock.unlock() continuation.resume(returning: cause) @@ -226,20 +272,15 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { } fileprivate func acknowledgeCompletion(leaseID: UUID) { - let waiters: [CheckedContinuation] - let cause: TerminalCause? 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) - lock.unlock() - for waiter in waiters { - waiter.resume(returning: cause) + if responseEnd == .notExpected, terminalCause == nil { + responseEnd = .acknowledged } - connection?.requestDidClose(self) + lock.unlock() + finishIfPossible() return } lock.unlock() @@ -248,18 +289,46 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { fileprivate func snapshot() -> RequestSnapshot { lock.lock() 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 workState { + case .reserved: .reserved + case .installed: .installed + case .running, .acknowledged: .running + } } + let responseEnd = responseEnd lock.unlock() - return .init(id: id, admissionOrdinal: admissionOrdinal, phase: phase) + return .init( + id: id, + admissionOrdinal: admissionOrdinal, + phase: phase, + responseEnd: responseEnd + ) + } + + private func finishIfPossible() { + let waiters: [CheckedContinuation] + let cause: TerminalCause? + lock.lock() + guard didClose == false, + case .acknowledged = workState, + responseEnd == .acknowledged || responseEnd == .closed else { + lock.unlock() + return + } + didClose = true + cause = terminalCause + waiters = closeWaiters + closeWaiters.removeAll(keepingCapacity: false) + lock.unlock() + for waiter in waiters { + waiter.resume(returning: cause) + } + connection?.requestDidClose(self) } } @@ -278,6 +347,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { private var nextRequestOrdinal: UInt64 = 0 private var requests: [UUID: RequestOperation] = [:] private var closeAcknowledged = false + private var closeAcknowledgementWaiters: [CheckedContinuation] = [] private var closeWaiters: [CheckedContinuation] = [] fileprivate init( @@ -324,6 +394,20 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { beginClosing(.transportFailure(message), signalResourceClose: true) } + package func closeAfterResponse() async { + beginClosing(.responseComplete, signalResourceClose: true) + await withCheckedContinuation { continuation in + lock.lock() + if closeAcknowledged { + lock.unlock() + continuation.resume() + } else { + closeAcknowledgementWaiters.append(continuation) + lock.unlock() + } + } + } + package func waitUntilClosed() async { await withCheckedContinuation { continuation in lock.lock() @@ -413,8 +497,11 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { private func acknowledgePeerClose() { let requests: [RequestOperation] + let acknowledgementWaiters: [CheckedContinuation] lock.lock() closeAcknowledged = true + acknowledgementWaiters = closeAcknowledgementWaiters + closeAcknowledgementWaiters.removeAll(keepingCapacity: false) switch phase { case .accepting, .admissionClosed: phase = .closing(.peerClosed) @@ -423,6 +510,9 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { requests = [] } lock.unlock() + for waiter in acknowledgementWaiters { + waiter.resume() + } for request in requests { request.beginClosing(.peerClosed) } diff --git a/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift b/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift index 594fb27..ec1466e 100644 --- a/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift +++ b/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift @@ -117,7 +117,7 @@ struct CodexReviewMCPHTTPServerTests { ], ], ]) - let response = await server.handleHTTPRequest(HTTPRequest( + let response = await server.validationResponseForTesting(HTTPRequest( method: "POST", headers: [ HTTPHeaderName.host: "review.local:9417", @@ -127,7 +127,7 @@ struct CodexReviewMCPHTTPServerTests { body: initializeBody, path: "/mcp" )) - let denied = await server.handleHTTPRequest(HTTPRequest( + let denied = await server.validationResponseForTesting(HTTPRequest( method: "POST", headers: [ HTTPHeaderName.host: "other.local:9417", @@ -138,10 +138,8 @@ struct CodexReviewMCPHTTPServerTests { path: "/mcp" )) - #expect(response.statusCode == 200) - #expect(response.headers[HTTPHeaderName.sessionID]?.isEmpty == false) - #expect(denied.statusCode == 421) - try await server.stop() + #expect(response == nil) + #expect(denied?.statusCode == 421) } @Test func streamableHTTPClassifiesAddressInUseBindError() { @@ -712,6 +710,273 @@ struct CodexReviewMCPHTTPServerTests { } } + @Test func networkFiniteStreamExhaustionFinishesItsActiveRequest() async throws { + try await withHTTPServer(store: CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) + )) { server in + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + _ = try await postJSONRPCData( + endpoint: endpoint, + sessionID: sessionID, + bodyData: makeToolsListBody(id: 29) + ) + #expect(await server.sessionActiveRequestCountForTesting(sessionID: sessionID) == 0) + } + } + + @Test func networkOpenStreamDisconnectFinishesItsActiveRequest() async throws { + let server = makeHTTPServer(configuration: .init( + host: "127.0.0.1", + port: 0, + streamHeartbeatInterval: .milliseconds(10) + )) + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + let connection = try await RawHTTPConnection.connect(to: endpoint) + try await connection.send(rawHTTPRequest( + endpoint: endpoint, + method: "GET", + sessionID: sessionID, + headers: [("Accept", "text/event-stream, application/json")] + )) + _ = try await connection.readResponseHead() + #expect(await server.sessionActiveRequestCountForTesting(sessionID: sessionID) == 1) + connection.reset() + #expect(await waitUntil(timeout: .seconds(2)) { + await server.sessionActiveRequestCountForTesting(sessionID: sessionID) == 0 + }) + try await server.stop() + } + + @Test func stopJoinsHeldFiniteResponseSource() async throws { + let server = makeHTTPServer() + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + await server.holdNextFiniteSourceCompletionForTesting() + let response = Task { + try await postJSONRPCData( + endpoint: endpoint, + sessionID: sessionID, + bodyData: makeToolsListBody(id: 30) + ) + } + await server.waitUntilFiniteSourceCompletionIsHeldForTesting() + #expect(try #require(await server.networkResourceSnapshotForTesting()) + .connections.flatMap(\.requests).count == 1) + + let stop = Task { try await server.stop() } + #expect(await waitUntil(timeout: .seconds(2)) { + await server.networkResourceSnapshotForTesting()?.phase == .closing(.serverStop) + }) + #expect(await server.eventLoopGroupShutdownCountForTesting() == 0) + await server.releaseFiniteSourceCompletionForTesting() + _ = try? await response.value + try await stop.value + #expect(await server.eventLoopGroupShutdownCountForTesting() == 1) + } + + @Test func heldPhysicalBodyWriteAllowsOneSourceRead() async throws { + let server = makeHTTPServer() + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + await server.holdNextResponseBodyWriteForTesting() + let response = Task { + try await postJSONRPCData( + endpoint: endpoint, + sessionID: sessionID, + bodyData: makeToolsListBody(id: 31) + ) + } + await server.waitUntilResponseBodyWriteIsHeldForTesting() + #expect(await server.responseSourceReadCountForTesting() == 1) + await server.releaseResponseBodyWriteForTesting() + _ = try await response.value + #expect(await server.responseSourceReadCountForTesting() >= 2) + try await server.stop() + } + + @Test func repeatedHeartbeatsCannotOvertakePendingBody() async { + #expect(await CodexReviewMCPHTTPServer.responseRendezvousPrioritizesBodyForTesting()) + } + + @Test func stopJoinsWriterBeforeEventLoopShutdown() async throws { + let server = makeHTTPServer() + 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, _) = try await URLSession.shared.bytes(for: request) + + let stop = Task { try await server.stop() } + await server.waitUntilWriterCompletionIsHeldForTesting() + #expect(await server.eventLoopGroupShutdownCountForTesting() == 0) + await server.releaseWriterCompletionForTesting() + try await stop.value + withExtendedLifetime(bytes) {} + #expect(await server.eventLoopGroupShutdownCountForTesting() == 1) + } + + @Test func loneNonPersistentRequestsSendCompleteResponseThenEOF() async throws { + let server = makeHTTPServer() + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + for (version, connectionHeader, id) in [ + ("HTTP/1.0", Optional.none, 32), + ("HTTP/1.1", Optional("close"), 33), + ] { + let connection = try await RawHTTPConnection.connect(to: endpoint) + try await connection.send(rawHTTPRequest( + endpoint: endpoint, + version: version, + sessionID: sessionID, + headers: connectionHeader.map { [("Connection", $0)] } ?? [], + body: makeToolsListBody(id: id) + )) + let head = try await connection.readResponseHead() + let body = try await connection.readUntilEOF() + connection.close() + #expect(head.contains(" 200 ")) + #expect(try decodeSSEJSON(from: body)["id"] as? Int == id) + if version == "HTTP/1.1" { + #expect(body.suffix(5) == Data("0\r\n\r\n".utf8)) + } + } + try await server.stop() + } + + @Test func channelCloseCancelsOpenSSEWithoutLateEventLoopWork() async throws { + let server = makeHTTPServer(configuration: .init( + host: "127.0.0.1", + port: 0, + streamHeartbeatInterval: .milliseconds(10) + )) + try await server.start() + let sessionID = try await initializeSession(endpoint: await server.url) + try await openAndCloseRawEventStream(endpoint: await server.url, sessionID: sessionID) + #expect(await waitUntil(timeout: .seconds(2)) { + await server.networkResourceSnapshotForTesting()?.connections + .flatMap(\.requests).isEmpty == true + }) + try await server.stop() + #expect(await server.eventLoopGroupShutdownCountForTesting() == 1) + } + + @Test func acknowledgedResponseEndWinsConcurrentStop() async throws { + let server = makeHTTPServer() + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + await server.holdNextResponseEndAcknowledgementForTesting() + let response = Task { + try await postJSONRPCData( + endpoint: endpoint, + sessionID: sessionID, + bodyData: makeToolsListBody(id: 34) + ) + } + await server.waitUntilResponseEndAcknowledgementIsHeldForTesting() + let before = try #require(await server.networkResourceSnapshotForTesting()? + .connections.flatMap(\.requests).first) + #expect(before.responseEnd == .acknowledged) + #expect(before.terminalCause == nil) + + let stop = Task { try await server.stop() } + #expect(await waitUntil(timeout: .seconds(2)) { + await server.networkResourceSnapshotForTesting()?.phase == .closing(.serverStop) + }) + let after = try #require(await server.networkResourceSnapshotForTesting()? + .connections.flatMap(\.requests).first) + #expect(after.responseEnd == .acknowledged) + #expect(after.terminalCause == nil) + await server.releaseResponseEndAcknowledgementForTesting() + _ = try? await response.value + try await stop.value + } + + @Test func clientDisconnectDrainsOwnedFiniteSource() async throws { + let server = makeHTTPServer() + try await server.start() + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + await server.holdNextFiniteSourceCompletionForTesting() + await server.holdNextWriterCompletionForTesting() + let connection = try await RawHTTPConnection.connect(to: endpoint) + try await connection.send(rawHTTPRequest( + endpoint: endpoint, + sessionID: sessionID, + body: makeToolsListBody(id: 35) + )) + await server.waitUntilFiniteSourceCompletionIsHeldForTesting() + _ = try await connection.readResponseHead() + connection.reset() + #expect(try #require(await server.networkResourceSnapshotForTesting()) + .connections.flatMap(\.requests).count == 1) + await server.releaseFiniteSourceCompletionForTesting() + await server.waitUntilWriterCompletionIsHeldForTesting() + let closing = try #require(await server.networkResourceSnapshotForTesting()? + .connections.flatMap(\.requests).first) + #expect(closing.responseEnd == .closed) + #expect(closing.terminalCause == .peerClosed || closing.terminalCause.map { + if case .transportFailure = $0 { true } else { false } + } == true) + await server.releaseWriterCompletionForTesting() + #expect(await waitUntil(timeout: .seconds(2)) { + await server.networkResourceSnapshotForTesting()?.connections + .flatMap(\.requests).isEmpty == true + }) + try await server.stop() + } + + @Test func nonStreamRequestRemainsOwnedUntilFinalEndWrite() async throws { + let server = makeHTTPServer() + try await server.start() + let endpoint = await server.url + await server.holdNextResponseEndWriteForTesting() + let connection = try await RawHTTPConnection.connect(to: endpoint) + try await connection.send(rawHTTPRequest( + endpoint: endpoint, + sessionID: nil, + body: makeToolsListBody(id: 36) + )) + await server.waitUntilResponseEndWriteIsHeldForTesting() + let pending = try #require(await server.networkResourceSnapshotForTesting()? + .connections.flatMap(\.requests).first) + #expect(pending.responseEnd == .pending) + await server.releaseResponseEndWriteForTesting() + #expect(await waitUntil(timeout: .seconds(2)) { + await server.networkResourceSnapshotForTesting()?.connections + .flatMap(\.requests).isEmpty == true + }) + #expect(try await connection.readResponseHead().contains(" 400 ")) + connection.close() + try await server.stop() + } + + @Test func acceptedChildCloseAcknowledgementPrecedesEventLoopShutdown() async throws { + let server = makeHTTPServer() + try await server.start() + let connection = try await RawHTTPConnection.connect(to: await server.url) + await server.holdNextStopCompletionForTesting() + let stop = Task { try await server.stop() } + await server.waitUntilStopCompletionIsHeldForTesting() + _ = try await connection.readUntilEOF() + #expect(await server.networkResourceSnapshotForTesting()?.isClosed == true) + #expect(await server.eventLoopGroupShutdownCountForTesting() == 0) + await server.releaseStopCompletionForTesting() + try await stop.value + connection.close() + #expect(await server.eventLoopGroupShutdownCountForTesting() == 1) + } + @Test func streamableHTTPCallsReviewStartWithCustomTarget() async throws { let backend = FakeCodexReviewBackend() let store = CodexReviewStore.makeTestingStore( @@ -1821,6 +2086,21 @@ struct CodexReviewMCPHTTPServerTests { } } + private func makeHTTPServer( + configuration: CodexReviewMCPHTTPServer.Configuration = .init( + host: "127.0.0.1", + port: 0 + ) + ) -> CodexReviewMCPHTTPServer { + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: FakeCodexReviewBackend()) + ) + return CodexReviewMCPHTTPServer( + adapter: CodexReviewMCPServer(store: store), + configuration: configuration + ) + } + private func withHTTPServer( store: CodexReviewStore, configuration: CodexReviewMCPHTTPServer.Configuration = .init( @@ -2287,6 +2567,27 @@ private final class RawHTTPConnection: @unchecked Sendable { } } + func reset() { + let shouldClose = lock.withLock { + guard isClosed == false else { return false } + isClosed = true + return true + } + if shouldClose { + var option = linger(l_onoff: 1, l_linger: 0) + _ = withUnsafePointer(to: &option) { + Darwin.setsockopt( + descriptor, + SOL_SOCKET, + SO_LINGER, + $0, + socklen_t(MemoryLayout.size) + ) + } + Darwin.close(descriptor) + } + } + deinit { close() }