Skip to content

Commit fe09a97

Browse files
authored
Merge pull request #42 from apple/agnosticdev/UpdateUDP
Updates to Swift UDP
2 parents c22eefb + 93ccf9e commit fe09a97

5 files changed

Lines changed: 245 additions & 21 deletions

File tree

Sources/SwiftNetwork/Protocols/Checksum.swift

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -57,6 +57,23 @@ enum ChecksumError: Error {
5757
case invalidBuffer
5858
}
5959

60+
struct ChecksumFlags: OptionSet {
61+
let rawValue: UInt8
62+
static let partial = ChecksumFlags(rawValue: 0x01)
63+
static let zeroInvert = ChecksumFlags(rawValue: 0x02)
64+
static let ip = ChecksumFlags(rawValue: 0x04)
65+
static let tcpIPv4 = ChecksumFlags(rawValue: 0x08)
66+
static let udpIPv4 = ChecksumFlags(rawValue: 0x10)
67+
static let tcpIPv6 = ChecksumFlags(rawValue: 0x20)
68+
static let udpIPv6 = ChecksumFlags(rawValue: 0x40)
69+
}
70+
71+
struct InterfaceChecksumFlags: OptionSet {
72+
let rawValue: UInt32
73+
static let udpIPv4 = InterfaceChecksumFlags(rawValue: 0x0000_0004)
74+
static let udpIPv6 = InterfaceChecksumFlags(rawValue: 0x0000_0040)
75+
}
76+
6077
@available(Network 0.1.0, *)
6178
extension IPv6Address {
6279
func checksum() -> UInt32 {
@@ -141,6 +158,11 @@ extension Frame {
141158
return ((~value) & 0xffff)
142159
}
143160

161+
mutating func setInternetChecksum(flags: ChecksumFlags, startOffset: UInt16, checksumOffset: UInt16) -> Bool {
162+
checksumOffloadFlags |= flags.rawValue
163+
return false
164+
}
165+
144166
mutating func finalizeIPChecksum(checksumOffset: Int, zeroInvert: Bool) throws(ChecksumError) {
145167
let unclaimedLength = self.unclaimedLength
146168
if unclaimedLength == 0 {

Sources/SwiftNetwork/Protocols/HarnessProtocols.swift

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -237,6 +237,10 @@ public class UpperHarness<LinkageType: InboundDataLinkage>: UpperHarnessProtocol
237237
return invokeGetMetadata() as? ProtocolMetadata<P>
238238
}
239239

240+
final public func getMetrics(requestedNetworkMetric: RequestedNetworkMetrics) -> NetworkMetrics? {
241+
lower.invokeGetMetrics(reference, requestedNetworkMetric: requestedNetworkMetric)
242+
}
243+
240244
public func setApplicationError(_ applicationError: UInt64, applicationErrorReason: String) {
241245
if let metadata: ProtocolMetadata<QUICProtocol> = self.getMetadata() {
242246
metadata.perProtocolMetadata?.quicConnectionMetadata?.applicationError = applicationError

Sources/SwiftNetwork/Protocols/UDPProtocol.swift

Lines changed: 77 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,7 @@ public struct UDPProtocol: NetworkProtocol {
4848
static public let noMetadata = UDPOptions(rawValue: 1 << 1)
4949
static public let ignoreInboundChecksum = UDPOptions(rawValue: 1 << 2)
5050
static public let useQUICStats = UDPOptions(rawValue: 1 << 3)
51+
static public let fullChecksumOffload = UDPOptions(rawValue: 1 << 4)
5152

5253
#if NETWORK_PRIVATE
5354
var privateStorage = UDPProtocolOptionsPrivateStorage()
@@ -97,7 +98,9 @@ public struct UDPProtocol: NetworkProtocol {
9798
var eventManager = ProtocolEventManager()
9899

99100
// Only called by newProtocolInstance()
100-
fileprivate static func registerNewUDP(on context: NetworkContext) -> ProtocolInstanceReference {
101+
fileprivate static func registerNewUDP(
102+
on context: NetworkContext,
103+
) -> ProtocolInstanceReference {
101104
let udp = UDPInstance(context: context)
102105
let registeredIndex = context.registerUDPInstance(udp)
103106
context.udpInstances[registeredIndex].udpInstanceIndex = registeredIndex
@@ -121,6 +124,9 @@ public struct UDPProtocol: NetworkProtocol {
121124

122125
var serviceClass = Parameters.ServiceClass.bestEffort
123126
var maximumDatagramSize: Int = 0
127+
var isIPv4: Bool {
128+
flags.contains(.isIPv4)
129+
}
124130

125131
struct Flags: OptionSet {
126132
init(rawValue: Self.RawValue) {
@@ -153,6 +159,13 @@ public struct UDPProtocol: NetworkProtocol {
153159
if self.maximumDatagramSize > UDPProtocol.headerLength {
154160
self.maximumDatagramSize -= UDPProtocol.headerLength
155161
}
162+
let udpChecksumOffload: UInt32 = path.hardwareChecksumFlags
163+
if isIPv4 && ((udpChecksumOffload & InterfaceChecksumFlags.udpIPv4.rawValue) != 0)
164+
|| !isIPv4 && ((udpChecksumOffload & InterfaceChecksumFlags.udpIPv6.rawValue) != 0)
165+
{
166+
self.flags.insert(.fullChecksumOffload)
167+
self.flags.remove(.partialChecksumOffload)
168+
}
156169
}
157170

158171
mutating func setup(
@@ -202,6 +215,21 @@ public struct UDPProtocol: NetworkProtocol {
202215
if udpOptions.ignoreInboundChecksum {
203216
self.flags.insert(.ignoreInboundChecksum)
204217
}
218+
if udpOptions.fullChecksumOffload {
219+
self.flags.insert(.fullChecksumOffload)
220+
}
221+
#if !NETWORK_EMBEDDED
222+
if let transport = parameters.defaultStack.transport {
223+
if transport.options == udpOptions, udpOptions.useQUICStats {
224+
self.flags.insert(.upperTransportIsQUIC)
225+
} else if let quicOptions = transport.options,
226+
quicOptions.matches(identifier: QUICStreamProtocol.identifier)
227+
|| quicOptions.matches(identifier: QUICConnectionProtocol.identifier)
228+
{
229+
self.flags.insert(.upperTransportIsQUIC)
230+
}
231+
}
232+
#endif
205233
}
206234
}
207235

@@ -213,7 +241,7 @@ public struct UDPProtocol: NetworkProtocol {
213241
let inbound = (inboundChecksum != nil)
214242
let existingChecksum: UInt16 = inboundChecksum ?? 0
215243
var checksumValue: UInt16 = 0
216-
if self.flags.contains(.isIPv4) {
244+
if isIPv4 {
217245
checksumValue = Checksum.ipv4PseudoHeader(
218246
source: inbound ? ipv4Remote : ipv4Local,
219247
dest: inbound ? ipv4Local : ipv4Remote,
@@ -295,7 +323,7 @@ public struct UDPProtocol: NetworkProtocol {
295323
return .removeFrameAndContinue
296324
}
297325

298-
guard self.flags.contains(.isIPv4) || checksum != 0 else {
326+
guard isIPv4 || checksum != 0 else {
299327
log.error("Received an IPv6 packet with zero checksum")
300328
frame.finalize(success: false)
301329
return .removeFrameAndContinue
@@ -391,7 +419,7 @@ public struct UDPProtocol: NetworkProtocol {
391419
frame.serviceClass = self.serviceClass
392420
}
393421

394-
if !self.flags.contains(.isIPv4) || !self.flags.contains(.noChecksum) {
422+
if !isIPv4 || !self.flags.contains(.noChecksum) {
395423
// Always insert pseudo header checksum
396424
let checksumValue = pseudoHeaderChecksum(inboundChecksum: nil, length: length)
397425

@@ -407,23 +435,41 @@ public struct UDPProtocol: NetworkProtocol {
407435
return .removeFrameAndContinue
408436
}
409437

410-
let finalizedChecksum = false
411438
if self.flags.contains(.fullChecksumOffload) {
412-
// TODO: Checksum offload
413-
} else if self.flags.contains(.partialChecksumOffload) {
414-
// TODO: Checksum offload
439+
let csumFlags: ChecksumFlags =
440+
isIPv4
441+
? [.udpIPv4, .zeroInvert]
442+
: [.udpIPv6, .zeroInvert]
443+
frame.checksumOffloadFlags |= csumFlags.rawValue
415444
}
416445

417-
if !finalizedChecksum {
418-
do throws(ChecksumError) {
419-
try frame.finalizeIPChecksum(checksumOffset: checksumOffset, zeroInvert: true)
420-
} catch {
421-
log.error("Failed to finalize UDP checksum")
422-
frame.finalize(success: false)
423-
return .removeFrameAndContinue
446+
if !self.flags.contains(.fullChecksumOffload) {
447+
var checksumOffloadDisabled = !self.flags.contains(.partialChecksumOffload)
448+
if !checksumOffloadDisabled {
449+
let csumStart = UInt16(isIPv4 ? IPProtocol.ipv4HeaderLength : IPProtocol.ipv6HeaderLength)
450+
let csumWrite = csumStart + UInt16(checksumOffset)
451+
let csumFlags: ChecksumFlags = [.partial, .zeroInvert]
452+
if !frame.setInternetChecksum(
453+
flags: csumFlags,
454+
startOffset: csumStart,
455+
checksumOffset: csumWrite
456+
) {
457+
checksumOffloadDisabled = true
458+
}
459+
}
460+
461+
if checksumOffloadDisabled {
462+
do throws(ChecksumError) {
463+
try frame.finalizeIPChecksum(checksumOffset: checksumOffset, zeroInvert: true)
464+
} catch {
465+
log.error("Failed to finalize UDP checksum")
466+
frame.finalize(success: false)
467+
return .removeFrameAndContinue
468+
}
424469
}
425470
}
426471
}
472+
transmitByteCount += length - UDPProtocol.headerLength
427473

428474
return .continueIterating
429475
}
@@ -434,6 +480,11 @@ public struct UDPProtocol: NetworkProtocol {
434480
#if !NETWORK_EMBEDDED
435481
var metadata: AbstractProtocolMetadata? { nil }
436482
#endif
483+
484+
func updateDataTransferSnapshot(_ snapshot: inout DataTransferSnapshot) {
485+
snapshot.receivedTransportByteCount = UInt64(receiveByteCount)
486+
snapshot.sentTransportByteCount = UInt64(transmitByteCount)
487+
}
437488
}
438489

439490
public init() {}
@@ -506,4 +557,15 @@ extension ProtocolOptions<UDPProtocol> {
506557
}
507558
}
508559
}
560+
561+
public var fullChecksumOffload: Bool {
562+
get { perProtocolOptions!.contains(.fullChecksumOffload) }
563+
set {
564+
if newValue {
565+
perProtocolOptions!.insert(.fullChecksumOffload)
566+
} else {
567+
perProtocolOptions!.remove(.fullChecksumOffload)
568+
}
569+
}
570+
}
509571
}

Tests/SwiftNetworkTests/SwiftNetworkMultiplexingTests.swift

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,10 @@ import XCTest
2525
@available(Network 0.1.0, *)
2626
final class SwiftNetworkMultiplexingTests: NetTestCase {
2727

28+
// 10.0.0.20
29+
static let localIPv4Address: [UInt8] = [0x0a, 0x00, 0x00, 0x14]
30+
// 10.0.0.117
31+
static let remoteIPv4Address: [UInt8] = [0x0a, 0x00, 0x00, 0x75]
2832
static let outputMessage: [UInt8] = [0x0a, 0x0b, 0x0c, 0x0d]
2933
static let inputMessage: [UInt8] = [0x0d, 0x0c, 0x0b, 0x0a, 0x01]
3034

@@ -33,8 +37,8 @@ final class SwiftNetworkMultiplexingTests: NetTestCase {
3337
let context = parameters.context
3438
let path = PathProperties(parameters: parameters)
3539

36-
let localEndpoint = Endpoint(address: IPv4Address(SwiftNetworkUDPTests.localIPv4Address)!, port: 1234)
37-
let remoteEndpoint = Endpoint(address: IPv4Address(SwiftNetworkUDPTests.localIPv4Address)!, port: 2345)
40+
let localEndpoint = Endpoint(address: IPv4Address(SwiftNetworkMultiplexingTests.localIPv4Address)!, port: 1234)
41+
let remoteEndpoint = Endpoint(address: IPv4Address(SwiftNetworkMultiplexingTests.localIPv4Address)!, port: 2345)
3842

3943
var instance: TestMultiplexingProtocol? = nil
4044
var upperHarness1: DatagramUpperHarness?
@@ -217,8 +221,8 @@ final class SwiftNetworkMultiplexingTests: NetTestCase {
217221
let context = parameters.context
218222
let path = PathProperties(parameters: parameters)
219223

220-
let localEndpoint = Endpoint(address: IPv4Address(SwiftNetworkUDPTests.localIPv4Address)!, port: 1234)
221-
let remoteEndpoint = Endpoint(address: IPv4Address(SwiftNetworkUDPTests.localIPv4Address)!, port: 2345)
224+
let localEndpoint = Endpoint(address: IPv4Address(SwiftNetworkMultiplexingTests.localIPv4Address)!, port: 1234)
225+
let remoteEndpoint = Endpoint(address: IPv4Address(SwiftNetworkMultiplexingTests.localIPv4Address)!, port: 2345)
222226

223227
// Use a high number of streams to ensure that we don't have poor scaling
224228
let upperHarnessCount = 1000

0 commit comments

Comments
 (0)