lib/wayland/src/runtime/client.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const sys = @import("sys");
  3 const wayland = @import("../root.zig");
  4 const catalog_model = @import("catalog.zig");
  5 const creation = @import("creation.zig");
  6 const dispatch = @import("dispatch.zig");
  7 const event = @import("event.zig");
  8 const global_state = @import("global.zig");
  9 const lifecycle = @import("lifecycle.zig");
 10 const object_model = @import("object.zig");
 11 const request_model = @import("request.zig");
 12 
 13 pub const default_client_object_count = object_model.default_client_object_count;
 14 pub const default_server_object_count = object_model.default_server_object_count;
 15 pub const default_retained_global_count = global_state.default_retained_global_count;
 16 pub const default_global_interface_name_byte_count =
 17     global_state.default_interface_name_byte_count_per_global;
 18 
 19 pub const Limits = struct {
 20     transport: wayland.TransportLimits = .{},
 21     client_object_count: usize = object_model.default_client_object_count,
 22     server_object_count: usize = object_model.default_server_object_count,
 23     retained_global_count: usize = global_state.default_retained_global_count,
 24     global_interface_name_byte_count: usize =
 25         global_state.default_interface_name_byte_count_per_global,
 26 };
 27 
 28 pub const CapacityError = wayland.TransportCapacityError ||
 29     wayland.protocol.value.EncoderCapacityError ||
 30     event.CapacityError ||
 31     wayland.ids.CapacityError ||
 32     object_model.CapacityError ||
 33     global_state.CapacityError ||
 34     error{CapacityOverflow};
 35 
 36 pub const Capacity = struct {
 37     transport: wayland.TransportCapacity,
 38     request_encoder: wayland.protocol.value.EncoderCapacity,
 39     event_storage: event.Capacity,
 40     event_creations: creation.Capacity,
 41     client_ids: wayland.ids.Capacity,
 42     objects: object_model.Capacity,
 43     globals: global_state.Capacity,
 44     total_requested_bytes: usize,
 45 
 46     pub fn derive(limits: Limits) CapacityError!Capacity {
 47         const transport = try wayland.TransportCapacity.derive(limits.transport);
 48         const request_encoder = try wayland.protocol.value.EncoderCapacity.derive(.{
 49             .payload_byte_count = @min(
 50                 transport.outbound_byte_count - wayland.wire.header_size,
 51                 wayland.protocol.value.maximum_payload_size,
 52             ),
 53             .descriptor_count = @min(
 54                 transport.outbound_descriptor_count,
 55                 sys.ancillary.maximum_descriptors,
 56             ),
 57         });
 58         const event_storage = try event.Capacity.derive(.{
 59             .payload_byte_count = @min(
 60                 transport.inbound_byte_count - wayland.wire.header_size,
 61                 wayland.protocol.value.maximum_payload_size,
 62             ),
 63             .descriptor_count = @min(
 64                 transport.inbound_descriptor_count,
 65                 event.maximum_descriptor_count,
 66             ),
 67         });
 68         const event_creations = creation.default_capacity;
 69         const client_ids = try wayland.ids.Capacity.derive(.{
 70             .client_id_count = limits.client_object_count,
 71         });
 72         const objects = try object_model.Capacity.derive(.{
 73             .client_object_count = limits.client_object_count,
 74             .server_object_count = limits.server_object_count,
 75         });
 76         const globals = try global_state.Capacity.derive(.{
 77             .retained_global_count = limits.retained_global_count,
 78             .interface_name_byte_count_per_global = limits.global_interface_name_byte_count,
 79         });
 80         const total_requested_bytes = try totalRequestedBytes(&.{
 81             transport.total_requested_bytes,
 82             request_encoder.total_requested_bytes,
 83             event_storage.total_requested_bytes,
 84             event_creations.total_requested_bytes,
 85             client_ids.total_requested_bytes,
 86             objects.total_requested_bytes,
 87             globals.total_requested_bytes,
 88         });
 89         return .{
 90             .transport = transport,
 91             .request_encoder = request_encoder,
 92             .event_storage = event_storage,
 93             .event_creations = event_creations,
 94             .client_ids = client_ids,
 95             .objects = objects,
 96             .globals = globals,
 97             .total_requested_bytes = total_requested_bytes,
 98         };
 99     }
100 };
101 
102 fn totalRequestedBytes(parts: []const usize) error{CapacityOverflow}!usize {
103     var total: usize = 0;
104     for (parts) |part| {
105         total = std.math.add(usize, total, part) catch return error.CapacityOverflow;
106     }
107     return total;
108 }
109 
110 pub const StorageStatus = struct {
111     transport: wayland.TransportStatus = .{},
112     request_encoder: wayland.protocol.value.EncoderStatus = .{},
113     event_storage: event.Status = .{},
114     event_creations: creation.Status = .{},
115     client_ids: wayland.ids.Status = .{},
116     objects: object_model.Status = .{},
117     globals: global_state.Status = .{},
118 };
119 
120 const InitialStorage = struct {
121     encoder: wayland.protocol.value.Encoder,
122     event_storage: event.Storage,
123     event_creations: creation.Storage,
124     ids: wayland.ids.Pool,
125     objects: object_model.Table,
126     globals: global_state.Set,
127 
128     fn init(
129         session_allocator: std.mem.Allocator,
130         capacity: Capacity,
131     ) std.mem.Allocator.Error!InitialStorage {
132         var encoder = try wayland.protocol.value.Encoder.init(
133             session_allocator,
134             capacity.request_encoder,
135         );
136         errdefer encoder.deinit();
137         var event_storage = try event.Storage.init(
138             session_allocator,
139             capacity.event_storage,
140         );
141         errdefer event_storage.deinit();
142         var event_creations = try creation.Storage.init(
143             session_allocator,
144             capacity.event_creations,
145         );
146         errdefer event_creations.deinit();
147         var ids = try wayland.ids.Pool.init(session_allocator, capacity.client_ids);
148         errdefer ids.deinit();
149         var objects = try object_model.Table.init(session_allocator, capacity.objects);
150         errdefer objects.deinit();
151         objects.addDisplay(
152             catalog_model.Catalog.standard().require("wl_display") catch unreachable,
153         ) catch unreachable;
154         return .{
155             .encoder = encoder,
156             .event_storage = event_storage,
157             .event_creations = event_creations,
158             .ids = ids,
159             .objects = objects,
160             .globals = try global_state.Set.init(session_allocator, capacity.globals),
161         };
162     }
163 
164     fn deinit(self: *InitialStorage) void {
165         self.globals.deinit();
166         self.objects.deinit();
167         self.ids.deinit();
168         self.event_creations.deinit();
169         self.event_storage.deinit();
170         self.encoder.deinit();
171         self.* = undefined;
172     }
173 };
174 
175 pub const Client = struct {
176     transport: wayland.Transport,
177     catalog: catalog_model.Catalog,
178     ids: wayland.ids.Pool,
179     objects: object_model.Table,
180     globals: global_state.Set,
181     encoder: wayland.protocol.value.Encoder,
182     event_storage: event.Storage,
183     event_creations: creation.Storage,
184     fatal_state: ?event.FatalView = null,
185     terminal_error: ?anyerror = null,
186 
187     pub fn connect(
188         session_allocator: std.mem.Allocator,
189         capacity: Capacity,
190     ) !Client {
191         var initial_storage = try InitialStorage.init(session_allocator, capacity);
192         errdefer initial_storage.deinit();
193         const transport = try wayland.connect(session_allocator, capacity.transport);
194         return initOwnedParts(transport, &initial_storage);
195     }
196 
197     pub fn connectNamed(
198         session_allocator: std.mem.Allocator,
199         display_name: []const u8,
200         capacity: Capacity,
201     ) !Client {
202         var initial_storage = try InitialStorage.init(session_allocator, capacity);
203         errdefer initial_storage.deinit();
204         const transport = try wayland.connectNamed(
205             session_allocator,
206             display_name,
207             capacity.transport,
208         );
209         return initOwnedParts(transport, &initial_storage);
210     }
211 
212     pub fn initOwned(
213         session_allocator: std.mem.Allocator,
214         owned_descriptor: sys.fd.Descriptor,
215         capacity: Capacity,
216     ) !Client {
217         var initial_storage = InitialStorage.init(session_allocator, capacity) catch |err| {
218             sys.fd.close(owned_descriptor);
219             return err;
220         };
221         errdefer initial_storage.deinit();
222         const transport = try wayland.Transport.initOwned(
223             session_allocator,
224             owned_descriptor,
225             capacity.transport,
226         );
227         return initOwnedParts(transport, &initial_storage);
228     }
229 
230     fn initOwnedParts(
231         transport: wayland.Transport,
232         initial_storage: *InitialStorage,
233     ) Client {
234         const catalog = catalog_model.Catalog.standard();
235         const client: Client = .{
236             .transport = transport,
237             .catalog = catalog,
238             .ids = initial_storage.ids,
239             .objects = initial_storage.objects,
240             .globals = initial_storage.globals,
241             .encoder = initial_storage.encoder,
242             .event_storage = initial_storage.event_storage,
243             .event_creations = initial_storage.event_creations,
244         };
245         initial_storage.* = undefined;
246         return client;
247     }
248 
249     pub fn deinit(self: *Client) void {
250         self.event_creations.deinit();
251         self.event_storage.deinit();
252         self.encoder.deinit();
253         self.globals.deinit();
254         self.objects.deinit();
255         self.ids.deinit();
256         self.transport.deinit();
257         self.* = undefined;
258     }
259 
260     pub fn fatal(self: *const Client) ?event.FatalView {
261         const state = self.fatal_state orelse return null;
262         return .{
263             .object_id = state.object_id,
264             .code = state.code,
265             .message = state.message,
266         };
267     }
268 
269     pub fn descriptor(self: *const Client) sys.fd.Descriptor {
270         return self.transport.descriptor;
271     }
272 
273     pub fn storageStatus(self: *const Client) StorageStatus {
274         return .{
275             .transport = self.transport.status(),
276             .request_encoder = self.encoder.status(),
277             .event_storage = self.event_storage.status(),
278             .event_creations = self.event_creations.status(),
279             .client_ids = self.ids.status(),
280             .objects = self.objects.status(),
281             .globals = self.globals.status(),
282         };
283     }
284 
285     pub fn object(self: *const Client, id: u32) ?object_model.Entry {
286         return self.objects.getLive(id);
287     }
288 
289     pub fn global(self: *const Client, registry_id: u32, name: u32) ?event.GlobalView {
290         const item = self.globals.get(registry_id, name) orelse return null;
291         return .{
292             .registry_id = item.registry_id,
293             .name = item.name,
294             .interface = item.interface,
295             .version = item.version,
296         };
297     }
298 
299     pub fn sync(self: *Client) !u32 {
300         var created: [1]u32 = undefined;
301         try self.request(
302             wayland.ids.display_id,
303             0,
304             &.{.{ .new_id = .fixed }},
305             &created,
306         );
307         return created[0];
308     }
309 
310     pub fn getRegistry(self: *Client) !u32 {
311         var created: [1]u32 = undefined;
312         try self.request(
313             wayland.ids.display_id,
314             1,
315             &.{.{ .new_id = .fixed }},
316             &created,
317         );
318         return created[0];
319     }
320 
321     pub fn bind(
322         self: *Client,
323         registry_id: u32,
324         name: u32,
325         interface: *const wayland.protocol.schema.Interface,
326         version: u32,
327     ) !u32 {
328         try self.ensureActive();
329         const registry_interface = try self.catalog.require("wl_registry");
330         const registry = try self.objects.requireLive(registry_id);
331         if (registry.interface != registry_interface) return error.NotRegistry;
332         const canonical = try self.catalog.canonical(interface);
333         const advertised = try self.globals.require(registry_id, name);
334         if (!std.mem.eql(u8, advertised.interface, canonical.name)) {
335             return error.InterfaceMismatch;
336         }
337         if (version == 0 or version > advertised.version or version > canonical.version) {
338             return error.InvalidVersion;
339         }
340 
341         var created: [1]u32 = undefined;
342         try self.request(
343             registry_id,
344             0,
345             &.{
346                 .{ .uint = name },
347                 .{ .new_id = .{ .dynamic = .{
348                     .interface = canonical.name,
349                     .version = version,
350                 } } },
351             },
352             &created,
353         );
354         return created[0];
355     }
356 
357     pub fn request(
358         self: *Client,
359         object_id: u32,
360         opcode: u16,
361         values: []const request_model.Value,
362         created_ids: []u32,
363     ) !void {
364         try self.ensureActive();
365         const parent = try self.objects.requireLive(object_id);
366         const metadata = parent.interface.request(opcode) orelse {
367             return error.UnknownRequestOpcode;
368         };
369         if (!metadata.supportedBy(parent.version)) return error.UnsupportedRequestVersion;
370         if (metadata.destructor and parent.origin == .display) return error.CannotDestroyDisplay;
371 
372         const prepared = try request_model.prepare(
373             &self.ids,
374             &self.objects,
375             self.catalog,
376             &self.encoder,
377             parent,
378             metadata,
379             values,
380             created_ids,
381         );
382         errdefer request_model.rollback(&self.ids, &self.objects, prepared.created_ids);
383         try self.transport.queueOwned(
384             object_id,
385             metadata.opcode,
386             prepared.encoded.payload,
387             prepared.encoded.descriptors,
388         );
389         if (metadata.destructor) {
390             lifecycle.retire(
391                 &self.ids,
392                 &self.objects,
393                 object_id,
394                 parent,
395                 .retired_request,
396             );
397         }
398     }
399 
400     pub fn flush(self: *Client) !wayland.stream.FlushStatus {
401         try self.ensureActive();
402         return self.transport.flush() catch |err| {
403             self.terminal_error = err;
404             return err;
405         };
406     }
407 
408     pub fn step(self: *Client) !event.Step {
409         try self.ensureActive();
410         self.event_storage.reset();
411         return dispatch.step(self);
412     }
413 
414     fn ensureActive(self: *const Client) !void {
415         if (self.terminal_error) |err| return err;
416     }
417 };
418 
419 test "runtime capacity derives request and event scratch" {
420     const capacity = try Capacity.derive(.{
421         .transport = .{
422             .inbound_byte_count = 13,
423             .inbound_descriptor_count = 2,
424             .outbound_byte_count = 17,
425             .outbound_descriptor_count = 3,
426         },
427         .client_object_count = 5,
428         .server_object_count = 7,
429         .retained_global_count = 3,
430         .global_interface_name_byte_count = 11,
431     });
432     try std.testing.expectEqual(@as(usize, 13), capacity.transport.inbound_byte_count);
433     try std.testing.expectEqual(@as(usize, 2), capacity.transport.inbound_descriptor_count);
434     try std.testing.expectEqual(@as(usize, 17), capacity.transport.outbound_byte_count);
435     try std.testing.expectEqual(@as(usize, 3), capacity.transport.outbound_descriptor_count);
436     try std.testing.expectEqual(@as(usize, 9), capacity.request_encoder.payload_byte_count);
437     try std.testing.expectEqual(@as(usize, 3), capacity.request_encoder.descriptor_count);
438     try std.testing.expectEqual(@as(usize, 5), capacity.event_storage.payload_byte_count);
439     try std.testing.expectEqual(@as(usize, 1), capacity.event_storage.descriptor_count);
440     try std.testing.expectEqual(
441         creation.maximum_event_creation_count,
442         capacity.event_creations.event_creation_count,
443     );
444     try std.testing.expectEqual(
445         @sizeOf(creation.Creation),
446         capacity.event_creations.event_creation_bytes,
447     );
448     try std.testing.expectEqual(@as(usize, 5), capacity.client_ids.client_id_count);
449     try std.testing.expectEqual(@as(usize, 5), capacity.objects.client_object_count);
450     try std.testing.expectEqual(@as(usize, 7), capacity.objects.server_object_count);
451     try std.testing.expectEqual(@as(usize, 3), capacity.globals.retained_global_count);
452     try std.testing.expectEqual(
453         @as(usize, 11),
454         capacity.globals.interface_name_byte_count_per_global,
455     );
456     try std.testing.expectEqual(
457         capacity.transport.total_requested_bytes +
458             capacity.request_encoder.total_requested_bytes +
459             capacity.event_storage.total_requested_bytes +
460             capacity.event_creations.total_requested_bytes +
461             capacity.client_ids.total_requested_bytes +
462             capacity.objects.total_requested_bytes +
463             capacity.globals.total_requested_bytes,
464         capacity.total_requested_bytes,
465     );
466 }
467 
468 test "runtime request scratch caps at one protocol message" {
469     const capacity = try Capacity.derive(.{ .transport = .{
470         .outbound_byte_count = 2 * @as(usize, wayland.wire.maximum_message_size),
471         .outbound_descriptor_count = 2 * sys.ancillary.maximum_descriptors,
472     } });
473     try std.testing.expectEqual(
474         @as(usize, wayland.protocol.value.maximum_payload_size),
475         capacity.request_encoder.payload_byte_count,
476     );
477     try std.testing.expectEqual(
478         @as(usize, sys.ancillary.maximum_descriptors),
479         capacity.request_encoder.descriptor_count,
480     );
481 }
482 
483 test "owned descriptor closes when any cold storage acquisition fails" {
484     if (comptime @import("builtin").os.tag != .linux) return error.SkipZigTest;
485     for (0..15) |fail_index| {
486         const sockets = sys.fd.socketPairUnixStream(.{ .close_on_exec = true }) catch |err|
487             switch (err) {
488                 error.UnsupportedPlatform => return error.SkipZigTest,
489                 else => return err,
490             };
491         defer sys.fd.close(sockets[1]);
492         var failing = std.testing.FailingAllocator.init(std.testing.allocator, .{
493             .fail_index = fail_index,
494         });
495         try std.testing.expectError(
496             error.OutOfMemory,
497             Client.initOwned(
498                 failing.allocator(),
499                 sockets[0],
500                 try Capacity.derive(.{}),
501             ),
502         );
503         try std.testing.expect(!sys.fd.isOpen(sockets[0]));
504     }
505 }
506 
507 test "runtime storage status composes event rejections" {
508     if (comptime @import("builtin").os.tag != .linux) return error.SkipZigTest;
509     const sockets = sys.fd.socketPairUnixStream(.{ .close_on_exec = true }) catch |err|
510         switch (err) {
511             error.UnsupportedPlatform => return error.SkipZigTest,
512             else => return err,
513         };
514     defer sys.fd.close(sockets[1]);
515     var client = try Client.initOwned(
516         std.testing.allocator,
517         sockets[0],
518         try Capacity.derive(.{}),
519     );
520     defer client.deinit();
521 
522     try std.testing.expectError(
523         error.EventPayloadCapacityExceeded,
524         client.event_storage.admit(client.event_storage.payload.len + 1, 0),
525     );
526     try std.testing.expectError(
527         error.EventDescriptorCapacityExceeded,
528         client.event_storage.admit(0, client.event_storage.descriptors.len + 1),
529     );
530     const status = client.storageStatus().event_storage;
531     try std.testing.expectEqual(
532         @as(u64, 1),
533         status.event_payload_capacity_rejection_count,
534     );
535     try std.testing.expectEqual(
536         @as(u64, 1),
537         status.event_descriptor_capacity_rejection_count,
538     );
539 }
540 
541 test "runtime storage status composes ID object and global rejections" {
542     if (comptime @import("builtin").os.tag != .linux) return error.SkipZigTest;
543     const sockets = sys.fd.socketPairUnixStream(.{ .close_on_exec = true }) catch |err|
544         switch (err) {
545             error.UnsupportedPlatform => return error.SkipZigTest,
546             else => return err,
547         };
548     defer sys.fd.close(sockets[1]);
549     var client = try Client.initOwned(
550         std.testing.allocator,
551         sockets[0],
552         try Capacity.derive(.{
553             .client_object_count = 1,
554             .server_object_count = 1,
555             .retained_global_count = 1,
556             .global_interface_name_byte_count = 3,
557         }),
558     );
559     defer client.deinit();
560     const callback = try client.catalog.require("wl_callback");
561 
562     const client_id = try client.ids.allocate();
563     try client.objects.addClient(client_id, callback, 1);
564     try std.testing.expectError(error.ClientIdCapacityExceeded, client.ids.allocate());
565     try std.testing.expectError(
566         error.ClientObjectCapacityExceeded,
567         client.objects.addClient(client_id + 1, callback, 1),
568     );
569     try client.objects.addServer(wayland.ids.first_server_id, callback, 1);
570     try std.testing.expectError(
571         error.ServerObjectCapacityExceeded,
572         client.objects.addServer(wayland.ids.first_server_id + 1, callback, 1),
573     );
574     _ = try client.globals.add(2, 7, "abc", 1);
575     try std.testing.expectError(
576         error.GlobalCapacityExceeded,
577         client.globals.add(2, 8, "a", 1),
578     );
579     try client.globals.remove(2, 7);
580     try std.testing.expectError(
581         error.GlobalInterfaceNameCapacityExceeded,
582         client.globals.add(2, 8, "abcd", 1),
583     );
584     const status = client.storageStatus();
585     try std.testing.expectEqual(
586         @as(u64, 1),
587         status.client_ids.client_id_capacity_rejection_count,
588     );
589     try std.testing.expectEqual(
590         @as(u64, 1),
591         status.objects.client_object_capacity_rejection_count,
592     );
593     try std.testing.expectEqual(
594         @as(u64, 1),
595         status.objects.server_object_capacity_rejection_count,
596     );
597     try std.testing.expectEqual(
598         @as(u64, 1),
599         status.globals.global_capacity_rejection_count,
600     );
601     try std.testing.expectEqual(
602         @as(u64, 1),
603         status.globals.global_interface_name_capacity_rejection_count,
604     );
605 }