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 }