@@ -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}
0 commit comments