Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ import NIOHTTP1

/// NIO-based channel handler for unary requests made through the Connect library.
final class ConnectUnaryChannelHandler: NIOCore.ChannelInboundHandler, @unchecked Sendable {
private let contentLength: NIOHTTPClient.ContentLength
private let eventLoop: NIOCore.EventLoop
private let request: Connect.HTTPRequest<Data?>
private let onMetrics: (Connect.HTTPMetrics) -> Void
Expand All @@ -35,11 +36,13 @@ final class ConnectUnaryChannelHandler: NIOCore.ChannelInboundHandler, @unchecke
init(
request: Connect.HTTPRequest<Data?>,
eventLoop: NIOCore.EventLoop,
contentLength: NIOHTTPClient.ContentLength = .sent,
onMetrics: @escaping (Connect.HTTPMetrics) -> Void,
onResponse: @escaping (Connect.HTTPResponse) -> Void
) {
self.request = request
self.eventLoop = eventLoop
self.contentLength = contentLength
self.onMetrics = onMetrics
self.onResponse = onResponse
}
Expand Down Expand Up @@ -103,7 +106,7 @@ final class ConnectUnaryChannelHandler: NIOCore.ChannelInboundHandler, @unchecke
}

var nioHeaders = NIOHTTP1.HTTPHeaders()
if let messageLength = self.request.message?.count {
if case .sent = self.contentLength, let messageLength = self.request.message?.count {
nioHeaders.add(name: "Content-Length", value: "\(messageLength)")
}
nioHeaders.add(name: "Host", value: self.request.url.host!)
Expand Down
19 changes: 18 additions & 1 deletion Libraries/ConnectNIO/Public/NIOHTTPClient.swift
Original file line number Diff line number Diff line change
Expand Up @@ -31,10 +31,18 @@ open class NIOHTTPClient: Connect.HTTPClientInterface, @unchecked Sendable {
private let port: Int
private let timeout: TimeInterval?
private let useSSL: Bool
private let contentLength: ContentLength

private var pendingRequests = [(NIOHTTP2.NIOHTTP2Handler.StreamMultiplexer?) -> Void]()
private var state = State.disconnected

/// Whether unary requests should include a `Content-Length` header. It is optional over
/// HTTP/2, the only protocol this client speaks.
public enum ContentLength: Sendable {
case sent
case omitted
}

private enum State {
case disconnected
case connecting
Expand All @@ -52,13 +60,21 @@ open class NIOHTTPClient: Connect.HTTPClientInterface, @unchecked Sendable {
/// (e.g., `https://connectrpc.com:8080`), the host's port will be used.
/// - parameter timeout: Optional timeout after which to terminate requests/streams if no
/// activity has occurred in the request or response path.
public init(host: String, port: Int? = nil, timeout: TimeInterval? = nil) {
/// - parameter contentLength: Whether unary requests should include a `Content-Length`
/// header. Defaults to sending it.
public init(
host: String,
port: Int? = nil,
timeout: TimeInterval? = nil,
contentLength: ContentLength = .sent
) {
let baseURL = URL(string: host)!
let useSSL = baseURL.scheme?.lowercased() == "https"
self.host = baseURL.host!
self.port = port ?? baseURL.port ?? (useSSL ? 443 : 80)
self.timeout = timeout
self.useSSL = useSSL
self.contentLength = contentLength
}

/// Called before the first request/stream is initialized, and the result is stored for reuse
Expand Down Expand Up @@ -118,6 +134,7 @@ open class NIOHTTPClient: Connect.HTTPClientInterface, @unchecked Sendable {
let handler = ConnectUnaryChannelHandler(
request: request,
eventLoop: eventLoop,
contentLength: self.contentLength,
onMetrics: onMetrics,
onResponse: onResponse
)
Expand Down
1 change: 1 addition & 0 deletions Package.swift
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,7 @@ let package = Package(
dependencies: [
"Connect",
"ConnectMocks",
"ConnectNIO",
.product(name: "SwiftProtobuf", package: "swift-protobuf"),
],
path: "Tests/UnitTests/ConnectLibraryTests",
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
// Copyright 2022-2025 The Connect Authors
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.

import Connect
@testable import ConnectNIO
import Foundation
import NIOCore
import NIOEmbedded
import NIOHTTP1
import Testing

struct ConnectUnaryChannelHandlerTests {
@Test
func sendsContentLengthByDefault() throws {
let headers = try Self.outboundHeaders(contentLength: .sent)

#expect(headers.first(name: "Content-Length") == "5")
}

@Test
func omitsContentLengthWhenRequested() throws {
let headers = try Self.outboundHeaders(contentLength: .omitted)

#expect(!headers.contains(name: "Content-Length"))
}

@Test(arguments: [NIOHTTPClient.ContentLength.sent, .omitted])
func alwaysSendsHostAndCallerHeaders(contentLength: NIOHTTPClient.ContentLength) throws {
let headers = try Self.outboundHeaders(contentLength: contentLength)

#expect(headers.first(name: "Host") == "connectrpc.com")
#expect(headers.first(name: "x-custom-header") == "value")
}

private static func outboundHeaders(
contentLength: NIOHTTPClient.ContentLength
) throws -> NIOHTTP1.HTTPHeaders {
let channel = EmbeddedChannel()
let request = Connect.HTTPRequest<Data?>(
url: URL(string: "https://connectrpc.com/connectrpc.Service/Method")!,
headers: ["x-custom-header": ["value"]],
message: Data([0, 0, 0, 0, 0]),
method: .post,
trailers: nil,
idempotencyLevel: .unknown
)
let handler = ConnectUnaryChannelHandler(
request: request,
eventLoop: channel.eventLoop,
contentLength: contentLength,
onMetrics: { _ in },
onResponse: { _ in }
)
try channel.pipeline.syncOperations.addHandler(handler)
channel.pipeline.fireChannelActive()

let outbound = try channel.readOutbound(as: NIOHTTP1.HTTPClientRequestPart.self)
guard case .head(let head) = try #require(outbound) else {
Issue.record("Expected a request head to be written")
return NIOHTTP1.HTTPHeaders()
}
return head.headers
}
}