lib/sys/src/ancillary.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const builtin = @import("builtin");
  3 const capabilities = @import("capabilities.zig");
  4 
  5 const linux = std.os.linux;
  6 const posix = std.posix;
  7 
  8 pub const required_capabilities = capabilities.noLibc(&.{ .descriptors, .networking });
  9 pub const Descriptor = posix.fd_t;
 10 pub const maximum_descriptors = 253;
 11 
 12 pub const SendError = error{
 13     UnsupportedPlatform,
 14     EmptyMessage,
 15     TooManyDescriptors,
 16     WouldBlock,
 17     BrokenPipe,
 18     ConnectionResetByPeer,
 19     InvalidDescriptor,
 20     MessageTooLarge,
 21     SendFailed,
 22 };
 23 
 24 pub const ReceiveError = error{
 25     UnsupportedPlatform,
 26     EmptyBuffer,
 27     WouldBlock,
 28     ConnectionResetByPeer,
 29     DescriptorTruncated,
 30     MalformedControl,
 31     ReceiveFailed,
 32 };
 33 
 34 pub const Received = struct {
 35     byte_count: usize,
 36     descriptor_count: usize,
 37 };
 38 
 39 pub fn send(
 40     socket: Descriptor,
 41     bytes: []const u8,
 42     descriptors: []const Descriptor,
 43 ) SendError!usize {
 44     if (comptime builtin.os.tag != .linux) return error.UnsupportedPlatform;
 45     if (bytes.len == 0) return error.EmptyMessage;
 46     if (descriptors.len > maximum_descriptors) return error.TooManyDescriptors;
 47 
 48     var control: [controlSpace(maximum_descriptors)]u8 align(@alignOf(linux.cmsghdr)) = @splat(0);
 49     const control_len = encodeDescriptors(&control, descriptors);
 50     var iovec: posix.iovec_const = .{ .base = bytes.ptr, .len = bytes.len };
 51     const message: linux.msghdr_const = .{
 52         .name = null,
 53         .namelen = 0,
 54         .iov = (&iovec)[0..1],
 55         .iovlen = 1,
 56         .control = if (control_len == 0) null else &control,
 57         .controllen = control_len,
 58         .flags = 0,
 59     };
 60 
 61     while (true) {
 62         const result = linux.sendmsg(socket, &message, linux.MSG.NOSIGNAL);
 63         switch (linux.errno(result)) {
 64             .SUCCESS => return @intCast(result),
 65             .INTR => continue,
 66             .AGAIN => return error.WouldBlock,
 67             .PIPE => return error.BrokenPipe,
 68             .CONNRESET => return error.ConnectionResetByPeer,
 69             .BADF => return error.InvalidDescriptor,
 70             .MSGSIZE => return error.MessageTooLarge,
 71             else => return error.SendFailed,
 72         }
 73     }
 74 }
 75 
 76 pub fn receive(
 77     socket: Descriptor,
 78     bytes: []u8,
 79     descriptors: *[maximum_descriptors]Descriptor,
 80 ) ReceiveError!Received {
 81     if (comptime builtin.os.tag != .linux) return error.UnsupportedPlatform;
 82     if (bytes.len == 0) return error.EmptyBuffer;
 83 
 84     var control: [controlSpace(maximum_descriptors)]u8 align(@alignOf(linux.cmsghdr)) = @splat(0);
 85     var iovec: posix.iovec = .{ .base = bytes.ptr, .len = bytes.len };
 86     var message: linux.msghdr = .{
 87         .name = null,
 88         .namelen = 0,
 89         .iov = (&iovec)[0..1],
 90         .iovlen = 1,
 91         .control = &control,
 92         .controllen = control.len,
 93         .flags = 0,
 94     };
 95 
 96     while (true) {
 97         const result = linux.recvmsg(socket, &message, linux.MSG.CMSG_CLOEXEC);
 98         switch (linux.errno(result)) {
 99             .SUCCESS => {
100                 var descriptor_count: usize = 0;
101                 parseDescriptors(control[0..message.controllen], descriptors, &descriptor_count) catch |err| {
102                     closeAll(descriptors[0..descriptor_count]);
103                     return err;
104                 };
105                 if (message.flags & linux.MSG.CTRUNC != 0) {
106                     closeAll(descriptors[0..descriptor_count]);
107                     return error.DescriptorTruncated;
108                 }
109                 return .{
110                     .byte_count = @intCast(result),
111                     .descriptor_count = descriptor_count,
112                 };
113             },
114             .INTR => continue,
115             .AGAIN => return error.WouldBlock,
116             .CONNRESET => return error.ConnectionResetByPeer,
117             else => return error.ReceiveFailed,
118         }
119     }
120 }
121 
122 fn encodeDescriptors(
123     control: *[controlSpace(maximum_descriptors)]u8,
124     descriptors: []const Descriptor,
125 ) usize {
126     if (descriptors.len == 0) return 0;
127 
128     const header: *linux.cmsghdr = @ptrCast(@alignCast(control));
129     header.* = .{
130         .len = controlLength(descriptors.len),
131         .level = linux.SOL.SOCKET,
132         .type = linux.SCM.RIGHTS,
133     };
134     const bytes = std.mem.sliceAsBytes(descriptors);
135     @memcpy(control[controlDataOffset()..][0..bytes.len], bytes);
136     return controlSpace(descriptors.len);
137 }
138 
139 fn parseDescriptors(
140     control: []align(@alignOf(linux.cmsghdr)) const u8,
141     descriptors: *[maximum_descriptors]Descriptor,
142     descriptor_count: *usize,
143 ) ReceiveError!void {
144     var offset: usize = 0;
145     while (offset + @sizeOf(linux.cmsghdr) <= control.len) {
146         const header: *const linux.cmsghdr = @ptrCast(@alignCast(control.ptr + offset));
147         if (header.len < controlDataOffset() or header.len > control.len - offset) {
148             return error.MalformedControl;
149         }
150 
151         if (header.level == linux.SOL.SOCKET and header.type == linux.SCM.RIGHTS) {
152             const data_len = header.len - controlDataOffset();
153             if (data_len % @sizeOf(Descriptor) != 0) return error.MalformedControl;
154             const count = data_len / @sizeOf(Descriptor);
155             if (descriptor_count.* + count > descriptors.len) return error.DescriptorTruncated;
156             const source = control[offset + controlDataOffset() ..][0..data_len];
157             const destination = descriptors[descriptor_count.*..][0..count];
158             @memcpy(std.mem.sliceAsBytes(destination), source);
159             descriptor_count.* += count;
160         }
161 
162         const next = offset + controlAlign(header.len);
163         if (next <= offset or next > control.len) return error.MalformedControl;
164         offset = next;
165     }
166 }
167 
168 fn closeAll(descriptors: []const Descriptor) void {
169     for (descriptors) |descriptor| _ = linux.close(descriptor);
170 }
171 
172 fn controlAlign(size: usize) usize {
173     const alignment = @sizeOf(usize);
174     const mask: usize = alignment - 1;
175     return (size + mask) & ~mask;
176 }
177 
178 fn controlDataOffset() usize {
179     return controlAlign(@sizeOf(linux.cmsghdr));
180 }
181 
182 fn controlLength(descriptor_count: usize) usize {
183     return controlDataOffset() + descriptor_count * @sizeOf(Descriptor);
184 }
185 
186 fn controlSpace(descriptor_count: usize) usize {
187     return controlDataOffset() + controlAlign(descriptor_count * @sizeOf(Descriptor));
188 }
189 
190 test "ancillary descriptor control sizing follows native alignment" {
191     try std.testing.expect(controlLength(1) <= controlSpace(1));
192     try std.testing.expectEqual(@as(usize, 0), controlSpace(0) % @sizeOf(usize));
193     try std.testing.expectEqual(@as(usize, 0), controlSpace(maximum_descriptors) % @sizeOf(usize));
194 }