diff --git a/build.zig b/build.zig index 80be26d..7794f15 100644 --- a/build.zig +++ b/build.zig @@ -19,6 +19,35 @@ pub fn build(b: *std.Build) void { }); const tests = b.addTest(.{ .root_module = mod }); + + // add interop test fixtures (lazy — only fetched when running tests) + if (b.lazyDependency("atproto-interop-tests", .{})) |interop| { + const interop_files = .{ + // syntax fixtures + .{ "tid_syntax_valid", "syntax/tid_syntax_valid.txt" }, + .{ "tid_syntax_invalid", "syntax/tid_syntax_invalid.txt" }, + .{ "did_syntax_valid", "syntax/did_syntax_valid.txt" }, + .{ "did_syntax_invalid", "syntax/did_syntax_invalid.txt" }, + .{ "handle_syntax_valid", "syntax/handle_syntax_valid.txt" }, + .{ "handle_syntax_invalid", "syntax/handle_syntax_invalid.txt" }, + .{ "nsid_syntax_valid", "syntax/nsid_syntax_valid.txt" }, + .{ "nsid_syntax_invalid", "syntax/nsid_syntax_invalid.txt" }, + .{ "recordkey_syntax_valid", "syntax/recordkey_syntax_valid.txt" }, + .{ "recordkey_syntax_invalid", "syntax/recordkey_syntax_invalid.txt" }, + .{ "aturi_syntax_valid", "syntax/aturi_syntax_valid.txt" }, + .{ "aturi_syntax_invalid", "syntax/aturi_syntax_invalid.txt" }, + // crypto fixtures + .{ "signature_fixtures", "crypto/signature-fixtures.json" }, + // mst fixtures + .{ "mst_key_heights", "mst/key_heights.json" }, + }; + inline for (interop_files) |entry| { + tests.root_module.addAnonymousImport(entry[0], .{ + .root_source_file = interop.path(entry[1]), + }); + } + } + const run_tests = b.addRunArtifact(tests); const test_step = b.step("test", "run unit tests"); diff --git a/build.zig.zon b/build.zig.zon index 6ac00c4..733934e 100644 --- a/build.zig.zon +++ b/build.zig.zon @@ -1,6 +1,6 @@ .{ .name = .zat, - .version = "0.1.7", + .version = "0.1.8", .fingerprint = 0x8da9db57ee82fbe4, .minimum_zig_version = "0.15.0", .dependencies = .{ @@ -8,6 +8,11 @@ .url = "https://github.com/karlseguin/websocket.zig/archive/97fefafa59cc78ce177cff540b8685cd7f699276.tar.gz", .hash = "websocket-0.1.0-ZPISdRlzAwBB_Bz2UMMqxYqF6YEVTIBoFsbzwPUJTHIc", }, + .@"atproto-interop-tests" = .{ + .url = "https://github.com/bluesky-social/atproto-interop-tests/archive/35bb5638ab1e5ce71fb88a0c95953fc557ef1925.tar.gz", + .hash = "N-V-__8AAIp5AQCe4JjGmPl8CplTkCis8PF1qvn7QX6GwDfu", + .lazy = true, + }, }, .paths = .{ "build.zig", diff --git a/src/internal/interop_tests.zig b/src/internal/interop_tests.zig new file mode 100644 index 0000000..4ec36e4 --- /dev/null +++ b/src/internal/interop_tests.zig @@ -0,0 +1,275 @@ +//! interop tests against bluesky-social/atproto-interop-tests fixtures +//! +//! validates zat's parsers and crypto against the official test vectors. + +const std = @import("std"); + +// types under test +const Tid = @import("tid.zig").Tid; +const Did = @import("did.zig").Did; +const Handle = @import("handle.zig").Handle; +const Nsid = @import("nsid.zig").Nsid; +const Rkey = @import("rkey.zig").Rkey; +const AtUri = @import("at_uri.zig").AtUri; + +// crypto +const jwt = @import("jwt.zig"); +const multibase = @import("multibase.zig"); +const multicodec = @import("multicodec.zig"); + +// === helpers === + +fn LineIterator(comptime sentinel: ?u8) type { + return struct { + inner: std.mem.SplitIterator(u8, .scalar), + + const Self = @This(); + + fn init(data: []const u8) Self { + // strip trailing sentinel if present (some files end with \n) + const trimmed = if (sentinel) |s| + if (data.len > 0 and data[data.len - 1] == s) data[0 .. data.len - 1] else data + else + data; + return .{ .inner = std.mem.splitScalar(u8, trimmed, '\n') }; + } + + fn next(self: *Self) ?[]const u8 { + while (self.inner.next()) |line| { + // skip blank lines and comments + if (line.len == 0) continue; + if (line[0] == '#') continue; + // strip trailing \r for windows line endings + const trimmed = if (line.len > 0 and line[line.len - 1] == '\r') + line[0 .. line.len - 1] + else + line; + if (trimmed.len == 0) continue; + return trimmed; + } + return null; + } + }; +} + +fn testLinesSentinel(comptime data: [:0]const u8) LineIterator(0) { + return LineIterator(0).init(data); +} + +/// run syntax validation tests for a parser type +fn syntaxTest( + comptime valid_data: [:0]const u8, + comptime invalid_data: [:0]const u8, + comptime parseFn: anytype, +) !void { + // test valid lines + var valid_lines = testLinesSentinel(valid_data); + var valid_count: usize = 0; + while (valid_lines.next()) |line| { + if (parseFn(line) == null) { + std.debug.print("FAIL: expected valid, got null for: '{s}'\n", .{line}); + return error.ExpectedValid; + } + valid_count += 1; + } + if (valid_count == 0) return error.NoTestCases; + + // test invalid lines + var invalid_lines = testLinesSentinel(invalid_data); + var invalid_count: usize = 0; + while (invalid_lines.next()) |line| { + if (parseFn(line) != null) { + std.debug.print("FAIL: expected null, got valid for: '{s}'\n", .{line}); + return error.ExpectedInvalid; + } + invalid_count += 1; + } + if (invalid_count == 0) return error.NoTestCases; +} + +// === tier 1: syntax validation === + +test "interop: tid syntax" { + try syntaxTest( + @embedFile("tid_syntax_valid"), + @embedFile("tid_syntax_invalid"), + Tid.parse, + ); +} + +test "interop: did syntax" { + try syntaxTest( + @embedFile("did_syntax_valid"), + @embedFile("did_syntax_invalid"), + Did.parse, + ); +} + +test "interop: handle syntax" { + try syntaxTest( + @embedFile("handle_syntax_valid"), + @embedFile("handle_syntax_invalid"), + Handle.parse, + ); +} + +test "interop: nsid syntax" { + try syntaxTest( + @embedFile("nsid_syntax_valid"), + @embedFile("nsid_syntax_invalid"), + Nsid.parse, + ); +} + +test "interop: rkey syntax" { + try syntaxTest( + @embedFile("recordkey_syntax_valid"), + @embedFile("recordkey_syntax_invalid"), + Rkey.parse, + ); +} + +test "interop: aturi syntax" { + try syntaxTest( + @embedFile("aturi_syntax_valid"), + @embedFile("aturi_syntax_invalid"), + AtUri.parse, + ); +} + +// === tier 2: crypto signature verification === + +fn base64StdDecode(allocator: std.mem.Allocator, input: []const u8) ![]u8 { + // try standard (padded) first, fall back to no-pad + const decoder = if (input.len > 0 and input[input.len - 1] == '=') + &std.base64.standard.Decoder + else + &std.base64.standard_no_pad.Decoder; + + const size = decoder.calcSizeForSlice(input) catch return error.InvalidBase64; + const output = try allocator.alloc(u8, size); + errdefer allocator.free(output); + decoder.decode(output, input) catch return error.InvalidBase64; + return output; +} + +test "interop: crypto signature verification" { + const allocator = std.testing.allocator; + + const fixture_json = @embedFile("signature_fixtures"); + const parsed = try std.json.parseFromSlice(std.json.Value, allocator, fixture_json, .{}); + defer parsed.deinit(); + + const fixtures = parsed.value.array.items; + var tested: usize = 0; + + for (fixtures) |fixture| { + const obj = fixture.object; + + const comment = if (obj.get("comment")) |v| switch (v) { + .string => |s| s, + else => "?", + } else "?"; + + const message_b64 = obj.get("messageBase64").?.string; + const algorithm = obj.get("algorithm").?.string; + const pub_key_did = obj.get("publicKeyDid").?.string; + const sig_b64 = obj.get("signatureBase64").?.string; + const valid = obj.get("validSignature").?.bool; + + // extract multibase key from did:key (strip "did:key:" prefix) + const did_key_prefix = "did:key:"; + if (!std.mem.startsWith(u8, pub_key_did, did_key_prefix)) return error.InvalidDidKey; + const multibase_key = pub_key_did[did_key_prefix.len..]; + + // decode message and signature + const message = try base64StdDecode(allocator, message_b64); + defer allocator.free(message); + + const sig_bytes = base64StdDecode(allocator, sig_b64) catch |err| { + // DER-encoded sigs may fail to decode at expected length — that's fine for invalid + if (!valid) { + tested += 1; + continue; + } + return err; + }; + defer allocator.free(sig_bytes); + + // decode public key from multibase+multicodec (did:key format) + const key_bytes = try multibase.decode(allocator, multibase_key); + defer allocator.free(key_bytes); + + const parsed_key = try multicodec.parsePublicKey(key_bytes); + + // verify signature + const verify_result = if (std.mem.eql(u8, algorithm, "ES256K")) + jwt.verifySecp256k1(message, sig_bytes, parsed_key.raw) + else if (std.mem.eql(u8, algorithm, "ES256")) + jwt.verifyP256(message, sig_bytes, parsed_key.raw) + else + error.UnsupportedAlgorithm; + + if (valid) { + verify_result catch |err| { + std.debug.print("FAIL: expected valid signature but got {s}: {s}\n", .{ @errorName(err), comment }); + return error.ExpectedValidSignature; + }; + } else { + if (verify_result) |_| { + std.debug.print("FAIL: expected invalid signature but verified OK: {s}\n", .{comment}); + return error.ExpectedInvalidSignature; + } else |_| {} + } + + tested += 1; + } + + // should have tested all 6 fixtures + try std.testing.expect(tested == fixtures.len); +} + +// === tier 3: MST key heights === + +/// compute MST tree depth for a record key. +/// depth = count leading zero bits in SHA-256(key), divided by 2, rounded down. +fn mstKeyHeight(key: []const u8) u32 { + var digest: [32]u8 = undefined; + std.crypto.hash.sha2.Sha256.hash(key, &digest, .{}); + var leading_zeros: u32 = 0; + for (digest) |byte| { + if (byte == 0) { + leading_zeros += 8; + } else { + leading_zeros += @clz(byte); + break; + } + } + return leading_zeros / 2; +} + +test "interop: mst key heights" { + const allocator = std.testing.allocator; + + const fixture_json = @embedFile("mst_key_heights"); + const parsed = try std.json.parseFromSlice(std.json.Value, allocator, fixture_json, .{}); + defer parsed.deinit(); + + const fixtures = parsed.value.array.items; + var tested: usize = 0; + + for (fixtures) |fixture| { + const obj = fixture.object; + const key = obj.get("key").?.string; + const expected_height: u32 = @intCast(obj.get("height").?.integer); + + const actual = mstKeyHeight(key); + if (actual != expected_height) { + std.debug.print("FAIL: key '{s}': expected height {d}, got {d}\n", .{ key, expected_height, actual }); + return error.WrongHeight; + } + tested += 1; + } + + try std.testing.expect(tested > 0); +} diff --git a/src/root.zig b/src/root.zig index 4689bc7..96d3023 100644 --- a/src/root.zig +++ b/src/root.zig @@ -44,3 +44,10 @@ pub const car = @import("internal/car.zig"); pub const firehose = @import("internal/firehose.zig"); pub const FirehoseClient = firehose.FirehoseClient; pub const FirehoseEvent = firehose.Event; + +// interop tests (test-only, references resolved by build.zig lazy dependency) +comptime { + if (@import("builtin").is_test) { + _ = @import("internal/interop_tests.zig"); + } +}