From e3335149a86a630c8e4e7c92628ba94931d373f8 Mon Sep 17 00:00:00 2001 From: Rui Paulo Date: Fri, 2 Oct 2026 16:10:04 -0700 Subject: [PATCH] CongestionControl: replace enum with struct and operation dispatch - Convert `CongestionControl` from an enum into a struct holding all controllers, dispatching through `CongestionControlOperation` and `CongestionControlQuery` protocols to avoid per-call copy-and-reassign. - Add placeholder initializers to each controller and make the type non-Optional on `QUICPath`, simplifying all call sites. - Generalize `inherit(from:)` to accept any `CongestionControlProtocol` and centralize re-creation and handoff via `reset` and `switchTo`. - Add `CongestionControlTests` covering dispatch, switching, and reset. --- .../SwiftNetwork/QUIC/CongestionControl.swift | 562 ++++++++++-------- Sources/SwiftNetwork/QUIC/Cubic.swift | 21 +- Sources/SwiftNetwork/QUIC/Ledbat.swift | 29 +- Sources/SwiftNetwork/QUIC/Prague.swift | 27 +- Sources/SwiftNetwork/QUIC/QUICPath.swift | 150 ++--- Tests/QUICTests/CongestionControlTests.swift | 301 ++++++++++ 6 files changed, 696 insertions(+), 394 deletions(-) create mode 100644 Tests/QUICTests/CongestionControlTests.swift diff --git a/Sources/SwiftNetwork/QUIC/CongestionControl.swift b/Sources/SwiftNetwork/QUIC/CongestionControl.swift index ebe1929c..88994080 100644 --- a/Sources/SwiftNetwork/QUIC/CongestionControl.swift +++ b/Sources/SwiftNetwork/QUIC/CongestionControl.swift @@ -24,69 +24,179 @@ internal import Logging internal import os #endif +// An operation on whichever congestion controller is active. @available(Network 0.1.0, *) -enum CongestionControl { - case cubic(algorithm: Cubic) +protocol CongestionControlOperation { + associatedtype Result + func callAsFunction(_ controller: inout Controller) -> Result +} + +// A read of whichever congestion controller is active. +@available(Network 0.1.0, *) +protocol CongestionControlQuery { + associatedtype Result + func callAsFunction(_ controller: Controller) -> Result +} + +@available(Network 0.1.0, *) +struct CongestionControl { + enum Algorithm: UInt8 { + case cubic + #if !NETWORK_EMBEDDED + case ledbat + case prague + #endif + + var name: String { + switch self { + case .cubic: return "CUBIC" + #if !NETWORK_EMBEDDED + case .ledbat: return "LEDBAT" + case .prague: return "PRAGUE" + #endif + } + } + } + + private(set) var algorithm: Algorithm + private var cubic: Cubic #if !NETWORK_EMBEDDED - case ledbat(algorithm: Ledbat) - case prague(algorithm: Prague) + private var ledbat: Ledbat + private var prague: Prague #endif - var congestionWindow: UInt64 { - switch self { - case .cubic(let cubic): - return cubic.congestionWindow + // Placeholder initializer to allow non-Optional types. + // CongestionControl is supposed to be initialized later. + init() { + let logPrefixer = LogPrefixer() + self.algorithm = .cubic + self.cubic = Cubic(placeholder: logPrefixer) #if !NETWORK_EMBEDDED - case .ledbat(let ledbat): - return ledbat.congestionWindow - case .prague(let prague): - return prague.congestionWindow + self.ledbat = Ledbat(placeholder: logPrefixer) + self.prague = Prague(placeholder: logPrefixer) + #endif + } + + init( + algorithm: Algorithm, + pacer: inout Pacer, + mss: Int, + qlog: QLog? = nil, + logPrefixer: LogPrefixer + ) { + self.algorithm = algorithm + self.cubic = Cubic(placeholder: logPrefixer) + #if !NETWORK_EMBEDDED + self.ledbat = Ledbat(placeholder: logPrefixer) + self.prague = Prague(placeholder: logPrefixer) + #endif + // Reset and properly initialize the algorithm. + reset(pacer: &pacer, mss: mss, qlog: qlog, logPrefixer: logPrefixer) + } + + // Re-creates the active controller from scratch. + mutating func reset(pacer: inout Pacer, mss: Int, qlog: QLog?, logPrefixer: LogPrefixer) { + switch algorithm { + case .cubic: + cubic = Cubic(pacer: &pacer, mss: mss, qlog: qlog, logPrefixer: logPrefixer) + #if !NETWORK_EMBEDDED + case .ledbat: + ledbat = Ledbat(mss: mss, qlog: qlog, logPrefixer: logPrefixer) + case .prague: + prague = Prague(pacer: &pacer, mss: mss, qlog: qlog, logPrefixer: logPrefixer) #endif } } - var availableCongestionWindow: UInt64 { - switch self { - case .cubic(let cubic): - return cubic.availableCongestionWindow + #if !NETWORK_EMBEDDED + private func handOff(to next: inout Next, mss: Int, qlog: QLog?) { + switch algorithm { + case .cubic: next.inherit(from: cubic, mss: mss, qlog: qlog) + case .ledbat: next.inherit(from: ledbat, mss: mss, qlog: qlog) + case .prague: next.inherit(from: prague, mss: mss, qlog: qlog) + } + } + + // Switches to another algorithm, handing it the outgoing controller's state. + mutating func switchTo( + _ newAlgorithm: Algorithm, + pacer: inout Pacer, + mss: Int, + qlog: QLog?, + logPrefixer: LogPrefixer + ) { + guard newAlgorithm != algorithm else { return } + switch newAlgorithm { + case .cubic: + var next = Cubic(pacer: &pacer, mss: mss, qlog: qlog, logPrefixer: logPrefixer) + handOff(to: &next, mss: mss, qlog: qlog) + cubic = next + case .ledbat: + var next = Ledbat(mss: mss, qlog: qlog, logPrefixer: logPrefixer) + handOff(to: &next, mss: mss, qlog: qlog) + ledbat = next + case .prague: + var next = Prague(pacer: &pacer, mss: mss, qlog: qlog, logPrefixer: logPrefixer) + handOff(to: &next, mss: mss, qlog: qlog) + prague = next + } + algorithm = newAlgorithm + } + #endif + + // Dispatches a mutating operation on the algorithm. + @inline(always) + private mutating func perform( + _ operation: Operation + ) -> Operation.Result { + switch algorithm { + case .cubic: return operation(&cubic) #if !NETWORK_EMBEDDED - case .ledbat(let ledbat): - return ledbat.availableCongestionWindow - case .prague(let prague): - return prague.availableCongestionWindow + case .ledbat: return operation(&ledbat) + case .prague: return operation(&prague) #endif } } - func canSend(packetLength: Int) -> Bool { - switch self { - case .cubic(let cubic): - return cubic.canSend(packetLength: packetLength) + // Dispatches a query operation on the algorithm. + @inline(always) + private func inspect(_ query: Query) -> Query.Result { + switch algorithm { + case .cubic: return query(cubic) #if !NETWORK_EMBEDDED - case .ledbat(let ledbat): - return ledbat.canSend(packetLength: packetLength) - case .prague(let prague): - return prague.canSend(packetLength: packetLength) + case .ledbat: return query(ledbat) + case .prague: return query(prague) #endif } } - mutating func persistentCongestion(mss: Int, qlog: QLog? = nil) { - switch self { - case .cubic(var cubic): - cubic.persistentCongestion(mss: mss, qlog: qlog) - self = .cubic(algorithm: cubic) + var name: String { algorithm.name } + + var congestionWindow: UInt64 { inspect(Read.CongestionWindow()) } + var availableCongestionWindow: UInt64 { inspect(Read.AvailableCongestionWindow()) } + var bytesInFlight: UInt64 { inspect(Read.BytesInFlight()) } + func canSend(packetLength: Int) -> Bool { inspect(Read.CanSend(packetLength: packetLength)) } + + // `inout` arguments can't be stored in a query, so this one dispatches by hand. + func filloutDataTransferSnapshot(dataTransferSnapshot: inout DataTransferSnapshot) { + switch algorithm { + case .cubic: + cubic.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) #if !NETWORK_EMBEDDED - case .ledbat(var ledbat): - ledbat.persistentCongestion(mss: mss, qlog: qlog) - self = .ledbat(algorithm: ledbat) - case .prague(var prague): - prague.persistentCongestion(mss: mss, qlog: qlog) - self = .prague(algorithm: prague) + case .ledbat: + ledbat.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) + case .prague: + prague.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) #endif } } + @inline(always) + mutating func persistentCongestion(mss: Int, qlog: QLog? = nil) { + perform(Op.PersistentCongestion(mss: mss, qlog: qlog)) + } + + // `RTT` is `~Copyable`, so it can't be stored in an operation; this one dispatches by hand. mutating func ackEnd( rtt: borrowing RTT, path: QUICPath?, @@ -95,53 +205,29 @@ enum CongestionControl { now: NetworkClock.Instant, qlog: QLog? = nil ) { - switch self { - case .cubic(algorithm: var cubic): + switch algorithm { + case .cubic: cubic.ackEnd(rtt: rtt, path: path, mss: mss, packetsLost: packetsLost, now: now, qlog: qlog) - self = .cubic(algorithm: cubic) #if !NETWORK_EMBEDDED - case .ledbat(algorithm: var ledbat): + case .ledbat: ledbat.ackEnd(rtt: rtt, path: path, mss: mss, packetsLost: packetsLost, now: now, qlog: qlog) - self = .ledbat(algorithm: ledbat) - case .prague(algorithm: var prague): + case .prague: prague.ackEnd(rtt: rtt, path: path, mss: mss, packetsLost: packetsLost, now: now, qlog: qlog) - self = .prague(algorithm: prague) #endif } } + @inline(always) mutating func packetSent(bytesSent: Int, qlog: QLog? = nil) { - switch self { - case .cubic(algorithm: var cubic): - cubic.packetSent(bytesSent: bytesSent, qlog: qlog) - self = .cubic(algorithm: cubic) - #if !NETWORK_EMBEDDED - case .ledbat(algorithm: var ledbat): - ledbat.packetSent(bytesSent: bytesSent, qlog: qlog) - self = .ledbat(algorithm: ledbat) - case .prague(algorithm: var prague): - prague.packetSent(bytesSent: bytesSent, qlog: qlog) - self = .prague(algorithm: prague) - #endif - } + perform(Op.PacketSent(bytesSent: bytesSent, qlog: qlog)) } + @inline(always) mutating func packetsAcked(bytesAcked: Int, sentTime: NetworkClock.Instant) { - switch self { - case .cubic(algorithm: var cubic): - cubic.packetsAcked(bytesAcked: bytesAcked, sentTime: sentTime) - self = .cubic(algorithm: cubic) - #if !NETWORK_EMBEDDED - case .ledbat(algorithm: var ledbat): - ledbat.packetsAcked(bytesAcked: bytesAcked, sentTime: sentTime) - self = .ledbat(algorithm: ledbat) - case .prague(algorithm: var prague): - prague.packetsAcked(bytesAcked: bytesAcked, sentTime: sentTime) - self = .prague(algorithm: prague) - #endif - } + perform(Op.PacketsAcked(bytesAcked: bytesAcked, sentTime: sentTime)) } + @inline(always) mutating func packetsLost( path: QUICPath?, bytesLost: Int, @@ -150,32 +236,8 @@ enum CongestionControl { smoothedRTT: NetworkDuration, now: NetworkClock.Instant ) -> Bool { - switch self { - case .cubic(algorithm: var cubic): - let reducedCongestionWindow = cubic.packetLost( - path: path, - bytesLost: bytesLost, - largestLostSentTime: largestLostSentTime, - mss: mss, - smoothedRTT: smoothedRTT, - now: now - ) - self = .cubic(algorithm: cubic) - return reducedCongestionWindow - #if !NETWORK_EMBEDDED - case .ledbat(algorithm: var ledbat): - let reducedCongestionWindow = ledbat.packetLost( - path: path, - bytesLost: bytesLost, - largestLostSentTime: largestLostSentTime, - mss: mss, - smoothedRTT: smoothedRTT, - now: now - ) - self = .ledbat(algorithm: ledbat) - return reducedCongestionWindow - case .prague(algorithm: var prague): - let reducedCongestionWindow = prague.packetLost( + perform( + Op.PacketsLost( path: path, bytesLost: bytesLost, largestLostSentTime: largestLostSentTime, @@ -183,12 +245,10 @@ enum CongestionControl { smoothedRTT: smoothedRTT, now: now ) - self = .prague(algorithm: prague) - return reducedCongestionWindow - #endif - } + ) } + @inline(always) mutating func processECN( path: QUICPath?, ceCount: Int, @@ -201,24 +261,8 @@ enum CongestionControl { now: NetworkClock.Instant, qlog: QLog? = nil ) { - switch self { - case .cubic(algorithm: var cubic): - cubic.processECN( - path: path, - ceCount: ceCount, - packetsAcked: packetsAcked, - largestSentPN: largestSentPN, - largestAckedPN: largestAckedPN, - largestAckedSentTime: largestAckedSentTime, - mss: mss, - smoothedRTT: smoothedRTT, - now: now, - qlog: qlog - ) - self = .cubic(algorithm: cubic) - #if !NETWORK_EMBEDDED - case .ledbat(algorithm: var ledbat): - ledbat.processECN( + perform( + Op.ProcessECN( path: path, ceCount: ceCount, packetsAcked: packetsAcked, @@ -230,141 +274,181 @@ enum CongestionControl { now: now, qlog: qlog ) - self = .ledbat(algorithm: ledbat) - case .prague(algorithm: var prague): - prague.processECN( - path: path, - ceCount: ceCount, - packetsAcked: packetsAcked, - largestSentPN: largestSentPN, - largestAckedPN: largestAckedPN, - largestAckedSentTime: largestAckedSentTime, - mss: mss, - smoothedRTT: smoothedRTT, - now: now, - qlog: qlog - ) - self = .prague(algorithm: prague) - #endif - } + ) } + @inline(always) mutating func packetDiscarded(bytesSent: Int, qlog: QLog? = nil) { - switch self { - case .cubic(var cubic): - cubic.packetDiscarded(bytesSent: bytesSent, qlog: qlog) - self = .cubic(algorithm: cubic) - #if !NETWORK_EMBEDDED - case .ledbat(var ledbat): - ledbat.packetDiscarded(bytesSent: bytesSent, qlog: qlog) - self = .ledbat(algorithm: ledbat) - case .prague(var prague): - prague.packetDiscarded(bytesSent: bytesSent, qlog: qlog) - self = .prague(algorithm: prague) - #endif - } + perform(Op.PacketDiscarded(bytesSent: bytesSent, qlog: qlog)) } + @inline(always) mutating func ackBegin() { - switch self { - case .cubic(algorithm: var cubic): - cubic.ackBegin() - self = .cubic(algorithm: cubic) - #if !NETWORK_EMBEDDED - case .ledbat(algorithm: var ledbat): - ledbat.ackBegin() - self = .ledbat(algorithm: ledbat) - case .prague(algorithm: var prague): - prague.ackBegin() - self = .prague(algorithm: prague) - #endif - } + perform(Op.AckBegin()) } - var bytesInFlight: UInt64 { - switch self { - case .cubic(algorithm: let cubic): - return cubic.bytesInFlight - #if !NETWORK_EMBEDDED - case .ledbat(algorithm: let ledbat): - return ledbat.bytesInFlight - case .prague(algorithm: let prague): - return prague.bytesInFlight - #endif - } + @inline(always) + mutating func spuriousRetransmit(qlog: QLog? = nil) { + perform(Op.SpuriousRetransmit(qlog: qlog)) } - var name: String { - switch self { - case .cubic(algorithm: _): - return "CUBIC" - #if !NETWORK_EMBEDDED - case .ledbat(algorithm: _): - return "LEDBAT" - case .prague(algorithm: _): - return "PRAGUE" - #endif - } + @inline(always) + mutating func mssChanged(mss: Int) { + perform(Op.MSSChanged(mss: mss)) } - mutating func spuriousRetransmit(qlog: QLog? = nil) { - switch self { - case .cubic(algorithm: var cubic): - cubic.spuriousRetransmit(qlog: qlog) - self = .cubic(algorithm: cubic) - #if !NETWORK_EMBEDDED - case .ledbat(algorithm: var ledbat): - ledbat.spuriousRetransmit(qlog: qlog) - self = .ledbat(algorithm: ledbat) - case .prague(algorithm: var prague): - prague.spuriousRetransmit(qlog: qlog) - self = .prague(algorithm: prague) - #endif - } + @inline(always) + mutating func idleTimeout(mss: Int) { + perform(Op.IdleTimeout(mss: mss)) } +} - mutating func mssChanged(mss: Int) { - switch self { - case .cubic(algorithm: var cubic): - cubic.mssChanged(mss: mss, qlog: nil) - self = .cubic(algorithm: cubic) - #if !NETWORK_EMBEDDED - case .ledbat(algorithm: var ledbat): - ledbat.mssChanged(mss: mss, qlog: nil) - self = .ledbat(algorithm: ledbat) - case .prague(algorithm: var prague): - prague.mssChanged(mss: mss, qlog: nil) - self = .prague(algorithm: prague) - #endif +// MARK: - Operations and queries + +@available(Network 0.1.0, *) +extension CongestionControl { + fileprivate enum Read { + struct CongestionWindow: CongestionControlQuery { + func callAsFunction(_ controller: Controller) -> UInt64 { + controller.congestionWindow + } } - } - mutating func idleTimeout(mss: Int) { - switch self { - case .cubic(algorithm: var cubic): - cubic.idleTimeout(mss: mss, qlog: nil) - self = .cubic(algorithm: cubic) - #if !NETWORK_EMBEDDED - case .ledbat(algorithm: var ledbat): - ledbat.idleTimeout(mss: mss, qlog: nil) - self = .ledbat(algorithm: ledbat) - case .prague(algorithm: var prague): - prague.idleTimeout(mss: mss, qlog: nil) - self = .prague(algorithm: prague) - #endif + struct AvailableCongestionWindow: CongestionControlQuery { + func callAsFunction(_ controller: Controller) -> UInt64 { + controller.availableCongestionWindow + } + } + + struct BytesInFlight: CongestionControlQuery { + func callAsFunction(_ controller: Controller) -> UInt64 { + controller.bytesInFlight + } + } + + struct CanSend: CongestionControlQuery { + let packetLength: Int + + func callAsFunction(_ controller: Controller) -> Bool { + controller.canSend(packetLength: packetLength) + } } } - func filloutDataTransferSnapshot(dataTransferSnapshot: inout DataTransferSnapshot) { - switch self { - case .cubic(algorithm: let cubic): - cubic.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) - #if !NETWORK_EMBEDDED - case .ledbat(algorithm: let ledbat): - ledbat.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) - case .prague(algorithm: let prague): - prague.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) - #endif + fileprivate enum Op { + struct PersistentCongestion: CongestionControlOperation { + let mss: Int + let qlog: QLog? + + func callAsFunction(_ controller: inout Controller) { + controller.persistentCongestion(mss: mss, qlog: qlog) + } + } + + struct PacketSent: CongestionControlOperation { + let bytesSent: Int + let qlog: QLog? + + func callAsFunction(_ controller: inout Controller) { + controller.packetSent(bytesSent: bytesSent, qlog: qlog) + } + } + + struct PacketsAcked: CongestionControlOperation { + let bytesAcked: Int + let sentTime: NetworkClock.Instant + + func callAsFunction(_ controller: inout Controller) { + controller.packetsAcked(bytesAcked: bytesAcked, sentTime: sentTime) + } + } + + struct PacketsLost: CongestionControlOperation { + let path: QUICPath? + let bytesLost: Int + let largestLostSentTime: NetworkClock.Instant + let mss: Int + let smoothedRTT: NetworkDuration + let now: NetworkClock.Instant + + func callAsFunction(_ controller: inout Controller) -> Bool { + controller.packetLost( + path: path, + bytesLost: bytesLost, + largestLostSentTime: largestLostSentTime, + mss: mss, + smoothedRTT: smoothedRTT, + now: now, + qlog: nil + ) + } + } + + struct ProcessECN: CongestionControlOperation { + let path: QUICPath? + let ceCount: Int + let packetsAcked: Int + let largestSentPN: Int64 + let largestAckedPN: Int64 + let largestAckedSentTime: NetworkClock.Instant + let mss: Int + let smoothedRTT: NetworkDuration + let now: NetworkClock.Instant + let qlog: QLog? + + func callAsFunction(_ controller: inout Controller) { + controller.processECN( + path: path, + ceCount: ceCount, + packetsAcked: packetsAcked, + largestSentPN: largestSentPN, + largestAckedPN: largestAckedPN, + largestAckedSentTime: largestAckedSentTime, + mss: mss, + smoothedRTT: smoothedRTT, + now: now, + qlog: qlog + ) + } + } + + struct PacketDiscarded: CongestionControlOperation { + let bytesSent: Int + let qlog: QLog? + + func callAsFunction(_ controller: inout Controller) { + controller.packetDiscarded(bytesSent: bytesSent, qlog: qlog) + } + } + + struct AckBegin: CongestionControlOperation { + func callAsFunction(_ controller: inout Controller) { + controller.ackBegin() + } + } + + struct SpuriousRetransmit: CongestionControlOperation { + let qlog: QLog? + + func callAsFunction(_ controller: inout Controller) { + controller.spuriousRetransmit(qlog: qlog) + } + } + + struct MSSChanged: CongestionControlOperation { + let mss: Int + + func callAsFunction(_ controller: inout Controller) { + controller.mssChanged(mss: mss, qlog: nil) + } + } + + struct IdleTimeout: CongestionControlOperation { + let mss: Int + + func callAsFunction(_ controller: inout Controller) { + controller.idleTimeout(mss: mss, qlog: nil) + } } } } @@ -388,7 +472,7 @@ protocol CongestionControlProtocol: PrefixedLoggable { var pipeAckIndex: Int { get set } mutating func inherit( - from: CongestionControl, + from other: some CongestionControlProtocol, mss: Int, qlog: QLog? ) @@ -403,7 +487,7 @@ protocol CongestionControlProtocol: PrefixedLoggable { ) mutating func spuriousRetransmit(qlog: QLog?) mutating func idleTimeout(mss: Int, qlog: QLog?) - /// Opens a recovery period at `now`, the time the loss was detected. + // Opens a recovery period at `now`, the time the loss was detected. mutating func enterRecovery(mss: Int, now: NetworkClock.Instant, qlog: QLog?) mutating func processECN( path: QUICPath?, @@ -543,7 +627,7 @@ extension CongestionControlProtocol { sentTime <= recoveryStartTime } - /// `sentTime` is when the packet went out, `now` when its loss was detected. + // `sentTime` is when the packet went out, `now` when its loss was detected. @discardableResult mutating func congestionEvent( sentTime: NetworkClock.Instant, diff --git a/Sources/SwiftNetwork/QUIC/Cubic.swift b/Sources/SwiftNetwork/QUIC/Cubic.swift index c57d5093..6599eac2 100644 --- a/Sources/SwiftNetwork/QUIC/Cubic.swift +++ b/Sources/SwiftNetwork/QUIC/Cubic.swift @@ -110,6 +110,10 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { logState(qlog: qlog, state: .slowStart, trigger: nil) } + init(placeholder logPrefixer: LogPrefixer) { + self.log = logPrefixer + } + private mutating func setK(mss: Int) { // K is the time period(s) that WCubic(t) function takes to increase // the current window size to WMax if there are no further @@ -457,24 +461,13 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { logUpdate(qlog: qlog) } - mutating func inherit(from: CongestionControl, mss: Int, qlog: QLog?) { + mutating func inherit(from other: some CongestionControlProtocol, mss: Int, qlog: QLog?) { // For Cubic, the old state will be stale and // as it will ramp up quickly in slow start, it // is best to start fresh. For congestion window // we can use the higher of last congestionWindow and initialCongestionWindow - switch from { - case .cubic(let cubic): - self.bytesInFlight = cubic.bytesInFlight - self.congestionWindow = max(cubic.congestionWindow, Cubic.initialCongestionWindow(mss)) - #if !NETWORK_EMBEDDED - case .ledbat(let ledbat): - self.bytesInFlight = ledbat.bytesInFlight - self.congestionWindow = max(ledbat.congestionWindow, Cubic.initialCongestionWindow(mss)) - case .prague(let prague): - self.bytesInFlight = prague.bytesInFlight - self.congestionWindow = max(prague.congestionWindow, Cubic.initialCongestionWindow(mss)) - #endif - } + self.bytesInFlight = other.bytesInFlight + self.congestionWindow = max(other.congestionWindow, Cubic.initialCongestionWindow(mss)) slowStartThreshold = UInt64.max resetInternal() logUpdate(qlog: qlog) diff --git a/Sources/SwiftNetwork/QUIC/Ledbat.swift b/Sources/SwiftNetwork/QUIC/Ledbat.swift index fa64b39e..2dbfbd75 100644 --- a/Sources/SwiftNetwork/QUIC/Ledbat.swift +++ b/Sources/SwiftNetwork/QUIC/Ledbat.swift @@ -66,6 +66,10 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { logUpdate(qlog: qlog) } + init(placeholder logPrefixer: LogPrefixer) { + self.log = logPrefixer + } + // GAIN is proportional to the ratio of base_delay // and TARGET delay, i.e., GAIN is smaller for bottlenecks // with small queues in order to ensure that LEDBAT yields @@ -336,24 +340,13 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { logUpdate(qlog: qlog) } - mutating func inherit(from: CongestionControl, mss: Int, qlog: QLog?) { - // LEDBAT has minimal state. We can continue using its own - // ssthresh from old state as it is somewhat stable. - // For congestion window, we can take the lower of - // its own cwnd and previous controller's cwnd. - switch from { - case .cubic(let cubic): - self.bytesInFlight = cubic.bytesInFlight - self.congestionWindow = min(cubic.congestionWindow, congestionWindow) - #if !NETWORK_EMBEDDED - case .ledbat(let ledbat): - self.bytesInFlight = ledbat.bytesInFlight - self.congestionWindow = min(ledbat.congestionWindow, congestionWindow) - case .prague(let prague): - self.bytesInFlight = prague.bytesInFlight - self.congestionWindow = min(prague.congestionWindow, congestionWindow) - #endif - } + mutating func inherit(from other: some CongestionControlProtocol, mss: Int, qlog: QLog?) { + // LEDBAT has minimal state, so `self` is expected to be freshly + // initialized and keeps its initial ssthresh. For the congestion + // window, take the lower of the initial cwnd and the previous + // controller's cwnd. + self.bytesInFlight = other.bytesInFlight + self.congestionWindow = min(other.congestionWindow, congestionWindow) logUpdate(qlog: qlog) resetInternal() } diff --git a/Sources/SwiftNetwork/QUIC/Prague.swift b/Sources/SwiftNetwork/QUIC/Prague.swift index d8bb2621..153bf293 100644 --- a/Sources/SwiftNetwork/QUIC/Prague.swift +++ b/Sources/SwiftNetwork/QUIC/Prague.swift @@ -140,6 +140,10 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { logState(qlog: qlog, state: .slowStart, trigger: nil) } + init(placeholder logPrefixer: LogPrefixer) { + self.log = logPrefixer + } + /// Computes the cubic K factor for the current congestion window. /// /// `K` is the time period(s) that the `W_cubic(t)` function takes to increase @@ -680,30 +684,13 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { logUpdate(qlog: qlog) } - mutating func inherit(from: CongestionControl, mss: Int, qlog: QLog?) { + mutating func inherit(from other: some CongestionControlProtocol, mss: Int, qlog: QLog?) { // For Prague, the old state will be stale and // as it will ramp up quickly in slow start, it // is best to start fresh. For congestion window // we can use the higher of last cwnd and INITIAL_CWND - switch from { - case .cubic(let cubic): - self.bytesInFlight = cubic.bytesInFlight - self.congestionWindow = max(cubic.congestionWindow, Prague.initialCongestionWindow(mss)) - #if !NETWORK_EMBEDDED - case .ledbat(let ledbat): - self.bytesInFlight = ledbat.bytesInFlight - self.congestionWindow = max( - ledbat.congestionWindow, - Prague.initialCongestionWindow(mss) - ) - case .prague(let prague): - self.bytesInFlight = prague.bytesInFlight - self.congestionWindow = max( - prague.congestionWindow, - Prague.initialCongestionWindow(mss) - ) - #endif - } + self.bytesInFlight = other.bytesInFlight + self.congestionWindow = max(other.congestionWindow, Prague.initialCongestionWindow(mss)) slowStartThreshold = UInt64.max scaledAlpha = Prague.maxAlpha << Prague.gShift resetInternal() diff --git a/Sources/SwiftNetwork/QUIC/QUICPath.swift b/Sources/SwiftNetwork/QUIC/QUICPath.swift index c2c71f70..d9a99993 100644 --- a/Sources/SwiftNetwork/QUIC/QUICPath.swift +++ b/Sources/SwiftNetwork/QUIC/QUICPath.swift @@ -156,7 +156,7 @@ public final class QUICPath: MultiplexingDatagramPath< var bdp = BandwidthDelayProduct() - private var congestionControl: CongestionControl? + private var congestionControl = CongestionControl() var pacer: Pacer @@ -371,13 +371,12 @@ public final class QUICPath: MultiplexingDatagramPath< let pacerEnabled = (pacePackets || QUICPreferences.shared.pacePackets) self.pacer = Pacer(enabled: pacerEnabled) - self.congestionControl = .cubic( - algorithm: Cubic( - pacer: &self.pacer, - mss: self.initialMSS, - qlog: parentProtocol.qLog, - logPrefixer: self.log - ) + self.congestionControl = CongestionControl( + algorithm: .cubic, + pacer: &self.pacer, + mss: self.initialMSS, + qlog: parentProtocol.qLog, + logPrefixer: self.log ) self.spinValue = parentProtocol.initialSpinValue @@ -422,42 +421,16 @@ public final class QUICPath: MultiplexingDatagramPath< } func resetCongestionControl() { - switch self.congestionControl { - case .cubic: - self.congestionControl = .cubic( - algorithm: Cubic( - pacer: &self.pacer, - mss: self.initialMSS, - qlog: parentProtocol.qLog, - logPrefixer: self.log - ) - ) - #if !NETWORK_EMBEDDED - case .ledbat: - self.congestionControl = .ledbat( - algorithm: Ledbat( - mss: self.initialMSS, - qlog: parentProtocol.qLog, - logPrefixer: self.log - ) - ) - case .prague: - self.congestionControl = .prague( - algorithm: Prague( - pacer: &self.pacer, - mss: self.initialMSS, - qlog: parentProtocol.qLog, - logPrefixer: self.log - ) - ) - #endif - case .none: - break - } + congestionControl.reset( + pacer: &pacer, + mss: initialMSS, + qlog: parentProtocol.qLog, + logPrefixer: log + ) } func idleTimeoutCongestionControl() { - self.congestionControl?.idleTimeout(mss: mss) + self.congestionControl.idleTimeout(mss: mss) } func setupL4SState(l4sEnabled: Bool?) { @@ -478,51 +451,22 @@ public final class QUICPath: MultiplexingDatagramPath< func markAsBackground(_ background: Bool) { #if !NETWORK_EMBEDDED // Use LEDBAT for background cases - switch self.congestionControl { - case .cubic: - if !background { return } // Nothing to do, already not background - var ledbat = Ledbat( - mss: self.initialMSS, - qlog: parentProtocol.qLog, - logPrefixer: self.log - ) - ledbat.inherit( - from: self.congestionControl!, - mss: self.initialMSS, - qlog: parentProtocol.qLog - ) - self.congestionControl = .ledbat(algorithm: ledbat) - case .ledbat: - if background { return } // Nothing to do, already background - // Inherit from the current LEDBAT so bytes in flight (and the window) carry over. - var cubic = Cubic( - pacer: &self.pacer, - mss: self.initialMSS, - qlog: parentProtocol.qLog, - logPrefixer: self.log - ) - cubic.inherit( - from: self.congestionControl!, - mss: self.initialMSS, - qlog: parentProtocol.qLog - ) - self.congestionControl = .cubic(algorithm: cubic) - case .prague: - if !background { return } // Nothing to do, already not background - var ledbat = Ledbat( - mss: self.initialMSS, - qlog: parentProtocol.qLog, - logPrefixer: self.log - ) - ledbat.inherit( - from: self.congestionControl!, - mss: self.initialMSS, - qlog: parentProtocol.qLog - ) - self.congestionControl = .ledbat(algorithm: ledbat) - case .none: - break + let target: CongestionControl.Algorithm + if background { + target = .ledbat + } else if congestionControl.algorithm == .ledbat { + target = .cubic + } else { + return // Nothing to do, already not background } + // The new controller inherits bytes in flight (and, by its own rule, the window). + congestionControl.switchTo( + target, + pacer: &pacer, + mss: initialMSS, + qlog: parentProtocol.qLog, + logPrefixer: log + ) #endif } @@ -702,22 +646,22 @@ public final class QUICPath: MultiplexingDatagramPath< extension QUICPath { @inline(always) var congestionControlWindow: UInt64 { - congestionControl?.congestionWindow ?? 0 + congestionControl.congestionWindow } @inline(always) var congestionControlAvailableCongestionWindow: UInt64 { - congestionControl?.availableCongestionWindow ?? 0 + congestionControl.availableCongestionWindow } @inline(always) func congestionControlCanSend(packetLength: Int) -> Bool { - congestionControl?.canSend(packetLength: packetLength) ?? false + congestionControl.canSend(packetLength: packetLength) } @inline(always) func congestionControlPersistentCongestion(mss: Int, qlog: QLog? = nil) { - congestionControl?.persistentCongestion(mss: mss, qlog: qlog) + congestionControl.persistentCongestion(mss: mss, qlog: qlog) } @inline(always) @@ -728,7 +672,7 @@ extension QUICPath { packetsLost: Bool, qlog: QLog? = nil ) { - congestionControl?.ackEnd( + congestionControl.ackEnd( rtt: rtt, path: self, mss: mss, @@ -740,12 +684,12 @@ extension QUICPath { @inline(always) func congestionControlPacketsSent(bytesSent: Int, qlog: QLog? = nil) { - congestionControl?.packetSent(bytesSent: bytesSent, qlog: qlog) + congestionControl.packetSent(bytesSent: bytesSent, qlog: qlog) } @inline(always) func congestionControlPacketsAcked(bytesAcked: Int, sentTime: NetworkClock.Instant) { - congestionControl?.packetsAcked(bytesAcked: bytesAcked, sentTime: sentTime) + congestionControl.packetsAcked(bytesAcked: bytesAcked, sentTime: sentTime) } @inline(always) @@ -757,49 +701,49 @@ extension QUICPath { ) -> Bool { // Loss accounting doesn't repace this path, so there is no path to hand down. let unpacedPath: QUICPath? = nil - return congestionControl?.packetsLost( + return congestionControl.packetsLost( path: unpacedPath, bytesLost: bytesLost, largestLostSentTime: largestLostSentTime, mss: mss, smoothedRTT: smoothedRTT, now: parentProtocol.now - ) ?? false + ) } @inline(always) func congestionControlPacketDiscarded(bytesSent: Int, qlog: QLog? = nil) { - congestionControl?.packetDiscarded(bytesSent: bytesSent, qlog: qlog) + congestionControl.packetDiscarded(bytesSent: bytesSent, qlog: qlog) } @inline(always) func congestionControlAckBegin() { - congestionControl?.ackBegin() + congestionControl.ackBegin() } @inline(always) var congestionControlBytesInFlight: UInt64 { - congestionControl?.bytesInFlight ?? 0 + congestionControl.bytesInFlight } @inline(always) var congestionControlName: String { - congestionControl?.name ?? "none" + congestionControl.name } @inline(always) func congestionControlSpuriousRetransmit(qlog: QLog? = nil) { - congestionControl?.spuriousRetransmit(qlog: qlog) + congestionControl.spuriousRetransmit(qlog: qlog) } @inline(always) func congestionControlMSSChanged(mss: Int) { - congestionControl?.mssChanged(mss: mss) + congestionControl.mssChanged(mss: mss) } @inline(always) func congestionControlIdleTimeout(mss: Int) { - congestionControl?.idleTimeout(mss: mss) + congestionControl.idleTimeout(mss: mss) } @inline(always) @@ -815,7 +759,7 @@ extension QUICPath { ) { // ECN accounting doesn't repace this path, so there is no path to hand down. let unpacedPath: QUICPath? = nil - congestionControl?.processECN( + congestionControl.processECN( path: unpacedPath, ceCount: ceCount, packetsAcked: packetsAcked, @@ -831,7 +775,7 @@ extension QUICPath { @inline(always) func congestionControlFilloutDataTransferSnapshot(snapshot: inout DataTransferSnapshot) { - congestionControl?.filloutDataTransferSnapshot(dataTransferSnapshot: &snapshot) + congestionControl.filloutDataTransferSnapshot(dataTransferSnapshot: &snapshot) } } diff --git a/Tests/QUICTests/CongestionControlTests.swift b/Tests/QUICTests/CongestionControlTests.swift new file mode 100644 index 00000000..87d262a7 --- /dev/null +++ b/Tests/QUICTests/CongestionControlTests.swift @@ -0,0 +1,301 @@ +//===----------------------------------------------------------------------===// +// +// This source file is part of the Swift open source project +// +// Copyright (c) 2026 Apple Inc. and the Swift project authors +// Licensed under Apache License v2.0 +// +// See LICENSE.txt for license information +// See CONTRIBUTORS.txt for the list of Swift project authors +// +// SPDX-License-Identifier: Apache-2.0 +// +//===----------------------------------------------------------------------===// + +#if !NETWORK_NO_SWIFT_QUIC + +import XCTest + +#if canImport(SwiftNetwork) +@_spi(Essentials) @_spi(ProtocolProvider) @testable import SwiftNetwork +#elseif canImport(Network) +@_spi(Essentials) @_spi(ProtocolProvider) @testable import Network +#endif + +#if canImport(SwiftNetworkTestHarness) +@_spi(TestHarness) @_spi(Essentials) @_spi(ProtocolProvider) import SwiftNetworkTestHarness +#endif + +/// Tests the `CongestionControl` container: dispatch to the active controller, switching +/// between controllers, and resetting. The controllers themselves are covered by +/// `CubicTests`, `LedbatTests` and `PragueTests`. +@available(Network 0.1.0, *) +final class CongestionControlTests: XCTestCase { + let mss = Constants.initialMSS + let logPrefixer = LogPrefixer("[CongestionControlTests]") + var pacer = Pacer(enabled: false) + var rtt: RTT! + // These tests drive the controllers directly, with no path to pace. + let noPath: QUICPath? = nil + + let cubicInitialWindow = UInt64(12000) + let ledbatInitialWindow = UInt64(2400) + let pragueInitialWindow = UInt64(12000) + + override func setUp() { + rtt = RTT(logPrefixer: logPrefixer) + rtt.smoothedRTT = .microseconds(0) + rtt.adjustedRTT = .milliseconds(100) + rtt.baseRTT = .milliseconds(100) + } + + private func makeCubic() -> CongestionControl { + CongestionControl(algorithm: .cubic, pacer: &pacer, mss: mss, logPrefixer: logPrefixer) + } + + private func switchTo(_ algorithm: CongestionControl.Algorithm, _ cc: inout CongestionControl) { + cc.switchTo(algorithm, pacer: &pacer, mss: mss, qlog: nil, logPrefixer: logPrefixer) + } + + private func slowStartThreshold(_ cc: CongestionControl) -> UInt64 { + var snapshot = DataTransferSnapshot() + cc.filloutDataTransferSnapshot(dataTransferSnapshot: &snapshot) + return snapshot.transportSlowStartThreshold + } + + /// Grows the window by sending and acknowledging `count` packets in slow start. + private func growWindow(_ cc: inout CongestionControl, packets count: Int) { + let time = NetworkClock.Instant.testBase + for _ in 0..