diff --git a/Sources/SwiftNetwork/QUIC/CongestionControl.swift b/Sources/SwiftNetwork/QUIC/CongestionControl.swift index ebe1929..e6c4e66 100644 --- a/Sources/SwiftNetwork/QUIC/CongestionControl.swift +++ b/Sources/SwiftNetwork/QUIC/CongestionControl.swift @@ -24,65 +24,346 @@ internal import Logging internal import os #endif +/// The state share by all algorithms Cubic, Ledbats, and Prague +/// +/// This is the single, authoritative copy of this state that is used throughout each algorithm @available(Network 0.1.0, *) -enum CongestionControl { - case cubic(algorithm: Cubic) - #if !NETWORK_EMBEDDED - case ledbat(algorithm: Ledbat) - case prague(algorithm: Prague) - #endif +struct CongestionControlState { + var congestionWindow = UInt64(0) + var bytesInFlight = UInt64(0) + var packetsAcked = UInt64(0) + var packetsMarked = UInt64(0) + var ecnCECounter = 0 + var largestSentPN = Int64(0) + var slowStartThreshold = UInt64.max + var prevSlowStartThreshold = UInt64.max + var recoveryStartTime = NetworkClock.Instant.zero + var bytesAcked = UInt64(0) + var pipeAckSamples = [UInt64(0)] + var pipeAckValue = UInt64(0) + var pipeAckSampleEnd = NetworkClock.Instant.zero + var pipeAckAcked = UInt64(0) + var pipeAckIndex = 0 - var congestionWindow: UInt64 { - switch self { - case .cubic(let cubic): - return cubic.congestionWindow + var congestionWindowValidationSamples: Int { + 3 + } + + var availableCongestionWindow: UInt64 { + if congestionWindow > bytesInFlight { + return congestionWindow - bytesInFlight + } else { + return 0 + } + } + + func congestionWindowValidated(log: LogPrefixer) -> Bool { + if pipeAckValue == 0 { + return true + } + // In slow-start, congestionWindow increases aggressively and pipeack + // might be lagging behind. Thus, give it 4x the space. + if congestionWindow < slowStartThreshold { + let congestionWindowLimit = congestionWindow >> 2 + if pipeAckValue < congestionWindowLimit { + log.datapath( + "Congestion window not validated in slow-start, pipeack: \(pipeAckValue), congestionWindow \(congestionWindow)" + ) + return false + } + } else { + let congestionWindowLimit = congestionWindow >> 1 + if pipeAckValue < congestionWindowLimit { + log.datapath( + "Congestion window not validated in congestion-avoidance, pipeack: \(pipeAckValue), congestionWindow: \(congestionWindow) " + ) + return false + } + } + + return true + } + + func lossFlightSize(log: LogPrefixer) -> UInt64 { + if !congestionWindowValidated(log: log) { + return max(pipeAckValue, bytesInFlight) + } else { + return congestionWindow + } + } + + mutating func incrementBytesInFlight(_ bytesSent: Int, log: LogPrefixer) { + bytesInFlight += UInt64(bytesSent) + log.datapath("Bytes in flight updated to \(bytesInFlight)") + + QUICSignpost.bytesInFlight(bytesInFlight: Int(bytesInFlight)) + } + + mutating func decrementBytesInFlight(_ bytes: UInt64, log: LogPrefixer) { + let result = bytesInFlight.subtractingReportingOverflow(bytes) + if result.overflow { + log.fault("Undeflow, \(bytes) decremented from \(bytesInFlight)") + bytesInFlight = 0 + } else { + bytesInFlight = result.partialValue + } + log.datapath("Bytes in flight updated to \(bytesInFlight)") + QUICSignpost.bytesInFlight(bytesInFlight: Int(bytesInFlight)) + } + + func packetInRecovery(sentTime: NetworkClock.Instant) -> Bool { + sentTime <= recoveryStartTime + } + + mutating func ackBegin() { + bytesAcked = 0 + } + + mutating func packetsAcked(bytesAcked: Int, sentTime: NetworkClock.Instant, log: LogPrefixer) { + let bytesAcked = UInt64(bytesAcked) + decrementBytesInFlight(bytesAcked, log: log) + if packetInRecovery(sentTime: sentTime) { + // Dont update the congestion window + log.datapath("Packet was sent before recovery, ignore") + return + } + // Congestion window is updated later in ackEnd + self.bytesAcked += bytesAcked + } + + mutating func packetSent(bytesSent: Int, log: LogPrefixer, qlog: QLog? = nil) { + incrementBytesInFlight(bytesSent, log: log) + logUpdate(log: log, qlog: qlog) + } + + mutating func packetDiscarded(bytesSent: Int, log: LogPrefixer, qlog: QLog? = nil) { + decrementBytesInFlight(UInt64(bytesSent), log: log) + logUpdate(log: log, qlog: qlog) + } + + mutating func mssChanged(mss: Int, log: LogPrefixer, qlog: QLog? = nil) { + congestionWindow = max(congestionWindow, UInt64(mss)) + logUpdate(log: log, qlog: qlog) + } + + // Compute if 1RTT or 1 round has elapsed by measuring if the + // packet sent after this instant has been acknowledged. + // largest_sent_pn is set at the start of a round + func rttElapsed(largestSentPN: Int64, largestAckedPN: Int64) -> Bool { + // A packet with pn higher than largest sent pn at + // the start of the round has been acknowledged + (largestSentPN == 0) || (largestAckedPN > largestSentPN) + } + + mutating func initPipeAckSamples() { + pipeAckSamples = Array(repeating: 0, count: congestionWindowValidationSamples) + pipeAckIndex = 0 + pipeAckValue = 0 + } + + mutating func setPipeAckSample(sample: UInt64) { + pipeAckSamples[pipeAckIndex] = sample + pipeAckIndex &+= 1 + pipeAckIndex = pipeAckIndex % congestionWindowValidationSamples + } + + mutating func pipeAckNewRound(target: NetworkClock.Instant) { + pipeAckSampleEnd = target == .zero ? .init(microseconds: 1) : target + pipeAckAcked = 0 + } + + mutating func updatePipeAckSamples() { + setPipeAckSample(sample: pipeAckAcked) + pipeAckValue = pipeAckAcked + for index in 0.. pipeAckValue { + pipeAckValue = pipeAckSamples[index] + } + } + } + + mutating func revalidateCongestionWindow( + smoothedRTT: NetworkDuration, + now: NetworkClock.Instant, + log: LogPrefixer + ) -> Bool { + if pipeAckSampleEnd == .zero { + pipeAckNewRound(target: now.advanced(by: smoothedRTT)) + } + pipeAckAcked += bytesAcked + // A full period passed? Update our pipeack samples + if now > pipeAckSampleEnd { + let period = pipeAckSampleEnd.duration(to: now) + if period > smoothedRTT { + // More than 1 RTT of inactivity, we need to set samples to 0 + setPipeAckSample(sample: 0) + if period > smoothedRTT * 2 { + // Reset the next sample as well + setPipeAckSample(sample: 0) + } + } + updatePipeAckSamples() + pipeAckNewRound(target: now + smoothedRTT) + } + return congestionWindowValidated(log: log) + } + + func canSend(packetLength: Int, log: LogPrefixer) -> Bool { + if availableCongestionWindow >= packetLength { + log.datapath( + "Can send packet because bytesInFlight \(bytesInFlight) + packetLength \(packetLength) <= congestionWindow \(congestionWindow)" + ) + return true + } else { + log.datapath( + "Congestion limited because bytesInFlight \(bytesInFlight) + packetLength \(packetLength) > congestionWindow \(congestionWindow)" + ) + QUICSignpost.congestionWindowLimited( + bytesInFlight: Int(bytesInFlight), + congestionWindow: Int(congestionWindow) + ) + return false + } + } + + func logUpdate(log: LogPrefixer, qlog: QLog?) { + if congestionWindow != UInt64.max { + log.datapath("Congestion window set to \(congestionWindow) bytes, bytes in flight \(bytesInFlight)") + QUICSignpost.congestionWindow(congestionWindow: Int(congestionWindow)) + } + #if QlogOutput + if let qlog { + qlog.congestionControlUpdated( + congestionWindow: congestionWindow, + bytesInFlight: bytesInFlight, + slowStartThresh: slowStartThreshold + ) + } + #endif + } +} + +/// Owns the shared congestion control state plus the active algorithm +@available(Network 0.1.0, *) +struct CongestionControl: ~Copyable { + enum Algorithm { + case cubic(algorithm: Cubic) #if !NETWORK_EMBEDDED - case .ledbat(let ledbat): - return ledbat.congestionWindow - case .prague(let prague): - return prague.congestionWindow + case ledbat(algorithm: Ledbat) + case prague(algorithm: Prague) #endif - } } - var availableCongestionWindow: UInt64 { - switch self { + var state: CongestionControlState + var log: LogPrefixer + var algorithm: Algorithm + + init(state: CongestionControlState = CongestionControlState(), algorithm: Algorithm) { + self.state = state + self.algorithm = algorithm + switch algorithm { case .cubic(let cubic): - return cubic.availableCongestionWindow + self.log = cubic.log #if !NETWORK_EMBEDDED case .ledbat(let ledbat): - return ledbat.availableCongestionWindow + self.log = ledbat.log case .prague(let prague): - return prague.availableCongestionWindow + self.log = prague.log #endif } } - func canSend(packetLength: Int) -> Bool { - switch self { - case .cubic(let cubic): - return cubic.canSend(packetLength: packetLength) + static func createCubic( + pacer: inout Pacer, + mss: Int, + qlog: QLog? = nil, + logPrefixer: LogPrefixer + ) -> CongestionControl { + var state = CongestionControlState() + let cubic = Cubic(state: &state, pacer: &pacer, mss: mss, qlog: qlog, logPrefixer: logPrefixer) + return CongestionControl(state: state, algorithm: .cubic(algorithm: cubic)) + } + + #if !NETWORK_EMBEDDED + static func createLedbat(mss: Int, qlog: QLog? = nil, logPrefixer: LogPrefixer) -> CongestionControl { + var state = CongestionControlState() + let ledbat = Ledbat(state: &state, mss: mss, qlog: qlog, logPrefixer: logPrefixer) + return CongestionControl(state: state, algorithm: .ledbat(algorithm: ledbat)) + } + + static func createPrague( + pacer: inout Pacer, + mss: Int, + qlog: QLog? = nil, + logPrefixer: LogPrefixer + ) -> CongestionControl { + var state = CongestionControlState() + let prague = Prague(state: &state, pacer: &pacer, mss: mss, qlog: qlog, logPrefixer: logPrefixer) + return CongestionControl(state: state, algorithm: .prague(algorithm: prague)) + } + #endif + + var congestionWindow: UInt64 { + state.congestionWindow + } + + var availableCongestionWindow: UInt64 { + state.availableCongestionWindow + } + + var bytesInFlight: UInt64 { + state.bytesInFlight + } + + var name: String { + switch algorithm { + case .cubic: + return "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 "LEDBAT" + case .prague: + return "PRAGUE" #endif } } + func canSend(packetLength: Int) -> Bool { + state.canSend(packetLength: packetLength, log: log) + } + + mutating func packetSent(bytesSent: Int, qlog: QLog? = nil) { + state.packetSent(bytesSent: bytesSent, log: log, qlog: qlog) + } + + mutating func packetsAcked(bytesAcked: Int, sentTime: NetworkClock.Instant) { + state.packetsAcked(bytesAcked: bytesAcked, sentTime: sentTime, log: log) + } + + mutating func packetDiscarded(bytesSent: Int, qlog: QLog? = nil) { + state.packetDiscarded(bytesSent: bytesSent, log: log, qlog: qlog) + } + + mutating func ackBegin() { + state.ackBegin() + } + + mutating func mssChanged(mss: Int) { + state.mssChanged(mss: mss, log: log, qlog: nil) + } + mutating func persistentCongestion(mss: Int, qlog: QLog? = nil) { - switch self { + switch algorithm { case .cubic(var cubic): - cubic.persistentCongestion(mss: mss, qlog: qlog) - self = .cubic(algorithm: cubic) + cubic.persistentCongestion(state: &state, mss: mss, qlog: qlog) + algorithm = .cubic(algorithm: cubic) #if !NETWORK_EMBEDDED case .ledbat(var ledbat): - ledbat.persistentCongestion(mss: mss, qlog: qlog) - self = .ledbat(algorithm: ledbat) + ledbat.persistentCongestion(state: &state, mss: mss, qlog: qlog) + algorithm = .ledbat(algorithm: ledbat) case .prague(var prague): - prague.persistentCongestion(mss: mss, qlog: qlog) - self = .prague(algorithm: prague) + prague.persistentCongestion(state: &state, mss: mss, qlog: qlog) + algorithm = .prague(algorithm: prague) #endif } } @@ -95,53 +376,22 @@ enum CongestionControl { now: NetworkClock.Instant, qlog: QLog? = nil ) { - switch self { - case .cubic(algorithm: var 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): - ledbat.ackEnd(rtt: rtt, path: path, mss: mss, packetsLost: packetsLost, now: now, qlog: qlog) - self = .ledbat(algorithm: ledbat) - case .prague(algorithm: var prague): - prague.ackEnd(rtt: rtt, path: path, mss: mss, packetsLost: packetsLost, now: now, qlog: qlog) - self = .prague(algorithm: prague) - #endif - } - } - - 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 - } - } - - 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) + switch algorithm { + case .cubic(var cubic): + cubic.ackEnd(state: &state, rtt: rtt, path: path, mss: mss, packetsLost: packetsLost, now: now, qlog: qlog) + algorithm = .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) + case .ledbat(var ledbat): + ledbat.ackEnd(state: &state, rtt: rtt, path: path, mss: mss, packetsLost: packetsLost, now: now, qlog: qlog) + algorithm = .ledbat(algorithm: ledbat) + case .prague(var prague): + prague.ackEnd(state: &state, rtt: rtt, path: path, mss: mss, packetsLost: packetsLost, now: now, qlog: qlog) + algorithm = .prague(algorithm: prague) #endif } } + @discardableResult mutating func packetsLost( path: QUICPath?, bytesLost: Int, @@ -150,9 +400,10 @@ enum CongestionControl { smoothedRTT: NetworkDuration, now: NetworkClock.Instant ) -> Bool { - switch self { - case .cubic(algorithm: var cubic): + switch algorithm { + case .cubic(var cubic): let reducedCongestionWindow = cubic.packetLost( + state: &state, path: path, bytesLost: bytesLost, largestLostSentTime: largestLostSentTime, @@ -160,11 +411,12 @@ enum CongestionControl { smoothedRTT: smoothedRTT, now: now ) - self = .cubic(algorithm: cubic) + algorithm = .cubic(algorithm: cubic) return reducedCongestionWindow #if !NETWORK_EMBEDDED - case .ledbat(algorithm: var ledbat): + case .ledbat(var ledbat): let reducedCongestionWindow = ledbat.packetLost( + state: &state, path: path, bytesLost: bytesLost, largestLostSentTime: largestLostSentTime, @@ -172,10 +424,11 @@ enum CongestionControl { smoothedRTT: smoothedRTT, now: now ) - self = .ledbat(algorithm: ledbat) + algorithm = .ledbat(algorithm: ledbat) return reducedCongestionWindow - case .prague(algorithm: var prague): + case .prague(var prague): let reducedCongestionWindow = prague.packetLost( + state: &state, path: path, bytesLost: bytesLost, largestLostSentTime: largestLostSentTime, @@ -183,7 +436,7 @@ enum CongestionControl { smoothedRTT: smoothedRTT, now: now ) - self = .prague(algorithm: prague) + algorithm = .prague(algorithm: prague) return reducedCongestionWindow #endif } @@ -201,9 +454,10 @@ enum CongestionControl { now: NetworkClock.Instant, qlog: QLog? = nil ) { - switch self { - case .cubic(algorithm: var cubic): + switch algorithm { + case .cubic(var cubic): cubic.processECN( + state: &state, path: path, ceCount: ceCount, packetsAcked: packetsAcked, @@ -215,10 +469,11 @@ enum CongestionControl { now: now, qlog: qlog ) - self = .cubic(algorithm: cubic) + algorithm = .cubic(algorithm: cubic) #if !NETWORK_EMBEDDED - case .ledbat(algorithm: var ledbat): + case .ledbat(var ledbat): ledbat.processECN( + state: &state, path: path, ceCount: ceCount, packetsAcked: packetsAcked, @@ -230,9 +485,10 @@ enum CongestionControl { now: now, qlog: qlog ) - self = .ledbat(algorithm: ledbat) - case .prague(algorithm: var prague): + algorithm = .ledbat(algorithm: ledbat) + case .prague(var prague): prague.processECN( + state: &state, path: path, ceCount: ceCount, packetsAcked: packetsAcked, @@ -244,126 +500,52 @@ enum CongestionControl { now: now, qlog: qlog ) - self = .prague(algorithm: prague) + algorithm = .prague(algorithm: prague) #endif } } - mutating func packetDiscarded(bytesSent: Int, qlog: QLog? = nil) { - switch self { + mutating func spuriousRetransmit(qlog: QLog? = nil) { + switch algorithm { case .cubic(var cubic): - cubic.packetDiscarded(bytesSent: bytesSent, qlog: qlog) - self = .cubic(algorithm: cubic) + cubic.spuriousRetransmit(state: &state, qlog: qlog) + algorithm = .cubic(algorithm: cubic) #if !NETWORK_EMBEDDED case .ledbat(var ledbat): - ledbat.packetDiscarded(bytesSent: bytesSent, qlog: qlog) - self = .ledbat(algorithm: ledbat) + ledbat.spuriousRetransmit(state: &state, qlog: qlog) + algorithm = .ledbat(algorithm: ledbat) case .prague(var prague): - prague.packetDiscarded(bytesSent: bytesSent, qlog: qlog) - self = .prague(algorithm: prague) - #endif - } - } - - 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 - } - } - - 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 - } - } - - var name: String { - switch self { - case .cubic(algorithm: _): - return "CUBIC" - #if !NETWORK_EMBEDDED - case .ledbat(algorithm: _): - return "LEDBAT" - case .prague(algorithm: _): - return "PRAGUE" - #endif - } - } - - 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 - } - } - - 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) + prague.spuriousRetransmit(state: &state, qlog: qlog) + algorithm = .prague(algorithm: prague) #endif } } mutating func idleTimeout(mss: Int) { - switch self { - case .cubic(algorithm: var cubic): - cubic.idleTimeout(mss: mss, qlog: nil) - self = .cubic(algorithm: cubic) + switch algorithm { + case .cubic(var cubic): + cubic.idleTimeout(state: &state, mss: mss, qlog: nil) + algorithm = .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) + case .ledbat(var ledbat): + ledbat.idleTimeout(state: &state, mss: mss, qlog: nil) + algorithm = .ledbat(algorithm: ledbat) + case .prague(var prague): + prague.idleTimeout(state: &state, mss: mss, qlog: nil) + algorithm = .prague(algorithm: prague) #endif } } func filloutDataTransferSnapshot(dataTransferSnapshot: inout DataTransferSnapshot) { - switch self { - case .cubic(algorithm: let cubic): - cubic.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) + switch algorithm { + case .cubic(let cubic): + cubic.filloutDataTransferSnapshot(state: state, dataTransferSnapshot: &dataTransferSnapshot) #if !NETWORK_EMBEDDED - case .ledbat(algorithm: let ledbat): - ledbat.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) - case .prague(algorithm: let prague): - prague.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) + case .ledbat(let ledbat): + ledbat.filloutDataTransferSnapshot(state: state, dataTransferSnapshot: &dataTransferSnapshot) + case .prague(let prague): + prague.filloutDataTransferSnapshot(state: state, dataTransferSnapshot: &dataTransferSnapshot) #endif } } @@ -371,29 +553,15 @@ enum CongestionControl { @available(Network 0.1.0, *) protocol CongestionControlProtocol: PrefixedLoggable { - var congestionWindow: UInt64 { get set } - var bytesInFlight: UInt64 { get set } - var packetsAcked: UInt64 { get set } - var packetsMarked: UInt64 { get set } - var ecnCECounter: Int { get set } - var largestSentPN: Int64 { get set } - var slowStartThreshold: UInt64 { get set } - var prevSlowStartThreshold: UInt64 { get set } - var recoveryStartTime: NetworkClock.Instant { get set } - var bytesAcked: UInt64 { get set } - var pipeAckSamples: [UInt64] { get set } - var pipeAckValue: UInt64 { get set } - var pipeAckSampleEnd: NetworkClock.Instant { get set } - var pipeAckAcked: UInt64 { get set } - var pipeAckIndex: Int { get set } - mutating func inherit( - from: CongestionControl, + from: CongestionControlState, + state: inout CongestionControlState, mss: Int, qlog: QLog? ) - mutating func reset(mss: Int, qlog: QLog?) + mutating func reset(state: inout CongestionControlState, mss: Int, qlog: QLog?) mutating func ackEnd( + state: inout CongestionControlState, rtt: borrowing RTT, path: QUICPath?, mss: Int, @@ -401,11 +569,12 @@ protocol CongestionControlProtocol: PrefixedLoggable { now: NetworkClock.Instant, qlog: QLog? ) - mutating func spuriousRetransmit(qlog: QLog?) - mutating func idleTimeout(mss: Int, qlog: QLog?) + mutating func spuriousRetransmit(state: inout CongestionControlState, qlog: QLog?) + mutating func idleTimeout(state: inout CongestionControlState, mss: Int, qlog: QLog?) /// Opens a recovery period at `now`, the time the loss was detected. - mutating func enterRecovery(mss: Int, now: NetworkClock.Instant, qlog: QLog?) + mutating func enterRecovery(state: inout CongestionControlState, mss: Int, now: NetworkClock.Instant, qlog: QLog?) mutating func processECN( + state: inout CongestionControlState, path: QUICPath?, ceCount: Int, packetsAcked: Int, @@ -418,6 +587,7 @@ protocol CongestionControlProtocol: PrefixedLoggable { qlog: QLog? ) mutating func packetLost( + state: inout CongestionControlState, path: QUICPath?, bytesLost: Int, largestLostSentTime: NetworkClock.Instant, @@ -426,243 +596,66 @@ protocol CongestionControlProtocol: PrefixedLoggable { now: NetworkClock.Instant, qlog: QLog? ) -> Bool - mutating func linkFlowControl( - largestAckSentTime: NetworkClock.Instant, - mss: Int, - now: NetworkClock.Instant, - qlog: QLog? - ) - mutating func persistentCongestion(mss: Int, qlog: QLog?) - mutating func mssChanged(mss: Int, qlog: QLog?) - mutating func packetDiscarded(bytesSent: Int, qlog: QLog?) - mutating func ackBegin() - mutating func packetSent(bytesSent: Int, qlog: QLog?) - mutating func packetsAcked(bytesAcked: Int, sentTime: NetworkClock.Instant) + mutating func persistentCongestion(state: inout CongestionControlState, mss: Int, qlog: QLog?) + func filloutDataTransferSnapshot(state: CongestionControlState, dataTransferSnapshot: inout DataTransferSnapshot) } @available(Network 0.1.0, *) extension CongestionControlProtocol { - var congestionWindowValidationSamples: Int { - 3 + mutating func packetSent(state: inout CongestionControlState, bytesSent: Int, qlog: QLog? = nil) { + state.packetSent(bytesSent: bytesSent, log: log, qlog: qlog) } - var availableCongestionWindow: UInt64 { - if congestionWindow > bytesInFlight { - return congestionWindow - bytesInFlight - } else { - return 0 - } + mutating func packetDiscarded(state: inout CongestionControlState, bytesSent: Int, qlog: QLog? = nil) { + state.packetDiscarded(bytesSent: bytesSent, log: log, qlog: qlog) } - var congestionWindowValidated: Bool { - if pipeAckValue == 0 { - return true - } - // In slow-start, congestionWindow increases aggressively and pipeack - // might be lagging behind. Thus, give it 4x the space. - if congestionWindow < slowStartThreshold { - let congestionWindowLimit = congestionWindow >> 2 - if pipeAckValue < congestionWindowLimit { - log.datapath( - "Congestion window not validated in slow-start, pipeack: \(pipeAckValue), congestionWindow \(congestionWindow)" - ) - return false - } - } else { - let congestionWindowLimit = congestionWindow >> 1 - if pipeAckValue < congestionWindowLimit { - log.datapath( - "Congestion window not validated in congestion-avoidance, pipeack: \(pipeAckValue), congestionWindow: \(congestionWindow) " - ) - return false - } - } - - return true + mutating func ackBegin(state: inout CongestionControlState) { + state.ackBegin() } - var lossFlightSize: UInt64 { - if !congestionWindowValidated { - return max(pipeAckValue, bytesInFlight) - } else { - return congestionWindow - } + mutating func packetsAcked(state: inout CongestionControlState, bytesAcked: Int, sentTime: NetworkClock.Instant) { + state.packetsAcked(bytesAcked: bytesAcked, sentTime: sentTime, log: log) } - mutating private func incrementBytesInFlight(_ bytesSent: Int) { - bytesInFlight += UInt64(bytesSent) - log.datapath("Bytes in flight updated to \(bytesInFlight)") - - QUICSignpost.bytesInFlight(bytesInFlight: Int(bytesInFlight)) + mutating func mssChanged(state: inout CongestionControlState, mss: Int, qlog: QLog? = nil) { + state.mssChanged(mss: mss, log: log, qlog: qlog) } - mutating func decrementBytesInFlight(_ bytes: UInt64) { - let result = bytesInFlight.subtractingReportingOverflow(bytes) - if result.overflow { - log.fault("Undeflow, \(bytes) decremented from \(bytesInFlight)") - bytesInFlight = 0 - } else { - bytesInFlight = result.partialValue - } - log.datapath("Bytes in flight updated to \(bytesInFlight)") - QUICSignpost.bytesInFlight(bytesInFlight: Int(bytesInFlight)) - } - - mutating func mssChanged(mss: Int, qlog: QLog? = nil) { - congestionWindow = max(congestionWindow, UInt64(mss)) - logUpdate(qlog: qlog) - } - - mutating func packetSent(bytesSent: Int, qlog: QLog? = nil) { - incrementBytesInFlight(bytesSent) - logUpdate(qlog: qlog) - } - - mutating func packetDiscarded(bytesSent: Int, qlog: QLog? = nil) { - decrementBytesInFlight(UInt64(bytesSent)) - logUpdate(qlog: qlog) - } - - mutating func ackBegin() { - bytesAcked = 0 - } - - mutating func packetsAcked(bytesAcked: Int, sentTime: NetworkClock.Instant) { - let bytesAcked = UInt64(bytesAcked) - decrementBytesInFlight(bytesAcked) - if packetInRecovery(sentTime: sentTime) { - // Dont update the congestion window - log.datapath("Packet was sent before recovery, ignore") - return - } - // Congestion window is updated later in ackEnd - self.bytesAcked += bytesAcked - } - - func packetInRecovery(sentTime: NetworkClock.Instant) -> Bool { - sentTime <= recoveryStartTime + func canSend(state: CongestionControlState, packetLength: Int) -> Bool { + state.canSend(packetLength: packetLength, log: log) } /// `sentTime` is when the packet went out, `now` when its loss was detected. @discardableResult mutating func congestionEvent( + state: inout CongestionControlState, sentTime: NetworkClock.Instant, mss: Int, now: NetworkClock.Instant, qlog: QLog? = nil ) -> Bool { // If the packet was sent before recovery started, do nothing - if packetInRecovery(sentTime: sentTime) { return false } + if state.packetInRecovery(sentTime: sentTime) { return false } // Enter recovery if the packet was sent // after start of the previous recovery period - enterRecovery(mss: mss, now: now, qlog: qlog) + enterRecovery(state: &state, mss: mss, now: now, qlog: qlog) return true } mutating func linkFlowControl( + state: inout CongestionControlState, largestAckSentTime: NetworkClock.Instant, mss: Int, now: NetworkClock.Instant, qlog: QLog? = nil ) { - congestionEvent(sentTime: largestAckSentTime, mss: mss, now: now, qlog: qlog) + congestionEvent(state: &state, sentTime: largestAckSentTime, mss: mss, now: now, qlog: qlog) log.debug( - "Link was flow controlled, reduced congestion window is \(congestionWindow) bytes" + "Link was flow controlled, reduced congestion window is \(state.congestionWindow) bytes" ) } - // Compute if 1RTT or 1 round has elapsed by measuring if the - // packet sent after this instant has been acknowledged. - // largest_sent_pn is set at the start of a round - func rttElapsed(largestSentPN: Int64, largestAckedPN: Int64) -> Bool { - // A packet with pn higher than largest sent pn at - // the start of the round has been acknowledged - (largestSentPN == 0) || (largestAckedPN > largestSentPN) - } - - mutating func initPipeAckSamples() { - pipeAckSamples = Array(repeating: 0, count: congestionWindowValidationSamples) - pipeAckIndex = 0 - pipeAckValue = 0 - } - - mutating func setPipeAckSample(sample: UInt64) { - pipeAckSamples[pipeAckIndex] = sample - pipeAckIndex &+= 1 - pipeAckIndex = pipeAckIndex % congestionWindowValidationSamples - } - - mutating func pipeAckNewRound(target: NetworkClock.Instant) { - pipeAckSampleEnd = target == .zero ? .init(microseconds: 1) : target - pipeAckAcked = 0 - } - - mutating func updatePipeAckSamples() { - setPipeAckSample(sample: pipeAckAcked) - pipeAckValue = pipeAckAcked - for index in 0.. pipeAckValue { - pipeAckValue = pipeAckSamples[index] - } - } - } - - mutating func revalidateCongestionWindow(smoothedRTT: NetworkDuration, now: NetworkClock.Instant) -> Bool { - if pipeAckSampleEnd == .zero { - pipeAckNewRound(target: now.advanced(by: smoothedRTT)) - } - pipeAckAcked += bytesAcked - // A full period passed? Update our pipeack samples - if now > pipeAckSampleEnd { - let period = pipeAckSampleEnd.duration(to: now) - if period > smoothedRTT { - // More than 1 RTT of inactivity, we need to set samples to 0 - setPipeAckSample(sample: 0) - if period > smoothedRTT * 2 { - // Reset the next sample as well - setPipeAckSample(sample: 0) - } - } - updatePipeAckSamples() - pipeAckNewRound(target: now + smoothedRTT) - } - return congestionWindowValidated - } - - func canSend(packetLength: Int) -> Bool { - if availableCongestionWindow >= packetLength { - log.datapath( - "Can send packet because bytesInFlight \(bytesInFlight) + packetLength \(packetLength) <= congestionWindow \(congestionWindow)" - ) - return true - } else { - log.datapath( - "Congestion limited because bytesInFlight \(bytesInFlight) + packetLength \(packetLength) > congestionWindow \(congestionWindow)" - ) - QUICSignpost.congestionWindowLimited( - bytesInFlight: Int(bytesInFlight), - congestionWindow: Int(congestionWindow) - ) - return false - } - } - - func logUpdate(qlog: QLog?) { - if congestionWindow != UInt64.max { - log.datapath("Congestion window set to \(congestionWindow) bytes, bytes in flight \(bytesInFlight)") - QUICSignpost.congestionWindow(congestionWindow: Int(congestionWindow)) - } - #if QlogOutput - if let qlog { - qlog.congestionControlUpdated( - congestionWindow: congestionWindow, - bytesInFlight: bytesInFlight, - slowStartThresh: slowStartThreshold - ) - } - #endif - } - func logState( qlog: QLog? = nil, state: QLogCongestionState, diff --git a/Sources/SwiftNetwork/QUIC/Cubic.swift b/Sources/SwiftNetwork/QUIC/Cubic.swift index c57d509..cf7c80a 100644 --- a/Sources/SwiftNetwork/QUIC/Cubic.swift +++ b/Sources/SwiftNetwork/QUIC/Cubic.swift @@ -54,22 +54,6 @@ extension CubicLikeProtocol { struct Cubic: CongestionControlProtocol, CubicLikeProtocol { var log: LogPrefixer - var congestionWindow = UInt64(0) - var bytesInFlight = UInt64(0) - var packetsAcked = UInt64(0) - var packetsMarked = UInt64(0) - var ecnCECounter = 0 - var largestSentPN = Int64(0) - var slowStartThreshold = UInt64.max - var prevSlowStartThreshold = UInt64.max - var recoveryStartTime = NetworkClock.Instant.zero - var bytesAcked = UInt64(0) - var pipeAckSamples = [UInt64(0)] - var pipeAckValue = UInt64(0) - var pipeAckSampleEnd = NetworkClock.Instant.zero - var pipeAckAcked = UInt64(0) - var pipeAckIndex = 0 - var K: Double = 0 var numCongestionEvents = 0 var totalAcked = UInt64(0) /* total bytes acked for cubic */ @@ -96,21 +80,22 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { UInt64(min(10 * mss, max(2 * mss, 14720))) } - init(pacer: inout Pacer, mss: Int, qlog: QLog? = nil, logPrefixer: LogPrefixer) { + init(state: inout CongestionControlState, pacer: inout Pacer, mss: Int, qlog: QLog? = nil, logPrefixer: LogPrefixer) + { self.log = logPrefixer - reset(mss: mss, qlog: qlog) + reset(state: &state, mss: mss, qlog: qlog) if pacer.enabled { let startupRate = - congestionWindow * System.Time.USEC_PER_SEC / UInt64(pacingInitialRTT.microseconds) + state.congestionWindow * System.Time.USEC_PER_SEC / UInt64(pacingInitialRTT.microseconds) let startupBurstSize = UInt64(mss) pacer.setInitialState(startupRate, UInt32(truncatingIfNeeded: startupBurstSize)) pacer.reset() } - logUpdate(qlog: qlog) + state.logUpdate(log: log, qlog: qlog) logState(qlog: qlog, state: .slowStart, trigger: nil) } - private mutating func setK(mss: Int) { + private mutating func setK(state: CongestionControlState, 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 // congestion events. Compute the cubic K using, @@ -119,12 +104,12 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { K = 0 return } - guard maxCongestionWindow > congestionWindow else { + guard maxCongestionWindow > state.congestionWindow else { K = 0 - if maxCongestionWindow < congestionWindow { + if maxCongestionWindow < state.congestionWindow { // Log a fault for underflow cases. Don't log if it would just be zero. let maxCongestionWindow = maxCongestionWindow - let congestionWindow = congestionWindow + let congestionWindow = state.congestionWindow Logger.proto.fault( "Max congestion window \(maxCongestionWindow) should be greater than congestion window \(congestionWindow)" ) @@ -132,7 +117,7 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { return } - K = Double(maxCongestionWindow - congestionWindow) / Cubic.cFactor + K = Double(maxCongestionWindow - state.congestionWindow) / Cubic.cFactor K = K / Double(mss) #if !NETWORK_EMBEDDED K = cbrt(K) @@ -141,7 +126,12 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { #endif } - private mutating func getTarget(mss: Int, smoothedRTT: NetworkDuration, now: NetworkClock.Instant) -> UInt64 { + private mutating func getTarget( + state: CongestionControlState, + mss: Int, + smoothedRTT: NetworkDuration, + now: NetworkClock.Instant + ) -> UInt64 { if epochStart == .zero { // If we exit slow start without any packet // loss, CUBIC switches to CA where t is the elapsed @@ -149,15 +139,15 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { // epochStart to now here. epochStart = now // Set originPoint for the start of epoch - if congestionWindow < maxCongestionWindow { + if state.congestionWindow < maxCongestionWindow { originPoint = maxCongestionWindow - setK(mss: mss) + setK(state: state, mss: mss) } else { - originPoint = congestionWindow + originPoint = state.congestionWindow K = 0 } // Reset tcpCongestionWindow to be in sync with cubic - tcpCongestionWindow = congestionWindow + tcpCongestionWindow = state.congestionWindow tcpTotalAcked = 0 } // Compute target cubic window W(t+RTT) for the next RTT using, @@ -186,6 +176,7 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { // Handle an in-sequence ACK in congestion avoidance phase private mutating func processAckCongestionAvoidance( + state: inout CongestionControlState, bytesAcked: UInt64, smoothedRTT: NetworkDuration, mss: Int, @@ -193,21 +184,21 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { ) { totalAcked += bytesAcked // compute W(t+RTT) - let WCubicNext = getTarget(mss: mss, smoothedRTT: smoothedRTT, now: now) + let WCubicNext = getTarget(state: state, mss: mss, smoothedRTT: smoothedRTT, now: now) updateTCPWindow(bytesAcked: bytesAcked, mss: mss) - if congestionWindow < WCubicNext { + if state.congestionWindow < WCubicNext { // Either concave or convex region // Total increase in 1RTT is (W(t+RTT) - congestionWindow). // To get increase per ACK, multiply by (bytesAcked / congestionWindow) let incr = - Double((WCubicNext - congestionWindow)) - * (Double(totalAcked) / Double(congestionWindow)) - congestionWindow += min(UInt64(incr), Cubic.initialCongestionWindow(mss)) + Double((WCubicNext - state.congestionWindow)) + * (Double(totalAcked) / Double(state.congestionWindow)) + state.congestionWindow += min(UInt64(incr), Cubic.initialCongestionWindow(mss)) totalAcked = 0 } - if congestionWindow < tcpCongestionWindow { + if state.congestionWindow < tcpCongestionWindow { // TCP friendly region - congestionWindow = tcpCongestionWindow + state.congestionWindow = tcpCongestionWindow // When the congestionWindow is set based on TF region, // we should reset the totalAcked counter as we // have already used bytes acked equivalent to @@ -218,11 +209,11 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { // Set WMax to congestionWindow to keep updating our current estimate of WMax // as we are probing for new limits at the start of connection if numCongestionEvents == 0 { - maxCongestionWindow = congestionWindow + maxCongestionWindow = state.congestionWindow } } - private mutating func updatePacerState(path: QUICPath?, smoothedRTT: NetworkDuration) { + private func updatePacerState(state: CongestionControlState, path: QUICPath?, smoothedRTT: NetworkDuration) { guard let path, path.pacer.enabled else { return } @@ -234,7 +225,8 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { smoothedRTT.microseconds == 0 ? pacingInitialRTT.microseconds : smoothedRTT.microseconds // Use 200% rate when in slow start - let pacedWindow = congestionWindow < slowStartThreshold ? congestionWindow * 2 : congestionWindow + let pacedWindow = + state.congestionWindow < state.slowStartThreshold ? state.congestionWindow * 2 : state.congestionWindow let rateInBytesPerSecond = pacedWindow * System.Time.USEC_PER_SEC / UInt64(smoothedRTTInMicroseconds) let burstSize = rateInBytesPerSecond >> burstQueueShift @@ -245,6 +237,7 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { @discardableResult mutating func packetLost( + state: inout CongestionControlState, path: QUICPath?, bytesLost: Int, largestLostSentTime: NetworkClock.Instant, @@ -253,28 +246,34 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { now: NetworkClock.Instant, qlog: QLog? = nil ) -> Bool { - decrementBytesInFlight(UInt64(bytesLost)) + state.decrementBytesInFlight(UInt64(bytesLost), log: log) let reducedCongestionWindow = congestionEvent( + state: &state, sentTime: largestLostSentTime, mss: mss, now: now, qlog: qlog ) - updatePacerState(path: path, smoothedRTT: smoothedRTT) + updatePacerState(state: state, path: path, smoothedRTT: smoothedRTT) return reducedCongestionWindow } - mutating func enterRecovery(mss: Int, now: NetworkClock.Instant, qlog: QLog? = nil) { - log.datapath("Entering Recovery: current cwin=\(congestionWindow)") - recoveryStartTime = now + mutating func enterRecovery( + state: inout CongestionControlState, + mss: Int, + now: NetworkClock.Instant, + qlog: QLog? = nil + ) { + log.datapath("Entering Recovery: current cwin=\(state.congestionWindow)") + state.recoveryStartTime = now lastMaxCongestionWindow = maxCongestionWindow - maxCongestionWindow = congestionWindow - congestionWindow = UInt64(Double(lossFlightSize) * Cubic.beta) - if _slowPath(congestionWindow < Cubic.minCongestionWindow(mss)) { - congestionWindow = UInt64(Cubic.minCongestionWindow(mss)) + maxCongestionWindow = state.congestionWindow + state.congestionWindow = UInt64(Double(state.lossFlightSize(log: log)) * Cubic.beta) + if _slowPath(state.congestionWindow < Cubic.minCongestionWindow(mss)) { + state.congestionWindow = UInt64(Cubic.minCongestionWindow(mss)) } - prevSlowStartThreshold = slowStartThreshold - slowStartThreshold = congestionWindow + state.prevSlowStartThreshold = state.slowStartThreshold + state.slowStartThreshold = state.congestionWindow // If Fast Convergence is supported, release more bandwidth // if saturation point is getting reduced due to new flows if maxCongestionWindow < lastMaxCongestionWindow { @@ -288,20 +287,21 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { // Compute epoch period K(s) that the window will take to increase // to last_max again after backoff due to loss. // Note that K = 0 if we enter congestion avoidance without loss. - setK(mss: mss) + setK(state: state, mss: mss) // Set the start of current congestion avoidance and the origin point epochStart = now originPoint = maxCongestionWindow // Reset tcpCongestionWindow to be in sync with cubic - tcpCongestionWindow = congestionWindow + tcpCongestionWindow = state.congestionWindow tcpTotalAcked = 0 numCongestionEvents += 1 - initPipeAckSamples() - logUpdate(qlog: qlog) + state.initPipeAckSamples() + state.logUpdate(log: log, qlog: qlog) logState(qlog: qlog, state: .recovery, trigger: nil) } mutating func ackEnd( + state: inout CongestionControlState, rtt: borrowing RTT, path: QUICPath?, mss: Int, @@ -314,37 +314,39 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { // this ACK processing return } - if bytesAcked == 0 { + if state.bytesAcked == 0 { // When we are in recovery period or received new CE counts return } let smoothedRTT = rtt.smoothedRTT - if !revalidateCongestionWindow(smoothedRTT: smoothedRTT, now: now) { - bytesAcked = 0 + if !state.revalidateCongestionWindow(smoothedRTT: smoothedRTT, now: now, log: log) { + state.bytesAcked = 0 return } - if congestionWindow < slowStartThreshold { - congestionWindow += min( - bytesAcked, + if state.congestionWindow < state.slowStartThreshold { + state.congestionWindow += min( + state.bytesAcked, Cubic.slowStartCongestionWindow(mss) ) } else { processAckCongestionAvoidance( - bytesAcked: bytesAcked, + state: &state, + bytesAcked: state.bytesAcked, smoothedRTT: smoothedRTT, mss: mss, now: now ) } // Should be a minimum of 2*MSS - if _slowPath(congestionWindow < Cubic.minCongestionWindow(mss)) { - congestionWindow = Cubic.minCongestionWindow(mss) + if _slowPath(state.congestionWindow < Cubic.minCongestionWindow(mss)) { + state.congestionWindow = Cubic.minCongestionWindow(mss) } - updatePacerState(path: path, smoothedRTT: smoothedRTT) - logUpdate(qlog: qlog) + updatePacerState(state: state, path: path, smoothedRTT: smoothedRTT) + state.logUpdate(log: log, qlog: qlog) } mutating func processECN( + state: inout CongestionControlState, path: QUICPath?, ceCount: Int, packetsAcked: Int, @@ -356,128 +358,122 @@ struct Cubic: CongestionControlProtocol, CubicLikeProtocol { now: NetworkClock.Instant, qlog: QLog? = nil ) { - if _slowPath(ceCount < ecnCECounter) { + if _slowPath(ceCount < state.ecnCECounter) { log.fault( - "New CE count \(ceCount) can't be less than current CE count \(ecnCECounter)" + "New CE count \(ceCount) can't be less than current CE count \(state.ecnCECounter)" ) } // Update packets acked and marked on every ACK, even if it // is not used by CUBIC. This state is relevant // for Prague and needs to be updated by other CCs // esp. LEDBAT. - packetsMarked = UInt64(ceCount) - self.packetsAcked = UInt64(packetsAcked) + state.packetsMarked = UInt64(ceCount) + state.packetsAcked = UInt64(packetsAcked) - if ceCount == ecnCECounter { + if ceCount == state.ecnCECounter { // No change in CE return } log.datapath( - "\(bytesAcked) bytes were ACKed with \(ecnCECounter) packets newly CE marked" + "\(state.bytesAcked) bytes were ACKed with \(state.ecnCECounter) packets newly CE marked" ) /* Update CE count even if we are already in CWR */ - ecnCECounter = ceCount + state.ecnCECounter = ceCount // Received an ACK with new CE counts, reset bytesAcked so // we that we don't increase congestionWindow during ackEnd - bytesAcked = 0 + state.bytesAcked = 0 - if !rttElapsed( - largestSentPN: self.largestSentPN, + if !state.rttElapsed( + largestSentPN: state.largestSentPN, largestAckedPN: largestAckedPN ) { // Haven't elapsed one RTT yet from last CWR return } - congestionEvent(sentTime: largestAckedSentTime, mss: mss, now: now, qlog: qlog) + congestionEvent(state: &state, sentTime: largestAckedSentTime, mss: mss, now: now, qlog: qlog) // Update pacer state as congestionWindow has changed - updatePacerState(path: path, smoothedRTT: smoothedRTT) + updatePacerState(state: state, path: path, smoothedRTT: smoothedRTT) // Start new round for CWR - self.largestSentPN = largestSentPN + state.largestSentPN = largestSentPN } - mutating func idleTimeout(mss: Int, qlog: QLog? = nil) { + mutating func idleTimeout(state: inout CongestionControlState, mss: Int, qlog: QLog? = nil) { // We want to ideally begin with slow start after idle period. // Set it to the larger of its current value, MAX (congestionWindow * Beta, IW) - slowStartThreshold = Cubic.idleTimeout( - slowStartThreshold: slowStartThreshold, - congestionWindow: congestionWindow, + state.slowStartThreshold = Cubic.idleTimeout( + slowStartThreshold: state.slowStartThreshold, + congestionWindow: state.congestionWindow, mss: mss ) // Set congestionWindow to initial congestion window - congestionWindow = min(congestionWindow, Cubic.initialCongestionWindow(mss)) - logUpdate(qlog: qlog) - resetInternal() + state.congestionWindow = min(state.congestionWindow, Cubic.initialCongestionWindow(mss)) + state.logUpdate(log: log, qlog: qlog) + resetInternal(state: &state) } - mutating func persistentCongestion(mss: Int, qlog: QLog? = nil) { - slowStartThreshold = Cubic.persistentCongestion( - congestionWindow: congestionWindow, + mutating func persistentCongestion(state: inout CongestionControlState, mss: Int, qlog: QLog? = nil) { + state.slowStartThreshold = Cubic.persistentCongestion( + congestionWindow: state.congestionWindow, mss: mss ) // Set the minimum congestion window - congestionWindow = Cubic.minCongestionWindow(mss) - logUpdate(qlog: qlog) + state.congestionWindow = Cubic.minCongestionWindow(mss) + state.logUpdate(log: log, qlog: qlog) logState(qlog: qlog, state: .slowStart, trigger: .persistentCongestion) } - mutating func spuriousRetransmit(qlog: QLog? = nil) { - guard maxCongestionWindow > 0 && prevSlowStartThreshold > 0 else { return } + mutating func spuriousRetransmit(state: inout CongestionControlState, qlog: QLog? = nil) { + guard maxCongestionWindow > 0 && state.prevSlowStartThreshold > 0 else { return } // Revert to the state before loss was detected - congestionWindow = max(maxCongestionWindow, congestionWindow) - slowStartThreshold = prevSlowStartThreshold - logUpdate(qlog: qlog) + state.congestionWindow = max(maxCongestionWindow, state.congestionWindow) + state.slowStartThreshold = state.prevSlowStartThreshold + state.logUpdate(log: log, qlog: qlog) } - private mutating func resetInternal() { - recoveryStartTime = .zero - prevSlowStartThreshold = UInt64.max + private mutating func resetInternal(state: inout CongestionControlState) { + state.recoveryStartTime = .zero + state.prevSlowStartThreshold = UInt64.max numCongestionEvents = 0 K = 0 totalAcked = 0 epochStart = .zero originPoint = 0 lastMaxCongestionWindow = 0 - maxCongestionWindow = congestionWindow + maxCongestionWindow = state.congestionWindow // CWV state - pipeAckSampleEnd = .zero - initPipeAckSamples() + state.pipeAckSampleEnd = .zero + state.initPipeAckSamples() } - func filloutDataTransferSnapshot(dataTransferSnapshot: inout DataTransferSnapshot) { - dataTransferSnapshot.transportCongestionWindow = congestionWindow - dataTransferSnapshot.transportSlowStartThreshold = slowStartThreshold + func filloutDataTransferSnapshot(state: CongestionControlState, dataTransferSnapshot: inout DataTransferSnapshot) { + dataTransferSnapshot.transportCongestionWindow = state.congestionWindow + dataTransferSnapshot.transportSlowStartThreshold = state.slowStartThreshold } - mutating func reset(mss: Int, qlog: QLog? = nil) { - congestionWindow = Cubic.initialCongestionWindow(mss) - slowStartThreshold = UInt64.max - resetInternal() - logUpdate(qlog: qlog) + mutating func reset(state: inout CongestionControlState, mss: Int, qlog: QLog? = nil) { + state.congestionWindow = Cubic.initialCongestionWindow(mss) + state.slowStartThreshold = UInt64.max + resetInternal(state: &state) + state.logUpdate(log: log, qlog: qlog) } - mutating func inherit(from: CongestionControl, mss: Int, qlog: QLog?) { + mutating func inherit( + from: CongestionControlState, + state: inout CongestionControlState, + mss: Int, + qlog: QLog? = nil + ) { // 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 - } - slowStartThreshold = UInt64.max - resetInternal() - logUpdate(qlog: qlog) + state.bytesInFlight = from.bytesInFlight + state.congestionWindow = max(from.congestionWindow, Cubic.initialCongestionWindow(mss)) + state.slowStartThreshold = UInt64.max + resetInternal(state: &state) + state.logUpdate(log: log, qlog: qlog) } } #endif diff --git a/Sources/SwiftNetwork/QUIC/Ledbat.swift b/Sources/SwiftNetwork/QUIC/Ledbat.swift index fa64b39..36e2dd7 100644 --- a/Sources/SwiftNetwork/QUIC/Ledbat.swift +++ b/Sources/SwiftNetwork/QUIC/Ledbat.swift @@ -28,22 +28,6 @@ internal import os struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { let log: LogPrefixer - var congestionWindow = UInt64(0) - var bytesInFlight = UInt64(0) - var packetsAcked = UInt64(0) - var packetsMarked = UInt64(0) - var ecnCECounter = 0 - var largestSentPN = Int64(0) - var slowStartThreshold = UInt64(0) - var prevSlowStartThreshold = UInt64(0) - var recoveryStartTime = NetworkClock.Instant.zero - var bytesAcked = UInt64(0) - var pipeAckSamples = [UInt64(0)] - var pipeAckValue = UInt64(0) - var pipeAckSampleEnd = NetworkClock.Instant.zero - var pipeAckAcked = UInt64(0) - var pipeAckIndex = 0 - private var prevCongestionWindow = UInt64(0) private var slowDownTimestamp = NetworkClock.Instant.zero private var slowDownBegin = NetworkClock.Instant.zero @@ -58,12 +42,12 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { UInt64(min(2 * mss, Ledbat.defaultCongestionWindow)) } - init(mss: Int, qlog: QLog? = nil, logPrefixer: LogPrefixer) { + init(state: inout CongestionControlState, mss: Int, qlog: QLog? = nil, logPrefixer: LogPrefixer) { self.log = logPrefixer - congestionWindow = Ledbat.initialCongestionWindow(mss) - slowStartThreshold = UInt64.max - reset(mss: mss, qlog: qlog) - logUpdate(qlog: qlog) + state.congestionWindow = Ledbat.initialCongestionWindow(mss) + state.slowStartThreshold = UInt64.max + reset(state: &state, mss: mss, qlog: qlog) + state.logUpdate(log: log, qlog: qlog) } // GAIN is proportional to the ratio of base_delay @@ -82,6 +66,7 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { @discardableResult mutating func packetLost( + state: inout CongestionControlState, path: QUICPath?, bytesLost: Int, largestLostSentTime: NetworkClock.Instant, @@ -90,8 +75,9 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { now: NetworkClock.Instant, qlog: QLog? = nil ) -> Bool { - decrementBytesInFlight(UInt64(bytesLost)) + state.decrementBytesInFlight(UInt64(bytesLost), log: log) let reducedCongestionWindow = congestionEvent( + state: &state, sentTime: largestLostSentTime, mss: mss, now: now, @@ -100,21 +86,27 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { return reducedCongestionWindow } - mutating func enterRecovery(mss: Int, now: NetworkClock.Instant, qlog: QLog? = nil) { - recoveryStartTime = now - prevCongestionWindow = congestionWindow - congestionWindow = UInt64(Double(lossFlightSize) * Ledbat.beta) - if _slowPath(congestionWindow < Ledbat.minCongestionWindow(mss)) { - congestionWindow = Ledbat.minCongestionWindow(mss) + mutating func enterRecovery( + state: inout CongestionControlState, + mss: Int, + now: NetworkClock.Instant, + qlog: QLog? = nil + ) { + state.recoveryStartTime = now + prevCongestionWindow = state.congestionWindow + state.congestionWindow = UInt64(Double(state.lossFlightSize(log: log)) * Ledbat.beta) + if _slowPath(state.congestionWindow < Ledbat.minCongestionWindow(mss)) { + state.congestionWindow = Ledbat.minCongestionWindow(mss) } - prevSlowStartThreshold = slowStartThreshold - slowStartThreshold = congestionWindow - initPipeAckSamples() - logUpdate(qlog: qlog) + state.prevSlowStartThreshold = state.slowStartThreshold + state.slowStartThreshold = state.congestionWindow + state.initPipeAckSamples() + state.logUpdate(log: log, qlog: qlog) logState(qlog: qlog, state: .recovery, trigger: nil) } mutating func ackEnd( + state: inout CongestionControlState, rtt: borrowing RTT, path: QUICPath?, mss: Int, @@ -127,13 +119,13 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { // this ACK processing return } - if bytesAcked == 0 { + if state.bytesAcked == 0 { // When we are in recovery period or received new CE counts return } let smoothedRTT = rtt.smoothedRTT - if !revalidateCongestionWindow(smoothedRTT: smoothedRTT, now: now) { - bytesAcked = 0 + if !state.revalidateCongestionWindow(smoothedRTT: smoothedRTT, now: now, log: log) { + state.bytesAcked = 0 return } let baseRTT = rtt.baseRTT @@ -154,9 +146,9 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { } if now < slowDownTimestamp + smoothedRTT * 2 { // Set cwnd to 2 packets and return - if congestionWindow > Ledbat.minCongestionWindow(mss) { - slowStartThreshold = congestionWindow - congestionWindow = Ledbat.minCongestionWindow(mss) + if state.congestionWindow > Ledbat.minCongestionWindow(mss) { + state.slowStartThreshold = state.congestionWindow + state.congestionWindow = Ledbat.minCongestionWindow(mss) } return } @@ -169,11 +161,11 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { // slow start, during CA, window growth will be bound // by ssthresh. let slowStartTarget = Ledbat.target * 0.75 - if congestionWindow < slowStartThreshold + if state.congestionWindow < state.slowStartThreshold && (numSlowDownEvents > 0 || qDelay < slowStartTarget) { - congestionWindow += UInt64( - gain(baseRTT) * Double(min(bytesAcked, Ledbat.slowStartCongestionWindow(mss))) + state.congestionWindow += UInt64( + gain(baseRTT) * Double(min(state.bytesAcked, Ledbat.slowStartCongestionWindow(mss))) ) // Reset the exit time if slowDownTimestamp != .zero { @@ -186,7 +178,7 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { if slowDownTimestamp == .zero { // On exit slow start due to higher queuing delay, cap // the ssthresh - slowStartThreshold = min(slowStartThreshold, congestionWindow) + state.slowStartThreshold = min(state.slowStartThreshold, state.congestionWindow) if numSlowDownEvents > 0 && slowDownEnd == .zero { // Set the slowdown end immediately after the // previous slowdown event @@ -203,8 +195,8 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { } // Additive increase -> W += GAIN (per RTT) if qDelay < Ledbat.target { - let tempIncrement = gain(baseRTT) * Double(bytesAcked) - congestionWindow += UInt64(tempIncrement * Double(mss) / Double(congestionWindow)) + let tempIncrement = gain(baseRTT) * Double(state.bytesAcked) + state.congestionWindow += UInt64(tempIncrement * Double(mss) / Double(state.congestionWindow)) } else { // Multiplicative decrease -> // W -= min(W * (qdelay/target - 1), W/2) (per RTT) @@ -216,22 +208,23 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { Double(qDelay.microseconds) / Double(Ledbat.target.microseconds) - 1, 0.5 ) - congestionWindow -= UInt64(tempMin * Double(min(bytesAcked, congestionWindow))) + state.congestionWindow -= UInt64(tempMin * Double(min(state.bytesAcked, state.congestionWindow))) // MD during Congestion Avoidance, limit ssthresh to // current cwnd - slowStartThreshold = min(slowStartThreshold, congestionWindow) + state.slowStartThreshold = min(state.slowStartThreshold, state.congestionWindow) } } // Should be a minimum of 2*MSS - if _slowPath(congestionWindow < Ledbat.minCongestionWindow(mss)) { - congestionWindow = Ledbat.minCongestionWindow(mss) + if _slowPath(state.congestionWindow < Ledbat.minCongestionWindow(mss)) { + state.congestionWindow = Ledbat.minCongestionWindow(mss) // ssthresh should be at least 2*MSS as well - slowStartThreshold = max(slowStartThreshold, congestionWindow) + state.slowStartThreshold = max(state.slowStartThreshold, state.congestionWindow) } - logUpdate(qlog: qlog) + state.logUpdate(log: log, qlog: qlog) } mutating func processECN( + state: inout CongestionControlState, path: QUICPath?, ceCount: Int, packetsAcked: Int, @@ -243,62 +236,62 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { now: NetworkClock.Instant, qlog: QLog? = nil ) { - if _slowPath(ceCount < ecnCECounter) { + if _slowPath(ceCount < state.ecnCECounter) { log.fault( - "New CE count \(ceCount) can't be less than current CE count \(ecnCECounter)" + "New CE count \(ceCount) can't be less than current CE count \(state.ecnCECounter)" ) } // Update packets acked and marked on every ACK, even if it // is not used by LEDBAT. This state is relevant // for Prague and needs to be updated by other CCs // esp. LEDBAT. - packetsMarked = UInt64(ceCount) - self.packetsAcked = UInt64(packetsAcked) + state.packetsMarked = UInt64(ceCount) + state.packetsAcked = UInt64(packetsAcked) - if ceCount == ecnCECounter { + if ceCount == state.ecnCECounter { // No change in CE return } log.datapath( - "\(bytesAcked) bytes were ACKed with \(ecnCECounter) packets newly CE marked" + "\(state.bytesAcked) bytes were ACKed with \(state.ecnCECounter) packets newly CE marked" ) // Update CE count even if we are already in CWR - ecnCECounter = ceCount + state.ecnCECounter = ceCount // Received an ACK with new CE counts, reset bytes_acked so // we that we don't increase cwnd during ack_end - bytesAcked = 0 - if !rttElapsed(largestSentPN: self.largestSentPN, largestAckedPN: largestAckedPN) { + state.bytesAcked = 0 + if !state.rttElapsed(largestSentPN: state.largestSentPN, largestAckedPN: largestAckedPN) { /* Haven't elapsed one RTT yet from last CWR */ return } - congestionEvent(sentTime: largestAckedSentTime, mss: mss, now: now) + congestionEvent(state: &state, sentTime: largestAckedSentTime, mss: mss, now: now) // Start new round for CWR - self.largestSentPN = largestSentPN + state.largestSentPN = largestSentPN } - mutating func spuriousRetransmit(qlog: QLog? = nil) { - guard prevCongestionWindow > 0 && prevSlowStartThreshold > 0 else { return } + mutating func spuriousRetransmit(state: inout CongestionControlState, qlog: QLog? = nil) { + guard prevCongestionWindow > 0 && state.prevSlowStartThreshold > 0 else { return } // Revert to the state before loss was detected - congestionWindow = max(prevCongestionWindow, congestionWindow) - slowStartThreshold = prevSlowStartThreshold - logUpdate(qlog: qlog) + state.congestionWindow = max(prevCongestionWindow, state.congestionWindow) + state.slowStartThreshold = state.prevSlowStartThreshold + state.logUpdate(log: log, qlog: qlog) } - mutating func persistentCongestion(mss: Int, qlog: QLog? = nil) { + mutating func persistentCongestion(state: inout CongestionControlState, mss: Int, qlog: QLog? = nil) { // Set the minimum congestion window let newCWND = Ledbat.minCongestionWindow(mss) - slowStartThreshold = max(UInt64(Double(congestionWindow) * Ledbat.beta), newCWND) - congestionWindow = newCWND - logUpdate(qlog: qlog) + state.slowStartThreshold = max(UInt64(Double(state.congestionWindow) * Ledbat.beta), newCWND) + state.congestionWindow = newCWND + state.logUpdate(log: log, qlog: qlog) logState(qlog: qlog, state: .slowStart, trigger: .persistentCongestion) } - private mutating func resetInternal() { - recoveryStartTime = .zero - prevSlowStartThreshold = 0 + private mutating func resetInternal(state: inout CongestionControlState) { + state.recoveryStartTime = .zero + state.prevSlowStartThreshold = 0 prevCongestionWindow = 0 slowDownTimestamp = .zero @@ -307,55 +300,49 @@ struct Ledbat: CongestionControlProtocol, CubicLikeProtocol { numSlowDownEvents = 0 // CWV state - pipeAckSampleEnd = .zero - initPipeAckSamples() + state.pipeAckSampleEnd = .zero + state.initPipeAckSamples() } - mutating func idleTimeout(mss: Int, qlog: QLog? = nil) { + mutating func idleTimeout(state: inout CongestionControlState, mss: Int, qlog: QLog? = nil) { // We want to ideally begin with slow start after idle period. // Set it to the larger of its current value, MAX (cwnd * Beta, IW) - slowStartThreshold = max( - slowStartThreshold, - max(UInt64(Double(congestionWindow) * Ledbat.beta), Ledbat.initialCongestionWindow(mss)) + state.slowStartThreshold = max( + state.slowStartThreshold, + max(UInt64(Double(state.congestionWindow) * Ledbat.beta), Ledbat.initialCongestionWindow(mss)) ) // Set cwnd to initial cwnd - congestionWindow = min(congestionWindow, Ledbat.initialCongestionWindow(mss)) - resetInternal() - logUpdate(qlog: qlog) + state.congestionWindow = min(state.congestionWindow, Ledbat.initialCongestionWindow(mss)) + resetInternal(state: &state) + state.logUpdate(log: log, qlog: qlog) } - func filloutDataTransferSnapshot(dataTransferSnapshot: inout DataTransferSnapshot) { - dataTransferSnapshot.transportCongestionWindow = congestionWindow - dataTransferSnapshot.transportSlowStartThreshold = slowStartThreshold + func filloutDataTransferSnapshot(state: CongestionControlState, dataTransferSnapshot: inout DataTransferSnapshot) { + dataTransferSnapshot.transportCongestionWindow = state.congestionWindow + dataTransferSnapshot.transportSlowStartThreshold = state.slowStartThreshold } - mutating func reset(mss: Int, qlog: QLog? = nil) { - congestionWindow = Ledbat.initialCongestionWindow(mss) - slowStartThreshold = UInt64.max - resetInternal() - logUpdate(qlog: qlog) + mutating func reset(state: inout CongestionControlState, mss: Int, qlog: QLog? = nil) { + state.congestionWindow = Ledbat.initialCongestionWindow(mss) + state.slowStartThreshold = UInt64.max + resetInternal(state: &state) + state.logUpdate(log: log, qlog: qlog) } - mutating func inherit(from: CongestionControl, mss: Int, qlog: QLog?) { + mutating func inherit( + from: CongestionControlState, + state: inout CongestionControlState, + mss: Int, + qlog: QLog? = nil + ) { // 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 - } - logUpdate(qlog: qlog) - resetInternal() + state.bytesInFlight = from.bytesInFlight + state.congestionWindow = min(from.congestionWindow, state.congestionWindow) + state.logUpdate(log: log, qlog: qlog) + resetInternal(state: &state) } } #endif diff --git a/Sources/SwiftNetwork/QUIC/Prague.swift b/Sources/SwiftNetwork/QUIC/Prague.swift index d8bb262..2c20857 100644 --- a/Sources/SwiftNetwork/QUIC/Prague.swift +++ b/Sources/SwiftNetwork/QUIC/Prague.swift @@ -42,22 +42,6 @@ enum RTTControlType: UInt8 { struct Prague: CongestionControlProtocol, CubicLikeProtocol { let log: LogPrefixer - var congestionWindow = UInt64(0) - var bytesInFlight = UInt64(0) - var packetsAcked = UInt64(0) - var packetsMarked = UInt64(0) - var ecnCECounter = 0 - var largestSentPN = Int64(0) - var slowStartThreshold = UInt64(0) - var prevSlowStartThreshold = UInt64(0) - var recoveryStartTime = NetworkClock.Instant.zero - var bytesAcked = UInt64(0) - var pipeAckSamples = [UInt64(0)] - var pipeAckValue = UInt64(0) - var pipeAckSampleEnd = NetworkClock.Instant.zero - var pipeAckAcked = UInt64(0) - var pipeAckIndex = 0 - private static let alphaShift = 20 private static let gShift = 4 private static let congestionWindowShift = 20 @@ -123,20 +107,21 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { UInt64(min(10 * mss, max(2 * mss, 14720))) } - init(pacer: inout Pacer, mss: Int, qlog: QLog? = nil, logPrefixer: LogPrefixer) { + init(state: inout CongestionControlState, pacer: inout Pacer, mss: Int, qlog: QLog? = nil, logPrefixer: LogPrefixer) + { self.log = logPrefixer - congestionWindow = Prague.initialCongestionWindow(mss) - slowStartThreshold = UInt64.max + state.congestionWindow = Prague.initialCongestionWindow(mss) + state.slowStartThreshold = UInt64.max scaledAlpha = Prague.maxAlpha << Prague.gShift - resetInternal() + resetInternal(state: &state) if pacer.enabled { let startupRate = - congestionWindow * System.Time.USEC_PER_SEC / UInt64(pacingInitialRTT.microseconds) + state.congestionWindow * System.Time.USEC_PER_SEC / UInt64(pacingInitialRTT.microseconds) let startupBurstSize = UInt64(mss) pacer.setInitialState(startupRate, UInt32(truncatingIfNeeded: startupBurstSize)) pacer.reset() } - logUpdate(qlog: qlog) + state.logUpdate(log: log, qlog: qlog) logState(qlog: qlog, state: .slowStart, trigger: nil) } @@ -146,12 +131,12 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { /// the current window size to `W_max` if there are no further /// congestion events. Computes the cubic `K` using /// `K = cubic_root(W_max(1-ß)/C)`. - private mutating func setCubicK(mss: Int) { + private mutating func setCubicK(state: CongestionControlState, mss: Int) { guard cubicMaxCongestionWindow != 0 else { cubicK = 0 return } - var K = Double(cubicMaxCongestionWindow - congestionWindow) / Prague.cFactor + var K = Double(cubicMaxCongestionWindow - state.congestionWindow) / Prague.cFactor K = K / Double(mss) #if !NETWORK_EMBEDDED K = cbrt(K) @@ -162,6 +147,7 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { } private mutating func getCubicTarget( + state: CongestionControlState, mss: Int, smoothedRTT: NetworkDuration, now: NetworkClock.Instant @@ -172,15 +158,15 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { // So, set epoch_start to now here. cubicEpochStart = now // Set origin_point for the start of epoch - if congestionWindow < cubicMaxCongestionWindow { + if state.congestionWindow < cubicMaxCongestionWindow { cubicOriginPoint = cubicMaxCongestionWindow - setCubicK(mss: mss) + setCubicK(state: state, mss: mss) } else { - cubicOriginPoint = congestionWindow + cubicOriginPoint = state.congestionWindow cubicK = 0 } // Reset reno_cwnd to be in sync with cubic - renoCongestionWindow = congestionWindow + renoCongestionWindow = state.congestionWindow renoAcked = 0 } @@ -212,6 +198,7 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { /// Handles an ACK in the congestion-avoidance phase after packet loss. private mutating func cubicProcessAckCA( + state: inout CongestionControlState, bytesAcked: UInt64, smoothedRTT: NetworkDuration, mss: Int, @@ -220,23 +207,23 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { cubicAcked += bytesAcked // compute W(t+RTT) - let wCubicNext = getCubicTarget(mss: mss, smoothedRTT: smoothedRTT, now: now) + let wCubicNext = getCubicTarget(state: state, mss: mss, smoothedRTT: smoothedRTT, now: now) updateRenoCongestionWindow(bytesAcked: bytesAcked, mss: mss) - if congestionWindow < wCubicNext { + if state.congestionWindow < wCubicNext { // Either concave or convex region // Total increase in 1RTT is (W(t+RTT) - cwnd). // To get increase per ACK, multiply by (bytes_acked / cwnd) let incr = - Double(wCubicNext - congestionWindow) - * Double(cubicAcked) / Double(congestionWindow) - congestionWindow += min(UInt64(incr), Prague.initialCongestionWindow(mss)) + Double(wCubicNext - state.congestionWindow) + * Double(cubicAcked) / Double(state.congestionWindow) + state.congestionWindow += min(UInt64(incr), Prague.initialCongestionWindow(mss)) cubicAcked = 0 } - if congestionWindow < renoCongestionWindow { + if state.congestionWindow < renoCongestionWindow { // TCP friendly region - congestionWindow = renoCongestionWindow + state.congestionWindow = renoCongestionWindow // When the cwnd is set based on Reno-Friendly region, // we should reset the cubic_acked counter as we // have already used bytes acked equivalent to @@ -247,7 +234,7 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { // Set W_max to cwnd to keep updating our current estimate of W_max // as we are probing for new limits at the start of connection if numCongestionEventsLoss == 0 { - cubicMaxCongestionWindow = congestionWindow + cubicMaxCongestionWindow = state.congestionWindow } } @@ -294,13 +281,13 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { } /// Handles an ACK in the congestion-avoidance phase after the decrease caused by CE. - private mutating func pragueCAAfterCE(bytesAcked: UInt64, mss: Int) { + private mutating func pragueCAAfterCE(state: inout CongestionControlState, bytesAcked: UInt64, mss: Int) { var increase = bytesAcked * UInt64(mss) * alphaAI - increase = (increase + (congestionWindow >> 1)) / congestionWindow - congestionWindow += increase >> Prague.congestionWindowShift + increase = (increase + (state.congestionWindow >> 1)) / state.congestionWindow + state.congestionWindow += increase >> Prague.congestionWindowShift } - private func updatePacerState(path: QUICPath?, smoothedRTT: NetworkDuration) { + private func updatePacerState(state: CongestionControlState, path: QUICPath?, smoothedRTT: NetworkDuration) { guard let path, path.pacer.enabled else { return } @@ -312,7 +299,8 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { smoothedRTT.microseconds == 0 ? pacingInitialRTT.microseconds : smoothedRTT.microseconds // Use 200% rate when in slow start - let pacedWindow = congestionWindow < slowStartThreshold ? congestionWindow * 2 : congestionWindow + let pacedWindow = + state.congestionWindow < state.slowStartThreshold ? state.congestionWindow * 2 : state.congestionWindow let rateInBytesPerSecond = pacedWindow * System.Time.USEC_PER_SEC / UInt64(smoothedRTTInMicroseconds) let burstSize = rateInBytesPerSecond >> burstQueueShift @@ -321,22 +309,23 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { path.pacer.setBurstSize(burstSize: UInt32(truncatingIfNeeded: burstSize)) } - private func packetInRecovery(sentTime: NetworkClock.Instant) -> Bool { - sentTime <= recoveryStartTime - } - - mutating func enterRecovery(mss: Int, now: NetworkClock.Instant, qlog: QLog? = nil) { - log.datapath("Entering Recovery: current cwin=\(congestionWindow)") - recoveryStartTime = now + mutating func enterRecovery( + state: inout CongestionControlState, + mss: Int, + now: NetworkClock.Instant, + qlog: QLog? = nil + ) { + log.datapath("Entering Recovery: current cwin=\(state.congestionWindow)") + state.recoveryStartTime = now cubicLastMaxCongestionWindow = cubicMaxCongestionWindow - cubicMaxCongestionWindow = congestionWindow + cubicMaxCongestionWindow = state.congestionWindow - congestionWindow = UInt64(Double(lossFlightSize) * Prague.beta) - if _slowPath(congestionWindow < Prague.minCongestionWindow(mss)) { - congestionWindow = Prague.minCongestionWindow(mss) + state.congestionWindow = UInt64(Double(state.lossFlightSize(log: log)) * Prague.beta) + if _slowPath(state.congestionWindow < Prague.minCongestionWindow(mss)) { + state.congestionWindow = Prague.minCongestionWindow(mss) } - prevSlowStartThreshold = slowStartThreshold - slowStartThreshold = congestionWindow + state.prevSlowStartThreshold = state.slowStartThreshold + state.slowStartThreshold = state.congestionWindow // If Fast Convergence is supported, release more bandwidth // if saturation point is getting reduced due to new flows @@ -353,23 +342,24 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { // Compute epoch period K(s) that the window will take to increase // to last_max again after backoff due to loss. // Note that K = 0 if we enter CA without loss. - setCubicK(mss: mss) + setCubicK(state: state, mss: mss) // Set the start of current CA and the origin point cubicEpochStart = now cubicOriginPoint = cubicMaxCongestionWindow // Reset renoCongestionWindow to be in sync with Prague - renoCongestionWindow = congestionWindow + renoCongestionWindow = state.congestionWindow renoAcked = 0 numCongestionEventsLoss += 1 reducedDueToCE = false - initPipeAckSamples() - logUpdate(qlog: qlog) + state.initPipeAckSamples() + state.logUpdate(log: log, qlog: qlog) logState(qlog: qlog, state: .recovery, trigger: nil) } /// Enters CWR for one RTT after receiving an ACK with new CE counts. private mutating func pragueCWR( + state: inout CongestionControlState, largestAckedSentTime: NetworkClock.Instant, mss: Int, qlog: QLog? = nil @@ -377,7 +367,7 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { numCongestionEventsCE += 1 // If the packet was sent before recovery started, do nothing - if packetInRecovery(sentTime: largestAckedSentTime) { + if state.packetInRecovery(sentTime: largestAckedSentTime) { return } @@ -388,29 +378,30 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { // increase cwnd during ack_end, even in CWR state. // // On entering CWR, cwnd = cwnd * (1 - DCTCP.alpha) / 2 - let reduction = (congestionWindow * alpha) >> (Prague.alphaShift + 1) - congestionWindow -= reduction + let reduction = (state.congestionWindow * alpha) >> (Prague.alphaShift + 1) + state.congestionWindow -= reduction // Should be at least 2 MSS - if _slowPath(congestionWindow < Prague.minCongestionWindow(mss)) { - congestionWindow = Prague.minCongestionWindow(mss) + if _slowPath(state.congestionWindow < Prague.minCongestionWindow(mss)) { + state.congestionWindow = Prague.minCongestionWindow(mss) } - slowStartThreshold = congestionWindow + state.slowStartThreshold = state.congestionWindow reducedDueToCE = true - logUpdate(qlog: qlog) + state.logUpdate(log: log, qlog: qlog) logState(qlog: qlog, state: .cwr, trigger: nil) } /// Updates alpha after receiving acknowledgments. private mutating func pragueUpdateAlpha( + state: inout CongestionControlState, largestSentPN: Int64, largestAckedPN: Int64, packetsMarked: UInt64, packetsAcked: UInt64 ) { - if !rttElapsed(largestSentPN: largestSentPNForAlpha, largestAckedPN: largestAckedPN) { + if !state.rttElapsed(largestSentPN: largestSentPNForAlpha, largestAckedPN: largestAckedPN) { // One RTT hasn't elapsed yet, don't update alpha log.datapath("One RTT hasn't elapsed, not updating alpha") return @@ -419,12 +410,12 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { var newlyMarked: UInt64 = 0 var newlyAcked: UInt64 = 0 - if packetsMarked > self.packetsMarked { - newlyMarked = packetsMarked - self.packetsMarked + if packetsMarked > state.packetsMarked { + newlyMarked = packetsMarked - state.packetsMarked } - if packetsAcked > self.packetsAcked { - newlyAcked = packetsAcked - self.packetsAcked + if packetsAcked > state.packetsAcked { + newlyAcked = packetsAcked - state.packetsAcked } else { log.error("No new packets were ACK'ed, we shouldn't be called") } @@ -456,27 +447,29 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { // New round for alpha largestSentPNForAlpha = largestSentPN - self.packetsMarked = packetsMarked - self.packetsAcked = packetsAcked + state.packetsMarked = packetsMarked + state.packetsAcked = packetsAcked } private mutating func pragueCongestionEvent( + state: inout CongestionControlState, sentTime: NetworkClock.Instant, mss: Int, now: NetworkClock.Instant, qlog: QLog? = nil ) -> Bool { // If the packet was sent before recovery started, do nothing - if packetInRecovery(sentTime: sentTime) { + if state.packetInRecovery(sentTime: sentTime) { return false } - enterRecovery(mss: mss, now: now, qlog: qlog) + enterRecovery(state: &state, mss: mss, now: now, qlog: qlog) return true } @discardableResult mutating func packetLost( + state: inout CongestionControlState, path: QUICPath?, bytesLost: Int, largestLostSentTime: NetworkClock.Instant, @@ -485,18 +478,20 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { now: NetworkClock.Instant, qlog: QLog? = nil ) -> Bool { - decrementBytesInFlight(UInt64(bytesLost)) + state.decrementBytesInFlight(UInt64(bytesLost), log: log) let reducedCongestionWindow = pragueCongestionEvent( + state: &state, sentTime: largestLostSentTime, mss: mss, now: now, qlog: qlog ) - updatePacerState(path: path, smoothedRTT: smoothedRTT) + updatePacerState(state: state, path: path, smoothedRTT: smoothedRTT) return reducedCongestionWindow } mutating func ackEnd( + state: inout CongestionControlState, rtt: borrowing RTT, path: QUICPath?, mss: Int, @@ -510,36 +505,43 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { return } - if bytesAcked == 0 { + if state.bytesAcked == 0 { // When we are in recovery period or received new CE counts return } let smoothedRTT = rtt.smoothedRTT - if !revalidateCongestionWindow(smoothedRTT: smoothedRTT, now: now) { - bytesAcked = 0 + if !state.revalidateCongestionWindow(smoothedRTT: smoothedRTT, now: now, log: log) { + state.bytesAcked = 0 return } - if congestionWindow < slowStartThreshold { - congestionWindow += min(bytesAcked, Prague.slowStartCongestionWindow(mss)) + if state.congestionWindow < state.slowStartThreshold { + state.congestionWindow += min(state.bytesAcked, Prague.slowStartCongestionWindow(mss)) } else { if reducedDueToCE { - pragueCAAfterCE(bytesAcked: bytesAcked, mss: mss) + pragueCAAfterCE(state: &state, bytesAcked: state.bytesAcked, mss: mss) } else { - cubicProcessAckCA(bytesAcked: bytesAcked, smoothedRTT: smoothedRTT, mss: mss, now: now) + cubicProcessAckCA( + state: &state, + bytesAcked: state.bytesAcked, + smoothedRTT: smoothedRTT, + mss: mss, + now: now + ) } } // Should be a minimum of 2*MSS - if _slowPath(congestionWindow < Prague.minCongestionWindow(mss)) { - congestionWindow = Prague.minCongestionWindow(mss) + if _slowPath(state.congestionWindow < Prague.minCongestionWindow(mss)) { + state.congestionWindow = Prague.minCongestionWindow(mss) } - updatePacerState(path: path, smoothedRTT: smoothedRTT) - logUpdate(qlog: qlog) + updatePacerState(state: state, path: path, smoothedRTT: smoothedRTT) + state.logUpdate(log: log, qlog: qlog) } mutating func processECN( + state: inout CongestionControlState, path: QUICPath?, ceCount: Int, packetsAcked: Int, @@ -551,16 +553,17 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { now: NetworkClock.Instant, qlog: QLog? = nil ) { - if _slowPath(ceCount < ecnCECounter) { + if _slowPath(ceCount < state.ecnCECounter) { log.fault( - "New CE count \(ceCount) can't be less than current CE count \(ecnCECounter)" + "New CE count \(ceCount) can't be less than current CE count \(state.ecnCECounter)" ) } // Update alpha of fraction of marked packets, // even when there are no new CE counts - if packetsAcked > Int(self.packetsAcked) { + if packetsAcked > Int(state.packetsAcked) { pragueUpdateAlpha( + state: &state, largestSentPN: largestSentPN, largestAckedPN: largestAckedPN, packetsMarked: UInt64(ceCount), @@ -568,27 +571,27 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { ) } - if ceCount == ecnCECounter { + if ceCount == state.ecnCECounter { // No change in CE return } log.datapath( - "\(bytesAcked) bytes were ACKed with \(ceCount - ecnCECounter) packets newly CE marked" + "\(state.bytesAcked) bytes were ACKed with \(ceCount - state.ecnCECounter) packets newly CE marked" ) // Received an ACK with new CE counts, subtract CE marked bytes // from bytes_acked, so that we use only unmarked bytes to // increase cwnd during ack_end - let ceBytes = UInt64(ceCount - ecnCECounter) * UInt64(mss) - if bytesAcked > ceBytes { - bytesAcked -= ceBytes + let ceBytes = UInt64(ceCount - state.ecnCECounter) * UInt64(mss) + if state.bytesAcked > ceBytes { + state.bytesAcked -= ceBytes } else { - bytesAcked = 0 + state.bytesAcked = 0 } // Update CE count even if we are already in CWR - ecnCECounter = ceCount + state.ecnCECounter = ceCount // Update AIMD alpha as SRTT might have changed if rttControl == .rateEquivalence { @@ -597,43 +600,43 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { pragueAIAlphaScalable(sRTT: smoothedRTT) } - if !rttElapsed(largestSentPN: self.largestSentPN, largestAckedPN: largestAckedPN) { + if !state.rttElapsed(largestSentPN: state.largestSentPN, largestAckedPN: largestAckedPN) { // Haven't elapsed one RTT yet from last CWR log.datapath("Haven't elapsed one RTT yet from last CWR") return } // Enter or stay in CWR if new counts are received - pragueCWR(largestAckedSentTime: largestAckedSentTime, mss: mss, qlog: qlog) + pragueCWR(state: &state, largestAckedSentTime: largestAckedSentTime, mss: mss, qlog: qlog) // Update pacer state as cwnd has changed - updatePacerState(path: path, smoothedRTT: smoothedRTT) + updatePacerState(state: state, path: path, smoothedRTT: smoothedRTT) // Start new round for CWR - self.largestSentPN = largestSentPN + state.largestSentPN = largestSentPN } - mutating func spuriousRetransmit(qlog: QLog? = nil) { - guard cubicMaxCongestionWindow > 0 && prevSlowStartThreshold > 0 else { return } + mutating func spuriousRetransmit(state: inout CongestionControlState, qlog: QLog? = nil) { + guard cubicMaxCongestionWindow > 0 && state.prevSlowStartThreshold > 0 else { return } // Revert to the state before loss was detected - congestionWindow = max(cubicMaxCongestionWindow, congestionWindow) - slowStartThreshold = prevSlowStartThreshold - logUpdate(qlog: qlog) + state.congestionWindow = max(cubicMaxCongestionWindow, state.congestionWindow) + state.slowStartThreshold = state.prevSlowStartThreshold + state.logUpdate(log: log, qlog: qlog) } - mutating func persistentCongestion(mss: Int, qlog: QLog? = nil) { + mutating func persistentCongestion(state: inout CongestionControlState, mss: Int, qlog: QLog? = nil) { // Set the minimum congestion window let newCWND = Prague.minCongestionWindow(mss) - slowStartThreshold = max(UInt64(Double(congestionWindow) * Prague.beta), newCWND) - congestionWindow = newCWND - logUpdate(qlog: qlog) + state.slowStartThreshold = max(UInt64(Double(state.congestionWindow) * Prague.beta), newCWND) + state.congestionWindow = newCWND + state.logUpdate(log: log, qlog: qlog) logState(qlog: qlog, state: .slowStart, trigger: .persistentCongestion) } - private mutating func resetInternal() { - recoveryStartTime = .zero - prevSlowStartThreshold = 0 + private mutating func resetInternal(state: inout CongestionControlState) { + state.recoveryStartTime = .zero + state.prevSlowStartThreshold = 0 numCongestionEventsLoss = 0 numCongestionEventsCE = 0 @@ -643,71 +646,59 @@ struct Prague: CongestionControlProtocol, CubicLikeProtocol { cubicEpochStart = .zero cubicOriginPoint = 0 cubicLastMaxCongestionWindow = 0 - cubicMaxCongestionWindow = congestionWindow + cubicMaxCongestionWindow = state.congestionWindow // Prague state rttControl = .rateEquivalence alphaAI = 1 << Prague.congestionWindowShift // CWV state - pipeAckSampleEnd = .zero - initPipeAckSamples() + state.pipeAckSampleEnd = .zero + state.initPipeAckSamples() } - mutating func idleTimeout(mss: Int, qlog: QLog? = nil) { + mutating func idleTimeout(state: inout CongestionControlState, mss: Int, qlog: QLog? = nil) { // We want to ideally begin with slow start after idle period. // Set it to the larger of its current value, MAX (cwnd * Beta, IW) - slowStartThreshold = max( - slowStartThreshold, - max(UInt64(Double(congestionWindow) * Prague.beta), Prague.initialCongestionWindow(mss)) + state.slowStartThreshold = max( + state.slowStartThreshold, + max(UInt64(Double(state.congestionWindow) * Prague.beta), Prague.initialCongestionWindow(mss)) ) // Set cwnd to initial cwnd - congestionWindow = min(congestionWindow, Prague.initialCongestionWindow(mss)) - logUpdate(qlog: qlog) - resetInternal() + state.congestionWindow = min(state.congestionWindow, Prague.initialCongestionWindow(mss)) + state.logUpdate(log: log, qlog: qlog) + resetInternal(state: &state) } - func filloutDataTransferSnapshot(dataTransferSnapshot: inout DataTransferSnapshot) { - dataTransferSnapshot.transportCongestionWindow = congestionWindow - dataTransferSnapshot.transportSlowStartThreshold = slowStartThreshold + func filloutDataTransferSnapshot(state: CongestionControlState, dataTransferSnapshot: inout DataTransferSnapshot) { + dataTransferSnapshot.transportCongestionWindow = state.congestionWindow + dataTransferSnapshot.transportSlowStartThreshold = state.slowStartThreshold } - mutating func reset(mss: Int, qlog: QLog? = nil) { - congestionWindow = Prague.initialCongestionWindow(mss) - slowStartThreshold = UInt64.max + mutating func reset(state: inout CongestionControlState, mss: Int, qlog: QLog? = nil) { + state.congestionWindow = Prague.initialCongestionWindow(mss) + state.slowStartThreshold = UInt64.max scaledAlpha = Prague.maxAlpha << Prague.gShift - resetInternal() - logUpdate(qlog: qlog) + resetInternal(state: &state) + state.logUpdate(log: log, qlog: qlog) } - mutating func inherit(from: CongestionControl, mss: Int, qlog: QLog?) { + mutating func inherit( + from: CongestionControlState, + state: inout CongestionControlState, + mss: Int, + qlog: QLog? = nil + ) { // 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 - } - slowStartThreshold = UInt64.max + state.bytesInFlight = from.bytesInFlight + state.congestionWindow = max(from.congestionWindow, Prague.initialCongestionWindow(mss)) + state.slowStartThreshold = UInt64.max scaledAlpha = Prague.maxAlpha << Prague.gShift - resetInternal() - logUpdate(qlog: qlog) + resetInternal(state: &state) + state.logUpdate(log: log, qlog: qlog) } } #endif diff --git a/Sources/SwiftNetwork/QUIC/QUICPath.swift b/Sources/SwiftNetwork/QUIC/QUICPath.swift index 021c2c2..4c6a3da 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 @@ -337,6 +337,14 @@ public final class QUICPath: MultiplexingDatagramPath< required init(parent: QUICConnection, in eventContext: inout NetworkContext.EventContext) { self.rtt = RTT(logPrefixer: parent.logPrefixer) self.pacer = Pacer() + // Overwritten with the real mss/qlog once `setup()` runs; RX/TX aren't allowed on a + // path until then, so this placeholder is never observed. + self.congestionControl = CongestionControl.createCubic( + pacer: &self.pacer, + mss: 0, + qlog: nil, + logPrefixer: parent.logPrefixer + ) super.init(parent: parent, in: &eventContext) } @@ -371,13 +379,11 @@ 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 = .createCubic( + pacer: &self.pacer, + mss: self.initialMSS, + qlog: parentProtocol.qLog, + logPrefixer: self.log ) self.spinValue = parentProtocol.initialSpinValue @@ -438,42 +444,34 @@ public final class QUICPath: MultiplexingDatagramPath< } func resetCongestionControl() { - switch self.congestionControl { + switch self.congestionControl.algorithm { case .cubic: - self.congestionControl = .cubic( - algorithm: Cubic( - pacer: &self.pacer, - mss: self.initialMSS, - qlog: parentProtocol.qLog, - logPrefixer: self.log - ) + self.congestionControl = .createCubic( + 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 - ) + self.congestionControl = .createLedbat( + 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 - ) + self.congestionControl = .createPrague( + pacer: &self.pacer, + mss: self.initialMSS, + qlog: parentProtocol.qLog, + logPrefixer: self.log ) #endif - case .none: - break } } func idleTimeoutCongestionControl() { - self.congestionControl?.idleTimeout(mss: mss) + self.congestionControl.idleTimeout(mss: mss) } func setupL4SState(l4sEnabled: Bool?) { @@ -494,50 +492,69 @@ public final class QUICPath: MultiplexingDatagramPath< func markAsBackground(_ background: Bool) { #if !NETWORK_EMBEDDED // Use LEDBAT for background cases - switch self.congestionControl { + switch self.congestionControl.algorithm { case .cubic: if !background { return } // Nothing to do, already not background + let oldState = self.congestionControl.state + var newState = CongestionControlState() var ledbat = Ledbat( + state: &newState, mss: self.initialMSS, qlog: parentProtocol.qLog, logPrefixer: self.log ) ledbat.inherit( - from: self.congestionControl!, + from: oldState, + state: &newState, mss: self.initialMSS, qlog: parentProtocol.qLog ) - self.congestionControl = .ledbat(algorithm: ledbat) + self.congestionControl = CongestionControl( + state: newState, + algorithm: .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. + let oldState = self.congestionControl.state + var newState = CongestionControlState() var cubic = Cubic( + state: &newState, pacer: &self.pacer, mss: self.initialMSS, qlog: parentProtocol.qLog, logPrefixer: self.log ) cubic.inherit( - from: self.congestionControl!, + from: oldState, + state: &newState, mss: self.initialMSS, qlog: parentProtocol.qLog ) - self.congestionControl = .cubic(algorithm: cubic) + self.congestionControl = CongestionControl( + state: newState, + algorithm: .cubic(algorithm: cubic) + ) case .prague: if !background { return } // Nothing to do, already not background + let oldState = self.congestionControl.state + var newState = CongestionControlState() var ledbat = Ledbat( + state: &newState, mss: self.initialMSS, qlog: parentProtocol.qLog, logPrefixer: self.log ) ledbat.inherit( - from: self.congestionControl!, + from: oldState, + state: &newState, mss: self.initialMSS, qlog: parentProtocol.qLog ) - self.congestionControl = .ledbat(algorithm: ledbat) - case .none: - break + self.congestionControl = CongestionControl( + state: newState, + algorithm: .ledbat(algorithm: ledbat) + ) } #endif } @@ -718,22 +735,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) @@ -744,7 +761,7 @@ extension QUICPath { packetsLost: Bool, qlog: QLog? = nil ) { - congestionControl?.ackEnd( + congestionControl.ackEnd( rtt: rtt, path: self, mss: mss, @@ -756,12 +773,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) @@ -773,49 +790,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) @@ -831,7 +848,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, @@ -847,7 +864,7 @@ extension QUICPath { @inline(always) func congestionControlFilloutDataTransferSnapshot(snapshot: inout DataTransferSnapshot) { - congestionControl?.filloutDataTransferSnapshot(dataTransferSnapshot: &snapshot) + congestionControl.filloutDataTransferSnapshot(dataTransferSnapshot: &snapshot) } } diff --git a/Sources/SwiftNetwork/QUIC/Recovery.swift b/Sources/SwiftNetwork/QUIC/Recovery.swift index 6dea9ad..dfe8091 100644 --- a/Sources/SwiftNetwork/QUIC/Recovery.swift +++ b/Sources/SwiftNetwork/QUIC/Recovery.swift @@ -129,7 +129,7 @@ struct Recovery: ~Copyable, PrefixedLoggable, NonCopyableTimerUser { if frontNumber == target { return 0 } if frontNumber > target { return nil } - var right = entryCount - 1 + var right = entryCount &- 1 let backNumber = outstandingPackets[right].packet.number.value if backNumber == target { return right } if backNumber < target { return nil } @@ -137,24 +137,24 @@ struct Recovery: ~Copyable, PrefixedLoggable, NonCopyableTimerUser { // Both ends are strictly inside the range now, so narrow the window // using the density bounds described above. var left = 1 - let fromFront = target - frontNumber + let fromFront = target &- frontNumber if fromFront < Int64(right) { right = Int(fromFront) } - let fromBack = backNumber - target - if fromBack < Int64(entryCount - 1 - left) { - left = entryCount - 1 - Int(fromBack) + let fromBack = backNumber &- target + if fromBack < Int64(entryCount &- 1 &- left) { + left = entryCount &- 1 &- Int(fromBack) } while left <= right { - let middle = left + (right - left) / 2 + let middle = left &+ (right &- left) / 2 let middleNumber = outstandingPackets[middle].packet.number.value if middleNumber == target { return middle } else if middleNumber < target { - left = middle + 1 + left = middle &+ 1 } else { - right = middle - 1 + right = middle &- 1 } } return nil @@ -172,7 +172,7 @@ struct Recovery: ~Copyable, PrefixedLoggable, NonCopyableTimerUser { ? outstandingPackets.removeFirst() : outstandingPackets.remove(at: index) if removedEntry.packet.largerPacket, largerPacketCount > 0 { - largerPacketCount -= 1 + largerPacketCount &-= 1 } return removedEntry } @@ -543,10 +543,14 @@ struct Recovery: ~Copyable, PrefixedLoggable, NonCopyableTimerUser { newlyECTAcked += 1 } - guard let sentPath = connection.path(for: ackedEntry.packet.sentPath) else { + let sentPath: QUICPath + if ackedEntry.packet.sentPath == path.pathIdentifier { + sentPath = path + } else if let lookedUpPath = connection.path(for: ackedEntry.packet.sentPath) { + sentPath = lookedUpPath + } else { continue } - if ackedEntry.lostTime != .zero { let sRTT = sentPath.rtt.smoothedRTT let latestRTT = sentPath.rtt.latestRTT diff --git a/Tests/QUICTests/CubicTests.swift b/Tests/QUICTests/CubicTests.swift index e36641a..2bc9f1d 100644 --- a/Tests/QUICTests/CubicTests.swift +++ b/Tests/QUICTests/CubicTests.swift @@ -32,6 +32,7 @@ final class CubicTests: XCTestCase { var rtt: RTT! let mss = Constants.initialMSS var cubic: Cubic! + var state = CongestionControlState() // These tests drive the algorithm directly, with no path to pace. let noPath: QUICPath? = nil var pacer: Pacer = Pacer(enabled: true) @@ -39,41 +40,43 @@ final class CubicTests: XCTestCase { override func setUp() { let logPrefixer = LogPrefixer("[CubicTests]") - cubic = Cubic(pacer: &pacer, mss: mss, logPrefixer: logPrefixer) + state = CongestionControlState() + cubic = Cubic(state: &state, pacer: &pacer, mss: mss, logPrefixer: logPrefixer) rtt = RTT(logPrefixer: logPrefixer) } func testCubicMSS() { // Test MSS > congestion window - XCTAssertEqual(cubic.availableCongestionWindow, defaultCongestionWindow) - cubic.mssChanged(mss: 65000) - XCTAssertEqual(cubic.availableCongestionWindow, 65000) - cubic.reset(mss: Constants.initialMSS) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) + cubic.mssChanged(state: &state, mss: 65000) + XCTAssertEqual(state.availableCongestionWindow, 65000) + cubic.reset(state: &state, mss: Constants.initialMSS) // Test MSS < congestion window - XCTAssertEqual(cubic.availableCongestionWindow, defaultCongestionWindow) - cubic.mssChanged(mss: 10) - XCTAssertEqual(cubic.availableCongestionWindow, defaultCongestionWindow) - cubic.reset(mss: Constants.initialMSS) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) + cubic.mssChanged(state: &state, mss: 10) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) + cubic.reset(state: &state, mss: Constants.initialMSS) } func testCubicReset() { - XCTAssertEqual(cubic.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) let time = NetworkClock.Instant.testBase - cubic.packetSent(bytesSent: 1000) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(cubic.availableCongestionWindow, 13000) - cubic.reset(mss: Constants.initialMSS) - XCTAssertEqual(cubic.availableCongestionWindow, defaultCongestionWindow) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 13000) + cubic.reset(state: &state, mss: Constants.initialMSS) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) } func testCubicLostPackets() { - XCTAssertEqual(cubic.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) // "Send" some packets and declare them lost let time = NetworkClock.Instant.testBase - cubic.packetSent(bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) cubic.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -81,29 +84,30 @@ final class CubicTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - XCTAssertEqual(cubic.availableCongestionWindow, 8400) + XCTAssertEqual(state.availableCongestionWindow, 8400) // See if we can send another packet - XCTAssertTrue(cubic.canSend(packetLength: 1000)) + XCTAssertTrue(cubic.canSend(state: state, packetLength: 1000)) } func testCubicSlowStart() { rtt.smoothedRTT = .microseconds(0) // "Send" some packets and declare one of them lost var time = NetworkClock.Instant.testBase - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) cubic.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -111,35 +115,36 @@ final class CubicTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - XCTAssertEqual(cubic.availableCongestionWindow, 11900) + XCTAssertEqual(state.availableCongestionWindow, 11900) // Make sure that another successful packet doesn't cause us to continue slow start time = NetworkClock.Instant.testBase.advanced(by: .microseconds(100)) - cubic.packetSent(bytesSent: 1000) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(cubic.availableCongestionWindow, 11953) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 11953) } func testCubicECN() { rtt.smoothedRTT = .microseconds(0) - XCTAssertEqual(cubic.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) // Test that CE counts will reduce the congestion window immediately and move CUBIC to Congestion avoidance var time = NetworkClock.Instant.testBase - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) cubic.processECN( + state: &state, path: noPath, ceCount: 1, packetsAcked: 6, @@ -150,34 +155,35 @@ final class CubicTests: XCTestCase { smoothedRTT: rtt.smoothedRTT, now: time ) - cubic.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(cubic.availableCongestionWindow, 8400) + cubic.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 8400) time = NetworkClock.Instant.testBase.advanced(by: .microseconds(100)) - cubic.packetSent(bytesSent: 1000) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) // congestion window grows during congestion avoidance - XCTAssertEqual(cubic.availableCongestionWindow, 8475) + XCTAssertEqual(state.availableCongestionWindow, 8475) } func testCubicECNEnterCWR() { rtt.smoothedRTT = .microseconds(0) - XCTAssertEqual(cubic.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) // Test that CE counts will reduce congestion window, enter congestion window recovery and after that we don't decrease congestion window for 1RTT even we receive new CE counts var time = NetworkClock.Instant.testBase - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) cubic.processECN( + state: &state, path: noPath, ceCount: 1, packetsAcked: 4, @@ -189,12 +195,13 @@ final class CubicTests: XCTestCase { now: time ) // availableCongestionWindow = congestionWindow - bytesInFlight = 8400 - 2000 = 6400 - XCTAssertEqual(cubic.availableCongestionWindow, 6400) + XCTAssertEqual(state.availableCongestionWindow, 6400) time = NetworkClock.Instant.testBase.advanced(by: .microseconds(100)) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) cubic.processECN( + state: &state, path: noPath, ceCount: 2, packetsAcked: 6, @@ -205,9 +212,9 @@ final class CubicTests: XCTestCase { smoothedRTT: rtt.smoothedRTT, now: time ) - cubic.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + cubic.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) // congestion window is the same 8400, bytes in flight has reduced to 0 - XCTAssertEqual(cubic.availableCongestionWindow, 8400) + XCTAssertEqual(state.availableCongestionWindow, 8400) } // CE marks reported through the path must reach the path's congestion controller, @@ -260,19 +267,20 @@ final class CubicTests: XCTestCase { rtt.smoothedRTT = .microseconds(0) // "Send" some packets and declare one of them lost var time = NetworkClock.Instant.testBase - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) cubic.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -280,45 +288,45 @@ final class CubicTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - cubic.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: true, now: time) - XCTAssertEqual(cubic.availableCongestionWindow, 8400) + cubic.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: true, now: time) + XCTAssertEqual(state.availableCongestionWindow, 8400) time = NetworkClock.Instant.testBase.advanced(by: .microseconds(100)) - cubic.packetSent(bytesSent: 1000) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(cubic.availableCongestionWindow, 8475) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 8475) } func testCubicIdleTimeout() { - XCTAssertEqual(cubic.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) let time = NetworkClock.Instant.testBase - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(cubic.availableCongestionWindow, 18000) - cubic.idleTimeout(mss: mss) - XCTAssertEqual(cubic.availableCongestionWindow, defaultCongestionWindow) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 18000) + cubic.idleTimeout(state: &state, mss: mss) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) } func testCubicPersistentCongestion() { - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.persistentCongestion(mss: mss) - XCTAssertEqual(cubic.availableCongestionWindow, 0) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.persistentCongestion(state: &state, mss: mss) + XCTAssertEqual(state.availableCongestionWindow, 0) } func testCubicCongestionLimited() { @@ -331,19 +339,20 @@ final class CubicTests: XCTestCase { // therefore give three reductions. var sentTime = NetworkClock.Instant.testBase var detectedAt = sentTime.advanced(by: .microseconds(100)) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1000, sentTime: sentTime) - cubic.packetsAcked(bytesAcked: 1000, sentTime: sentTime) - cubic.packetsAcked(bytesAcked: 1000, sentTime: sentTime) - cubic.packetsAcked(bytesAcked: 1000, sentTime: sentTime) - cubic.packetsAcked(bytesAcked: 1000, sentTime: sentTime) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: sentTime) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: sentTime) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: sentTime) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: sentTime) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: sentTime) cubic.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -351,14 +360,15 @@ final class CubicTests: XCTestCase { smoothedRTT: .microseconds(0), now: detectedAt ) - XCTAssertEqual(cubic.availableCongestionWindow, 8400) + XCTAssertEqual(state.availableCongestionWindow, 8400) sentTime = sentTime.advanced(by: .microseconds(1000)) detectedAt = sentTime.advanced(by: .microseconds(100)) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) cubic.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -367,6 +377,7 @@ final class CubicTests: XCTestCase { now: detectedAt ) cubic.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -375,6 +386,7 @@ final class CubicTests: XCTestCase { now: detectedAt ) cubic.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -383,6 +395,7 @@ final class CubicTests: XCTestCase { now: detectedAt ) cubic.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -392,9 +405,10 @@ final class CubicTests: XCTestCase { ) sentTime = sentTime.advanced(by: .microseconds(1000)) detectedAt = sentTime.advanced(by: .microseconds(100)) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) cubic.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -403,6 +417,7 @@ final class CubicTests: XCTestCase { now: detectedAt ) cubic.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -413,23 +428,24 @@ final class CubicTests: XCTestCase { // One reduction per round, three rounds: 12000 -> 8400 -> 5880 -> 4116, each step // `UInt64(Double(window) * Cubic.beta)`. Written out rather than as `pow(beta, 3)`, which // is 0.34299999999999997 and truncates to 4115; the reductions compound one at a time. - XCTAssertEqual(cubic.availableCongestionWindow, 4116) - XCTAssertFalse(cubic.canSend(packetLength: 10000)) + XCTAssertEqual(state.availableCongestionWindow, 4116) + XCTAssertFalse(cubic.canSend(state: state, packetLength: 10000)) } func testCubicPacketDiscard() { - cubic.packetSent(bytesSent: 1000) - cubic.packetDiscarded(bytesSent: 1000) - XCTAssertEqual(cubic.bytesInFlight, 0) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetDiscarded(state: &state, bytesSent: 1000) + XCTAssertEqual(state.bytesInFlight, 0) } func testCubicSpuriousRetransmit() { let time = NetworkClock.Instant.testBase - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) cubic.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -437,20 +453,21 @@ final class CubicTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - cubic.spuriousRetransmit() - XCTAssertEqual(cubic.availableCongestionWindow, 9000) + cubic.spuriousRetransmit(state: &state) + XCTAssertEqual(state.availableCongestionWindow, 9000) } // Tests that we can enter CA without any loss after idle period func testCubicCongestionAvoidance() { let time = NetworkClock.Instant.testBase - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) cubic.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -458,29 +475,29 @@ final class CubicTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - cubic.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: true, now: time) - XCTAssertEqual(cubic.availableCongestionWindow, 8400) - cubic.idleTimeout(mss: mss) - XCTAssertEqual(cubic.availableCongestionWindow, 8400) - cubic.packetSent(bytesSent: 1200) - cubic.packetSent(bytesSent: 1200) - cubic.packetSent(bytesSent: 1200) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1200, sentTime: time) - cubic.packetsAcked(bytesAcked: 1200, sentTime: time) - cubic.packetsAcked(bytesAcked: 1200, sentTime: time) - cubic.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + cubic.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: true, now: time) + XCTAssertEqual(state.availableCongestionWindow, 8400) + cubic.idleTimeout(state: &state, mss: mss) + XCTAssertEqual(state.availableCongestionWindow, 8400) + cubic.packetSent(state: &state, bytesSent: 1200) + cubic.packetSent(state: &state, bytesSent: 1200) + cubic.packetSent(state: &state, bytesSent: 1200) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1200, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1200, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1200, sentTime: time) + cubic.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) // Enter CA - XCTAssertEqual(cubic.availableCongestionWindow, 12000) + XCTAssertEqual(state.availableCongestionWindow, 12000) for _ in 0..<12 { - cubic.packetSent(bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) } - cubic.ackBegin() + cubic.ackBegin(state: &state) for _ in 0..<12 { - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) } - cubic.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(cubic.availableCongestionWindow, 13200) + cubic.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 13200) } @@ -488,20 +505,20 @@ final class CubicTests: XCTestCase { var dataTransferSnapshot = DataTransferSnapshot() XCTAssertEqual(dataTransferSnapshot.transportCongestionWindow, 0) XCTAssertEqual(dataTransferSnapshot.transportSlowStartThreshold, 0) - cubic.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) + cubic.filloutDataTransferSnapshot(state: state, dataTransferSnapshot: &dataTransferSnapshot) XCTAssertTrue(dataTransferSnapshot.transportCongestionWindow > 0) XCTAssertTrue(dataTransferSnapshot.transportSlowStartThreshold > 0) let existingCongestionWindow = dataTransferSnapshot.transportCongestionWindow let time = NetworkClock.Instant.testBase - cubic.packetSent(bytesSent: 1000) - cubic.packetSent(bytesSent: 1000) - cubic.ackBegin() - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.packetsAcked(bytesAcked: 1000, sentTime: time) - cubic.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - - cubic.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.packetSent(state: &state, bytesSent: 1000) + cubic.ackBegin(state: &state) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + cubic.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + + cubic.filloutDataTransferSnapshot(state: state, dataTransferSnapshot: &dataTransferSnapshot) XCTAssertEqual( dataTransferSnapshot.transportCongestionWindow, (existingCongestionWindow + 2000) diff --git a/Tests/QUICTests/LedbatTests.swift b/Tests/QUICTests/LedbatTests.swift index 6f8c45f..28a5087 100644 --- a/Tests/QUICTests/LedbatTests.swift +++ b/Tests/QUICTests/LedbatTests.swift @@ -31,68 +31,71 @@ final class LedbatTests: XCTestCase { var rtt: RTT! let mss = Constants.initialMSS var ledbat: Ledbat! + var state = CongestionControlState() // These tests drive the algorithm directly, with no path to pace. let noPath: QUICPath? = nil let defaultCongestionWindow = UInt64(2400) override func setUp() { let logPrefixer = LogPrefixer("[LedbatTests]") - ledbat = Ledbat(mss: mss, logPrefixer: logPrefixer) + state = CongestionControlState() + ledbat = Ledbat(state: &state, mss: mss, logPrefixer: logPrefixer) rtt = RTT(logPrefixer: logPrefixer) rtt.baseRTT = .milliseconds(100) } func testLedbatMSS() { // Test MSS > congestion window - XCTAssertEqual(ledbat.availableCongestionWindow, defaultCongestionWindow) - ledbat.mssChanged(mss: 65000) - XCTAssertEqual(ledbat.availableCongestionWindow, 65000) - ledbat.reset(mss: Constants.initialMSS) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) + ledbat.mssChanged(state: &state, mss: 65000) + XCTAssertEqual(state.availableCongestionWindow, 65000) + ledbat.reset(state: &state, mss: Constants.initialMSS) // Test MSS < congestion window - XCTAssertEqual(ledbat.availableCongestionWindow, defaultCongestionWindow) - ledbat.mssChanged(mss: 10) - XCTAssertEqual(ledbat.availableCongestionWindow, defaultCongestionWindow) - ledbat.reset(mss: Constants.initialMSS) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) + ledbat.mssChanged(state: &state, mss: 10) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) + ledbat.reset(state: &state, mss: Constants.initialMSS) } func testLedbatReset() { - XCTAssertEqual(ledbat.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) // SRTT = 100ms, Current RTT = 120ms rtt.adjustedRTT = .milliseconds(120) rtt.smoothedRTT = .milliseconds(100) let time = NetworkClock.Instant.testBase - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 2900) - ledbat.reset(mss: Constants.initialMSS) - XCTAssertEqual(ledbat.availableCongestionWindow, defaultCongestionWindow) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 2900) + ledbat.reset(state: &state, mss: Constants.initialMSS) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) } func testLedbatLostPackets() { - XCTAssertEqual(ledbat.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) // SRTT = 100ms, Current RTT = 120ms rtt.adjustedRTT = .milliseconds(120) rtt.smoothedRTT = .milliseconds(100) let time = NetworkClock.Instant.testBase // Send to increase cwnd - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 4400) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 4400) // Send a packet and declare them lost - ledbat.packetSent(bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) ledbat.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -100,37 +103,38 @@ final class LedbatTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - XCTAssertEqual(ledbat.availableCongestionWindow, 2400) + XCTAssertEqual(state.availableCongestionWindow, 2400) // See if we can send another packet - XCTAssertTrue(ledbat.canSend(packetLength: 1000)) + XCTAssertTrue(ledbat.canSend(state: state, packetLength: 1000)) } func testLedbatSlowStart() { - XCTAssertEqual(ledbat.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) // SRTT = 100ms, Current RTT = 120ms rtt.adjustedRTT = .milliseconds(120) rtt.smoothedRTT = .milliseconds(100) // Send some packets to increase cwnd var time = NetworkClock.Instant.testBase - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 4900) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 4900) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) ledbat.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -138,66 +142,67 @@ final class LedbatTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - XCTAssertEqual(ledbat.availableCongestionWindow, 2450) + XCTAssertEqual(state.availableCongestionWindow, 2450) // Additive increase during CA time = NetworkClock.Instant.testBase.advanced(by: .microseconds(100)) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 2939) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 2939) // Mulitplicative decrease during CA // Current RTT = 180ms rtt.adjustedRTT = .milliseconds(180) time = NetworkClock.Instant.testBase.advanced(by: .microseconds(100)) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 2606) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 2606) } func testLedbatECN() { // SRTT = 100ms, Current RTT = 120ms rtt.adjustedRTT = .milliseconds(120) rtt.smoothedRTT = .milliseconds(100) - XCTAssertEqual(ledbat.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) // Lets increase the window first to go higher than MIN_CWND var time = NetworkClock.Instant.testBase - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 5400) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 5400) time = NetworkClock.Instant.testBase.advanced(by: .microseconds(100)) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) ledbat.processECN( + state: &state, path: noPath, ceCount: 1, packetsAcked: 6, @@ -208,55 +213,56 @@ final class LedbatTests: XCTestCase { smoothedRTT: rtt.smoothedRTT, now: time ) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 2700) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 2700) time = NetworkClock.Instant.testBase.advanced(by: .microseconds(200)) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) // cwnd grows during congestion avoidance - XCTAssertEqual(ledbat.availableCongestionWindow, 2922) + XCTAssertEqual(state.availableCongestionWindow, 2922) } func testLedbatECNEnterCWR() { // SRTT = 100ms, base RTT = 100ms network RTT = 120ms rtt.adjustedRTT = .milliseconds(120) rtt.smoothedRTT = .milliseconds(100) - XCTAssertEqual(ledbat.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) // Lets increase the window first to go higher than MIN_CWND var time = NetworkClock.Instant.testBase - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 5400) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 5400) // Test that CE counts will reduce cwnd, enter CWR and after that we don't decrease cwnd for 1RTT even we receive new CE counts time = NetworkClock.Instant.testBase.advanced(by: .microseconds(100)) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) ledbat.processECN( + state: &state, path: noPath, ceCount: 1, packetsAcked: 4, @@ -267,15 +273,16 @@ final class LedbatTests: XCTestCase { smoothedRTT: rtt.smoothedRTT, now: time ) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) // allowed cwnd = cwnd - bytes_in_flight = 2700 - 2000 = 700 - XCTAssertEqual(ledbat.availableCongestionWindow, 700) + XCTAssertEqual(state.availableCongestionWindow, 700) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) ledbat.processECN( + state: &state, path: noPath, ceCount: 2, packetsAcked: 6, @@ -286,9 +293,9 @@ final class LedbatTests: XCTestCase { smoothedRTT: rtt.smoothedRTT, now: time ) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) // cwnd is same 2700, bytes in flight has reduced to 0 - XCTAssertEqual(ledbat.availableCongestionWindow, 2700) + XCTAssertEqual(state.availableCongestionWindow, 2700) } func testLedbatAckDuringRecovery() { @@ -297,19 +304,20 @@ final class LedbatTests: XCTestCase { rtt.smoothedRTT = .milliseconds(100) // "Send" some packets and declare one of them lost var time = NetworkClock.Instant.testBase - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) ledbat.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -317,51 +325,51 @@ final class LedbatTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: true, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 2400) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: true, now: time) + XCTAssertEqual(state.availableCongestionWindow, 2400) time = NetworkClock.Instant.testBase.advanced(by: .microseconds(100)) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 2650) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 2650) } func testLedbatIdleTimeout() { - XCTAssertEqual(ledbat.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) // SRTT = 100ms, base RTT = 100ms network RTT = 120ms rtt.adjustedRTT = .milliseconds(120) rtt.smoothedRTT = .milliseconds(100) let time = NetworkClock.Instant.testBase - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 5400) - ledbat.idleTimeout(mss: mss) - XCTAssertEqual(ledbat.availableCongestionWindow, defaultCongestionWindow) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 5400) + ledbat.idleTimeout(state: &state, mss: mss) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) } func testLedbatPersistentCongestion() { // SRTT = 100ms, base RTT = 100ms network RTT = 120ms rtt.adjustedRTT = .milliseconds(120) rtt.smoothedRTT = .milliseconds(100) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.persistentCongestion(mss: mss) - XCTAssertEqual(ledbat.availableCongestionWindow, 0) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.persistentCongestion(state: &state, mss: mss) + XCTAssertEqual(state.availableCongestionWindow, 0) } func testLedbatCongestionLimited() { @@ -384,22 +392,23 @@ final class LedbatTests: XCTestCase { // pipeack sample and `lossFlightSize` stays equal to the window. var sentTime = NetworkClock.Instant.testBase for _ in 0..<4 { - ledbat.packetSent(bytesSent: 12000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 12000, sentTime: sentTime) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: sentTime) + ledbat.packetSent(state: &state, bytesSent: 12000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 12000, sentTime: sentTime) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: sentTime) sentTime = sentTime.advanced(by: .milliseconds(1)) } - XCTAssertEqual(ledbat.availableCongestionWindow, 26400) + XCTAssertEqual(state.availableCongestionWindow, 26400) // Three rounds, four losses each, halving the window once per round. for expectedWindow in [UInt64(13200), 6600, 3300] { let detectedAt = sentTime.advanced(by: .microseconds(100)) for _ in 0..<4 { - ledbat.packetSent(bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) } for lossIndex in 0..<4 { let openedRecovery = ledbat.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -410,16 +419,16 @@ final class LedbatTests: XCTestCase { // Only the first loss opens a period; the rest were sent before it started. XCTAssertEqual(openedRecovery, lossIndex == 0) } - XCTAssertEqual(ledbat.availableCongestionWindow, expectedWindow) + XCTAssertEqual(state.availableCongestionWindow, expectedWindow) sentTime = sentTime.advanced(by: .milliseconds(1)) } - XCTAssertFalse(ledbat.canSend(packetLength: 4000)) + XCTAssertFalse(ledbat.canSend(state: state, packetLength: 4000)) } func testLedbatPacketDiscard() { - ledbat.packetSent(bytesSent: 1000) - ledbat.packetDiscarded(bytesSent: 1000) - XCTAssertEqual(ledbat.bytesInFlight, 0) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetDiscarded(state: &state, bytesSent: 1000) + XCTAssertEqual(state.bytesInFlight, 0) } func testLedbatSpuriousRetransmit() { @@ -427,13 +436,14 @@ final class LedbatTests: XCTestCase { rtt.adjustedRTT = .milliseconds(120) rtt.smoothedRTT = .milliseconds(100) let time = NetworkClock.Instant.testBase - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 2900) - ledbat.packetSent(bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 2900) + ledbat.packetSent(state: &state, bytesSent: 1000) ledbat.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -441,8 +451,8 @@ final class LedbatTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - ledbat.spuriousRetransmit() - XCTAssertEqual(ledbat.availableCongestionWindow, 2900) + ledbat.spuriousRetransmit(state: &state) + XCTAssertEqual(state.availableCongestionWindow, 2900) } // Tests that we can enter CA without any loss after idle period @@ -451,13 +461,14 @@ final class LedbatTests: XCTestCase { rtt.adjustedRTT = .milliseconds(120) rtt.smoothedRTT = .milliseconds(100) let time = NetworkClock.Instant.testBase - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) ledbat.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -465,27 +476,27 @@ final class LedbatTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: true, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 2400) - ledbat.idleTimeout(mss: mss) - XCTAssertEqual(ledbat.availableCongestionWindow, 2400) - ledbat.packetSent(bytesSent: 1200) - ledbat.packetSent(bytesSent: 1200) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1200, sentTime: time) - ledbat.packetsAcked(bytesAcked: 1200, sentTime: time) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: true, now: time) + XCTAssertEqual(state.availableCongestionWindow, 2400) + ledbat.idleTimeout(state: &state, mss: mss) + XCTAssertEqual(state.availableCongestionWindow, 2400) + ledbat.packetSent(state: &state, bytesSent: 1200) + ledbat.packetSent(state: &state, bytesSent: 1200) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1200, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1200, sentTime: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) // Enter CA - XCTAssertEqual(ledbat.availableCongestionWindow, 3000) + XCTAssertEqual(state.availableCongestionWindow, 3000) for _ in 0..<3 { - ledbat.packetSent(bytesSent: 1000) + ledbat.packetSent(state: &state, bytesSent: 1000) } - ledbat.ackBegin() + ledbat.ackBegin(state: &state) for _ in 0..<3 { - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) } - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(ledbat.availableCongestionWindow, 3600) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 3600) } @@ -493,19 +504,19 @@ final class LedbatTests: XCTestCase { var dataTransferSnapshot = DataTransferSnapshot() XCTAssertEqual(dataTransferSnapshot.transportCongestionWindow, 0) XCTAssertEqual(dataTransferSnapshot.transportSlowStartThreshold, 0) - ledbat.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) + ledbat.filloutDataTransferSnapshot(state: state, dataTransferSnapshot: &dataTransferSnapshot) XCTAssertTrue(dataTransferSnapshot.transportCongestionWindow > 0) XCTAssertTrue(dataTransferSnapshot.transportSlowStartThreshold > 0) rtt.adjustedRTT = .milliseconds(120) rtt.smoothedRTT = .milliseconds(100) let time = NetworkClock.Instant.testBase - ledbat.packetSent(bytesSent: 1000) - ledbat.ackBegin() - ledbat.packetsAcked(bytesAcked: 1000, sentTime: time) - ledbat.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + ledbat.packetSent(state: &state, bytesSent: 1000) + ledbat.ackBegin(state: &state) + ledbat.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + ledbat.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - ledbat.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) + ledbat.filloutDataTransferSnapshot(state: state, dataTransferSnapshot: &dataTransferSnapshot) XCTAssertEqual(dataTransferSnapshot.transportCongestionWindow, 2900) } diff --git a/Tests/QUICTests/PragueTests.swift b/Tests/QUICTests/PragueTests.swift index 544dde3..c6deb9a 100644 --- a/Tests/QUICTests/PragueTests.swift +++ b/Tests/QUICTests/PragueTests.swift @@ -32,6 +32,7 @@ final class PragueTests: XCTestCase { var rtt: RTT! let mss = Constants.initialMSS var prague: Prague! + var state = CongestionControlState() // These tests drive the algorithm directly, with no path to pace. let noPath: QUICPath? = nil var pacer: Pacer = Pacer(enabled: true) @@ -39,41 +40,43 @@ final class PragueTests: XCTestCase { override func setUp() { let logPrefixer = LogPrefixer("[PragueTests]") - prague = Prague(pacer: &pacer, mss: mss, logPrefixer: logPrefixer) + state = CongestionControlState() + prague = Prague(state: &state, pacer: &pacer, mss: mss, logPrefixer: logPrefixer) rtt = RTT(logPrefixer: logPrefixer) } func testPragueMSS() { /* Test MSS > congestion window */ - XCTAssertEqual(prague.availableCongestionWindow, defaultCongestionWindow) - prague.mssChanged(mss: 65000) - XCTAssertEqual(prague.availableCongestionWindow, 65000) - prague.reset(mss: Constants.initialMSS) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) + prague.mssChanged(state: &state, mss: 65000) + XCTAssertEqual(state.availableCongestionWindow, 65000) + prague.reset(state: &state, mss: Constants.initialMSS) /* Test MSS < congestion window */ - XCTAssertEqual(prague.availableCongestionWindow, defaultCongestionWindow) - prague.mssChanged(mss: 10) - XCTAssertEqual(prague.availableCongestionWindow, defaultCongestionWindow) - prague.reset(mss: Constants.initialMSS) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) + prague.mssChanged(state: &state, mss: 10) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) + prague.reset(state: &state, mss: Constants.initialMSS) } func testPragueReset() { - XCTAssertEqual(prague.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) let time = NetworkClock.Instant.testBase - prague.packetSent(bytesSent: 1000) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(prague.availableCongestionWindow, 13000) - prague.reset(mss: Constants.initialMSS) - XCTAssertEqual(prague.availableCongestionWindow, defaultCongestionWindow) + prague.packetSent(state: &state, bytesSent: 1000) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 13000) + prague.reset(state: &state, mss: Constants.initialMSS) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) } func testPragueLostPackets() { - XCTAssertEqual(prague.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) /* "Send" some packets and declare them lost */ let time = NetworkClock.Instant.testBase - prague.packetSent(bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) prague.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -81,29 +84,30 @@ final class PragueTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - XCTAssertEqual(prague.availableCongestionWindow, 8400) + XCTAssertEqual(state.availableCongestionWindow, 8400) /* See if we can send another packet */ - XCTAssertTrue(prague.canSend(packetLength: 1000)) + XCTAssertTrue(prague.canSend(state: state, packetLength: 1000)) } func testPragueSlowStart() { rtt.smoothedRTT = .microseconds(10) /* "Send" some packets and declare one of them lost */ var time = NetworkClock.Instant.testBase - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) prague.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -111,35 +115,36 @@ final class PragueTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - XCTAssertEqual(prague.availableCongestionWindow, 11900) + XCTAssertEqual(state.availableCongestionWindow, 11900) /* Make sure that another successful packet doesn't cause us to continue slow start */ time = NetworkClock.Instant.testBase.advanced(by: .microseconds(100)) - prague.packetSent(bytesSent: 1000) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(prague.availableCongestionWindow, 11953) + prague.packetSent(state: &state, bytesSent: 1000) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 11953) } func testPragueECN() { rtt.smoothedRTT = .milliseconds(15) - XCTAssertEqual(prague.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) /* Test that CE counts will reduce the congestion window immediately and move Prague to Congestion avoidance */ var time = NetworkClock.Instant.testBase - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) prague.processECN( + state: &state, path: noPath, ceCount: 1, packetsAcked: 6, @@ -150,36 +155,37 @@ final class PragueTests: XCTestCase { smoothedRTT: rtt.smoothedRTT, now: time ) - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) // cwnd after reduction = 6313 and after AI increase for 5 unmarked packets = 6826 - XCTAssertEqual(prague.availableCongestionWindow, 6826) + XCTAssertEqual(state.availableCongestionWindow, 6826) time = NetworkClock.Instant.testBase.advanced(by: .microseconds(100)) - prague.packetSent(bytesSent: 1000) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + prague.packetSent(state: &state, bytesSent: 1000) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) /* congestion window grows during congestion avoidance */ - XCTAssertEqual(prague.availableCongestionWindow, 6924) + XCTAssertEqual(state.availableCongestionWindow, 6924) } func testPragueECNEnterCWR() { rtt.smoothedRTT = .milliseconds(15) - XCTAssertEqual(prague.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) /* Test that CE counts will reduce congestion window, enter CWR and after that we don't decrease congestion window for 1RTT even we receive new CE counts */ let time = NetworkClock.Instant.testBase - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) prague.processECN( + state: &state, path: noPath, ceCount: 1, packetsAcked: 4, @@ -190,15 +196,16 @@ final class PragueTests: XCTestCase { smoothedRTT: rtt.smoothedRTT, now: time ) - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) // cwnd after decrease = 6282, after AI increase = 6582 // allowed cwnd = cwnd - bytes_in_flight = 6582 - 2000 = 4582 - XCTAssertEqual(prague.availableCongestionWindow, 4582) + XCTAssertEqual(state.availableCongestionWindow, 4582) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) prague.processECN( + state: &state, path: noPath, ceCount: 2, packetsAcked: 6, @@ -209,28 +216,29 @@ final class PragueTests: XCTestCase { smoothedRTT: rtt.smoothedRTT, now: time ) - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) // cwnd is 6664 after AI for 1 unmarked packet - XCTAssertEqual(prague.availableCongestionWindow, 6664) + XCTAssertEqual(state.availableCongestionWindow, 6664) } func testPragueAckDuringRecovery() { rtt.smoothedRTT = .microseconds(10) /* "Send" some packets and declare one of them lost */ var time = NetworkClock.Instant.testBase - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) prague.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -238,45 +246,45 @@ final class PragueTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: true, now: time) - XCTAssertEqual(prague.availableCongestionWindow, 8400) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: true, now: time) + XCTAssertEqual(state.availableCongestionWindow, 8400) time = NetworkClock.Instant.testBase.advanced(by: .microseconds(100)) - prague.packetSent(bytesSent: 1000) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(prague.availableCongestionWindow, 8475) + prague.packetSent(state: &state, bytesSent: 1000) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 8475) } func testPragueIdleTimeout() { - XCTAssertEqual(prague.availableCongestionWindow, defaultCongestionWindow) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) let time = NetworkClock.Instant.testBase - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(prague.availableCongestionWindow, 18000) - prague.idleTimeout(mss: mss) - XCTAssertEqual(prague.availableCongestionWindow, defaultCongestionWindow) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 18000) + prague.idleTimeout(state: &state, mss: mss) + XCTAssertEqual(state.availableCongestionWindow, defaultCongestionWindow) } func testPraguePersistentCongestion() { - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.persistentCongestion(mss: mss) - XCTAssertEqual(prague.availableCongestionWindow, 0) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.persistentCongestion(state: &state, mss: mss) + XCTAssertEqual(state.availableCongestionWindow, 0) } func testPragueCongestionLimited() { @@ -285,19 +293,20 @@ final class PragueTests: XCTestCase { // sending give three reductions. var sentTime = NetworkClock.Instant.testBase var detectedAt = sentTime.advanced(by: .microseconds(100)) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1000, sentTime: sentTime) - prague.packetsAcked(bytesAcked: 1000, sentTime: sentTime) - prague.packetsAcked(bytesAcked: 1000, sentTime: sentTime) - prague.packetsAcked(bytesAcked: 1000, sentTime: sentTime) - prague.packetsAcked(bytesAcked: 1000, sentTime: sentTime) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: sentTime) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: sentTime) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: sentTime) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: sentTime) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: sentTime) prague.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -305,14 +314,15 @@ final class PragueTests: XCTestCase { smoothedRTT: .microseconds(0), now: detectedAt ) - XCTAssertEqual(prague.availableCongestionWindow, 8400) + XCTAssertEqual(state.availableCongestionWindow, 8400) sentTime = sentTime.advanced(by: .microseconds(1000)) detectedAt = sentTime.advanced(by: .microseconds(100)) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) prague.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -321,6 +331,7 @@ final class PragueTests: XCTestCase { now: detectedAt ) prague.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -329,6 +340,7 @@ final class PragueTests: XCTestCase { now: detectedAt ) prague.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -337,6 +349,7 @@ final class PragueTests: XCTestCase { now: detectedAt ) prague.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -346,9 +359,10 @@ final class PragueTests: XCTestCase { ) sentTime = sentTime.advanced(by: .microseconds(1000)) detectedAt = sentTime.advanced(by: .microseconds(100)) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) prague.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -357,6 +371,7 @@ final class PragueTests: XCTestCase { now: detectedAt ) prague.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: sentTime, @@ -367,23 +382,24 @@ final class PragueTests: XCTestCase { // One reduction per round, three rounds: 12000 -> 8400 -> 5880 -> 4116, each step // `UInt64(Double(window) * Prague.beta)`. Written out rather than as `pow(beta, 3)`, which // is 0.34299999999999997 and truncates to 4115; the reductions compound one at a time. - XCTAssertEqual(prague.availableCongestionWindow, 4116) - XCTAssertFalse(prague.canSend(packetLength: 10000)) + XCTAssertEqual(state.availableCongestionWindow, 4116) + XCTAssertFalse(prague.canSend(state: state, packetLength: 10000)) } func testPraguePacketDiscard() { - prague.packetSent(bytesSent: 1000) - prague.packetDiscarded(bytesSent: 1000) - XCTAssertEqual(prague.bytesInFlight, 0) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetDiscarded(state: &state, bytesSent: 1000) + XCTAssertEqual(state.bytesInFlight, 0) } func testPragueSpuriousRetransmit() { let time = NetworkClock.Instant.testBase - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) prague.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -391,20 +407,21 @@ final class PragueTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - prague.spuriousRetransmit() - XCTAssertEqual(prague.availableCongestionWindow, 9000) + prague.spuriousRetransmit(state: &state) + XCTAssertEqual(state.availableCongestionWindow, 9000) } /* Tests that we can enter CA without any loss after idle period */ func testPragueCongestionAvoidance() { var time = NetworkClock.Instant.testBase - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) prague.packetLost( + state: &state, path: noPath, bytesLost: 1000, largestLostSentTime: time, @@ -412,50 +429,50 @@ final class PragueTests: XCTestCase { smoothedRTT: .microseconds(0), now: time ) - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: true, now: time) - XCTAssertEqual(prague.availableCongestionWindow, 8400) - prague.idleTimeout(mss: mss) - XCTAssertEqual(prague.availableCongestionWindow, 8400) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: true, now: time) + XCTAssertEqual(state.availableCongestionWindow, 8400) + prague.idleTimeout(state: &state, mss: mss) + XCTAssertEqual(state.availableCongestionWindow, 8400) time = NetworkClock.Instant.testBase - prague.packetSent(bytesSent: 1200) - prague.packetSent(bytesSent: 1200) - prague.packetSent(bytesSent: 1200) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1200, sentTime: time) - prague.packetsAcked(bytesAcked: 1200, sentTime: time) - prague.packetsAcked(bytesAcked: 1200, sentTime: time) - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + prague.packetSent(state: &state, bytesSent: 1200) + prague.packetSent(state: &state, bytesSent: 1200) + prague.packetSent(state: &state, bytesSent: 1200) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1200, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1200, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1200, sentTime: time) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) /* Enter CA */ - XCTAssertEqual(prague.availableCongestionWindow, 12000) + XCTAssertEqual(state.availableCongestionWindow, 12000) for _ in 0..<12 { - prague.packetSent(bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) } - prague.ackBegin() + prague.ackBegin(state: &state) for _ in 0..<12 { - prague.packetsAcked(bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) } - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - XCTAssertEqual(prague.availableCongestionWindow, 13200) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + XCTAssertEqual(state.availableCongestionWindow, 13200) } func testPragueDataTransferSnapshot() { var dataTransferSnapshot = DataTransferSnapshot() XCTAssertEqual(dataTransferSnapshot.transportCongestionWindow, 0) XCTAssertEqual(dataTransferSnapshot.transportSlowStartThreshold, 0) - prague.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) + prague.filloutDataTransferSnapshot(state: state, dataTransferSnapshot: &dataTransferSnapshot) XCTAssertTrue(dataTransferSnapshot.transportCongestionWindow > 0) XCTAssertTrue(dataTransferSnapshot.transportSlowStartThreshold > 0) let existingCongestionWindow = dataTransferSnapshot.transportCongestionWindow let time = NetworkClock.Instant.testBase - prague.packetSent(bytesSent: 1000) - prague.packetSent(bytesSent: 1000) - prague.ackBegin() - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.packetsAcked(bytesAcked: 1000, sentTime: time) - prague.ackEnd(rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) + prague.packetSent(state: &state, bytesSent: 1000) + prague.packetSent(state: &state, bytesSent: 1000) + prague.ackBegin(state: &state) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.packetsAcked(state: &state, bytesAcked: 1000, sentTime: time) + prague.ackEnd(state: &state, rtt: rtt, path: noPath, mss: mss, packetsLost: false, now: time) - prague.filloutDataTransferSnapshot(dataTransferSnapshot: &dataTransferSnapshot) + prague.filloutDataTransferSnapshot(state: state, dataTransferSnapshot: &dataTransferSnapshot) XCTAssertEqual( dataTransferSnapshot.transportCongestionWindow, (existingCongestionWindow + 2000)