lib/wayland/src/stream/outbox.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const sys = @import("sys");
  3 const wayland = @import("../root.zig");
  4 const storage = @import("storage.zig");
  5 
  6 const wire = wayland.wire;
  7 
  8 pub const FlushStatus = enum {
  9     drained,
 10     pending,
 11 };
 12 
 13 pub const Outbox = struct {
 14     session_allocator: std.mem.Allocator,
 15     bytes: []u8,
 16     descriptors: []sys.fd.Descriptor,
 17     byte_count: usize = 0,
 18     descriptor_count: usize = 0,
 19     byte_offset: usize = 0,
 20     descriptor_offset: usize = 0,
 21     byte_capacity_rejection_count: u64 = 0,
 22     descriptor_capacity_rejection_count: u64 = 0,
 23 
 24     pub fn init(
 25         session_allocator: std.mem.Allocator,
 26         capacity: storage.Capacity,
 27     ) std.mem.Allocator.Error!Outbox {
 28         const bytes = try session_allocator.alloc(u8, capacity.outbound_byte_count);
 29         errdefer session_allocator.free(bytes);
 30         return .{
 31             .session_allocator = session_allocator,
 32             .bytes = bytes,
 33             .descriptors = try session_allocator.alloc(
 34                 sys.fd.Descriptor,
 35                 capacity.outbound_descriptor_count,
 36             ),
 37         };
 38     }
 39 
 40     pub fn deinit(self: *Outbox) void {
 41         self.assertValid();
 42         closeAll(self.descriptors[self.descriptor_offset..self.descriptor_count]);
 43         self.session_allocator.free(self.bytes);
 44         if (self.descriptors.len != 0) self.session_allocator.free(self.descriptors);
 45         self.* = undefined;
 46     }
 47 
 48     pub fn queueOwned(
 49         self: *Outbox,
 50         object_id: u32,
 51         opcode: u16,
 52         payload: []const u8,
 53         descriptors: []const sys.fd.Descriptor,
 54     ) (wire.Error || storage.StorageError || error{TooManyDescriptors})!void {
 55         if (descriptors.len > sys.ancillary.maximum_descriptors) return error.TooManyDescriptors;
 56         const header = try wire.Header.init(object_id, opcode, payload.len);
 57         var encoded: [wire.header_size]u8 = undefined;
 58         try header.encode(&encoded);
 59 
 60         self.compact();
 61         const message_byte_count = encoded.len + payload.len;
 62         var failure: ?storage.StorageError = null;
 63         if (message_byte_count > self.bytes.len - self.byte_count) {
 64             self.byte_capacity_rejection_count +|= 1;
 65             failure = error.OutboundByteCapacityExceeded;
 66         }
 67         if (descriptors.len > self.descriptors.len - self.descriptor_count) {
 68             self.descriptor_capacity_rejection_count +|= 1;
 69             if (failure == null) failure = error.OutboundDescriptorCapacityExceeded;
 70         }
 71         if (failure) |err| return err;
 72         @memcpy(self.bytes[self.byte_count..][0..encoded.len], &encoded);
 73         self.byte_count += encoded.len;
 74         @memcpy(self.bytes[self.byte_count..][0..payload.len], payload);
 75         self.byte_count += payload.len;
 76         @memcpy(
 77             self.descriptors[self.descriptor_count..][0..descriptors.len],
 78             descriptors,
 79         );
 80         self.descriptor_count += descriptors.len;
 81         self.assertValid();
 82     }
 83 
 84     pub fn flush(self: *Outbox, socket: sys.fd.Descriptor) !FlushStatus {
 85         while (self.byte_offset < self.byte_count) {
 86             const pending_bytes = self.bytes[self.byte_offset..self.byte_count];
 87             const pending_descriptors = self.descriptors[self.descriptor_offset..self.descriptor_count];
 88             const descriptor_count = @min(pending_descriptors.len, sys.ancillary.maximum_descriptors);
 89             const bytes = if (pending_descriptors.len > descriptor_count)
 90                 pending_bytes[0..1]
 91             else
 92                 pending_bytes;
 93             const descriptors = pending_descriptors[0..descriptor_count];
 94             const sent = sys.ancillary.send(socket, bytes, descriptors) catch |err| switch (err) {
 95                 error.WouldBlock => return .pending,
 96                 else => return err,
 97             };
 98             if (sent == 0) return error.SendFailed;
 99             closeAll(descriptors);
100             self.descriptor_offset += descriptors.len;
101             self.byte_offset += sent;
102             self.assertValid();
103         }
104 
105         std.debug.assert(self.descriptor_offset == self.descriptor_count);
106         self.byte_count = 0;
107         self.descriptor_count = 0;
108         self.byte_offset = 0;
109         self.descriptor_offset = 0;
110         self.assertValid();
111         return .drained;
112     }
113 
114     pub fn pendingByteCount(self: *const Outbox) usize {
115         return self.byte_count - self.byte_offset;
116     }
117 
118     pub fn pendingDescriptorCount(self: *const Outbox) usize {
119         return self.descriptor_count - self.descriptor_offset;
120     }
121 
122     pub fn status(self: *const Outbox) storage.Status {
123         return .{
124             .outbound_byte_capacity_rejection_count = self.byte_capacity_rejection_count,
125             .outbound_descriptor_capacity_rejection_count = self.descriptor_capacity_rejection_count,
126         };
127     }
128 
129     fn compact(self: *Outbox) void {
130         storage.compactSlice(u8, self.bytes, &self.byte_count, &self.byte_offset);
131         storage.compactSlice(
132             sys.fd.Descriptor,
133             self.descriptors,
134             &self.descriptor_count,
135             &self.descriptor_offset,
136         );
137     }
138 
139     fn assertValid(self: *const Outbox) void {
140         std.debug.assert(self.byte_offset <= self.byte_count);
141         std.debug.assert(self.byte_count <= self.bytes.len);
142         std.debug.assert(self.descriptor_offset <= self.descriptor_count);
143         std.debug.assert(self.descriptor_count <= self.descriptors.len);
144     }
145 };
146 
147 fn closeAll(descriptors: []const sys.fd.Descriptor) void {
148     for (descriptors) |descriptor| sys.fd.close(descriptor);
149 }
150 
151 test "outbox retains partial writes under socket backpressure" {
152     const sockets = sys.fd.socketPairUnixStream(.{ .close_on_exec = true }) catch |err| switch (err) {
153         error.UnsupportedPlatform => return error.SkipZigTest,
154         else => return err,
155     };
156     defer sys.fd.close(sockets[0]);
157     defer sys.fd.close(sockets[1]);
158     try sys.fd.setNonBlocking(sockets[0]);
159 
160     const capacity = try storage.Capacity.derive(.{
161         .outbound_byte_count = 16 * @as(usize, wire.maximum_message_size),
162     });
163     var outbox = try Outbox.init(std.testing.allocator, capacity);
164     defer outbox.deinit();
165     var payload: [wire.maximum_message_size - wire.header_size]u8 = @splat(0xa5);
166     for (0..16) |index_value| {
167         try outbox.queueOwned(2, @intCast(index_value), &payload, &.{});
168     }
169     const initial = outbox.pendingByteCount();
170     try std.testing.expectEqual(FlushStatus.pending, try outbox.flush(sockets[0]));
171     try std.testing.expect(outbox.pendingByteCount() > 0);
172     try std.testing.expect(outbox.pendingByteCount() < initial);
173 }
174 
175 test "outbox retains owned descriptors while the socket would block" {
176     const sockets = sys.fd.socketPairUnixStream(.{ .close_on_exec = true }) catch |err| switch (err) {
177         error.UnsupportedPlatform => return error.SkipZigTest,
178         else => return err,
179     };
180     defer sys.fd.close(sockets[0]);
181     defer sys.fd.close(sockets[1]);
182     try sys.fd.setNonBlocking(sockets[0]);
183 
184     const fill: [4096]u8 = @splat(0xcc);
185     while (true) {
186         _ = sys.fd.write(sockets[0], &fill) catch |err| switch (err) {
187             error.WouldBlock => break,
188             else => return err,
189         };
190     }
191 
192     const pipe = try sys.fd.pipeWithOptions(.{ .close_on_exec = true });
193     defer sys.fd.close(pipe[0]);
194     var queued = false;
195     defer if (!queued) sys.fd.close(pipe[1]);
196     const capacity = try storage.Capacity.derive(.{ .outbound_descriptor_count = 1 });
197     var outbox = try Outbox.init(std.testing.allocator, capacity);
198     defer outbox.deinit();
199     try outbox.queueOwned(2, 0, &.{}, &.{pipe[1]});
200     queued = true;
201 
202     try std.testing.expectEqual(FlushStatus.pending, try outbox.flush(sockets[0]));
203     try std.testing.expectEqual(@as(usize, 1), outbox.pendingDescriptorCount());
204     try std.testing.expectEqual(@as(usize, 1), try sys.fd.write(pipe[1], "x"));
205     var received: [1]u8 = undefined;
206     try std.testing.expectEqual(@as(usize, 1), try sys.fd.read(pipe[0], &received));
207     try std.testing.expectEqual(@as(u8, 'x'), received[0]);
208 }
209 
210 test "outbox max plus one preserves retained bytes and saturates status" {
211     const capacity = try storage.Capacity.derive(.{ .outbound_byte_count = 8 });
212     var outbox = try Outbox.init(std.testing.allocator, capacity);
213     defer outbox.deinit();
214 
215     try outbox.queueOwned(2, 0, &.{}, &.{});
216     try std.testing.expectError(
217         error.OutboundByteCapacityExceeded,
218         outbox.queueOwned(3, 1, &.{}, &.{}),
219     );
220     try std.testing.expectEqual(@as(usize, 8), outbox.pendingByteCount());
221     try std.testing.expectEqual(@as(u64, 1), outbox.status().outbound_byte_capacity_rejection_count);
222 
223     outbox.byte_capacity_rejection_count = std.math.maxInt(u64);
224     try std.testing.expectError(
225         error.OutboundByteCapacityExceeded,
226         outbox.queueOwned(3, 1, &.{}, &.{}),
227     );
228     try std.testing.expectEqual(
229         std.math.maxInt(u64),
230         outbox.status().outbound_byte_capacity_rejection_count,
231     );
232 }
233 
234 test "outbox max plus one leaves rejected descriptor with caller" {
235     const first_pipe = try sys.fd.pipeWithOptions(.{ .close_on_exec = true });
236     defer sys.fd.close(first_pipe[0]);
237     const second_pipe = try sys.fd.pipeWithOptions(.{ .close_on_exec = true });
238     defer sys.fd.close(second_pipe[0]);
239     defer sys.fd.close(second_pipe[1]);
240 
241     const capacity = try storage.Capacity.derive(.{
242         .outbound_byte_count = 2 * wire.header_size,
243         .outbound_descriptor_count = 1,
244     });
245     var outbox = try Outbox.init(std.testing.allocator, capacity);
246     defer outbox.deinit();
247     try outbox.queueOwned(2, 0, &.{}, &.{first_pipe[1]});
248     try std.testing.expectError(
249         error.OutboundDescriptorCapacityExceeded,
250         outbox.queueOwned(3, 1, &.{}, &.{second_pipe[1]}),
251     );
252     try std.testing.expectEqual(@as(usize, 8), outbox.pendingByteCount());
253     try std.testing.expectEqual(@as(usize, 1), outbox.pendingDescriptorCount());
254     try std.testing.expectEqual(
255         @as(u64, 1),
256         outbox.status().outbound_descriptor_capacity_rejection_count,
257     );
258 }