Skip to content

Commit 158b138

Browse files
committed
mpt: add support for prefix scans and proofs
Replace Tree.Prove with a single Tree.Path method that returns the raw path proof for a key or prefix lookup, plus top-level ProveLookup and ProvePrefix functions that reduce a path to a key or prefix claim and its proof, verified by VerifyLookup (renamed from Verify) and the new VerifyPrefix, which share a verifyPath helper. This keeps the derivation logic in one place instead of duplicating it across the mem and disk Tree implementations. Implement Scan for both implementations. Extend testdata/mkverify.go and testdata/verify.txt with prefix proof test vectors, and add testdata/mkprefix.go, which generates testdata/prefix.txt: the full matrix of prefixes tried against a range of trees (empty, every prefix of a set of keys chosen to exercise the key padding rules, and 200 SHA256-derived keys), so that other implementations can check Scan, ProvePrefix, and VerifyPrefix against the same vectors. TestPrefix now reads that file instead of generating the matrix itself, and TestMalformedPath is new.
1 parent f5aef32 commit 158b138

10 files changed

Lines changed: 8638 additions & 168 deletions

File tree

mpt/disk.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -201,7 +201,7 @@ func memOpen(file1, file2, disk File, op string) (_ Tree, err error) {
201201
if err != nil {
202202
return nil, err
203203
}
204-
h := emptyTreeHash()
204+
h := emptyTreeHash
205205
if err := t.mutate(mem[hdrHash:], h[:]); err != nil {
206206
return nil, err
207207
}

mpt/dmem.go

Lines changed: 138 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@ import (
88
"bytes"
99
"encoding/binary"
1010
"fmt"
11+
"iter"
1112
)
1213

1314
// hash returns the hash for the given tree node.
@@ -338,44 +339,43 @@ func (t *diskTree) predict(s []node, a addr, pbit int, list []KeyVal) ([]node, [
338339
return s, list, nil
339340
}
340341

341-
// Prove returns a proof of the presence or absence of key in t.
342-
func (t *diskTree) Prove(key Key) (val Val, ok bool, proof Proof, err error) {
342+
// Path returns the path proof for key in t.
343+
func (t *diskTree) Path(key Key) (Proof, error) {
343344
t.mmu.RLock()
344345
defer t.mmu.RUnlock()
345346

346347
if t.err != nil {
347-
return Val{}, false, nil, t.err
348+
return nil, t.err
348349
}
349350
if t.hdr().dirty() {
350-
return Val{}, false, nil, ErrModifiedTree
351+
return nil, ErrModifiedTree
351352
}
352353
root, err := t.node(t.hdr().root())
353354
if err != nil {
354-
return Val{}, false, nil, err
355+
return nil, err
355356
}
356357
if root == nil {
357-
return Val{}, false, Proof{}, nil
358+
return Proof{}, nil
358359
}
359-
return root.prove(t, -1, key)
360+
return root.path(t, -1, key)
360361
}
361362

362-
func (n *diskNode) prove(t *diskTree, pbit int, key Key) (val Val, ok bool, proof Proof, err error) {
363+
// path returns the path proof for key in the subtree rooted at n.
364+
// pbit is the parent bit depth, controlling whether n is viewed as a leaf.
365+
func (n *diskNode) path(t *diskTree, pbit int, key Key) (Proof, error) {
363366
nbit := n.bit()
364367
if nbit <= pbit {
365368
// view n as leaf
366369
nkey, nval, err := n.keyVal(t)
367370
if err != nil {
368-
return Val{}, false, nil, err
369-
}
370-
if bytes.Equal(nkey, key) {
371-
return nval, true, Proof{}, nil
371+
return nil, err
372372
}
373373
var p Proof
374374
p = binary.AppendUvarint(p, uint64(len(nkey)))
375-
p = append(p, nkey[:]...)
375+
p = append(p, nkey...)
376376
p = binary.AppendUvarint(p, uint64(len(nval)))
377-
p = append(p, nval[:]...)
378-
return Val{}, false, p, nil
377+
p = append(p, nval...)
378+
return p, nil
379379
}
380380

381381
childAddr, sibAddr := n.left(), n.right()
@@ -384,24 +384,141 @@ func (n *diskNode) prove(t *diskTree, pbit int, key Key) (val Val, ok bool, proo
384384
}
385385
child, err := t.node(childAddr)
386386
if err != nil {
387-
return Val{}, false, nil, err
387+
return nil, err
388388
}
389389
sib, err := t.node(sibAddr)
390390
if err != nil {
391-
return Val{}, false, nil, err
391+
return nil, err
392392
}
393393
sibHash, err := sib.hash(t, nbit)
394394
if err != nil {
395-
return Val{}, false, nil, err
395+
return nil, err
396396
}
397397

398-
val, ok, proof, err = child.prove(t, nbit, key)
398+
proof, err := child.path(t, nbit, key)
399399
if err != nil {
400-
return
400+
return nil, err
401401
}
402402
proof = binary.AppendUvarint(proof, uint64(nbit))
403403
proof = append(proof, sibHash[:]...)
404-
return
404+
return proof, nil
405+
}
406+
407+
// Scan returns an iterator over key-val pairs whose keys start with prefix.
408+
func (t *diskTree) Scan(prefix []byte) iter.Seq2[KeyVal, error] {
409+
return func(yield func(KeyVal, error) bool) {
410+
t.mmu.RLock()
411+
defer t.mmu.RUnlock()
412+
413+
if t.err != nil {
414+
yield(KeyVal{}, t.err)
415+
return
416+
}
417+
n, pbit, err := t.subtree(prefix)
418+
if err != nil {
419+
yield(KeyVal{}, err)
420+
return
421+
}
422+
if n == nil {
423+
return
424+
}
425+
n.scan(t, pbit, yield)
426+
}
427+
}
428+
429+
// subtree returns the root of the subtree holding every key that starts
430+
// with prefix, along with its parent's bit depth, or nil if the tree
431+
// holds no such key.
432+
func (t *diskTree) subtree(prefix []byte) (*diskNode, int, error) {
433+
n, err := t.node(t.hdr().root())
434+
if err != nil || n == nil {
435+
return nil, 0, err
436+
}
437+
438+
// Look up prefix as if it were a key, stopping at the first node
439+
// that splits at a bit index at or past the end of the prefix:
440+
// every key below that node agrees with the others there,
441+
// so either all of them start with prefix or none do.
442+
pbits := 8 * len(prefix)
443+
pbit := -1
444+
for n.bit() > pbit && n.bit() < pbits {
445+
nbit := n.bit()
446+
a := n.left()
447+
if bit(prefix, nbit) != 0 {
448+
a = n.right()
449+
}
450+
if n, err = t.child(a); err != nil {
451+
return nil, 0, err
452+
}
453+
pbit = nbit
454+
}
455+
456+
// Check one key to decide for all of them.
457+
left, err := n.leftmost(t, pbit)
458+
if err != nil {
459+
return nil, 0, err
460+
}
461+
key, err := left.key(t)
462+
if err != nil {
463+
return nil, 0, err
464+
}
465+
if !key.HasPrefix(prefix) {
466+
return nil, 0, nil
467+
}
468+
return n, pbit, nil
469+
}
470+
471+
// child returns the node at address a, which must not be a nil address.
472+
func (t *diskTree) child(a addr) (*diskNode, error) {
473+
n, err := t.node(a)
474+
if err != nil {
475+
return nil, err
476+
}
477+
if n == nil {
478+
return nil, t.broken(errCorrupt)
479+
}
480+
return n, nil
481+
}
482+
483+
// leftmost returns the leaf holding the smallest key
484+
// in the subtree rooted at n.
485+
func (n *diskNode) leftmost(t *diskTree, pbit int) (*diskNode, error) {
486+
for n.bit() > pbit {
487+
pbit = n.bit()
488+
next, err := t.child(n.left())
489+
if err != nil {
490+
return nil, err
491+
}
492+
n = next
493+
}
494+
return n, nil
495+
}
496+
497+
// scan yields the key-val pairs in the subtree rooted at n, in key order,
498+
// reporting whether iteration should continue.
499+
func (n *diskNode) scan(t *diskTree, pbit int, yield func(KeyVal, error) bool) bool {
500+
nbit := n.bit()
501+
if nbit <= pbit {
502+
// view n as leaf
503+
key, val, err := n.keyVal(t)
504+
if err != nil {
505+
yield(KeyVal{}, err)
506+
return false
507+
}
508+
return yield(KeyVal{key, val}, nil)
509+
}
510+
511+
left, err := t.child(n.left())
512+
if err != nil {
513+
yield(KeyVal{}, err)
514+
return false
515+
}
516+
right, err := t.child(n.right())
517+
if err != nil {
518+
yield(KeyVal{}, err)
519+
return false
520+
}
521+
return left.scan(t, nbit, yield) && right.scan(t, nbit, yield)
405522
}
406523

407524
func (t *diskTree) check() {

mpt/mem.go

Lines changed: 84 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ import (
99
"encoding/binary"
1010
"errors"
1111
"fmt"
12+
"iter"
1213
)
1314

1415
// A memTree is an in-memory [Tree].
@@ -42,7 +43,7 @@ func (n *memNode) bit() int {
4243
// NewMemTree returns a new in-memory [Tree].
4344
func NewMemTree() Tree {
4445
t := &memTree{
45-
hash: emptyTreeHash(),
46+
hash: emptyTreeHash,
4647
exact: true,
4748
}
4849
return t
@@ -245,33 +246,32 @@ func (t *memTree) predict(s []node, n *memNode, pbit int, list []KeyVal) ([]node
245246
return s, list
246247
}
247248

248-
// Prove returns a proof of the presence or absence of key in t.
249-
func (t *memTree) Prove(key Key) (val Val, ok bool, proof Proof, err error) {
249+
// Path returns the path proof for key in t.
250+
func (t *memTree) Path(key Key) (Proof, error) {
250251
if t.err != nil {
251-
return Val{}, false, nil, t.err
252+
return nil, t.err
252253
}
253254
if t.dirty {
254-
return Val{}, false, nil, ErrModifiedTree
255+
return nil, ErrModifiedTree
255256
}
256257
if t.root == nil {
257-
return Val{}, false, Proof{}, nil
258+
return Proof{}, nil
258259
}
259-
return t.root.prove(-1, key)
260+
return t.root.path(-1, key), nil
260261
}
261262

262-
func (n *memNode) prove(pbit int, key Key) (val Val, ok bool, proof Proof, err error) {
263+
// path returns the path proof for key in the subtree rooted at n.
264+
// pbit is the parent bit depth, controlling whether n is viewed as a leaf.
265+
func (n *memNode) path(pbit int, key Key) Proof {
263266
nbit := n.bit()
264267
if nbit <= pbit {
265268
// view n as leaf
266-
if bytes.Equal(n.key, key) {
267-
return n.val, true, Proof{}, nil
268-
}
269269
var p Proof
270270
p = binary.AppendUvarint(p, uint64(len(n.key)))
271-
p = append(p, n.key[:]...)
271+
p = append(p, n.key...)
272272
p = binary.AppendUvarint(p, uint64(len(n.val)))
273-
p = append(p, n.val[:]...)
274-
return Val{}, false, p, nil
273+
p = append(p, n.val...)
274+
return p
275275
}
276276

277277
var sib Hash
@@ -284,10 +284,78 @@ func (n *memNode) prove(pbit int, key Key) (val Val, ok bool, proof Proof, err e
284284
sib = n.left.hash(nbit)
285285
}
286286

287-
val, ok, proof, _ = child.prove(nbit, key)
287+
proof := child.path(nbit, key)
288288
proof = binary.AppendUvarint(proof, uint64(nbit))
289289
proof = append(proof, sib[:]...)
290-
return
290+
return proof
291+
}
292+
293+
// Scan returns an iterator over key-val pairs whose keys start with prefix.
294+
func (t *memTree) Scan(prefix []byte) iter.Seq2[KeyVal, error] {
295+
return func(yield func(KeyVal, error) bool) {
296+
if t.err != nil {
297+
yield(KeyVal{}, t.err)
298+
return
299+
}
300+
n, pbit := t.subtree(prefix)
301+
if n == nil {
302+
return
303+
}
304+
n.scan(pbit, yield)
305+
}
306+
}
307+
308+
// subtree returns the root of the subtree holding every key that starts
309+
// with prefix, along with its parent's bit depth, or nil if the tree
310+
// holds no such key.
311+
func (t *memTree) subtree(prefix []byte) (*memNode, int) {
312+
n := t.root
313+
if n == nil {
314+
return nil, 0
315+
}
316+
317+
// Look up prefix as if it were a key, stopping at the first node
318+
// that splits at a bit index at or past the end of the prefix:
319+
// every key below that node agrees with the others there,
320+
// so either all of them start with prefix or none do.
321+
pbits := 8 * len(prefix)
322+
pbit := -1
323+
for n.bit() > pbit && n.bit() < pbits {
324+
nbit := n.bit()
325+
if bit(prefix, nbit) == 0 {
326+
n = n.left
327+
} else {
328+
n = n.right
329+
}
330+
pbit = nbit
331+
}
332+
333+
// Check one key to decide for all of them.
334+
if !n.leftmost(pbit).key.HasPrefix(prefix) {
335+
return nil, 0
336+
}
337+
return n, pbit
338+
}
339+
340+
// leftmost returns the leaf holding the smallest key
341+
// in the subtree rooted at n.
342+
func (n *memNode) leftmost(pbit int) *memNode {
343+
for n.bit() > pbit {
344+
pbit = n.bit()
345+
n = n.left
346+
}
347+
return n
348+
}
349+
350+
// scan yields the key-val pairs in the subtree rooted at n, in key order,
351+
// reporting whether iteration should continue.
352+
func (n *memNode) scan(pbit int, yield func(KeyVal, error) bool) bool {
353+
nbit := n.bit()
354+
if nbit <= pbit {
355+
// view n as leaf
356+
return yield(KeyVal{n.key, n.val}, nil)
357+
}
358+
return n.left.scan(nbit, yield) && n.right.scan(nbit, yield)
291359
}
292360

293361
// check checks all the tree invariants, walking the entire tree.

0 commit comments

Comments
 (0)