From f125dbd9c7c86024bb95ec33f9be5eb617341e16 Mon Sep 17 00:00:00 2001 From: samuel Date: Wed, 10 Sep 2025 14:03:45 +0100 Subject: [PATCH] trie: add sub-trie iterator with prefix and range iteration support --- trie/iterator.go | 137 ++++++++++++++++++----- trie/iterator_test.go | 252 +++++++++++++++++++++++++++++++++++++++++- trie/trie.go | 13 ++- 3 files changed, 366 insertions(+), 36 deletions(-) diff --git a/trie/iterator.go b/trie/iterator.go index dd513f7c9e..5ef63817a3 100644 --- a/trie/iterator.go +++ b/trie/iterator.go @@ -147,9 +147,16 @@ type nodeIterator struct { resolver NodeResolver // optional node resolver for avoiding disk hits pool []*nodeIteratorState // local pool for iterator states - // Fields for subtree iteration + // Fields for subtree iteration (original byte keys) startKey []byte // Start key for subtree iteration (nil for full trie) stopKey []byte // Stop key for subtree iteration (nil for full trie) + + // Precomputed nibble paths for efficient comparison + startPath []byte // Precomputed hex path for startKey (without terminator) + stopPath []byte // Precomputed hex path for stopKey (without terminator) + + // Iteration mode + prefixMode bool // True if this is prefix iteration (use HasPrefix check) } // errIteratorEnd is stored in nodeIterator.err when iteration is done. @@ -301,25 +308,33 @@ func (it *nodeIterator) Next(descend bool) bool { return false } - // Check if we're still within the subtree boundaries - // Note: path is already hex-encoded by the iterator - if it.startKey != nil && len(path) > 0 { - startKeyHex := keybytesToHex(it.startKey) - // Remove terminator from startKey hex if present - if hasTerm(startKeyHex) { - startKeyHex = startKeyHex[:len(startKeyHex)-1] - } - if !bytes.HasPrefix(path, startKeyHex) { - it.err = errIteratorEnd - return false - } - } - if it.stopKey != nil && len(path) > 0 { - stopKeyHex := keybytesToHex(it.stopKey) - if hasTerm(stopKeyHex) { - stopKeyHex = stopKeyHex[:len(stopKeyHex)-1] - } - if bytes.Compare(path, stopKeyHex) >= 0 { + // Check if we're still within the subtree boundaries using precomputed paths + if it.startPath != nil && len(path) > 0 { + if it.prefixMode { + // For prefix iteration, use HasPrefix to ensure we stay within the prefix + if !bytes.HasPrefix(path, it.startPath) { + it.err = errIteratorEnd + return false + } + } else { + // For range iteration, ensure we don't return nodes before the lower bound. + // Advance the iterator until we reach a node at or after startPath. + for bytes.Compare(path, it.startPath) < 0 { + // Progress the iterator by pushing the current candidate, then peeking again. + it.push(state, parentIndex, path) + state, parentIndex, path, err = it.peek(descend) + it.err = err + if it.err != nil { + return false + } + if len(path) == 0 { + break + } + } + } + } + if it.stopPath != nil && len(path) > 0 { + if bytes.Compare(path, it.stopPath) >= 0 { it.err = errIteratorEnd return false } @@ -867,21 +882,85 @@ func (it *unionIterator) Error() error { } // NewSubtreeIterator creates an iterator that only traverses nodes within a subtree -// defined by the given startKey and stopKey. The startKey defines where iteration -// starts, and stopKey defines where it ends (exclusive). +// defined by the given startKey and stopKey. This supports general range iteration +// where startKey is inclusive and stopKey is exclusive. +// +// The iterator will only visit nodes whose keys k satisfy: startKey <= k < stopKey, +// where comparisons are performed in lexicographic order of byte keys (internally +// implemented via hex-nibble path comparisons for efficiency). +// +// If startKey is nil, iteration starts from the beginning. If stopKey is nil, +// iteration continues to the end of the trie. For prefix iteration, use the +// Trie.NodeIteratorWithPrefix method which handles prefix semantics correctly. func NewSubtreeIterator(trie *Trie, startKey, stopKey []byte) NodeIterator { + return newSubtreeIterator(trie, startKey, stopKey, false) +} + +// newPrefixIterator creates an iterator that only traverses nodes with the given prefix. +// This ensures that only keys starting with the prefix are visited. +func newPrefixIterator(trie *Trie, prefix []byte) NodeIterator { + stopKey := nextKey(prefix) + return newSubtreeIterator(trie, prefix, stopKey, true) +} + +// nextKey returns the next possible key after the given prefix. +// For example, "abc" -> "abd", "ab\xff" -> "ac", etc. +func nextKey(prefix []byte) []byte { + if len(prefix) == 0 { + return nil + } + // Make a copy to avoid modifying the original + next := make([]byte, len(prefix)) + copy(next, prefix) + + // Increment the last byte that isn't 0xff + for i := len(next) - 1; i >= 0; i-- { + if next[i] < 0xff { + next[i]++ + return next + } + // If it's 0xff, we need to carry over + // Trim trailing 0xff bytes + next = next[:i] + } + // If all bytes were 0xff, return nil (no upper bound) + return nil +} + +func newSubtreeIterator(trie *Trie, startKey, stopKey []byte, prefixMode bool) NodeIterator { + // Precompute nibble paths for efficient comparison + var startPath, stopPath []byte + if startKey != nil { + startPath = keybytesToHex(startKey) + if hasTerm(startPath) { + startPath = startPath[:len(startPath)-1] + } + } + if stopKey != nil { + stopPath = keybytesToHex(stopKey) + if hasTerm(stopPath) { + stopPath = stopPath[:len(stopPath)-1] + } + } + if trie.Hash() == types.EmptyRootHash { return &nodeIterator{ - trie: trie, - err: errIteratorEnd, - startKey: startKey, - stopKey: stopKey, + trie: trie, + err: errIteratorEnd, + startKey: startKey, + stopKey: stopKey, + startPath: startPath, + stopPath: stopPath, + prefixMode: prefixMode, } } it := &nodeIterator{ - trie: trie, - startKey: startKey, - stopKey: stopKey, + trie: trie, + startKey: startKey, + stopKey: stopKey, + startPath: startPath, + stopPath: stopPath, + prefixMode: prefixMode, } // Seek to the starting position if startKey is provided if startKey != nil && len(startKey) > 0 { diff --git a/trie/iterator_test.go b/trie/iterator_test.go index b0e7999ab0..f433c266d9 100644 --- a/trie/iterator_test.go +++ b/trie/iterator_test.go @@ -644,13 +644,12 @@ func TestSubtreeIterator(t *testing.T) { root, nodes := tr.Commit(false) db.Update(root, types.EmptyRootHash, trienode.NewWithNodeSet(nodes)) - // Test subtree iteration with prefix "do" + // Test prefix iteration using NewPrefixIterator prefix := []byte("do") - stop := []byte("e") // We need to re-open the trie from the committed state tr, _ = New(TrieID(root), db) - it := NewSubtreeIterator(tr, prefix, stop) + it := newPrefixIterator(tr, prefix) found := make(map[string]string) for it.Next(true) { @@ -782,6 +781,253 @@ func TestEmptyPrefixIterator(t *testing.T) { } } +// TestPrefixIteratorEdgeCases tests various edge cases for prefix iteration +func TestPrefixIteratorEdgeCases(t *testing.T) { + // Create a trie with test data + trie := NewEmpty(newTestDatabase(rawdb.NewMemoryDatabase(), rawdb.HashScheme)) + testData := map[string]string{ + "abc": "value1", + "abcd": "value2", + "abce": "value3", + "abd": "value4", + "dog": "value5", + "dog\xff": "value6", // Test with 0xff byte + "dog\xff\xff": "value7", // Multiple 0xff bytes + } + for key, value := range testData { + trie.Update([]byte(key), []byte(value)) + } + + // Test 1: Prefix not present in trie + t.Run("NonexistentPrefix", func(t *testing.T) { + iter, err := trie.NodeIteratorWithPrefix([]byte("xyz")) + if err != nil { + t.Fatalf("Failed to create iterator: %v", err) + } + count := 0 + for iter.Next(true) { + if iter.Leaf() { + count++ + } + } + if count != 0 { + t.Errorf("Expected 0 results for nonexistent prefix, got %d", count) + } + }) + + // Test 2: Prefix exactly equals an existing key + t.Run("ExactKeyPrefix", func(t *testing.T) { + iter, err := trie.NodeIteratorWithPrefix([]byte("abc")) + if err != nil { + t.Fatalf("Failed to create iterator: %v", err) + } + found := make(map[string]bool) + for iter.Next(true) { + if iter.Leaf() { + found[string(iter.LeafKey())] = true + } + } + // Should find "abc", "abcd", "abce" but not "abd" + if !found["abc"] || !found["abcd"] || !found["abce"] { + t.Errorf("Missing expected keys: got %v", found) + } + if found["abd"] { + t.Errorf("Found unexpected key 'abd' with prefix 'abc'") + } + }) + + // Test 3: Prefix with trailing 0xff + t.Run("TrailingFFPrefix", func(t *testing.T) { + iter, err := trie.NodeIteratorWithPrefix([]byte("dog\xff")) + if err != nil { + t.Fatalf("Failed to create iterator: %v", err) + } + found := make(map[string]bool) + for iter.Next(true) { + if iter.Leaf() { + found[string(iter.LeafKey())] = true + } + } + // Should find "dog\xff" and "dog\xff\xff" + if !found["dog\xff"] || !found["dog\xff\xff"] { + t.Errorf("Missing expected keys with 0xff: got %v", found) + } + if found["dog"] { + t.Errorf("Found unexpected key 'dog' with prefix 'dog\\xff'") + } + }) + + // Test 4: All 0xff case (edge case for nextKey) + t.Run("AllFFPrefix", func(t *testing.T) { + // Add a key with all 0xff bytes + allFF := []byte{0xff, 0xff} + trie.Update(allFF, []byte("all_ff_value")) + trie.Update(append(allFF, 0x00), []byte("all_ff_plus")) + + iter, err := trie.NodeIteratorWithPrefix(allFF) + if err != nil { + t.Fatalf("Failed to create iterator: %v", err) + } + count := 0 + for iter.Next(true) { + if iter.Leaf() { + count++ + } + } + // Should find at least the allFF key itself + if count < 1 { + t.Errorf("Expected at least 1 result for all-0xff prefix, got %d", count) + } + }) + + // Test 5: Empty prefix (should iterate entire trie) + t.Run("EmptyPrefix", func(t *testing.T) { + iter, err := trie.NodeIteratorWithPrefix([]byte{}) + if err != nil { + t.Fatalf("Failed to create iterator: %v", err) + } + count := 0 + for iter.Next(true) { + if iter.Leaf() { + count++ + } + } + // Should find all keys in the trie + expectedCount := len(testData) + 2 // +2 for the extra keys added in test 4 + if count != expectedCount { + t.Errorf("Expected %d results for empty prefix, got %d", expectedCount, count) + } + }) +} + +// TestGeneralRangeIteration tests NewSubtreeIterator with arbitrary start/stop ranges +func TestGeneralRangeIteration(t *testing.T) { + // Create a trie with test data + trie := NewEmpty(newTestDatabase(rawdb.NewMemoryDatabase(), rawdb.HashScheme)) + testData := map[string]string{ + "apple": "fruit1", + "apricot": "fruit2", + "banana": "fruit3", + "cherry": "fruit4", + "date": "fruit5", + "fig": "fruit6", + "grape": "fruit7", + } + for key, value := range testData { + trie.Update([]byte(key), []byte(value)) + } + + // Test range iteration from "banana" to "fig" (exclusive) + t.Run("RangeIteration", func(t *testing.T) { + iter := NewSubtreeIterator(trie, []byte("banana"), []byte("fig")) + found := make(map[string]bool) + for iter.Next(true) { + if iter.Leaf() { + found[string(iter.LeafKey())] = true + } + } + // Should find "banana", "cherry", "date" but not "fig" + if !found["banana"] || !found["cherry"] || !found["date"] { + t.Errorf("Missing expected keys in range: got %v", found) + } + if found["apple"] || found["apricot"] || found["fig"] || found["grape"] { + t.Errorf("Found unexpected keys outside range: got %v", found) + } + }) + + // Test with nil stopKey (iterate to end) + t.Run("NilStopKey", func(t *testing.T) { + iter := NewSubtreeIterator(trie, []byte("date"), nil) + found := make(map[string]bool) + for iter.Next(true) { + if iter.Leaf() { + found[string(iter.LeafKey())] = true + } + } + // Should find "date", "fig", "grape" + if !found["date"] || !found["fig"] || !found["grape"] { + t.Errorf("Missing expected keys from 'date' to end: got %v", found) + } + if found["apple"] || found["banana"] || found["cherry"] { + t.Errorf("Found unexpected keys before 'date': got %v", found) + } + }) + + // Test with nil startKey (iterate from beginning) + t.Run("NilStartKey", func(t *testing.T) { + iter := NewSubtreeIterator(trie, nil, []byte("cherry")) + found := make(map[string]bool) + for iter.Next(true) { + if iter.Leaf() { + found[string(iter.LeafKey())] = true + } + } + // Should find "apple", "apricot", "banana" but not "cherry" or later + if !found["apple"] || !found["apricot"] || !found["banana"] { + t.Errorf("Missing expected keys before 'cherry': got %v", found) + } + if found["cherry"] || found["date"] || found["fig"] || found["grape"] { + t.Errorf("Found unexpected keys at or after 'cherry': got %v", found) + } + }) +} + +// TestPrefixIteratorWithDescend tests prefix iteration with descend=false +func TestPrefixIteratorWithDescend(t *testing.T) { + // Create a trie with nested structure + trie := NewEmpty(newTestDatabase(rawdb.NewMemoryDatabase(), rawdb.HashScheme)) + testData := map[string]string{ + "a": "value_a", + "a/b": "value_ab", + "a/b/c": "value_abc", + "a/b/d": "value_abd", + "a/e": "value_ae", + "b": "value_b", + } + for key, value := range testData { + trie.Update([]byte(key), []byte(value)) + } + + // Test skipping subtrees with descend=false + t.Run("SkipSubtrees", func(t *testing.T) { + iter, err := trie.NodeIteratorWithPrefix([]byte("a")) + if err != nil { + t.Fatalf("Failed to create iterator: %v", err) + } + + // Count nodes at each level + nodesVisited := 0 + leafsFound := make(map[string]bool) + + // First call with descend=true to enter the "a" subtree + if !iter.Next(true) { + t.Fatal("Expected to find at least one node") + } + nodesVisited++ + + // Continue iteration, sometimes with descend=false + descendPattern := []bool{false, true, false, true, true} + for i := 0; iter.Next(descendPattern[i%len(descendPattern)]); i++ { + nodesVisited++ + if iter.Leaf() { + leafsFound[string(iter.LeafKey())] = true + } + } + + // We should still respect the prefix boundary even when skipping + for key := range leafsFound { + if !bytes.HasPrefix([]byte(key), []byte("a")) { + t.Errorf("Found key outside prefix when using descend=false: %s", key) + } + } + + // Should not have found "b" even if we skip some subtrees + if leafsFound["b"] { + t.Error("Iterator leaked outside prefix boundary with descend=false") + } + }) +} + func BenchmarkIterator(b *testing.B) { diskDb, srcDb, tr, _ := makeTestTrie(rawdb.HashScheme) root := tr.Hash() diff --git a/trie/trie.go b/trie/trie.go index 782fb04db4..02c62f6cf1 100644 --- a/trie/trie.go +++ b/trie/trie.go @@ -134,15 +134,20 @@ func (t *Trie) NodeIterator(start []byte) (NodeIterator, error) { return newNodeIterator(t, start), nil } -// NodeIteratorWithPrefix returns an iterator that returns nodes of the trie with a specific prefix. -// Iteration starts at the key after the given prefix and stops when leaving the subtree. +// NodeIteratorWithPrefix returns an iterator that returns nodes of the trie whose leaf keys +// start with the given prefix. Iteration includes all keys k where prefix <= k < nextKey(prefix), +// effectively returning only keys that have the prefix. The iteration stops once it would +// encounter a key that doesn't start with the prefix. +// +// For example, with prefix "dog", the iterator will return "dog", "dogcat", "dogfish" but +// not "dot" or "fog". An empty prefix iterates the entire trie. func (t *Trie) NodeIteratorWithPrefix(prefix []byte) (NodeIterator, error) { // Short circuit if the trie is already committed and not usable. if t.committed { return nil, ErrCommitted } - // Use NewSubtreeIterator with just a startKey and no stopKey boundary - return NewSubtreeIterator(t, prefix, nil), nil + // Use the dedicated prefix iterator which handles prefix checking correctly + return newPrefixIterator(t, prefix), nil } // MustGet is a wrapper of Get and will omit any encountered error but just