lib/quic/src/connection/ack.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const quic = @import("../root.zig");
  3 
  4 pub const Range = struct {
  5     smallest: u62,
  6     largest: u62,
  7 };
  8 
  9 pub const Ranges = struct {
 10     values: []Range,
 11     count: u16 = 0,
 12 
 13     pub fn init(values: []Range) Ranges {
 14         std.debug.assert(values.len <= std.math.maxInt(u16));
 15         for (values) |*value| value.* = .{ .smallest = 0, .largest = 0 };
 16         return .{ .values = values };
 17     }
 18 
 19     pub fn contains(self: *const Ranges, packet_number: u62) bool {
 20         for (0..self.values.len) |index| {
 21             if (index >= self.count) break;
 22             const value = self.values[index];
 23             if (packet_number < value.smallest) return false;
 24             if (packet_number <= value.largest) return true;
 25         }
 26         return false;
 27     }
 28 
 29     pub fn insert(self: *Ranges, packet_number: u62) bool {
 30         if (self.contains(packet_number)) return false;
 31         var at = self.insertionIndex(packet_number);
 32         const joins_left = self.joinsLeft(at, packet_number);
 33         const joins_right = self.joinsRight(at, packet_number);
 34         if (joins_left and joins_right) {
 35             self.values[at - 1].largest = self.values[at].largest;
 36             self.remove(at);
 37             return true;
 38         }
 39         if (joins_left) {
 40             self.values[at - 1].largest = packet_number;
 41             return true;
 42         }
 43         if (joins_right) {
 44             self.values[at].smallest = packet_number;
 45             return true;
 46         }
 47         if (self.count == self.values.len) {
 48             if (at == 0) return true;
 49             self.remove(0);
 50             at = self.insertionIndex(packet_number);
 51         }
 52         self.insertAt(at, .{ .smallest = packet_number, .largest = packet_number });
 53         return true;
 54     }
 55 
 56     pub fn encode(self: *const Ranges, delay: u62, output: *quic.cursor.Write) !void {
 57         if (self.count == 0) return error.EmptyRanges;
 58         std.debug.assert(self.count <= self.values.len);
 59         std.debug.assert(self.count - 1 <= quic.frame.ack_ranges_max);
 60         const newest = self.values[self.count - 1];
 61         try output.byte(0x02);
 62         _ = try quic.varint.write(newest.largest, output);
 63         _ = try quic.varint.write(delay, output);
 64         _ = try quic.varint.write(self.count - 1, output);
 65         _ = try quic.varint.write(newest.largest - newest.smallest, output);
 66         var previous = newest;
 67         for (0..self.values.len) |reverse| {
 68             if (reverse + 1 >= self.count) break;
 69             const index = self.count - 2 - reverse;
 70             const current = self.values[index];
 71             _ = try quic.varint.write(previous.smallest - current.largest - 2, output);
 72             _ = try quic.varint.write(current.largest - current.smallest, output);
 73             previous = current;
 74         }
 75     }
 76 
 77     fn insertionIndex(self: *const Ranges, packet_number: u62) usize {
 78         for (0..self.values.len) |index| {
 79             if (index >= self.count) return index;
 80             if (packet_number < self.values[index].smallest) return index;
 81         }
 82         return self.count;
 83     }
 84 
 85     fn joinsLeft(self: *const Ranges, at: usize, packet_number: u62) bool {
 86         if (at == 0) return false;
 87         const largest = self.values[at - 1].largest;
 88         if (largest == std.math.maxInt(u62)) return false;
 89         return largest + 1 == packet_number;
 90     }
 91 
 92     fn joinsRight(self: *const Ranges, at: usize, packet_number: u62) bool {
 93         if (at >= self.count) return false;
 94         if (packet_number == std.math.maxInt(u62)) return false;
 95         return packet_number + 1 == self.values[at].smallest;
 96     }
 97 
 98     fn insertAt(self: *Ranges, at: usize, value: Range) void {
 99         std.debug.assert(self.count < self.values.len);
100         var index: usize = self.count;
101         for (0..self.values.len) |_| {
102             if (index <= at) break;
103             self.values[index] = self.values[index - 1];
104             index -= 1;
105         }
106         self.values[at] = value;
107         self.count += 1;
108     }
109 
110     fn remove(self: *Ranges, at: usize) void {
111         std.debug.assert(at < self.count);
112         for (0..self.values.len) |offset| {
113             const index = at + offset;
114             if (index + 1 >= self.count) break;
115             self.values[index] = self.values[index + 1];
116         }
117         self.count -= 1;
118     }
119 };
120 
121 test "RFC 9000 section 13.2.3 ACK ranges merge and discard the oldest range" {
122     var values: [2]Range = undefined;
123     var ranges = Ranges.init(&values);
124     for ([_]u62{ 4, 2, 3, 8, 6 }) |packet_number| _ = ranges.insert(packet_number);
125     try std.testing.expectEqual(@as(u16, 2), ranges.count);
126     try std.testing.expectEqual(Range{ .smallest = 6, .largest = 6 }, ranges.values[0]);
127     try std.testing.expectEqual(Range{ .smallest = 8, .largest = 8 }, ranges.values[1]);
128 }
129 
130 test "RFC 9000 section 19.3 ACK encoding decodes to tracked ranges" {
131     var values: [3]Range = undefined;
132     var ranges = Ranges.init(&values);
133     for ([_]u62{ 1, 2, 5, 9, 10 }) |packet_number| _ = ranges.insert(packet_number);
134     var bytes: [64]u8 = undefined;
135     var output = quic.cursor.Write.init(&bytes);
136     try ranges.encode(7, &output);
137     var input = quic.cursor.Read.init(output.written());
138     const decoded = try quic.frame.decode(&input);
139     try std.testing.expectEqual(@as(u62, 10), decoded.ack.largest);
140     try std.testing.expectEqual(@as(u8, 2), decoded.ack.range_count);
141 }