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 }