lib/choir/src/serialization/binary/reader.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const binary = @import("root.zig");
3
4 pub fn Reader(comptime limits: binary.Limits) type {
5 return struct {
6 bytes: []const u8,
7 offset: usize = 0,
8 total_entries: usize = 0,
9
10 const Self = @This();
11
12 pub fn init(bytes: []const u8) binary.Error!Self {
13 if (bytes.len > limits.serialized_bytes) return error.LimitExceeded;
14 return .{ .bytes = bytes };
15 }
16
17 pub fn atEnd(self: Self) bool {
18 return self.offset == self.bytes.len;
19 }
20
21 pub fn readRaw(self: *Self, length: usize) binary.Error![]const u8 {
22 if (length > self.bytes.len - self.offset) return error.Truncated;
23 const value = self.bytes[self.offset..][0..length];
24 self.offset += length;
25 return value;
26 }
27
28 pub fn readInt(self: *Self, comptime T: type) binary.Error!T {
29 const value = try self.readRaw(@sizeOf(T));
30 return std.mem.readInt(T, value[0..@sizeOf(T)], .little);
31 }
32
33 pub fn readBool(self: *Self) binary.Error!bool {
34 return switch (try self.readInt(u8)) {
35 0 => false,
36 1 => true,
37 else => error.InvalidValue,
38 };
39 }
40
41 pub fn readTag(self: *Self, comptime T: type) binary.Error!T {
42 return try binary.decodeTag(T, try self.readInt(u8));
43 }
44
45 pub fn readCount(self: *Self) binary.Error!usize {
46 const count: usize = @intCast(try self.readInt(u32));
47 if (count > limits.collection_entries) return error.LimitExceeded;
48 if (count > self.bytes.len - self.offset) return error.Truncated;
49 self.total_entries = std.math.add(usize, self.total_entries, count) catch return error.LimitExceeded;
50 if (self.total_entries > limits.total_entries) return error.LimitExceeded;
51 return count;
52 }
53
54 pub fn readString(self: *Self) binary.Error![]const u8 {
55 return try self.readLengthPrefixed(limits.string_bytes);
56 }
57
58 pub fn readBlob(self: *Self) binary.Error![]const u8 {
59 return try self.readLengthPrefixed(limits.blob_bytes);
60 }
61
62 pub fn readOptionalString(self: *Self) binary.Error!?[]const u8 {
63 if (!try self.readBool()) return null;
64 return try self.readString();
65 }
66
67 pub fn readOptionalU16(self: *Self) binary.Error!?u16 {
68 if (!try self.readBool()) return null;
69 return try self.readInt(u16);
70 }
71
72 fn readLengthPrefixed(self: *Self, maximum: usize) binary.Error![]const u8 {
73 const length: usize = @intCast(try self.readInt(u32));
74 if (length > maximum) return error.LimitExceeded;
75 return try self.readRaw(length);
76 }
77 };
78 }
79
80 test "binary reader rejects malformed primitive and aggregate fields" {
81 const limits = binary.Limits{
82 .serialized_bytes = 128,
83 .string_bytes = 8,
84 .blob_bytes = 16,
85 .collection_entries = 4,
86 .total_entries = 8,
87 };
88 const TestReader = Reader(limits);
89
90 var invalid_bool = try TestReader.init(&.{2});
91 try std.testing.expectError(error.InvalidValue, invalid_bool.readBool());
92
93 var invalid_string_length: [@sizeOf(u32)]u8 = undefined;
94 std.mem.writeInt(u32, &invalid_string_length, limits.string_bytes + 1, .little);
95 var invalid_string = try TestReader.init(&invalid_string_length);
96 try std.testing.expectError(error.LimitExceeded, invalid_string.readString());
97
98 var invalid_blob_length: [@sizeOf(u32)]u8 = undefined;
99 std.mem.writeInt(u32, &invalid_blob_length, limits.blob_bytes + 1, .little);
100 var invalid_blob = try TestReader.init(&invalid_blob_length);
101 try std.testing.expectError(error.LimitExceeded, invalid_blob.readBlob());
102
103 const count_fields = limits.total_entries / limits.collection_entries + 1;
104 var aggregate_bytes: [limits.collection_entries + count_fields * @sizeOf(u32)]u8 = @splat(0);
105 for (0..count_fields) |index| {
106 const offset = index * @sizeOf(u32);
107 std.mem.writeInt(u32, aggregate_bytes[offset..][0..@sizeOf(u32)], limits.collection_entries, .little);
108 }
109 var aggregate = try TestReader.init(&aggregate_bytes);
110 for (0..count_fields - 1) |_| _ = try aggregate.readCount();
111 try std.testing.expectError(error.LimitExceeded, aggregate.readCount());
112 }