lib/trace/src/checkpoint.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 const Allocator = std.mem.Allocator;
4
5 pub const CaptureFn = *const fn (context: *anyopaque, allocator: Allocator) anyerror![]u8;
6 pub const PrepareRestoreFn = *const fn (context: *anyopaque, bytes: []const u8) anyerror!void;
7 pub const CommitRestoreFn = *const fn (context: *anyopaque) void;
8 pub const CancelRestoreFn = *const fn (context: *anyopaque) void;
9
10 pub const Provider = struct {
11 context: *anyopaque,
12 capture: CaptureFn,
13 prepare_restore: PrepareRestoreFn,
14 commit_restore: CommitRestoreFn,
15 cancel_restore: CancelRestoreFn,
16
17 pub fn captureAlloc(self: Provider, allocator: Allocator) ![]u8 {
18 return try self.capture(self.context, allocator);
19 }
20
21 pub fn prepareRestoreBytes(self: Provider, bytes: []const u8) !void {
22 try self.prepare_restore(self.context, bytes);
23 }
24
25 pub fn commitRestore(self: Provider) void {
26 self.commit_restore(self.context);
27 }
28
29 pub fn cancelRestore(self: Provider) void {
30 self.cancel_restore(self.context);
31 }
32 };
33
34 const ProviderTestState = struct {
35 value: [32]u8 = undefined,
36 value_len: usize = 0,
37 staged: [32]u8 = undefined,
38 staged_len: usize = 0,
39 prepared: bool = false,
40
41 fn init(bytes: []const u8) ProviderTestState {
42 var self: ProviderTestState = .{};
43 self.setValue(bytes);
44 return self;
45 }
46
47 fn setValue(self: *ProviderTestState, bytes: []const u8) void {
48 std.debug.assert(bytes.len <= self.value.len);
49 @memcpy(self.value[0..bytes.len], bytes);
50 self.value_len = bytes.len;
51 }
52
53 fn current(self: *const ProviderTestState) []const u8 {
54 return self.value[0..self.value_len];
55 }
56
57 fn capture(context: *anyopaque, allocator: Allocator) ![]u8 {
58 const self: *@This() = @ptrCast(@alignCast(context));
59 return try allocator.dupe(u8, self.current());
60 }
61
62 fn prepareRestore(context: *anyopaque, bytes: []const u8) !void {
63 const self: *@This() = @ptrCast(@alignCast(context));
64 if (bytes.len > self.staged.len) return error.TestProviderCapacityExceeded;
65 @memcpy(self.staged[0..bytes.len], bytes);
66 self.staged_len = bytes.len;
67 self.prepared = true;
68 }
69
70 fn commitRestore(context: *anyopaque) void {
71 const self: *@This() = @ptrCast(@alignCast(context));
72 std.debug.assert(self.prepared);
73 @memcpy(self.value[0..self.staged_len], self.staged[0..self.staged_len]);
74 self.value_len = self.staged_len;
75 self.prepared = false;
76 }
77
78 fn cancelRestore(context: *anyopaque) void {
79 const self: *@This() = @ptrCast(@alignCast(context));
80 self.prepared = false;
81 }
82 };
83
84 test "provider stages restore before an infallible commit" {
85 var state = ProviderTestState.init("before");
86 const provider = Provider{
87 .context = &state,
88 .capture = ProviderTestState.capture,
89 .prepare_restore = ProviderTestState.prepareRestore,
90 .commit_restore = ProviderTestState.commitRestore,
91 .cancel_restore = ProviderTestState.cancelRestore,
92 };
93
94 {
95 const bytes = try provider.captureAlloc(std.testing.allocator);
96 defer std.testing.allocator.free(bytes);
97 state.setValue("after");
98 try provider.prepareRestoreBytes(bytes);
99 try std.testing.expectEqualStrings("after", state.current());
100 provider.commitRestore();
101 }
102
103 try std.testing.expectEqualStrings("before", state.current());
104 }
105
106 test "provider cancel discards a staged restore" {
107 var state = ProviderTestState.init("before");
108 const provider = Provider{
109 .context = &state,
110 .capture = ProviderTestState.capture,
111 .prepare_restore = ProviderTestState.prepareRestore,
112 .commit_restore = ProviderTestState.commitRestore,
113 .cancel_restore = ProviderTestState.cancelRestore,
114 };
115
116 try provider.prepareRestoreBytes("after");
117 provider.cancelRestore();
118
119 try std.testing.expectEqualStrings("before", state.current());
120 try std.testing.expect(!state.prepared);
121 }