diff --git a/CHANGELOG.md b/CHANGELOG.md index d79edea..34291ad 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,13 @@ # changelog +## 0.1.9 + +- **feat**: merkle search tree (MST) — `mst.Mst` with `put`, `get`, `delete`, `rootCid` +- **feat**: ECDSA signing — `signSecp256k1`, `signP256` with low-S normalization (RFC 6979) +- **feat**: `did:key` construction — `multicodec.formatDidKey`, `multicodec.encodePublicKey` +- **feat**: multibase encoding — base58btc encode, base32lower encode/decode +- interop tests: MST common prefix (13 vectors), commit proofs (6 fixtures) + ## 0.1.8 - **fix**: NSID parser rejects TLD starting with digit (e.g. `1.0.0.127.record`) diff --git a/build.zig b/build.zig index 7794f15..6d09371 100644 --- a/build.zig +++ b/build.zig @@ -40,6 +40,8 @@ pub fn build(b: *std.Build) void { .{ "signature_fixtures", "crypto/signature-fixtures.json" }, // mst fixtures .{ "mst_key_heights", "mst/key_heights.json" }, + .{ "common_prefix", "mst/common_prefix.json" }, + .{ "commit_proofs", "firehose/commit-proof-fixtures.json" }, }; inline for (interop_files) |entry| { tests.root_module.addAnonymousImport(entry[0], .{ diff --git a/build.zig.zon b/build.zig.zon index 733934e..3a8b10f 100644 --- a/build.zig.zon +++ b/build.zig.zon @@ -1,6 +1,6 @@ .{ .name = .zat, - .version = "0.1.8", + .version = "0.1.9", .fingerprint = 0x8da9db57ee82fbe4, .minimum_zig_version = "0.15.0", .dependencies = .{ diff --git a/src/internal/interop_tests.zig b/src/internal/interop_tests.zig index 4ec36e4..feb8004 100644 --- a/src/internal/interop_tests.zig +++ b/src/internal/interop_tests.zig @@ -17,6 +17,10 @@ const jwt = @import("jwt.zig"); const multibase = @import("multibase.zig"); const multicodec = @import("multicodec.zig"); +// mst +const mst = @import("mst.zig"); +const cbor = @import("cbor.zig"); + // === helpers === fn LineIterator(comptime sentinel: ?u8) type { @@ -229,24 +233,7 @@ test "interop: crypto signature verification" { 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; -} +// === tier 3: MST === test "interop: mst key heights" { const allocator = std.testing.allocator; @@ -263,7 +250,7 @@ test "interop: mst key heights" { const key = obj.get("key").?.string; const expected_height: u32 = @intCast(obj.get("height").?.integer); - const actual = mstKeyHeight(key); + const actual = mst.keyHeight(key); if (actual != expected_height) { std.debug.print("FAIL: key '{s}': expected height {d}, got {d}\n", .{ key, expected_height, actual }); return error.WrongHeight; @@ -273,3 +260,113 @@ test "interop: mst key heights" { try std.testing.expect(tested > 0); } + +test "interop: mst common prefix" { + const allocator = std.testing.allocator; + + const fixture_json = @embedFile("common_prefix"); + 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 left = obj.get("left").?.string; + const right = obj.get("right").?.string; + const expected_len: usize = @intCast(obj.get("len").?.integer); + + const actual = mst.commonPrefixLen(left, right); + if (actual != expected_len) { + std.debug.print("FAIL: commonPrefixLen('{s}', '{s}'): expected {d}, got {d}\n", .{ left, right, expected_len, actual }); + return error.WrongPrefixLen; + } + tested += 1; + } + + try std.testing.expect(tested == 13); +} + +test "interop: mst commit proofs" { + const allocator = std.testing.allocator; + + const fixture_json = @embedFile("commit_proofs"); + 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| { + var arena = std.heap.ArenaAllocator.init(allocator); + defer arena.deinit(); + const a = arena.allocator(); + + const obj = fixture.object; + const comment = if (obj.get("comment")) |v| switch (v) { + .string => |s| s, + else => "?", + } else "?"; + + // parse leaf value CID + const leaf_value_str = obj.get("leafValue").?.string; + const leaf_cid = try mst.parseCidString(a, leaf_value_str); + + // build initial tree from keys + var tree = mst.Mst.init(a); + const keys = obj.get("keys").?.array.items; + for (keys) |key_val| { + try tree.put(key_val.string, leaf_cid); + } + + // verify root before commit + const root_before_str = obj.get("rootBeforeCommit").?.string; + const expected_before = try mst.parseCidString(a, root_before_str); + + const actual_before = try tree.rootCid(); + if (!std.mem.eql(u8, actual_before.raw, expected_before.raw)) { + std.debug.print("FAIL [{s}]: rootBeforeCommit mismatch\n", .{comment}); + std.debug.print(" expected: {s}\n", .{root_before_str}); + // print hex for debugging + std.debug.print(" expected raw ({d}): ", .{expected_before.raw.len}); + for (expected_before.raw) |b| std.debug.print("{x:0>2}", .{b}); + std.debug.print("\n actual raw ({d}): ", .{actual_before.raw.len}); + for (actual_before.raw) |b| std.debug.print("{x:0>2}", .{b}); + std.debug.print("\n", .{}); + return error.RootBeforeMismatch; + } + + // apply adds + const adds = obj.get("adds").?.array.items; + for (adds) |add_val| { + try tree.put(add_val.string, leaf_cid); + } + + // apply dels + const dels = obj.get("dels").?.array.items; + for (dels) |del_val| { + try tree.delete(del_val.string); + } + + // verify root after commit + const root_after_str = obj.get("rootAfterCommit").?.string; + const expected_after = try mst.parseCidString(a, root_after_str); + + const actual_after = try tree.rootCid(); + if (!std.mem.eql(u8, actual_after.raw, expected_after.raw)) { + std.debug.print("FAIL [{s}]: rootAfterCommit mismatch\n", .{comment}); + std.debug.print(" expected: {s}\n", .{root_after_str}); + std.debug.print(" expected raw ({d}): ", .{expected_after.raw.len}); + for (expected_after.raw) |b| std.debug.print("{x:0>2}", .{b}); + std.debug.print("\n actual raw ({d}): ", .{actual_after.raw.len}); + for (actual_after.raw) |b| std.debug.print("{x:0>2}", .{b}); + std.debug.print("\n", .{}); + return error.RootAfterMismatch; + } + + tested += 1; + } + + try std.testing.expect(tested == 6); +} diff --git a/src/internal/jwt.zig b/src/internal/jwt.zig index 72c58df..5c43af6 100644 --- a/src/internal/jwt.zig +++ b/src/internal/jwt.zig @@ -252,42 +252,55 @@ const p256_half_order: [32]u8 = .{ 0x79, 0xDC, 0xE5, 0x61, 0x7E, 0x31, 0x92, 0xA8, }; -pub fn verifySecp256k1(message: []const u8, sig_bytes: []const u8, public_key_raw: []const u8) !void { - const Scheme = crypto.sign.ecdsa.EcdsaSecp256k1Sha256; +/// ECDSA signature (r || s, 64 bytes) +pub const Signature = struct { + bytes: [64]u8, +}; - // parse signature (r || s, 64 bytes) - if (sig_bytes.len != 64) return error.InvalidSignature; - const sig = Scheme.Signature.fromBytes(sig_bytes[0..64].*); +/// sign a message with deterministic RFC 6979 ECDSA and low-S normalization +fn signEcdsa(comptime Scheme: type, comptime Curve: type, comptime half_order: [32]u8, message: []const u8, secret_key_bytes: []const u8) !Signature { + if (secret_key_bytes.len != 32) return error.InvalidSecretKey; + const sk = Scheme.SecretKey.fromBytes(secret_key_bytes[0..32].*) catch return error.InvalidSecretKey; + const kp = Scheme.KeyPair.fromSecretKey(sk) catch return error.InvalidSecretKey; - // reject high-S signatures (atproto requires low-S) - rejectHighS(secp256k1_half_order, sig.s) catch return error.SignatureVerificationFailed; + var sig = kp.sign(message, null) catch return error.SigningFailed; - // parse public key from SEC1 compressed format - if (public_key_raw.len != 33) return error.InvalidPublicKey; - const public_key = Scheme.PublicKey.fromSec1(public_key_raw) catch return error.InvalidPublicKey; + if (bigEndianGt(sig.s, half_order)) { + sig.s = Curve.scalar.neg(sig.s, .big) catch return error.SigningFailed; + } - // verify - sig.verify(message, public_key) catch return error.SignatureVerificationFailed; + return .{ .bytes = sig.toBytes() }; } -pub fn verifyP256(message: []const u8, sig_bytes: []const u8, public_key_raw: []const u8) !void { - const Scheme = crypto.sign.ecdsa.EcdsaP256Sha256; - - // parse signature (r || s, 64 bytes) +/// verify an ECDSA signature, rejecting high-S +fn verifyEcdsa(comptime Scheme: type, comptime half_order: [32]u8, message: []const u8, sig_bytes: []const u8, public_key_raw: []const u8) !void { if (sig_bytes.len != 64) return error.InvalidSignature; const sig = Scheme.Signature.fromBytes(sig_bytes[0..64].*); - // reject high-S signatures (atproto requires low-S) - rejectHighS(p256_half_order, sig.s) catch return error.SignatureVerificationFailed; + rejectHighS(half_order, sig.s) catch return error.SignatureVerificationFailed; - // parse public key from SEC1 compressed format if (public_key_raw.len != 33) return error.InvalidPublicKey; const public_key = Scheme.PublicKey.fromSec1(public_key_raw) catch return error.InvalidPublicKey; - // verify sig.verify(message, public_key) catch return error.SignatureVerificationFailed; } +pub fn signSecp256k1(message: []const u8, secret_key_bytes: []const u8) !Signature { + return signEcdsa(crypto.sign.ecdsa.EcdsaSecp256k1Sha256, crypto.ecc.Secp256k1, secp256k1_half_order, message, secret_key_bytes); +} + +pub fn signP256(message: []const u8, secret_key_bytes: []const u8) !Signature { + return signEcdsa(crypto.sign.ecdsa.EcdsaP256Sha256, crypto.ecc.P256, p256_half_order, message, secret_key_bytes); +} + +pub fn verifySecp256k1(message: []const u8, sig_bytes: []const u8, public_key_raw: []const u8) !void { + return verifyEcdsa(crypto.sign.ecdsa.EcdsaSecp256k1Sha256, secp256k1_half_order, message, sig_bytes, public_key_raw); +} + +pub fn verifyP256(message: []const u8, sig_bytes: []const u8, public_key_raw: []const u8) !void { + return verifyEcdsa(crypto.sign.ecdsa.EcdsaP256Sha256, p256_half_order, message, sig_bytes, public_key_raw); +} + // === tests === test "parse jwt structure" { @@ -369,3 +382,66 @@ test "reject signature with wrong key" { // should fail verification with wrong key try std.testing.expectError(error.SignatureVerificationFailed, jwt.verify(wrong_key)); } + +test "sign and verify round-trip - secp256k1" { + // generate a deterministic keypair using a fixed seed + const Scheme = crypto.sign.ecdsa.EcdsaSecp256k1Sha256; + const sk_bytes = [_]u8{ + 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, + 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f, 0x10, + 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, + 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e, 0x1f, 0x20, + }; + + const message = "hello atproto"; + const sig = try signSecp256k1(message, &sk_bytes); + + // verify low-S: s must be <= half_order + const s = sig.bytes[32..64].*; + try std.testing.expect(!bigEndianGt(s, secp256k1_half_order)); + + // verify with the corresponding public key + const sk = try Scheme.SecretKey.fromBytes(sk_bytes); + const kp = try Scheme.KeyPair.fromSecretKey(sk); + const pk_bytes = kp.public_key.toCompressedSec1(); + + try verifySecp256k1(message, &sig.bytes, &pk_bytes); +} + +test "sign and verify round-trip - P-256" { + const Scheme = crypto.sign.ecdsa.EcdsaP256Sha256; + const sk_bytes = [_]u8{ + 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, + 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f, 0x30, + 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, + 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f, 0x40, + }; + + const message = "hello atproto p256"; + const sig = try signP256(message, &sk_bytes); + + // verify low-S + const s = sig.bytes[32..64].*; + try std.testing.expect(!bigEndianGt(s, p256_half_order)); + + // verify with the corresponding public key + const sk = try Scheme.SecretKey.fromBytes(sk_bytes); + const kp = try Scheme.KeyPair.fromSecretKey(sk); + const pk_bytes = kp.public_key.toCompressedSec1(); + + try verifyP256(message, &sig.bytes, &pk_bytes); +} + +test "sign produces deterministic signatures" { + const sk_bytes = [_]u8{ + 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, + 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f, 0x10, + 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, + 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e, 0x1f, 0x20, + }; + const message = "deterministic test"; + + const sig1 = try signSecp256k1(message, &sk_bytes); + const sig2 = try signSecp256k1(message, &sk_bytes); + try std.testing.expectEqualSlices(u8, &sig1.bytes, &sig2.bytes); +} diff --git a/src/internal/mst.zig b/src/internal/mst.zig new file mode 100644 index 0000000..622cbcc --- /dev/null +++ b/src/internal/mst.zig @@ -0,0 +1,781 @@ +//! merkle search tree (MST) +//! +//! the AT Protocol repository data structure. a deterministic search tree +//! where each key's tree layer is derived from the leading zero bits of +//! SHA-256(key). keys are stored sorted within each node, with subtree +//! pointers interleaved between entries. +//! +//! see: https://atproto.com/specs/repository#mst-structure + +const std = @import("std"); +const cbor = @import("cbor.zig"); +const multibase = @import("multibase.zig"); +const Allocator = std.mem.Allocator; + +/// compute MST tree layer for a key. +/// layer = count leading zero bits in SHA-256(key), divided by 2, rounded down. +pub fn keyHeight(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; +} + +/// byte-level common prefix length between two strings +pub fn commonPrefixLen(a: []const u8, b: []const u8) usize { + const min_len = @min(a.len, b.len); + var i: usize = 0; + while (i < min_len) : (i += 1) { + if (a[i] != b[i]) break; + } + return i; +} + +/// parse a CID string (base32lower multibase, e.g. "bafyrei...") +pub fn parseCidString(allocator: Allocator, s: []const u8) !cbor.Cid { + if (s.len == 0) return error.InvalidCid; + // strip 'b' multibase prefix and decode base32lower + if (s[0] != 'b') return error.UnsupportedEncoding; + const raw = try multibase.base32lower.decode(allocator, s[1..]); + return .{ .raw = raw }; +} + +/// MST node. stores a left subtree pointer and a list of entries. +/// each entry has a key, CID value, and optional right subtree. +const Node = struct { + left: ?*Node, + entries: std.ArrayList(Entry), + + const Entry = struct { + key: []const u8, + value: cbor.Cid, + right: ?*Node, + }; + + fn init() Node { + return .{ + .left = null, + .entries = .{}, + }; + } +}; + +/// merkle search tree +pub const Mst = struct { + allocator: Allocator, + root: ?*Node, + root_layer: ?u32, + + pub fn init(allocator: Allocator) Mst { + return .{ + .allocator = allocator, + .root = null, + .root_layer = null, + }; + } + + /// insert or update a key-value pair + pub fn put(self: *Mst, key: []const u8, value: cbor.Cid) !void { + const height = keyHeight(key); + + if (self.root == null) { + // empty tree: create root at key's height + const node = try self.createNode(); + try node.entries.append(self.allocator, .{ + .key = try self.allocator.dupe(u8, key), + .value = value, + .right = null, + }); + self.root = node; + self.root_layer = height; + return; + } + + const root_layer = self.root_layer.?; + + if (height > root_layer) { + // key belongs above the current root — lift + self.root = try self.insertAbove(self.root.?, root_layer, key, value, height); + self.root_layer = height; + } else if (height == root_layer) { + // key belongs at root layer + self.root = try self.insertAtLayer(self.root.?, key, value, height); + } else { + // key belongs below — recurse into subtree + try self.insertBelow(self.root.?, root_layer, key, value, height); + } + } + + /// look up a key, returning its CID value if present + pub fn get(self: *const Mst, key: []const u8) ?cbor.Cid { + return findKey(self.root, self.root_layer orelse return null, key, keyHeight(key)); + } + + fn findKey(maybe_node: ?*Node, layer: u32, key: []const u8, height: u32) ?cbor.Cid { + const node = maybe_node orelse return null; + + if (height == layer) { + for (node.entries.items) |entry| { + const cmp = std.mem.order(u8, key, entry.key); + if (cmp == .eq) return entry.value; + if (cmp == .lt) return null; + } + return null; + } + + // height < layer: recurse into the subtree gap containing key + for (node.entries.items, 0..) |entry, i| { + if (std.mem.order(u8, key, entry.key) == .lt) { + const subtree = if (i == 0) node.left else node.entries.items[i - 1].right; + return findKey(subtree, layer - 1, key, height); + } + } + // after all entries + const last_right = if (node.entries.items.len > 0) + node.entries.items[node.entries.items.len - 1].right + else + node.left; + return findKey(last_right, layer - 1, key, height); + } + + /// delete a key from the tree + pub fn delete(self: *Mst, key: []const u8) !void { + if (self.root == null) return; + try self.deleteFromNode(self.root.?, self.root_layer.?, key); + // trim: if root has no entries and only left subtree, collapse + while (self.root) |root| { + if (root.entries.items.len == 0) { + if (root.left) |left| { + self.root = left; + if (self.root_layer.? > 0) { + self.root_layer = self.root_layer.? - 1; + } else { + self.root = null; + self.root_layer = null; + break; + } + } else { + self.root = null; + self.root_layer = null; + break; + } + } else break; + } + } + + fn deleteFromNode(self: *Mst, node: *Node, layer: u32, key: []const u8) !void { + const height = keyHeight(key); + + if (height == layer) { + // find and remove the entry + for (node.entries.items, 0..) |entry, i| { + if (std.mem.eql(u8, entry.key, key)) { + // merge left and right subtrees around the deleted entry + const left_sub = if (i == 0) node.left else node.entries.items[i - 1].right; + const right_sub = entry.right; + const merged = try self.mergeSubtrees(left_sub, right_sub); + + if (i == 0) { + node.left = merged; + } else { + node.entries.items[i - 1].right = merged; + } + + self.allocator.free(entry.key); + _ = node.entries.orderedRemove(i); + return; + } + } + return; // key not found + } + + // height < layer: recurse into the appropriate gap + if (node.entries.items.len == 0) { + if (node.left) |left| { + try self.deleteFromNode(left, layer - 1, key); + } + return; + } + + for (node.entries.items, 0..) |entry, i| { + if (std.mem.order(u8, key, entry.key) == .lt) { + const subtree = if (i == 0) &node.left else &node.entries.items[i - 1].right; + if (subtree.*) |sub| { + try self.deleteFromNode(sub, layer - 1, key); + } + return; + } + } + // after all entries + const last = &node.entries.items[node.entries.items.len - 1].right; + if (last.*) |sub| { + try self.deleteFromNode(sub, layer - 1, key); + } + } + + /// merge two subtrees that were separated by a deleted entry. + /// both nodes are at the same layer. concatenate their entries + /// and recursively merge if the junction creates adjacent children. + /// follows the Go reference `appendMerge` / `mergeNodes` algorithm. + fn mergeSubtrees(self: *Mst, left: ?*Node, right: ?*Node) !?*Node { + if (left == null) return right; + if (right == null) return left; + + const l = left.?; + const r = right.?; + + // create merged node: takes left's `left` pointer and all entries from both + const merged = try self.createNode(); + merged.left = l.left; + + // copy left entries + for (l.entries.items) |entry| { + try merged.entries.append(self.allocator, entry); + } + + // check junction: last entry of left's `right` vs right's `left` + if (merged.entries.items.len > 0) { + const last = &merged.entries.items[merged.entries.items.len - 1]; + if (last.right != null and r.left != null) { + // both sides of the junction are subtrees — recursively merge + last.right = try self.mergeSubtrees(last.right, r.left); + } else if (last.right == null and r.left != null) { + last.right = r.left; + } + // if last.right != null and r.left == null, keep last.right as-is + } else { + // left has no entries: junction is merged.left vs r.left + if (merged.left != null and r.left != null) { + merged.left = try self.mergeSubtrees(merged.left, r.left); + } else if (merged.left == null) { + merged.left = r.left; + } + } + + // copy right entries + for (r.entries.items) |entry| { + try merged.entries.append(self.allocator, entry); + } + + return merged; + } + + const MstError = Allocator.Error; + + /// compute the root CID of the tree + pub fn rootCid(self: *Mst) MstError!cbor.Cid { + return self.nodeCid(self.root); + } + + fn nodeCid(self: *Mst, maybe_node: ?*Node) MstError!cbor.Cid { + const encoded = try self.serializeNode(maybe_node); + defer self.allocator.free(encoded); + return cbor.Cid.forDagCbor(self.allocator, encoded); + } + + fn serializeNode(self: *Mst, maybe_node: ?*Node) MstError![]u8 { + const node = maybe_node orelse { + // empty node: { "l": null, "e": [] } + return cbor.encodeAlloc(self.allocator, .{ .map = &.{ + .{ .key = "e", .value = .{ .array = &.{} } }, + .{ .key = "l", .value = .null }, + } }); + }; + + // compute left subtree CID + const left_value: cbor.Value = if (node.left) |left| blk: { + const left_cid = try self.nodeCid(left); + break :blk .{ .cid = left_cid }; + } else .null; + + // build entry array with prefix compression + var entry_values: std.ArrayList(cbor.Value) = .{}; + defer entry_values.deinit(self.allocator); + + var prev_key: []const u8 = ""; + for (node.entries.items) |entry| { + const prefix_len = commonPrefixLen(prev_key, entry.key); + const suffix = entry.key[prefix_len..]; + + // right subtree CID + const tree_val: cbor.Value = if (entry.right) |right| blk: { + const right_cid = try self.nodeCid(right); + break :blk .{ .cid = right_cid }; + } else .null; + + // allocate map entries on heap (stack-local &.{...} would alias across iterations) + const map_entries = try self.allocator.alloc(cbor.Value.MapEntry, 4); + map_entries[0] = .{ .key = "k", .value = .{ .bytes = suffix } }; + map_entries[1] = .{ .key = "p", .value = .{ .unsigned = prefix_len } }; + map_entries[2] = .{ .key = "t", .value = tree_val }; + map_entries[3] = .{ .key = "v", .value = .{ .cid = entry.value } }; + + try entry_values.append(self.allocator, .{ .map = map_entries }); + + prev_key = entry.key; + } + + const entries_slice = try self.allocator.dupe(cbor.Value, entry_values.items); + defer self.allocator.free(entries_slice); + + return cbor.encodeAlloc(self.allocator, .{ .map = &.{ + .{ .key = "e", .value = .{ .array = entries_slice } }, + .{ .key = "l", .value = left_value }, + } }); + } + + // === internal helpers === + + fn createNode(self: *Mst) !*Node { + const node = try self.allocator.create(Node); + node.* = Node.init(); + return node; + } + + /// insert a key that belongs above the current root. + /// splits the tree at its own layer, wraps each half in parent nodes + /// to bridge the layer gap, then assembles the new root. + fn insertAbove(self: *Mst, node: *Node, node_layer: u32, key: []const u8, value: cbor.Cid, target_layer: u32) !*Node { + // 1. split the tree at its current layer around the key + const splits = try self.splitNode(node, key); + var left = splits.left; + var right = splits.right; + + // 2. wrap each half in parent layers (bridge the gap) + // "extraLayersToAdd = keyZeros - layer" + // "intentionally starting at 1, since first layer is taken care of by split" + const extra_layers = target_layer - node_layer; + var i: u32 = 1; + while (i < extra_layers) : (i += 1) { + if (left) |l| { + const parent = try self.createNode(); + parent.left = l; + left = parent; + } + if (right) |r| { + const parent = try self.createNode(); + parent.left = r; + right = parent; + } + } + + // 3. assemble new root: [left_tree, key_leaf, right_tree] + const new_root = try self.createNode(); + new_root.left = left; + try new_root.entries.append(self.allocator, .{ + .key = try self.allocator.dupe(u8, key), + .value = value, + .right = right, + }); + return new_root; + } + + /// insert a key at the same layer as the node + fn insertAtLayer(self: *Mst, node: *Node, key: []const u8, value: cbor.Cid, layer: u32) !*Node { + _ = layer; + // find insertion position + var insert_idx: usize = node.entries.items.len; + for (node.entries.items, 0..) |entry, i| { + const cmp = std.mem.order(u8, key, entry.key); + if (cmp == .eq) { + // update existing + node.entries.items[i].value = value; + return node; + } + if (cmp == .lt) { + insert_idx = i; + break; + } + } + + // split the subtree that spans the insertion gap + const gap_subtree = if (insert_idx == 0) node.left else node.entries.items[insert_idx - 1].right; + + var left_split: ?*Node = null; + var right_split: ?*Node = null; + + if (gap_subtree) |subtree| { + const splits = try self.splitNode(subtree, key); + left_split = splits.left; + right_split = splits.right; + } + + // update the pointer before the gap + if (insert_idx == 0) { + node.left = left_split; + } else { + node.entries.items[insert_idx - 1].right = left_split; + } + + // insert the new entry + try node.entries.insert(self.allocator, insert_idx, .{ + .key = try self.allocator.dupe(u8, key), + .value = value, + .right = right_split, + }); + + return node; + } + + /// insert a key below the current node's layer + fn insertBelow(self: *Mst, node: *Node, node_layer: u32, key: []const u8, value: cbor.Cid, target_height: u32) !void { + // find which gap the key falls into + for (node.entries.items, 0..) |entry, i| { + const cmp = std.mem.order(u8, key, entry.key); + if (cmp == .eq) { + // update existing + node.entries.items[i].value = value; + return; + } + if (cmp == .lt) { + // key goes in the gap before this entry + const subtree_ptr = if (i == 0) &node.left else &node.entries.items[i - 1].right; + try self.insertIntoGap(subtree_ptr, node_layer - 1, key, value, target_height); + return; + } + } + // key goes after all entries + const last_ptr = if (node.entries.items.len > 0) + &node.entries.items[node.entries.items.len - 1].right + else + &node.left; + try self.insertIntoGap(last_ptr, node_layer - 1, key, value, target_height); + } + + fn insertIntoGap(self: *Mst, subtree_ptr: *?*Node, gap_layer: u32, key: []const u8, value: cbor.Cid, target_height: u32) MstError!void { + if (target_height == gap_layer) { + // insert at this layer + if (subtree_ptr.*) |existing| { + subtree_ptr.* = try self.insertAtLayer(existing, key, value, gap_layer); + } else { + const new_node = try self.createNode(); + try new_node.entries.append(self.allocator, .{ + .key = try self.allocator.dupe(u8, key), + .value = value, + .right = null, + }); + subtree_ptr.* = new_node; + } + } else if (target_height > gap_layer) { + // need to lift — split and wrap + if (subtree_ptr.*) |existing| { + subtree_ptr.* = try self.insertAbove(existing, gap_layer, key, value, target_height); + } else { + const new_node = try self.createNode(); + try new_node.entries.append(self.allocator, .{ + .key = try self.allocator.dupe(u8, key), + .value = value, + .right = null, + }); + subtree_ptr.* = new_node; + } + } else { + // target_height < gap_layer: recurse deeper + if (subtree_ptr.*) |existing| { + try self.insertBelow(existing, gap_layer, key, value, target_height); + } else { + // create node at gap_layer and recurse + const new_node = try self.createNode(); + subtree_ptr.* = new_node; + try self.insertBelow(new_node, gap_layer, key, value, target_height); + } + } + } + + /// split a subtree around a key: everything < key goes left, everything >= key goes right. + /// follows the Go reference: find split point among leaf entries, then recursively + /// split the subtree in the gap if needed. + fn splitNode(self: *Mst, node: *Node, key: []const u8) !struct { left: ?*Node, right: ?*Node } { + // find the first entry >= key + var split_idx: usize = node.entries.items.len; + for (node.entries.items, 0..) |entry, i| { + if (std.mem.order(u8, key, entry.key) != .gt) { + split_idx = i; + break; + } + } + + // left gets entries [0..split_idx), right gets entries [split_idx..] + var left_node = try self.createNode(); + var right_node = try self.createNode(); + + // left node takes the original node's left subtree + left_node.left = node.left; + + // copy entries to left + for (node.entries.items[0..split_idx]) |entry| { + try left_node.entries.append(self.allocator, entry); + } + + // copy entries to right + for (node.entries.items[split_idx..]) |entry| { + try right_node.entries.append(self.allocator, entry); + } + + // the subtree between the last left entry and first right entry may need recursive splitting. + // in our representation: this is the right pointer of the last left entry (or left's left if no entries) + // for the right node, its "left" is initially null — we need to set it from the gap. + + // split the gap subtree between the two halves + if (left_node.entries.items.len > 0) { + const last_left = &left_node.entries.items[left_node.entries.items.len - 1]; + if (last_left.right) |gap_subtree| { + const sub_split = try self.splitNode(gap_subtree, key); + last_left.right = sub_split.left; + right_node.left = sub_split.right; + } + } else if (left_node.left != null and split_idx == 0) { + // all entries went right — the gap is the original node's left subtree + const sub_split = try self.splitNode(left_node.left.?, key); + left_node.left = sub_split.left; + right_node.left = sub_split.right; + } + + const left_result: ?*Node = if (left_node.entries.items.len > 0 or left_node.left != null) left_node else null; + const right_result: ?*Node = if (right_node.entries.items.len > 0 or right_node.left != null) right_node else null; + + return .{ .left = left_result, .right = right_result }; + } +}; + +// === tests === + +test "keyHeight" { + // values from interop test fixtures + try std.testing.expectEqual(@as(u32, 0), keyHeight("")); + try std.testing.expectEqual(@as(u32, 0), keyHeight("asdf")); + try std.testing.expectEqual(@as(u32, 1), keyHeight("blue")); + try std.testing.expectEqual(@as(u32, 0), keyHeight("2653ae71")); + try std.testing.expectEqual(@as(u32, 2), keyHeight("88bfafc7")); + try std.testing.expectEqual(@as(u32, 4), keyHeight("2a92d355")); + try std.testing.expectEqual(@as(u32, 6), keyHeight("884976f5")); + try std.testing.expectEqual(@as(u32, 4), keyHeight("app.bsky.feed.post/454397e440ec")); + try std.testing.expectEqual(@as(u32, 8), keyHeight("app.bsky.feed.post/9adeb165882c")); +} + +test "commonPrefixLen" { + try std.testing.expectEqual(@as(usize, 0), commonPrefixLen("", "")); + try std.testing.expectEqual(@as(usize, 3), commonPrefixLen("abc", "abc")); + try std.testing.expectEqual(@as(usize, 0), commonPrefixLen("", "abc")); + try std.testing.expectEqual(@as(usize, 2), commonPrefixLen("ab", "abc")); + try std.testing.expectEqual(@as(usize, 3), commonPrefixLen("abcde", "abc")); + try std.testing.expectEqual(@as(usize, 0), commonPrefixLen("abcde", "qbb")); +} + +test "put and get" { + const alloc = std.testing.allocator; + var arena = std.heap.ArenaAllocator.init(alloc); + defer arena.deinit(); + const a = arena.allocator(); + + var tree = Mst.init(a); + + const cid1 = try cbor.Cid.forDagCbor(a, "value1"); + const cid2 = try cbor.Cid.forDagCbor(a, "value2"); + + try tree.put("key1", cid1); + try tree.put("key2", cid2); + + const got1 = tree.get("key1") orelse return error.NotFound; + try std.testing.expectEqualSlices(u8, cid1.raw, got1.raw); + + const got2 = tree.get("key2") orelse return error.NotFound; + try std.testing.expectEqualSlices(u8, cid2.raw, got2.raw); + + try std.testing.expect(tree.get("nonexistent") == null); +} + +test "put and delete" { + const alloc = std.testing.allocator; + var arena = std.heap.ArenaAllocator.init(alloc); + defer arena.deinit(); + const a = arena.allocator(); + + var tree = Mst.init(a); + const cid = try cbor.Cid.forDagCbor(a, "value"); + + try tree.put("key1", cid); + try tree.put("key2", cid); + + try std.testing.expect(tree.get("key1") != null); + try tree.delete("key1"); + try std.testing.expect(tree.get("key1") == null); + try std.testing.expect(tree.get("key2") != null); +} + +test "rootCid is deterministic" { + const alloc = std.testing.allocator; + var arena = std.heap.ArenaAllocator.init(alloc); + defer arena.deinit(); + const a = arena.allocator(); + + const cid_val = try cbor.Cid.forDagCbor(a, "leaf"); + + // build tree 1 + var tree1 = Mst.init(a); + try tree1.put("a", cid_val); + try tree1.put("b", cid_val); + const root1 = try tree1.rootCid(); + + // build tree 2 (same keys, same order) + var tree2 = Mst.init(a); + try tree2.put("a", cid_val); + try tree2.put("b", cid_val); + const root2 = try tree2.rootCid(); + + try std.testing.expectEqualSlices(u8, root1.raw, root2.raw); +} + +test "empty tree rootCid matches reference" { + const alloc = std.testing.allocator; + var arena = std.heap.ArenaAllocator.init(alloc); + defer arena.deinit(); + const a = arena.allocator(); + + var tree = Mst.init(a); + const root = try tree.rootCid(); + try std.testing.expectEqual(@as(u64, 1), root.version().?); + + // known empty tree CID from Go reference implementation + const expected = try parseCidString(a, "bafyreie5737gdxlw5i64vzichcalba3z2v5n6icifvx5xytvske7mr3hpm"); + try std.testing.expectEqualSlices(u8, expected.raw, root.raw); +} + +test "single key rootCid matches reference" { + const alloc = std.testing.allocator; + var arena = std.heap.ArenaAllocator.init(alloc); + defer arena.deinit(); + const a = arena.allocator(); + + var tree = Mst.init(a); + // use a known CID value (the leaf CID from commit-proof fixtures) + const leaf_cid = try parseCidString(a, "bafyreie5cvv4h45feadgeuwhbcutmh6t2ceseocckahdoe6uat64zmz454"); + + // single layer-0 key + try tree.put("com.example.record/3jqfcqzm3fo2j", leaf_cid); + + const root = try tree.rootCid(); + const expected = try parseCidString(a, "bafyreibj4lsc3aqnrvphp5xmrnfoorvru4wynt6lwidqbm2623a6tatzdu"); + try std.testing.expectEqualSlices(u8, expected.raw, root.raw); +} + +test "single layer-2 key rootCid matches reference" { + const alloc = std.testing.allocator; + var arena = std.heap.ArenaAllocator.init(alloc); + defer arena.deinit(); + const a = arena.allocator(); + + var tree = Mst.init(a); + const leaf_cid = try parseCidString(a, "bafyreie5cvv4h45feadgeuwhbcutmh6t2ceseocckahdoe6uat64zmz454"); + + // single layer-2 key + try tree.put("com.example.record/3jqfcqzm3fx2j", leaf_cid); + + const root = try tree.rootCid(); + const expected = try parseCidString(a, "bafyreih7wfei65pxzhauoibu3ls7jgmkju4bspy4t2ha2qdjnzqvoy33ai"); + try std.testing.expectEqualSlices(u8, expected.raw, root.raw); +} + +test "5 key tree matches reference" { + const alloc = std.testing.allocator; + var arena = std.heap.ArenaAllocator.init(alloc); + defer arena.deinit(); + const a = arena.allocator(); + + var tree = Mst.init(a); + const leaf_cid = try parseCidString(a, "bafyreie5cvv4h45feadgeuwhbcutmh6t2ceseocckahdoe6uat64zmz454"); + + // 5 keys from Go test (note: last key has 4fc not 3ft) + const keys = [_][]const u8{ + "com.example.record/3jqfcqzm3fp2j", + "com.example.record/3jqfcqzm3fr2j", + "com.example.record/3jqfcqzm3fs2j", + "com.example.record/3jqfcqzm3ft2j", + "com.example.record/3jqfcqzm4fc2j", + }; + + for (keys) |key| { + try tree.put(key, leaf_cid); + } + + const root = try tree.rootCid(); + const expected = try parseCidString(a, "bafyreicmahysq4n6wfuxo522m6dpiy7z7qzym3dzs756t5n7nfdgccwq7m"); + try std.testing.expectEqualSlices(u8, expected.raw, root.raw); +} + +test "two deep split fixture" { + const alloc = std.testing.allocator; + var arena = std.heap.ArenaAllocator.init(alloc); + defer arena.deinit(); + const a = arena.allocator(); + + const leaf_cid = try parseCidString(a, "bafyreie5cvv4h45feadgeuwhbcutmh6t2ceseocckahdoe6uat64zmz454"); + + var tree = Mst.init(a); + const initial_keys = [_][]const u8{ + "A0/374913", "B1/986427", "C0/451630", + "E0/670489", "F1/085263", "G0/765327", + }; + for (initial_keys) |key| { + try tree.put(key, leaf_cid); + } + + const expected_before = try parseCidString(a, "bafyreicraprx2xwnico4tuqir3ozsxpz46qkcpox3obf5bagicqwurghpy"); + try std.testing.expectEqualSlices(u8, expected_before.raw, (try tree.rootCid()).raw); + + try tree.put("D2/269196", leaf_cid); + + const expected_after = try parseCidString(a, "bafyreihvay6pazw3dfa47u5d2tn3rd6pa57sr37bo5bqyvjuqc73ib65my"); + try std.testing.expectEqualSlices(u8, expected_after.raw, (try tree.rootCid()).raw); +} + +test "complex multi-op commit" { + const alloc = std.testing.allocator; + var arena = std.heap.ArenaAllocator.init(alloc); + defer arena.deinit(); + const a = arena.allocator(); + + const leaf_cid = try parseCidString(a, "bafyreie5cvv4h45feadgeuwhbcutmh6t2ceseocckahdoe6uat64zmz454"); + + var tree = Mst.init(a); + const initial_keys = [_][]const u8{ + "B0/601692", "C2/014073", "D0/952776", + "E2/819540", "F0/697858", "H0/131238", + }; + for (initial_keys) |key| { + try tree.put(key, leaf_cid); + } + + const expected_before = try parseCidString(a, "bafyreigr3plnts7dax6yokvinbhcqpyicdfgg6npvvyx6okc5jo55slfqi"); + try std.testing.expectEqualSlices(u8, expected_before.raw, (try tree.rootCid()).raw); + + // adds + try tree.put("A2/827942", leaf_cid); + try tree.put("G2/611528", leaf_cid); + // del + try tree.delete("C2/014073"); + + const expected_after = try parseCidString(a, "bafyreiftrcrbhrwmi37u4egedlg56gk3jeh3tvmqvwgowoifuklfysyx54"); + try std.testing.expectEqualSlices(u8, expected_after.raw, (try tree.rootCid()).raw); +} + +test "parseCidString" { + const alloc = std.testing.allocator; + var arena = std.heap.ArenaAllocator.init(alloc); + defer arena.deinit(); + const a = arena.allocator(); + + const cid = try parseCidString(a, "bafyreie5cvv4h45feadgeuwhbcutmh6t2ceseocckahdoe6uat64zmz454"); + try std.testing.expectEqual(@as(u64, 1), cid.version().?); + try std.testing.expectEqual(@as(u64, 0x71), cid.codec().?); + try std.testing.expectEqual(@as(u64, 0x12), cid.hashFn().?); + try std.testing.expectEqual(@as(usize, 32), cid.digest().?.len); +} diff --git a/src/internal/multibase.zig b/src/internal/multibase.zig index 28b7c76..0d5f3fe 100644 --- a/src/internal/multibase.zig +++ b/src/internal/multibase.zig @@ -1,7 +1,7 @@ -//! multibase decoder +//! multibase codec //! -//! decodes multibase-encoded strings (prefix + encoded data). -//! currently supports base58btc (z prefix) for DID document public keys. +//! encodes and decodes multibase-encoded strings (prefix + encoded data). +//! supports base58btc (z prefix) and base32lower (b prefix). //! //! see: https://github.com/multiformats/multibase @@ -10,10 +10,12 @@ const std = @import("std"); /// multibase encoding types pub const Encoding = enum { base58btc, // z prefix + base32lower, // b prefix pub fn fromPrefix(prefix: u8) ?Encoding { return switch (prefix) { 'z' => .base58btc, + 'b' => .base32lower, else => null, }; } @@ -28,10 +30,19 @@ pub fn decode(allocator: std.mem.Allocator, input: []const u8) ![]u8 { return switch (encoding) { .base58btc => try base58btc.decode(allocator, input[1..]), + .base32lower => try base32lower.decode(allocator, input[1..]), }; } -/// base58btc decoder (bitcoin alphabet) +/// encode raw bytes to a multibase string with the given encoding +pub fn encode(allocator: std.mem.Allocator, encoding: Encoding, data: []const u8) ![]u8 { + return switch (encoding) { + .base58btc => try base58btc.encode(allocator, data), + .base32lower => try base32lower.encode(allocator, data), + }; +} + +/// base58btc codec (bitcoin alphabet) pub const base58btc = struct { /// bitcoin base58 alphabet const alphabet = "123456789ABCDEFGHJKLMNPQRSTUVWXYZabcdefghijkmnopqrstuvwxyz"; @@ -45,6 +56,67 @@ pub const base58btc = struct { break :blk table; }; + /// encode bytes to base58btc string with 'z' multibase prefix + pub fn encode(allocator: std.mem.Allocator, input: []const u8) ![]u8 { + // count leading zero bytes → leading '1's + var leading_zeros: usize = 0; + for (input) |b| { + if (b != 0) break; + leading_zeros += 1; + } + + if (input.len == 0 or leading_zeros == input.len) { + // all zeros (or empty) + const result = try allocator.alloc(u8, 1 + leading_zeros); + result[0] = 'z'; // multibase prefix + @memset(result[1..], '1'); + return result; + } + + // load bytes into big integer (big-endian) + var acc = try std.math.big.int.Managed.init(allocator); + defer acc.deinit(); + + for (input) |b| { + try acc.shiftLeft(&acc, 8); + try acc.addScalar(&acc, b); + } + + // repeatedly divide by 58 to extract base58 digits + var digits: std.ArrayList(u8) = .{}; + defer digits.deinit(allocator); + + var divisor = try std.math.big.int.Managed.initSet(allocator, @as(u64, 58)); + defer divisor.deinit(); + + var quotient = try std.math.big.int.Managed.init(allocator); + defer quotient.deinit(); + + var remainder = try std.math.big.int.Managed.init(allocator); + defer remainder.deinit(); + + while (!acc.toConst().eqlZero()) { + try quotient.divFloor(&remainder, &acc, &divisor); + const digit: usize = @intCast(remainder.toConst().toInt(u64) catch 0); + try digits.append(allocator, alphabet[digit]); + try acc.copy(quotient.toConst()); + } + + // result: 'z' prefix + leading '1's + reversed digits + const total_len = 1 + leading_zeros + digits.items.len; + const result = try allocator.alloc(u8, total_len); + result[0] = 'z'; // multibase prefix + @memset(result[1 .. 1 + leading_zeros], '1'); + + // digits were accumulated LSB-first, reverse into result + const digit_slice = result[1 + leading_zeros ..]; + for (digits.items, 0..) |d, i| { + digit_slice[digits.items.len - 1 - i] = d; + } + + return result; + } + /// decode base58btc string to bytes pub fn decode(allocator: std.mem.Allocator, input: []const u8) ![]u8 { if (input.len == 0) return allocator.alloc(u8, 0); @@ -56,14 +128,7 @@ pub const base58btc = struct { leading_zeros += 1; } - // estimate output size: each base58 char represents ~5.86 bits - // use a simple overestimate: input.len bytes is more than enough - const max_output = input.len; - const result = try allocator.alloc(u8, max_output); - errdefer allocator.free(result); - // decode using big integer arithmetic - // accumulator = accumulator * 58 + digit var acc = try std.math.big.int.Managed.init(allocator); defer acc.deinit(); @@ -75,40 +140,28 @@ pub const base58btc = struct { for (input) |c| { const digit = decode_table[c]; - if (digit < 0) { - allocator.free(result); - return error.InvalidCharacter; - } + if (digit < 0) return error.InvalidCharacter; - // acc = acc * 58 + digit try temp.mul(&acc, &multiplier); try acc.copy(temp.toConst()); try acc.addScalar(&acc, @as(u8, @intCast(digit))); } - // convert big int to bytes (big-endian for base58) + // convert big int to bytes (big-endian) const limbs = acc.toConst().limbs; const limb_count = acc.len(); - // calculate byte size from limbs var byte_count: usize = 0; if (limb_count > 0 and !acc.toConst().eqlZero()) { - const bit_count = acc.toConst().bitCountAbs(); - byte_count = (bit_count + 7) / 8; + byte_count = (acc.toConst().bitCountAbs() + 7) / 8; } - // write bytes in big-endian order - var output_bytes = try allocator.alloc(u8, leading_zeros + byte_count); - errdefer allocator.free(output_bytes); - - // leading zeros - @memset(output_bytes[0..leading_zeros], 0); + const result = try allocator.alloc(u8, leading_zeros + byte_count); + @memset(result[0..leading_zeros], 0); // convert limbs to big-endian bytes if (byte_count > 0) { - const output_slice = output_bytes[leading_zeros..]; - - // limbs are in little-endian order, we need big-endian output + const output_slice = result[leading_zeros..]; var pos: usize = byte_count; for (limbs[0..limb_count]) |limb| { const limb_bytes = @sizeOf(@TypeOf(limb)); @@ -120,8 +173,87 @@ pub const base58btc = struct { } } - allocator.free(result); - return output_bytes; + return result; + } +}; + +/// base32lower codec (RFC 4648, lowercase, no padding) +pub const base32lower = struct { + const alphabet = "abcdefghijklmnopqrstuvwxyz234567"; + + const decode_table: [256]i8 = blk: { + var table: [256]i8 = .{-1} ** 256; + for (alphabet, 0..) |c, i| { + table[c] = @intCast(i); + } + break :blk table; + }; + + /// encode bytes to base32lower string with 'b' multibase prefix + pub fn encode(allocator: std.mem.Allocator, input: []const u8) ![]u8 { + if (input.len == 0) { + const result = try allocator.alloc(u8, 1); + result[0] = 'b'; + return result; + } + + // base32: 5 bytes → 8 chars + const out_len = (input.len * 8 + 4) / 5; // ceil(bits / 5) + const result = try allocator.alloc(u8, 1 + out_len); + result[0] = 'b'; // multibase prefix + + var bit_buf: u32 = 0; + var bits: u5 = 0; + var pos: usize = 1; + + for (input) |byte| { + bit_buf = (bit_buf << 8) | byte; + bits += 8; + while (bits >= 5) { + bits -= 5; + const idx: u5 = @truncate(bit_buf >> bits); + result[pos] = alphabet[idx]; + pos += 1; + } + } + + // remaining bits (left-aligned) + if (bits > 0) { + const idx: u5 = @truncate(bit_buf << (@as(u5, 5) - bits)); + result[pos] = alphabet[idx]; + pos += 1; + } + + return result[0..pos]; + } + + /// decode base32lower string (no multibase prefix) to bytes + pub fn decode(allocator: std.mem.Allocator, input: []const u8) ![]u8 { + if (input.len == 0) return allocator.alloc(u8, 0); + + const out_len = input.len * 5 / 8; + const result = try allocator.alloc(u8, out_len); + errdefer allocator.free(result); + + var bit_buf: u32 = 0; + var bits: u4 = 0; + var pos: usize = 0; + + for (input) |c| { + if (c == '=') break; // stop at padding + const digit = decode_table[c]; + if (digit < 0) return error.InvalidCharacter; + + bit_buf = (bit_buf << 5) | @as(u32, @intCast(digit)); + bits += 5; + if (bits >= 8) { + bits -= 8; + result[pos] = @truncate(bit_buf >> bits); + pos += 1; + } + } + + return allocator.realloc(result, pos); } }; @@ -188,3 +320,65 @@ test "base58btc decode real multibase key - secp256k1" { // compressed point prefix should be 0x02 or 0x03 try std.testing.expect(parsed.raw[0] == 0x02 or parsed.raw[0] == 0x03); } + +test "base58btc encode-decode round-trip" { + const alloc = std.testing.allocator; + + { + const original = "abc"; + const encoded = try base58btc.encode(alloc, original); + defer alloc.free(encoded); + // should have 'z' prefix + try std.testing.expectEqual(@as(u8, 'z'), encoded[0]); + + const decoded = try decode(alloc, encoded); + defer alloc.free(decoded); + try std.testing.expectEqualSlices(u8, original, decoded); + } + + // round-trip with leading zeros + { + const original = &[_]u8{ 0, 0, 0x01 }; + const encoded = try base58btc.encode(alloc, original); + defer alloc.free(encoded); + const decoded = try decode(alloc, encoded); + defer alloc.free(decoded); + try std.testing.expectEqualSlices(u8, original, decoded); + } +} + +test "base32lower encode-decode round-trip" { + const alloc = std.testing.allocator; + + { + const original = "hello"; + const encoded = try base32lower.encode(alloc, original); + defer alloc.free(encoded); + try std.testing.expectEqual(@as(u8, 'b'), encoded[0]); + + const decoded = try decode(alloc, encoded); + defer alloc.free(decoded); + try std.testing.expectEqualSlices(u8, original, decoded); + } + + // empty + { + const encoded = try base32lower.encode(alloc, ""); + defer alloc.free(encoded); + try std.testing.expectEqualStrings("b", encoded); + } +} + +test "base32lower decode bafyrei prefix" { + const alloc = std.testing.allocator; + // CIDv1 dag-cbor sha2-256 always starts with "bafyrei" in base32lower + // "bafyreie5cvv4h45feadgeuwhbcutmh6t2ceseocckahdoe6uat64zmz454" + const input = "afyreie5cvv4h45feadgeuwhbcutmh6t2ceseocckahdoe6uat64zmz454"; + const decoded = try base32lower.decode(alloc, input); + defer alloc.free(decoded); + // CIDv1: version=1(0x01), codec=dag-cbor(0x71), hash=sha2-256(0x12), len=32(0x20) + try std.testing.expectEqual(@as(u8, 0x01), decoded[0]); + try std.testing.expectEqual(@as(u8, 0x71), decoded[1]); + try std.testing.expectEqual(@as(u8, 0x12), decoded[2]); + try std.testing.expectEqual(@as(u8, 0x20), decoded[3]); +} diff --git a/src/internal/multicodec.zig b/src/internal/multicodec.zig index 838d154..647a456 100644 --- a/src/internal/multicodec.zig +++ b/src/internal/multicodec.zig @@ -51,6 +51,43 @@ pub fn parsePublicKey(data: []const u8) !PublicKey { return error.UnsupportedKeyType; } +/// encode a raw public key with multicodec prefix +pub fn encodePublicKey(allocator: std.mem.Allocator, key_type: KeyType, raw: []const u8) ![]u8 { + if (raw.len != 33) return error.InvalidKeyLength; + + const result = try allocator.alloc(u8, 2 + raw.len); + switch (key_type) { + .secp256k1 => { + result[0] = 0xe7; + result[1] = 0x01; + }, + .p256 => { + result[0] = 0x80; + result[1] = 0x24; + }, + } + @memcpy(result[2..], raw); + return result; +} + +/// format a raw public key as a did:key string +pub fn formatDidKey(allocator: std.mem.Allocator, key_type: KeyType, raw: []const u8) ![]u8 { + const multibase = @import("multibase.zig"); + + const mc_bytes = try encodePublicKey(allocator, key_type, raw); + defer allocator.free(mc_bytes); + + const multibase_str = try multibase.encode(allocator, .base58btc, mc_bytes); + defer allocator.free(multibase_str); + + // "did:key:" + multibase string (which already has 'z' prefix) + const prefix = "did:key:"; + const result = try allocator.alloc(u8, prefix.len + multibase_str.len); + @memcpy(result[0..prefix.len], prefix); + @memcpy(result[prefix.len..], multibase_str); + return result; +} + // === tests === test "parse secp256k1 key" { @@ -88,3 +125,61 @@ test "reject too short" { const data = [_]u8{0xe7}; try std.testing.expectError(error.TooShort, parsePublicKey(&data)); } + +test "encode-decode round-trip secp256k1" { + const alloc = std.testing.allocator; + var raw: [33]u8 = undefined; + raw[0] = 0x02; + @memset(raw[1..], 0xaa); + + const encoded = try encodePublicKey(alloc, .secp256k1, &raw); + defer alloc.free(encoded); + + const parsed = try parsePublicKey(encoded); + try std.testing.expectEqual(KeyType.secp256k1, parsed.key_type); + try std.testing.expectEqualSlices(u8, &raw, parsed.raw); +} + +test "did:key round-trip secp256k1" { + const alloc = std.testing.allocator; + const multibase = @import("multibase.zig"); + + var raw: [33]u8 = undefined; + raw[0] = 0x02; + @memset(raw[1..], 0xcc); + + const did_key_str = try formatDidKey(alloc, .secp256k1, &raw); + defer alloc.free(did_key_str); + + // should start with "did:key:z" + try std.testing.expect(std.mem.startsWith(u8, did_key_str, "did:key:z")); + + // parse back: strip "did:key:" prefix, decode multibase, parse multicodec + const multibase_str = did_key_str["did:key:".len..]; + const mc_bytes = try multibase.decode(alloc, multibase_str); + defer alloc.free(mc_bytes); + + const parsed = try parsePublicKey(mc_bytes); + try std.testing.expectEqual(KeyType.secp256k1, parsed.key_type); + try std.testing.expectEqualSlices(u8, &raw, parsed.raw); +} + +test "did:key round-trip p256" { + const alloc = std.testing.allocator; + const multibase = @import("multibase.zig"); + + var raw: [33]u8 = undefined; + raw[0] = 0x03; + @memset(raw[1..], 0xdd); + + const did_key_str = try formatDidKey(alloc, .p256, &raw); + defer alloc.free(did_key_str); + + const multibase_str = did_key_str["did:key:".len..]; + const mc_bytes = try multibase.decode(alloc, multibase_str); + defer alloc.free(mc_bytes); + + const parsed = try parsePublicKey(mc_bytes); + try std.testing.expectEqual(KeyType.p256, parsed.key_type); + try std.testing.expectEqualSlices(u8, &raw, parsed.raw); +} diff --git a/src/root.zig b/src/root.zig index 96d3023..34207df 100644 --- a/src/root.zig +++ b/src/root.zig @@ -27,6 +27,9 @@ pub const Jwt = @import("internal/jwt.zig").Jwt; pub const multibase = @import("internal/multibase.zig"); pub const multicodec = @import("internal/multicodec.zig"); +// mst +pub const mst = @import("internal/mst.zig"); + // sync / firehose const sync = @import("internal/sync.zig"); pub const CommitAction = sync.CommitAction;