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
158 changes: 145 additions & 13 deletions Sources/SwiftNetwork/Protocols/IPProtocol.swift
Original file line number Diff line number Diff line change
Expand Up @@ -918,7 +918,11 @@ public struct IPProtocol: NetworkProtocol {
}
}

mutating func writeOutboundFrames(_ frames: inout FrameArray) {
mutating func writeOutboundFrames(
_ frames: inout FrameArray,
lower: OutboundDatagramLinkage,
selfReference: ProtocolInstanceReference
) {
frames.iterateMutableFrames { frame in
guard frame.unclaim(fromStart: IPv4Instance.headerLength) else {
frame.finalize(success: false)
Expand Down Expand Up @@ -969,12 +973,23 @@ public struct IPProtocol: NetworkProtocol {
var fragmentationSucceeded = true
var fragmentFrames = FrameArray()
while cursor < payloadLength {
// Determine if last or hold large the chunk length is
// Determine if last or how large the chunk length is
let remaining = payloadLength - cursor
let isLast = remaining <= fragmentRoom
let chunkLength = isLast ? remaining : fragmentRoom
// Create the fragment frame with this chunk length
var fragmentFrame = Frame(count: IPv4Instance.headerLength + chunkLength)
let fragmentFrameSize = IPv4Instance.headerLength + chunkLength
guard
var allocatedFrames = try? lower.invokeGetDatagramsToSend(
selfReference,
maximumDatagramCount: 1,
minimumDatagramSize: fragmentFrameSize
),
var fragmentFrame = allocatedFrames.popFirst()
else {
fragmentationSucceeded = false
break
}
// MF bit is always set except for the last fragment
let ipOff = UInt16(isLast ? 0 : 0x2000) | UInt16(cursor / 8)
let fragmentTotalLength = UInt16(IPv4Instance.headerLength + chunkLength)
Expand Down Expand Up @@ -1663,11 +1678,15 @@ public struct IPProtocol: NetworkProtocol {
}
}

mutating func writeOutboundFrames(_ frames: inout FrameArray) {
frames.iterateMutableFrames { frame in
mutating func writeOutboundFrames(
_ frames: inout FrameArray,
lower: OutboundDatagramLinkage,
selfReference: ProtocolInstanceReference
) {
frames.iterateMutableFrames { (frame: inout Frame) -> FrameArray.FrameIterationResult in
_ = frame.unclaim(fromStart: IPv6Instance.headerLength)

let payloadLength = UInt16(frame.unclaimedLength - IPv6Instance.headerLength)
let payloadLength = frame.unclaimedLength - IPv6Instance.headerLength
let localAddressValue = self.localAddress.addressValue
let remoteAddressValue = self.remoteAddress.addressValue

Expand All @@ -1686,9 +1705,110 @@ public struct IPProtocol: NetworkProtocol {
if dscpValue != 0 {
flow |= UInt32(bigEndian: (UInt32(dscpValue) << 22) & 0x0fc0_0000) // IP6FLOW_DSCP_SHIFT
}

let enableFragmentation: Bool
if let fragmentationOverride = frame.fragmentationOverride {
enableFragmentation = fragmentationOverride
} else {
enableFragmentation = self.flags.enableFragmentation
}

// IPv6 header + Fragment Extension Header
let ipv6CompleteHeaderLength =
IPv6Instance.headerLength + IPv6Instance.fragmentExtensionHeaderLength
let mtu = self.pathProperties.mtu
var maxPayloadPerFragment = 0
if mtu > ipv6CompleteHeaderLength {
maxPayloadPerFragment = mtu - ipv6CompleteHeaderLength
}

// Handle fragmentation if payloadLength is greater than maxPayloadPerFragment and enableFragmentation is enabled
if enableFragmentation && maxPayloadPerFragment > 0 && payloadLength > maxPayloadPerFragment {
var randomNumber = SystemRandomNumberGenerator()
let fragmentID = UInt32(truncatingIfNeeded: randomNumber.next())
// Align fragment payload to blocks of 8 bytes - RFC 2460
let fragmentRoom = maxPayloadPerFragment - (maxPayloadPerFragment % 8)
guard fragmentRoom > 0 else {
frame.finalize(success: false)
return .removeFrameAndContinue
}
var cursor = 0
var fragmentationSucceeded = true
var fragmentFrames = FrameArray()
while cursor < payloadLength {
// Determine if last or how large the chunk length is
let remaining = payloadLength - cursor
let isLast = remaining <= fragmentRoom
let chunkLength = isLast ? remaining : fragmentRoom
// Fragment Extension Header + this chunk's payload.
let fragmentLength = UInt16(chunkLength + IPv6Instance.fragmentExtensionHeaderLength)
// Fragment offset flags
let offsetFlags = UInt16(cursor) | (isLast ? 0 : UInt16(IPv6Instance.ip6fMoreFragmentMask))
let fragmentFrameSize = ipv6CompleteHeaderLength + chunkLength
guard
var allocatedFrames = try? lower.invokeGetDatagramsToSend(
selfReference,
maximumDatagramCount: 1,
minimumDatagramSize: fragmentFrameSize
),
var fragmentFrame = allocatedFrames.popFirst()
else {
fragmentationSucceeded = false
break
}
let result = Serializer.serialize(&fragmentFrame, claim: false) {
write throws(SerializationError) in
// IPv6 base header
try write.uint32(flow)
try write.uint16NetworkByteOrder(fragmentLength)
try write.uint8(IPv6Instance.fragmentExtensionHeader)
try write.uint8(self.hopLimit)
try write.uint32(localAddressValue.0)
try write.uint32(localAddressValue.1)
try write.uint32(localAddressValue.2)
try write.uint32(localAddressValue.3)
try write.uint32(remoteAddressValue.0)
try write.uint32(remoteAddressValue.1)
try write.uint32(remoteAddressValue.2)
try write.uint32(remoteAddressValue.3)
// Fragment Extension Header
try write.uint8(self.ipProtocolNumber)
try write.uint8(0)
try write.uint16NetworkByteOrder(offsetFlags)
try write.uint32(fragmentID)
}
guard result.isValid else {
Logger.proto.error("Serializing IPv6 fragment failed with result: \(result)")
fragmentFrame.finalize(success: false)
fragmentationSucceeded = false
break
}
let copied = frame.copyInto(
&fragmentFrame,
atOffset: ipv6CompleteHeaderLength,
fromOffset: IPv6Instance.headerLength + cursor,
length: chunkLength
)
guard copied == chunkLength else {
fragmentFrame.finalize(success: false)
fragmentationSucceeded = false
break
}
self.counters.txPackets += 1
fragmentFrames.add(frame: fragmentFrame)
cursor += chunkLength
}
frame.finalize(success: fragmentationSucceeded)
if fragmentationSucceeded {
return .replaceWithFramesAndContinue(fragmentFrames)
}
fragmentFrames.finalizeAllFramesAsFailed()
return .removeFrameAndContinue
}
// Standard path
let result = Serializer.serialize(&frame, claim: false) { write throws(SerializationError) in
try write.uint32(flow)
try write.uint16NetworkByteOrder(payloadLength)
try write.uint16NetworkByteOrder(UInt16(payloadLength))
try write.uint8(self.ipProtocolNumber)
try write.uint8(self.hopLimit)
try write.uint32(localAddressValue.0)
Expand All @@ -1702,10 +1822,10 @@ public struct IPProtocol: NetworkProtocol {
}
if !result.isValid {
Logger.proto.error("Serializing IPv6 packet failed with result: \(result)")
return true
return .continueIterating
}
self.counters.txPackets += 1
return true
return .continueIterating
}
}
}
Expand Down Expand Up @@ -1883,7 +2003,14 @@ public struct IPProtocol: NetworkProtocol {
}

mutating func sendDatagrams(_ datagrams: consuming FrameArray) throws(NetworkError) {
IPInstance.processOutbound(&self.instanceType, datagrams: &datagrams)
let lower = self.lower
let selfReference = self.effectiveSelfReference
IPInstance.processOutbound(
&self.instanceType,
lower: lower,
selfReference: selfReference,
datagrams: &datagrams
)
try invokeSendDatagrams(datagrams)
}

Expand All @@ -1904,13 +2031,18 @@ public struct IPProtocol: NetworkProtocol {
}

@inline(__always)
private static func processOutbound(_ instanceType: inout IPInstanceType, datagrams: inout FrameArray) {
private static func processOutbound(
_ instanceType: inout IPInstanceType,
lower: OutboundDatagramLinkage,
selfReference: ProtocolInstanceReference,
datagrams: inout FrameArray
) {
switch instanceType {
case .ipv4(var instance):
instance.writeOutboundFrames(&datagrams)
instance.writeOutboundFrames(&datagrams, lower: lower, selfReference: selfReference)
instanceType = .ipv4(instance)
case .ipv6(var instance):
instance.writeOutboundFrames(&datagrams)
instance.writeOutboundFrames(&datagrams, lower: lower, selfReference: selfReference)
instanceType = .ipv6(instance)
}
}
Expand Down
125 changes: 125 additions & 0 deletions Tests/SwiftNetworkTests/SwiftNetworkIPTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -1026,6 +1026,131 @@ final class SwiftNetworkIPTests: NetTestCase {
wait(for: [expectation], timeout: 10.0)
}

func testIPv6OutboundFragmentation() {
// Send a 29-byte payload on a path with MTU set to 64 and fragmentation enabled.
// IPv6 base header is 40 bytes, Fragment Extension Header is 8 bytes (total overhead = 48).
// Max payload per fragment = 64 - 48 = 16 bytes, aligned to 8 = 16 bytes.
// Expected: 2 fragments — payload[0..<16] with MF=1, payload[16..<29] with MF=0.
let mtu = 64
let payload: [UInt8] = [
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08,
0x09, 0x0A, 0x0B, 0x0C, 0x0D, 0x0E, 0x0F, 0x10,
0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18,
0x19, 0x1A, 0x1B, 0x1C, 0x1D,
]
let parameters = Parameters()
let expectation = XCTestExpectation()
let context = parameters.context
context.async {
defer { expectation.fulfill() }
var path = PathProperties(parameters: parameters)
path.directInterface = Interface(
index: 1,
name: "lo0",
type: .loopback,
subtype: .other,
mtu: mtu
)
path.effectiveMTU = UInt32(mtu)

let reference = IPProtocol.instance(context: parameters.context)
let ipOptions = IPProtocol.options()
ipOptions.flags = IPProtocol.IPOptions.Flags(rawValue: ipOptions.flags.rawValue)
.union(.fragmentationEnabledOverridden)
.union(.fragmentationEnabled)
ipOptions.setLogID(prefix: "C", parent: "1", protocolLogIDNumber: 2)
ipOptions.setProtocolInstance(reference)
parameters.defaultStack.internet = .ip(ipOptions)

let udpOptions = UDPProtocol.options()
udpOptions.noMetadata = true
udpOptions.setLogID(prefix: "C", parent: "1", protocolLogIDNumber: 1)
parameters.defaultStack.transport = .udp(udpOptions)

let localEndpoint = Endpoint(address: IPv6Address(SwiftNetworkIPTests.localIPv6Address)!, port: 0)
let remoteEndpoint = Endpoint(address: IPv6Address(SwiftNetworkIPTests.remoteIPv6Address)!, port: 0)
let ipLinkage = OutboundDatagramLinkage(reference: reference)
guard
let upperHarness = DatagramUpperHarness(
identifier: "Client",
local: localEndpoint,
remote: remoteEndpoint,
parameters: parameters,
path: path,
context: parameters.context,
lowerProtocol: ipLinkage
)
else {
XCTFail("Failed to attach IP to upper harness")
return
}
let lowerHarness = DatagramLowerHarness(identifier: "Client", context: parameters.context)
do {
try reference.attachLowerDatagramProtocol(
lowerHarness.reference,
remote: remoteEndpoint,
local: localEndpoint,
parameters: parameters,
path: path
)
} catch {
XCTFail("Failed to attach IP to lower harness")
return
}
upperHarness.start { connected in XCTAssertTrue(connected) }
_ = upperHarness.write(payload)

var fragments: [[UInt8]] = []
while lowerHarness.hasOutboundPackets {
if let packet = lowerHarness.extractLastOutboundPacket() {
fragments.append(packet)
}
}
// Expect exactly 2 fragments: 16-byte payload + 13-byte payload
XCTAssertEqual(fragments.count, 2, "Expected exactly 2 IPv6 fragments")
guard fragments.count == 2,
let fragmentOne = fragments.first,
let fragmentTwo = fragments.last
else { return }

// Each fragment must be at least 48 bytes (40 IPv6 base header + 8 Fragment Extension Header)
let minFragmentSize = 40 + 8
XCTAssertGreaterThanOrEqual(fragmentOne.count, minFragmentSize)
XCTAssertGreaterThanOrEqual(fragmentTwo.count, minFragmentSize)
guard fragmentOne.count >= minFragmentSize, fragmentTwo.count >= minFragmentSize else { return }

XCTAssertEqual(fragmentOne[6], 0x2C, "Fragment 1 Next header must be IPPROTO_FRAGMENT (44)")
XCTAssertEqual(fragmentTwo[6], 0x2C, "Fragment 2 Next header must be IPPROTO_FRAGMENT (44)")

XCTAssertEqual(fragmentOne[40], 0x11, "Fragment 1 must be UDP (17)")
XCTAssertEqual(fragmentTwo[40], 0x11, "Fragment 2 must be UDP (17)")

// Both fragments must share the same identifier (bytes 44–47)
XCTAssertEqual(
Array(fragmentOne[44...47]),
Array(fragmentTwo[44...47]),
"All fragments must share the same IPv6 fragment identifier"
)
let offsetFlags1 = UInt16(fragmentOne[42]) << 8 | UInt16(fragmentOne[43])
XCTAssertNotEqual(offsetFlags1 & 0x0001, 0, "Fragment 1 must have MF=1") // More fragments
XCTAssertEqual(offsetFlags1 & 0xFFF8, 0, "Fragment 1 must have byte offset=0")
// Fragment 1 payload must be the first 16 bytes of the original payload
XCTAssertEqual(fragmentOne.count, minFragmentSize + 16, "Fragment 1 must be header + 16 payload bytes")
XCTAssertEqual(Array(fragmentOne[minFragmentSize...]), Array(payload[0..<16]))

let offsetFlags2 = UInt16(fragmentTwo[42]) << 8 | UInt16(fragmentTwo[43])
XCTAssertEqual(offsetFlags2 & 0x0001, 0, "Fragment 2 must have MF=0") // No more fragments
XCTAssertEqual(offsetFlags2 & 0xFFF8, 16, "Fragment 2 must have byte offset=16")
// Fragment 2 payload must be the remaining 13 bytes of the original payload
XCTAssertEqual(fragmentTwo.count, minFragmentSize + 13, "Fragment 2 must be header + 13 payload bytes")
XCTAssertEqual(Array(fragmentTwo[minFragmentSize...]), Array(payload[16...]))

upperHarness.stop()
upperHarness.teardown()
}
wait(for: [expectation], timeout: 10.0)
}

// Sets up a minimal IP harness to test different fragment and reassembly conditions
private func processIPFragment(
packets: [[UInt8]],
Expand Down
Loading