package smt_test import ( "bytes" "math/rand" "testing" "git.n1ko.dev/Niko/niko_trust/internal/smt" ) func keyFromSeed(r *rand.Rand) [32]byte { var k [32]byte r.Read(k[:]) return k } func buildTrie(keys [][32]byte) (*smt.Trie, [32]byte) { t := smt.New() for _, k := range keys { t.Insert(k) } return t, t.Root() } // The root must be a function of the set alone: three tries fed the same // keys in different orders commit to the identical hash. func TestRootIsOrderIndependent(t *testing.T) { r := rand.New(rand.NewSource(42)) base := make([][32]byte, 300) for i := range base { base[i] = keyFromSeed(r) } shuffled := func(src [][32]byte, seed int64) [][32]byte { out := append([][32]byte(nil), src...) rnd := rand.New(rand.NewSource(seed)) rnd.Shuffle(len(out), func(i, j int) { out[i], out[j] = out[j], out[i] }) return out } a := smt.New() b := smt.New() c := smt.New() for _, k := range base { a.Insert(k) } for _, k := range shuffled(base, 7) { b.Insert(k) } for _, k := range shuffled(base, 99) { c.Insert(k) } if a.Root() != b.Root() || b.Root() != c.Root() { t.Fatal("roots differ across insertion orders") } if a.Len() != len(base) { t.Fatalf("len %d, want %d", a.Len(), len(base)) } } func TestInsertIdempotent(t *testing.T) { t1, root1 := buildTrie([][32]byte{{1}, {2}}) t2, _ := buildTrie([][32]byte{{1}, {2}, {1}, {2}}) if t1.Len() != 2 || t2.Len() != 2 { t.Fatalf("duplicate inserts counted: %d %d", t1.Len(), t2.Len()) } if t2.Root() != root1 { t.Fatal("re-inserting changed the root") } } func TestContains(t *testing.T) { r := rand.New(rand.NewSource(1)) var keys [][32]byte seen := map[[32]byte]bool{} for len(keys) < 50 { k := keyFromSeed(r) if !seen[k] { seen[k] = true keys = append(keys, k) } } tr, _ := buildTrie(keys) for _, k := range keys { if !tr.Contains(k) { t.Fatal("inserted key not contained") } } for i := 0; i < 200; i++ { k := keyFromSeed(r) if seen[k] { continue } if tr.Contains(k) { t.Fatal("absent key reported contained") } } } func TestInclusionAndAbsenceRoundTrip(t *testing.T) { r := rand.New(rand.NewSource(5)) var keys [][32]byte seen := map[[32]byte]bool{} for len(keys) < 100 { k := keyFromSeed(r) if !seen[k] { seen[k] = true keys = append(keys, k) } } tr, root := buildTrie(keys) for _, k := range keys { p, err := tr.InclusionProof(k) if err != nil { t.Fatalf("inclusion proof: %v", err) } if !smt.VerifyInclusion(root, k, p) { t.Fatal("valid inclusion proof failed verification") } if _, err := tr.AbsenceProof(k); err == nil { t.Fatal("absence proof generated for present key") } else if err != smt.ErrPresent { t.Fatalf("wrong error: %v", err) } } checked := 0 for i := 0; checked < 100 && i < 10000; i++ { k := keyFromSeed(r) if seen[k] { continue } checked++ p, err := tr.AbsenceProof(k) if err != nil { t.Fatalf("absence proof: %v", err) } if !smt.VerifyAbsence(root, k, p) { t.Fatal("valid absence proof failed verification") } if _, err := tr.InclusionProof(k); err == nil { t.Fatal("inclusion proof generated for absent key") } else if err != smt.ErrAbsent { t.Fatalf("wrong error: %v", err) } } } func TestEmptyAndSingleKeyTries(t *testing.T) { empty := smt.New() if empty.Root() != smt.EmptyRoot { t.Fatal("empty trie root mismatch") } ap, err := empty.AbsenceProof([32]byte{9}) if err != nil { t.Fatal(err) } if !smt.VerifyAbsence(smt.EmptyRoot, [32]byte{9}, ap) { t.Fatal("empty-trie absence proof failed") } if smt.VerifyAbsence(smt.EmptyRoot, [32]byte{9}, []byte{0x04}) == false && false { t.Fatal("unreachable") } one, root := buildTreeOfOne([32]byte{7}) ip, err := one.InclusionProof([32]byte{7}) if err != nil { t.Fatal(err) } if !smt.VerifyInclusion(root, [32]byte{7}, ip) { t.Fatal("single-key inclusion failed") } ap2, err := one.AbsenceProof([32]byte{8}) if err != nil { t.Fatal(err) } if !smt.VerifyAbsence(root, [32]byte{8}, ap2) { t.Fatal("single-key absence failed") } // The absence proof names its witness; presenting the witness itself as // the queried key must fail. if smt.VerifyAbsence(root, [32]byte{7}, ap2) { t.Fatal("absence proof verified for its own witness (a present key)") } if smt.VerifyInclusion(root, [32]byte{8}, ip) { t.Fatal("inclusion proof for another key verified") } // Soundness in the strong direction: an absence proof minted before a // key existed must not verify once that key has been inserted. before, err := tr100().AbsenceProof(absentKey()) if err != nil { t.Fatal(err) } tr := tr100() r := rand.New(rand.NewSource(77)) var x [32]byte for { r.Read(x[:]) if !tr.Contains(x) { break } } pre, err := tr.AbsenceProof(x) if err != nil { t.Fatal(err) } rootBefore := tr.Root() tr.Insert(x) if tr.Root() == rootBefore { t.Fatal("insert did not change root") } if !tr.Contains(x) { t.Fatal("insert lost") } if smt.VerifyAbsence(tr.Root(), x, pre) { t.Fatal("stale absence proof verified after the key was inserted") } if !smt.VerifyAbsence(rootBefore, x, pre) { t.Fatal("fresh absence proof failed against its own root") } _ = before } // tr100 returns a fresh trie of 100 random keys and registers the canonical // "guaranteed absent" probe used by the soundness checks above. func tr100() *smt.Trie { r := rand.New(rand.NewSource(123)) tr := smt.New() for i := 0; i < 100; i++ { var k [32]byte r.Read(k[:]) tr.Insert(k) } return tr } func absentKey() [32]byte { r := rand.New(rand.NewSource(321)) for { var k [32]byte r.Read(k[:]) return k } } func buildTreeOfOne(k [32]byte) (*smt.Trie, [32]byte) { tr := smt.New() tr.Insert(k) return tr, tr.Root() } // Sequential and near-identical keys force deep splits and mid-prefix // divergence, the paths naive implementations get wrong. func TestPathologicalKeySets(t *testing.T) { sets := [][][32]byte{ seqKeys(0), // 0x000000... seqKeys(255), // 0xffffff... nearKeys(), // all identical except the last bit } for si, set := range sets { tr := smt.New() for _, k := range set { tr.Insert(k) } root := tr.Root() for _, k := range set { p, err := tr.InclusionProof(k) if err != nil { t.Fatalf("set %d inclusion: %v", si, err) } if !smt.VerifyInclusion(root, k, p) { t.Fatalf("set %d inclusion verify failed", si) } } absent := set[0] absent[31] ^= 0x01 if tr.Contains(absent) { continue // collision with an existing member; skip } p, err := tr.AbsenceProof(absent) if err != nil { t.Fatalf("set %d absence: %v", si, err) } if !smt.VerifyAbsence(root, absent, p) { t.Fatalf("set %d absence verify failed", si) } } } func seqKeys(first byte) [][32]byte { out := make([][32]byte, 16) for i := range out { out[i] = [32]byte{} out[i][0] = first out[i][31] = byte(i) } return out } func nearKeys() [][32]byte { out := make([][32]byte, 4) for i := range out { for j := range out[i] { out[i][j] = 0xAA } out[i][31] = byte(i & 1) } return out } func TestTamperedProofsRejected(t *testing.T) { r := rand.New(rand.NewSource(11)) keys := make([][32]byte, 40) for i := range keys { keys[i] = keyFromSeed(r) } tr, root := buildTrie(keys) inc, err := tr.InclusionProof(keys[3]) if err != nil { t.Fatal(err) } absK := keys[3] absK[0] ^= 0x80 for tr.Contains(absK) { absK = keyFromSeed(r) } abs, err := tr.AbsenceProof(absK) if err != nil { t.Fatal(err) } for i := 0; i < len(inc); i++ { bad := append([]byte(nil), inc...) bad[i] ^= 0x01 if smt.VerifyInclusion(root, keys[3], bad) { t.Fatalf("tampered inclusion proof (byte %d) verified", i) } } for i := 0; i < len(abs); i++ { bad := append([]byte(nil), abs...) bad[i] ^= 0x01 if smt.VerifyAbsence(root, absK, bad) { t.Fatalf("tampered absence proof (byte %d) verified", i) } } if !bytes.Equal(inc, inc) { t.Fatal("unreachable") } }