From 8c1d9ca5d64da805153ae5d655eea10b77c2fa0c Mon Sep 17 00:00:00 2001 From: Kazuki Nakashima <65545348+lynnswap@users.noreply.github.com> Date: Sat, 22 Aug 2026 20:24:57 +0900 Subject: [PATCH] fix(mcp): bound HTTP request admission --- .../CodexReviewMCPHTTPServer.swift | 312 ++++++++-- .../MCPHTTPNetworkResourceOwner.swift | 7 +- .../CodexReviewMCPHTTPServerTests.swift | 571 +++++++++++++++++- .../MCPHTTPNetworkResourceOwnerTests.swift | 29 + 4 files changed, 874 insertions(+), 45 deletions(-) diff --git a/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift b/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift index cadfbdf..f611868 100644 --- a/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift +++ b/Sources/CodexReviewMCPServer/CodexReviewMCPHTTPServer.swift @@ -84,13 +84,15 @@ package extension CodexReviewMCPHTTPServer { package var retryInterval: Int? package var streamHeartbeatInterval: Duration? package var boundedReviewWaitDuration: Duration + package var maximumRequestBodyBytes: Int package init( host: String = "localhost", port: Int = 9417, endpoint: String = "/mcp", sessionTimeout: TimeInterval = 3600, - retryInterval: Int? = 1000 + retryInterval: Int? = 1000, + maximumRequestBodyBytes: Int = 1_048_576 ) { self.init( host: host, @@ -99,7 +101,8 @@ package extension CodexReviewMCPHTTPServer { sessionTimeout: sessionTimeout, retryInterval: retryInterval, streamHeartbeatInterval: .seconds(30), - boundedReviewWaitDuration: .seconds(540) + boundedReviewWaitDuration: .seconds(540), + maximumRequestBodyBytes: maximumRequestBodyBytes ) } @@ -110,8 +113,13 @@ package extension CodexReviewMCPHTTPServer { sessionTimeout: TimeInterval = 3600, retryInterval: Int? = 1000, streamHeartbeatInterval: Duration?, - boundedReviewWaitDuration: Duration = .seconds(540) + boundedReviewWaitDuration: Duration = .seconds(540), + maximumRequestBodyBytes: Int = 1_048_576 ) { + precondition( + maximumRequestBodyBytes >= 0, + "MCP HTTP Configuration owns a nonnegative request-body byte limit." + ) self.host = host self.port = port self.endpoint = endpoint.hasPrefix("/") ? endpoint : "/\(endpoint)" @@ -119,6 +127,7 @@ package extension CodexReviewMCPHTTPServer { self.retryInterval = retryInterval self.streamHeartbeatInterval = streamHeartbeatInterval self.boundedReviewWaitDuration = boundedReviewWaitDuration + self.maximumRequestBodyBytes = maximumRequestBodyBytes } package func url(boundPort: Int? = nil) -> URL { @@ -400,6 +409,7 @@ package actor CodexReviewMCPHTTPServer { networkResources: MCPHTTPNetworkResourceOwner ) async -> StartingGenerationResult { let group = MultiThreadedEventLoopGroup(numberOfThreads: 1) + let maximumRequestBodyBytes = configuration.maximumRequestBodyBytes let bootstrap = ServerBootstrap(group: group) .serverChannelOption(ChannelOptions.backlog, value: 128) .serverChannelOption(ChannelOptions.socketOption(.so_reuseaddr), value: 1) @@ -410,7 +420,8 @@ package actor CodexReviewMCPHTTPServer { return channel.pipeline.configureHTTPServerPipeline().flatMap { channel.pipeline.addHandler(CodexReviewMCPHTTPHandler( server: self, - connection: connection + connection: connection, + maximumRequestBodyBytes: maximumRequestBodyBytes )) } } @@ -864,6 +875,19 @@ package actor CodexReviewMCPHTTPServer { } } + package func networkResourceSnapshotForTesting() -> MCPHTTPNetworkResourceOwner.Snapshot? { + switch lifecycleState { + case .starting(let operation): + operation.networkResources.snapshot() + case .running(let resources): + resources.networkResources.snapshot() + case .stopping(_, let resources?, _): + resources.networkResources.snapshot() + case .stopping, .stopped: + nil + } + } + package func eventLoopGroupShutdownCountForTesting() -> Int { eventLoopGroupShutdownCount } @@ -1185,6 +1209,103 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked typealias InboundIn = HTTPServerRequestPart typealias OutboundOut = HTTPServerResponsePart + private enum RequestBodyResult: Sendable { + case body(Data?) + case payloadTooLarge + case expectationFailed + case cancelled + } + + private final class RequestBodyReceipt: @unchecked Sendable { + private let maximumByteCount: Int + private let lock = NSLock() + private var body = Data() + private var result: RequestBodyResult? + private var waiter: CheckedContinuation? + + init(maximumByteCount: Int) { + self.maximumByteCount = maximumByteCount + } + + func receive(_ buffer: ByteBuffer) -> Bool { + let readableByteCount = buffer.readableBytes + guard readableByteCount > 0 else { + return false + } + + lock.lock() + guard result == nil else { + lock.unlock() + return false + } + guard readableByteCount <= maximumByteCount - body.count else { + lock.unlock() + return true + } + body.append(contentsOf: buffer.readableBytesView) + lock.unlock() + return false + } + + func finish() { + let outcome: RequestBodyResult + lock.lock() + guard result == nil else { + lock.unlock() + return + } + outcome = .body(body.isEmpty ? nil : body) + result = outcome + body.removeAll(keepingCapacity: false) + let continuation = waiter + waiter = nil + lock.unlock() + continuation?.resume(returning: outcome) + } + + func reject(_ outcome: RequestBodyResult) { + let continuation: CheckedContinuation? + lock.lock() + guard result == nil else { + lock.unlock() + return + } + result = outcome + body.removeAll(keepingCapacity: false) + continuation = waiter + waiter = nil + lock.unlock() + continuation?.resume(returning: outcome) + } + + func waitForResult() async -> RequestBodyResult { + await withTaskCancellationHandler { + await withCheckedContinuation { continuation in + lock.lock() + if let result { + lock.unlock() + continuation.resume(returning: result) + } else { + precondition( + waiter == nil, + "One request operation owns the body receipt waiter." + ) + waiter = continuation + lock.unlock() + } + } + } onCancel: { + self.reject(.cancelled) + } + } + } + + private enum RequestExpectation: Equatable { + case none + case continueRequest + case unsupported + } + private struct ResponsePartWriter: @unchecked Sendable { let handler: CodexReviewMCPHTTPHandler let context: ChannelHandlerContext @@ -1207,12 +1328,12 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked } private struct RequestState { - var head: HTTPRequestHead - var bodyBuffer: ByteBuffer + var bodyReceipt: RequestBodyReceipt } private let server: CodexReviewMCPHTTPServer private let connection: MCPHTTPNetworkResourceOwner.Connection + private let maximumRequestBodyBytes: Int private var requestState: RequestState? private var activeStreamTask: Task? private var activeStreamID: UUID? @@ -1220,43 +1341,132 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked init( server: CodexReviewMCPHTTPServer, - connection: MCPHTTPNetworkResourceOwner.Connection + connection: MCPHTTPNetworkResourceOwner.Connection, + maximumRequestBodyBytes: Int ) { self.server = server self.connection = connection + self.maximumRequestBodyBytes = maximumRequestBodyBytes } func channelRead(context: ChannelHandlerContext, data: NIOAny) { let part = unwrapInboundIn(data) switch part { case .head(let head): - requestState = RequestState( - head: head, - bodyBuffer: context.channel.allocator.buffer(capacity: 0) - ) - case .body(var buffer): - requestState?.bodyBuffer.writeBuffer(&buffer) - case .end: - guard let state = requestState else { - return + receiveRequestHead(head, context: context) + case .body(let buffer): + if requestState?.bodyReceipt.receive(buffer) == true { + connection.closeAdmission() + requestState?.bodyReceipt.reject(.payloadTooLarge) } + case .end: + let receipt = requestState?.bodyReceipt requestState = nil - guard let admittedRequest = connection.admitRequest() else { - context.close(promise: nil) + receipt?.finish() + } + } + + private func receiveRequestHead( + _ head: HTTPRequestHead, + context: ChannelHandlerContext + ) { + let expectation = requestExpectation(for: head) + let contentLengthExceedsLimit = contentLengthExceedsLimit(head) + let rejection: RequestBodyResult? + if contentLengthExceedsLimit { + rejection = .payloadTooLarge + } else if expectation == .unsupported { + rejection = .expectationFailed + } else { + rejection = nil + } + let shouldSendContinue = expectation == .continueRequest && rejection == nil + let finalForConnection = head.isKeepAlive == false || rejection != nil + guard let admittedRequest = connection.admitRequest( + finalForConnection: finalForConnection + ) else { + context.close(promise: nil) + return + } + + let bodyReceipt = RequestBodyReceipt(maximumByteCount: maximumRequestBodyBytes) + requestState = .init(bodyReceipt: bodyReceipt) + nonisolated(unsafe) let context = context + let task = Task { [self] in + defer { + bodyReceipt.reject(.cancelled) + admittedRequest.lease.acknowledgeCompletion() + } + guard await admittedRequest.lease.waitUntilStartIsAllowed() else { return } - nonisolated(unsafe) let context = context - let task = Task { [self] in - defer { - admittedRequest.lease.acknowledgeCompletion() - } - guard await admittedRequest.lease.waitUntilStartIsAllowed() else { + if shouldSendContinue { + do { + let responseHead = HTTPResponseHead(version: head.version, status: .continue) + try await writeResponsePart( + .head(responseHead), + context: context, + eventLoop: context.eventLoop + ) + } catch { + connection.transportFailed(error.localizedDescription) return } - await handleRequest(state: state, context: context) } - admittedRequest.lease.install(task) + let bodyResult = await bodyReceipt.waitForResult() + guard Task.isCancelled == false else { + return + } + switch bodyResult { + case .body(let body): + await handleRequest(head: head, body: body, context: context) + case .payloadTooLarge: + await writeRequestRejection( + status: .payloadTooLarge, + version: head.version, + context: context + ) + case .expectationFailed: + await writeRequestRejection( + status: .expectationFailed, + version: head.version, + context: context + ) + case .cancelled: + return + } + } + admittedRequest.lease.install(task) + + if let rejection { + bodyReceipt.reject(rejection) + } + } + + private func contentLengthExceedsLimit(_ head: HTTPRequestHead) -> Bool { + guard let rawValue = head.headers.first(name: "content-length") else { + return false } + guard let contentLength = UInt64(rawValue) else { + return true + } + return contentLength > UInt64(maximumRequestBodyBytes) + } + + private func requestExpectation(for head: HTTPRequestHead) -> RequestExpectation { + guard head.version.major == 1, head.version.minor >= 1 else { + return .none + } + let values = head.headers[canonicalForm: "expect"] + guard values.isEmpty == false else { + return .none + } + guard values.count == 1, + String(values[0]).caseInsensitiveCompare("100-continue") == .orderedSame + else { + return .unsupported + } + return .continueRequest } func channelReadComplete(context: ChannelHandlerContext) { @@ -1265,6 +1475,8 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked } func channelInactive(context: ChannelHandlerContext) { + requestState?.bodyReceipt.reject(.cancelled) + requestState = nil connection.peerClosed() finishActiveStream() context.fireChannelInactive() @@ -1272,6 +1484,8 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked func userInboundEventTriggered(context: ChannelHandlerContext, event: Any) { if case ChannelEvent.inputClosed = event { + requestState?.bodyReceipt.reject(.cancelled) + requestState = nil connection.peerClosed() finishActiveStream() context.close(promise: nil) @@ -1281,6 +1495,8 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked } func errorCaught(context: ChannelHandlerContext, error: any Error) { + requestState?.bodyReceipt.reject(.cancelled) + requestState = nil connection.transportFailed(error.localizedDescription) finishActiveStream() context.close(promise: nil) @@ -1295,10 +1511,10 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked } private func handleRequest( - state: RequestState, + head: HTTPRequestHead, + body: Data?, 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 { @@ -1310,14 +1526,14 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked return } - let request = makeHTTPRequest(from: state) + let request = makeHTTPRequest(head: head, body: body) let response = await server.handleTrackedHTTPRequest(request) await writeResponse(response, version: head.version, context: context) } - private func makeHTTPRequest(from state: RequestState) -> HTTPRequest { + private func makeHTTPRequest(head: HTTPRequestHead, body: Data?) -> HTTPRequest { var headers: [String: String] = [:] - for (name, value) in state.head.headers { + for (name, value) in head.headers { if let existing = headers[name] { headers[name] = "\(existing), \(value)" } else { @@ -1325,24 +1541,36 @@ private final class CodexReviewMCPHTTPHandler: ChannelInboundHandler, @unchecked } } - let body: Data? - if state.bodyBuffer.readableBytes > 0, - let bytes = state.bodyBuffer.getBytes(at: 0, length: state.bodyBuffer.readableBytes) - { - body = Data(bytes) - } else { - body = nil - } - - let path = String(state.head.uri.split(separator: "?").first ?? Substring(state.head.uri)) + let path = String(head.uri.split(separator: "?").first ?? Substring(head.uri)) return HTTPRequest( - method: state.head.method.rawValue, + method: head.method.rawValue, headers: headers, body: body, path: path ) } + private func writeRequestRejection( + status: HTTPResponseStatus, + version: HTTPVersion, + context: ChannelHandlerContext + ) async { + nonisolated(unsafe) let context = context + let eventLoop = context.eventLoop + var head = HTTPResponseHead(version: version, status: status) + head.headers.add(name: "Content-Length", value: "0") + head.headers.add(name: "Connection", value: "close") + do { + try await writeResponsePart(.head(head), context: context, eventLoop: eventLoop) + try await writeResponsePart(.end(nil), context: context, eventLoop: eventLoop) + } catch { + logger.error("MCP request rejection write failed: \(error.localizedDescription, privacy: .public)") + } + eventLoop.execute { + context.close(promise: nil) + } + } + private func writeResponse( _ trackedResponse: TrackedHTTPResponse, version: HTTPVersion, diff --git a/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift b/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift index 25a9e92..85351a1 100644 --- a/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift +++ b/Sources/CodexReviewMCPServer/MCPHTTPNetworkResourceOwner.swift @@ -296,7 +296,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { } } - package func admitRequest() -> AdmittedRequest? { + package func admitRequest(finalForConnection: Bool = false) -> AdmittedRequest? { lock.lock() guard phase == .accepting else { lock.unlock() @@ -309,6 +309,9 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { ) let lease = operation.makeLease() requests[operation.id] = operation + if finalForConnection { + phase = .admissionClosed + } lock.unlock() return .init(operation: operation, lease: lease) } @@ -334,7 +337,7 @@ package final class MCPHTTPNetworkResourceOwner: @unchecked Sendable { } } - fileprivate func closeAdmission() { + package func closeAdmission() { lock.lock() if phase == .accepting { phase = .admissionClosed diff --git a/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift b/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift index d949238..594fb27 100644 --- a/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift +++ b/Tests/CodexReviewMCPServerTests/CodexReviewMCPHTTPServerTests.swift @@ -389,6 +389,329 @@ struct CodexReviewMCPHTTPServerTests { try await server.stop() } + @Test func slowFirstPOSTKeepsPipelinedSecondMutationOutsideAdmission() async throws { + let backend = FakeCodexReviewBackend() + let startGate = AsyncGate() + await backend.holdStartReview(with: startGate) + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: backend), + idGenerator: .init(next: { "job-1" }) + ) + + try await withHTTPServer(store: store) { server in + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + let connection = try await RawHTTPConnection.connect(to: endpoint) + defer { connection.close() } + let first = try makeReviewStartBody(id: 2) + let second = try makeReviewStartBody(id: 3) + + try await connection.send( + rawHTTPRequest(endpoint: endpoint, sessionID: sessionID, body: first) + + rawHTTPRequest(endpoint: endpoint, sessionID: sessionID, body: second) + ) + try await backend.waitForStartReview(timeout: .seconds(2)) + + let snapshot = try #require(await server.networkResourceSnapshotForTesting()) + #expect(snapshot.connections.flatMap(\.requests).count == 1) + #expect(await backend.recordedCommands().filter { + if case .startReview = $0 { true } else { false } + }.count == 1) + + connection.close() + await startGate.open() + } + } + + @Test func openGETKeepsSameConnectionPOSTOutsideAdmission() async throws { + let backend = FakeCodexReviewBackend() + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: backend) + ) + + try await withHTTPServer(store: store) { server in + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + let connection = try await RawHTTPConnection.connect(to: endpoint) + defer { connection.close() } + let get = try rawHTTPRequest( + endpoint: endpoint, + method: "GET", + sessionID: sessionID, + headers: [("Accept", "text/event-stream, application/json")] + ) + let post = try rawHTTPRequest( + endpoint: endpoint, + sessionID: sessionID, + body: makeReviewStartBody(id: 2) + ) + + try await connection.send(get + post) + #expect(try await connection.readResponseHead().contains(" 200 ")) + + let snapshot = try #require(await server.networkResourceSnapshotForTesting()) + #expect(snapshot.connections.flatMap(\.requests).count == 1) + #expect(await backend.recordedCommands().isEmpty) + } + } + + @Test func sameSessionDifferentConnectionsRemainParallel() async throws { + let backend = FakeCodexReviewBackend() + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: backend), + idGenerator: .init(next: { "job-1" }) + ) + + try await withHTTPServer(store: store) { server in + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + let streamConnection = try await RawHTTPConnection.connect(to: endpoint) + defer { streamConnection.close() } + try await streamConnection.send(rawHTTPRequest( + endpoint: endpoint, + method: "GET", + sessionID: sessionID, + headers: [("Accept", "text/event-stream, application/json")] + )) + #expect(try await streamConnection.readResponseHead().contains(" 200 ")) + + async let response = postJSONRPCData( + endpoint: endpoint, + sessionID: sessionID, + bodyData: makeReviewStartBody(id: 2) + ) + try await backend.waitForStartReview(timeout: .seconds(2)) + await backend.yield(.completed(summary: "Done", result: "review text")) + _ = try await response + + #expect(await backend.recordedCommands().contains { + if case .startReview = $0 { true } else { false } + }) + } + } + + @Test func nonKeepAliveRequestRejectsPipelinedMutationBeforeDomainAdmission() async throws { + let backend = FakeCodexReviewBackend() + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: backend), + idGenerator: .init(next: { "job-1" }) + ) + + try await withHTTPServer(store: store) { server in + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + let connection = try await RawHTTPConnection.connect(to: endpoint) + defer { connection.close() } + let first = try rawHTTPRequest( + endpoint: endpoint, + sessionID: sessionID, + headers: [("Connection", "close")], + body: makeToolsListBody(id: 2) + ) + let mutation = try rawHTTPRequest( + endpoint: endpoint, + sessionID: sessionID, + body: makeReviewStartBody(id: 3) + ) + + try await connection.send(first + mutation) + #expect(try await connection.readResponseHead().contains(" 200 ")) + _ = try await connection.readUntilEOF() + + #expect(await backend.recordedCommands().isEmpty) + } + } + + @Test func oneConnectionOwnsOnlyOneOfManyPipelinedRequests() async throws { + let backend = FakeCodexReviewBackend() + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: backend) + ) + + try await withHTTPServer(store: store) { server in + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + let connection = try await RawHTTPConnection.connect(to: endpoint) + defer { connection.close() } + var requests = try rawHTTPRequest( + endpoint: endpoint, + method: "GET", + sessionID: sessionID, + headers: [("Accept", "text/event-stream, application/json")] + ) + for id in 2..<34 { + requests += try rawHTTPRequest( + endpoint: endpoint, + sessionID: sessionID, + body: makeToolsListBody(id: id) + ) + } + + try await connection.send(requests) + #expect(try await connection.readResponseHead().contains(" 200 ")) + + let snapshot = try #require(await server.networkResourceSnapshotForTesting()) + #expect(snapshot.connections.flatMap(\.requests).count == 1) + } + } + + @Test func requestBodyLimitAcceptsNAndRejectsKnownAndChunkedNPlusOne() async throws { + let limit = 512 + let backend = FakeCodexReviewBackend() + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: backend) + ) + let configuration = CodexReviewMCPHTTPServer.Configuration( + host: "127.0.0.1", + port: 0, + maximumRequestBodyBytes: limit + ) + + try await withHTTPServer(store: store, configuration: configuration) { server in + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + var exactBody = try makeToolsListBody(id: 2) + exactBody.append(Data(repeating: 0x20, count: limit - exactBody.count)) + + let exact = try await RawHTTPConnection.connect(to: endpoint) + try await exact.send(rawHTTPRequest( + endpoint: endpoint, + sessionID: sessionID, + body: exactBody + )) + #expect(try await exact.readResponseHead().contains(" 200 ")) + exact.close() + + let knownTooLarge = try await RawHTTPConnection.connect(to: endpoint) + try await knownTooLarge.send(rawHTTPRequest( + endpoint: endpoint, + sessionID: sessionID, + headers: [ + ("Content-Length", "\(limit + 1)"), + ("Expect", "100-continue"), + ] + )) + let knownHead = try await knownTooLarge.readResponseHead() + #expect(knownHead.contains(" 413 ")) + #expect(knownHead.lowercased().contains("connection: close")) + _ = try await knownTooLarge.readUntilEOF() + knownTooLarge.close() + + let largerThanInt = try await RawHTTPConnection.connect(to: endpoint) + try await largerThanInt.send(rawHTTPRequest( + endpoint: endpoint, + sessionID: sessionID, + headers: [("Content-Length", "9223372036854775808")] + )) + #expect(try await largerThanInt.readResponseHead().contains(" 413 ")) + _ = try await largerThanInt.readUntilEOF() + largerThanInt.close() + + let chunked = try await RawHTTPConnection.connect(to: endpoint) + var chunkedBody = Data("\(String(limit, radix: 16))\r\n".utf8) + chunkedBody.append(Data(repeating: 0x61, count: limit)) + chunkedBody.append(Data("\r\n1\r\nb\r\n".utf8)) + try await chunked.send(rawHTTPRequest( + endpoint: endpoint, + sessionID: sessionID, + headers: [("Transfer-Encoding", "chunked")], + body: chunkedBody + )) + let chunkedHead = try await chunked.readResponseHead() + #expect(chunkedHead.contains(" 413 ")) + #expect(chunkedHead.lowercased().contains("connection: close")) + _ = try await chunked.readUntilEOF() + chunked.close() + } + } + + @Test func expectContinueUsesTheStandardPipelineBeforeReadingTheBody() async throws { + let backend = FakeCodexReviewBackend() + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: backend) + ) + + try await withHTTPServer(store: store) { server in + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + let body = try makeToolsListBody(id: 2) + let connection = try await RawHTTPConnection.connect(to: endpoint) + defer { connection.close() } + try await connection.send(rawHTTPRequest( + endpoint: endpoint, + sessionID: sessionID, + headers: [ + ("Content-Length", "\(body.count)"), + ("Expect", "100-continue"), + ] + )) + + #expect(try await connection.readResponseHead().contains(" 100 ")) + try await connection.send(body) + #expect(try await connection.readResponseHead().contains(" 200 ")) + + let unsupported = try await RawHTTPConnection.connect(to: endpoint) + try await unsupported.send(rawHTTPRequest( + endpoint: endpoint, + sessionID: sessionID, + headers: [ + ("Content-Length", "\(body.count)"), + ("Expect", "custom-expectation"), + ] + )) + #expect(try await unsupported.readResponseHead().contains(" 417 ")) + _ = try await unsupported.readUntilEOF() + unsupported.close() + + let http10 = try await RawHTTPConnection.connect(to: endpoint) + try await http10.send(rawHTTPRequest( + endpoint: endpoint, + version: "HTTP/1.0", + sessionID: sessionID, + headers: [ + ("Content-Length", "\(body.count)"), + ("Expect", "100-continue"), + ], + body: body + )) + let http10Head = try await http10.readResponseHead() + #expect(http10Head.contains(" 200 ")) + #expect(http10Head.contains(" 100 ") == false) + http10.close() + } + } + + @Test func stopCancelsAHeadAdmittedBodyReceipt() async throws { + let backend = FakeCodexReviewBackend() + let store = CodexReviewStore.makeTestingStore( + backend: TestingCodexReviewStoreBackend(reviewBackend: backend) + ) + + try await withHTTPServer(store: store) { server in + let endpoint = await server.url + let sessionID = try await initializeSession(endpoint: endpoint) + let body = try makeToolsListBody(id: 2) + let connection = try await RawHTTPConnection.connect(to: endpoint) + defer { connection.close() } + try await connection.send(rawHTTPRequest( + endpoint: endpoint, + sessionID: sessionID, + headers: [ + ("Content-Length", "\(body.count)"), + ("Expect", "100-continue"), + ] + )) + #expect(try await connection.readResponseHead().contains(" 100 ")) + #expect(try #require(await server.networkResourceSnapshotForTesting()) + .connections.flatMap(\.requests).count == 1) + + try await server.stop() + + _ = try await connection.readUntilEOF() + #expect(await server.currentGenerationIDForTesting() == nil) + } + } + @Test func streamableHTTPCallsReviewStartWithCustomTarget() async throws { let backend = FakeCodexReviewBackend() let store = CodexReviewStore.makeTestingStore( @@ -1500,7 +1823,10 @@ struct CodexReviewMCPHTTPServerTests { private func withHTTPServer( store: CodexReviewStore, - configuration: CodexReviewMCPHTTPServer.Configuration = .init(port: 0), + configuration: CodexReviewMCPHTTPServer.Configuration = .init( + host: "127.0.0.1", + port: 0 + ), operation: (CodexReviewMCPHTTPServer) async throws -> T ) async throws -> T { let adapter = CodexReviewMCPServer(store: store) @@ -1561,6 +1887,71 @@ struct CodexReviewMCPHTTPServerTests { try JSONSerialization.data(withJSONObject: body) } + private func makeToolsListBody(id: Int) throws -> Data { + try makeJSONBody([ + "jsonrpc": "2.0", + "id": id, + "method": "tools/list", + ]) + } + + private func makeReviewStartBody(id: Int) throws -> Data { + try makeJSONBody([ + "jsonrpc": "2.0", + "id": id, + "method": "tools/call", + "params": [ + "name": "review_start", + "arguments": [ + "cwd": "/tmp/project", + "target": ["type": "uncommittedChanges"], + ], + ], + ]) + } + + private func rawHTTPRequest( + endpoint: URL, + method: String = "POST", + version: String = "HTTP/1.1", + sessionID: String?, + headers: [(String, String)] = [], + body: Data? = nil + ) throws -> Data { + let components = try #require(URLComponents(url: endpoint, resolvingAgainstBaseURL: false)) + let host = try #require(components.host) + let port = try #require(components.port) + var requestHeaders: [(String, String)] = [ + ("Host", "\(host):\(port)"), + ] + if method == "POST" { + requestHeaders.append(("Content-Type", "application/json")) + requestHeaders.append(("Accept", "text/event-stream, application/json")) + } + if let sessionID { + requestHeaders.append(("MCP-Session-Id", sessionID)) + } + requestHeaders.append(contentsOf: headers) + let hasFramingHeader = requestHeaders.contains { name, _ in + name.caseInsensitiveCompare("Content-Length") == .orderedSame + || name.caseInsensitiveCompare("Transfer-Encoding") == .orderedSame + } + if let body, hasFramingHeader == false { + requestHeaders.append(("Content-Length", "\(body.count)")) + } + + let serializedHeaders = requestHeaders + .map { "\($0.0): \($0.1)\r\n" } + .joined() + var request = Data( + "\(method) \(endpoint.path) \(version)\r\n\(serializedHeaders)\r\n".utf8 + ) + if let body { + request.append(body) + } + return request + } + private func canonicalJSON(_ value: Any) throws -> String { let data = try JSONSerialization.data(withJSONObject: value, options: [.sortedKeys]) return String(decoding: data, as: UTF8.self) @@ -1723,6 +2114,184 @@ struct CodexReviewMCPHTTPServerTests { } } +private final class RawHTTPConnection: @unchecked Sendable { + private let descriptor: Int32 + private let lock = NSLock() + private var bufferedInput = Data() + private var isClosed = false + + private init(descriptor: Int32) { + self.descriptor = descriptor + } + + static func connect(to endpoint: URL) async throws -> RawHTTPConnection { + let components = try #require(URLComponents(url: endpoint, resolvingAgainstBaseURL: false)) + let host = try #require(components.host) + let ipv4Host = host == "localhost" ? "127.0.0.1" : host + let port = try #require(components.port) + return try await Task.detached { + let descriptor = Darwin.socket(AF_INET, SOCK_STREAM, 0) + guard descriptor >= 0 else { + throw currentPOSIXError() + } + do { + var noSignal = Int32(1) + _ = withUnsafePointer(to: &noSignal) { + Darwin.setsockopt( + descriptor, + SOL_SOCKET, + SO_NOSIGPIPE, + $0, + socklen_t(MemoryLayout.size) + ) + } + var timeout = timeval(tv_sec: 5, tv_usec: 0) + _ = withUnsafePointer(to: &timeout) { + Darwin.setsockopt( + descriptor, + SOL_SOCKET, + SO_RCVTIMEO, + $0, + socklen_t(MemoryLayout.size) + ) + } + + 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 + guard inet_pton(AF_INET, ipv4Host, &address.sin_addr) == 1 else { + 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 { + throw currentPOSIXError() + } + return RawHTTPConnection(descriptor: descriptor) + } catch { + Darwin.close(descriptor) + throw error + } + }.value + } + + func send(_ bytes: Data) async throws { + let descriptor = self.descriptor + try await Task.detached { + try bytes.withUnsafeBytes { rawBuffer in + guard let baseAddress = rawBuffer.baseAddress else { + return + } + var sent = 0 + while sent < rawBuffer.count { + let count = Darwin.send( + descriptor, + baseAddress.advanced(by: sent), + rawBuffer.count - sent, + 0 + ) + if count < 0, errno == EINTR { + continue + } + guard count > 0 else { + throw currentPOSIXError() + } + sent += count + } + } + }.value + } + + func readResponseHead() async throws -> String { + let descriptor = self.descriptor + var bytes = lock.withLock { + defer { bufferedInput.removeAll(keepingCapacity: false) } + return bufferedInput + } + let terminator = Data("\r\n\r\n".utf8) + let result = try await Task.detached { + while let range = bytes.range(of: terminator) { + let head = Data(bytes[.. 0 else { + if count == 0 { + throw testError("Connection closed before the response head completed") + } + throw currentPOSIXError() + } + bytes.append(contentsOf: buffer.prefix(count)) + if let range = bytes.range(of: terminator) { + let head = Data(bytes[.. Data { + let descriptor = self.descriptor + var bytes = lock.withLock { + defer { bufferedInput.removeAll(keepingCapacity: false) } + return bufferedInput + } + return try await Task.detached { + while true { + var buffer = [UInt8](repeating: 0, count: 4096) + let count = Darwin.recv(descriptor, &buffer, buffer.count, 0) + if count < 0, errno == EINTR { + continue + } + if count == 0 { + return bytes + } + guard count > 0 else { + throw currentPOSIXError() + } + bytes.append(contentsOf: buffer.prefix(count)) + } + }.value + } + + func close() { + let shouldClose = lock.withLock { + guard isClosed == false else { + return false + } + isClosed = true + return true + } + if shouldClose { + Darwin.shutdown(descriptor, SHUT_RDWR) + Darwin.close(descriptor) + } + } + + deinit { + close() + } +} + private nonisolated func currentPOSIXError() -> NSError { NSError(domain: NSPOSIXErrorDomain, code: Int(errno)) } diff --git a/Tests/CodexReviewMCPServerTests/MCPHTTPNetworkResourceOwnerTests.swift b/Tests/CodexReviewMCPServerTests/MCPHTTPNetworkResourceOwnerTests.swift index e41d3a1..3bc0091 100644 --- a/Tests/CodexReviewMCPServerTests/MCPHTTPNetworkResourceOwnerTests.swift +++ b/Tests/CodexReviewMCPServerTests/MCPHTTPNetworkResourceOwnerTests.swift @@ -106,6 +106,35 @@ struct MCPHTTPNetworkResourceOwnerTests { await closing.waitUntilClosed() } + @Test func finalRequestAdmissionAtomicallyClosesFutureAdmission() async throws { + let owner = MCPHTTPNetworkResourceOwner(generationID: 9) + let resource = TestingConnectionResource() + let connection = try #require(owner.admitConnection(resource)) + + let admitted = try #require(connection.admitRequest(finalForConnection: true)) + + let snapshot = try #require(owner.snapshot().connections.first) + #expect(snapshot.phase == .admissionClosed) + #expect(snapshot.requests.map(\.id) == [admitted.operation.id]) + #expect(connection.admitRequest(finalForConnection: false) == nil) + + let task = Task { + defer { + admitted.lease.acknowledgeCompletion() + } + guard await admitted.lease.waitUntilStartIsAllowed() else { + return + } + } + admitted.lease.install(task) + await task.value + + let closing = owner.beginClosing(.serverStop) + await resource.waitUntilCloseIsSignalled() + resource.acknowledgeClose() + await closing.waitUntilClosed() + } + @Test func transportFailureWinsPeerCloseAndDrainsItsRequest() async throws { let owner = MCPHTTPNetworkResourceOwner(generationID: 4) let resource = TestingConnectionResource()