diff --git a/bench_test.go b/bench_test.go index fc799cd..befdd0d 100644 --- a/bench_test.go +++ b/bench_test.go @@ -5,16 +5,14 @@ import ( "bytes" "encoding/binary" "fmt" - "maps" + "math" rand "math/rand/v2" "os" - "reflect" "slices" "strings" "testing" "github.com/phiryll/kv" - "github.com/stretchr/testify/assert" ) // No benchmark can have a truly random element, random seeds must be constants! @@ -48,6 +46,45 @@ var ( benchStoreConfigs = createBenchStoreConfigs() ) +func randomBytes(n int, random *rand.Rand) []byte { + if n == 0 { + return []byte{} + } + k := (n-1)/8 + 1 + b := make([]byte, k*8) + for i := range k { + binary.BigEndian.PutUint64(b[i*8:], random.Uint64()) + } + return b[:n] +} + +func randomByte(random *rand.Rand) byte { + return byte(random.UintN(256)) +} + +// Returns a random key of with length chosen from a roughly normal distribution +// with the given mean. Lengths will range from 0 to 2*mean. +func randomKey(meanLen int, random *rand.Rand) []byte { + const bound = 4.0 // chosen experimentally + val := random.NormFloat64() + for val < -bound || val > +bound { + val = random.NormFloat64() + } + // val is in [-bound, +bound], translate that to [0, 2*mean] + val = (val + bound) * float64(meanLen) / bound + return randomBytes(int(math.Round(val)), random) +} + +func randomFixedLengthKey(keyLen int, random *rand.Rand) []byte { + return randomBytes(keyLen, random) +} + +func shuffle[S ~[]E, E any](slice S, random *rand.Rand) { + random.Shuffle(len(slice), func(i, j int) { + slice[i], slice[j] = slice[j], slice[i] + }) +} + func BenchmarkTraverser(b *testing.B) { benchTraverser(b, "kind=pre-order", kv.TestingPreOrder) benchTraverser(b, "kind=post-order", kv.TestingPostOrder) @@ -116,6 +153,8 @@ func BenchmarkChildBounds(b *testing.B) { bounds *Bounds keys keySet }{ + // No need to benchmark reverse bounds, the code is the same. + // no common prefix { From(nil).To(empty), @@ -127,14 +166,11 @@ func BenchmarkChildBounds(b *testing.B) { }, { From(nil).To(high), - keySet{ - empty, nextKey(empty), before, high[:1], high[:2], high[:3], - prevKey(high), high, nextKey(high), after, - }, + keySet{empty, high[:1], high[:2], high[:3], prevKey(high), high, nextKey(high), after}, }, { From(nil).To(nil), - keySet{empty, nextKey(empty), within}, + keySet{empty, nextKey(empty), after}, }, { From(empty).To(nextKey(empty)), @@ -142,36 +178,33 @@ func BenchmarkChildBounds(b *testing.B) { }, { From(empty).To(high), - keySet{ - empty, nextKey(empty), before, high[:1], high[:2], high[:3], - prevKey(high), high, nextKey(high), after, - }, + keySet{empty, nextKey(empty), high[:1], high[:2], high[:3], prevKey(high), high, nextKey(high), after}, }, { From(empty).To(nil), - keySet{empty, nextKey(empty), within}, + keySet{empty, nextKey(empty), after}, }, { From(nextKey(empty)).To(high), keySet{ - empty, nextKey(empty), nextKey(nextKey(empty)), before, high[:1], high[:2], high[:3], + empty, nextKey(empty), nextKey(nextKey(empty)), high[:1], high[:2], high[:3], prevKey(high), high, nextKey(high), after, }, }, { From(nextKey(empty)).To(nil), - keySet{empty, nextKey(empty), nextKey(nextKey(empty)), within}, + keySet{empty, nextKey(empty), nextKey(nextKey(empty)), after}, }, { From(low).To(high), keySet{ - empty, nextKey(empty), before, low[:1], low[:2], low[:3], prevKey(low), low, nextKey(low), - within, high[:1], high[:2], high[:3], prevKey(high), high, nextKey(high), after, + empty, nextKey(empty), low[:1], low[:2], low[:3], prevKey(low), low, nextKey(low), + high[:1], high[:2], high[:3], prevKey(high), high, nextKey(high), after, }, }, { From(low).To(nil), - keySet{empty, nextKey(empty), before, low[:1], low[:2], low[:3], prevKey(low), low, nextKey(low), after}, + keySet{empty, nextKey(empty), low[:1], low[:2], low[:3], prevKey(low), low, nextKey(low), after}, }, // 2 byte common prefix @@ -193,20 +226,11 @@ func BenchmarkChildBounds(b *testing.B) { }, } { b.Run(fmt.Sprintf("bounds=%s", tt.bounds), func(b *testing.B) { - forward := tt.bounds - reverse := From(tt.bounds.End).DownTo(tt.bounds.Begin) for _, k := range tt.keys { b.Run("key="+kv.KeyName(k), func(b *testing.B) { - b.Run("dir=forward", func(b *testing.B) { - for b.Loop() { - kv.TestingChildBounds(forward, k) - } - }) - b.Run("dir=reverse", func(b *testing.B) { - for b.Loop() { - kv.TestingChildBounds(reverse, k) - } - }) + for b.Loop() { + kv.TestingChildBounds(tt.bounds, k) + } }) } }) @@ -367,55 +391,10 @@ func createBenchStoreConfigs() []*storeConfig { return append(createBenchRandomStoreConfigs(), createBenchWordStoreConfigs()...) } -func TestBenchStoreConfigs(t *testing.T) { - t.Parallel() - for _, config := range benchStoreConfigs { - t.Run(config.name, func(t *testing.T) { - t.Parallel() - assert.Len(t, config.entries, config.size) - assert.Equal(t, len(config.present), len(config.absent)) - assert.Equal(t, 1, len(config.present[0])+len(config.absent[0])) - assert.Equal(t, 1<<8, len(config.present[1])+len(config.absent[1])) - assert.Equal(t, 1<<16, len(config.present[2])+len(config.absent[2])) - assert.Len(t, config.forward, 1<<16) - assert.Len(t, config.reverse, 1<<16) - - present := maps.Clone(config.entries) - for i := range len(config.present) { - if i > 2 { - assert.Len(t, config.absent[i], 1<<16) - } - for _, k := range config.absent[i] { - assert.Len(t, k, i) - _, ok := present[string(k)] - assert.False(t, ok) - } - for _, k := range config.present[i] { - assert.Len(t, k, i) - _, ok := present[string(k)] - assert.True(t, ok) - delete(present, string(k)) - } - } - assert.Empty(t, present) - }) - } -} - -func TestBenchStoreConfigRepeatability(t *testing.T) { - t.Parallel() - for i, config := range createBenchStoreConfigs() { - t.Run(config.name, func(t *testing.T) { - t.Parallel() - assert.True(t, reflect.DeepEqual(benchStoreConfigs[i], config)) - }) - } -} - // This helps to understand how factory() can impact other benchmarks which use it. func BenchmarkFactory(b *testing.B) { for _, def := range implDefs { - b.Run("impl="+def.name, func(b *testing.B) { + b.Run(def.name, func(b *testing.B) { for b.Loop() { _ = def.factory() } @@ -461,7 +440,7 @@ func BenchmarkSparse(b *testing.B) { } shuffle(keys, random) for _, def := range implDefs { - b.Run("impl="+def.name, func(b *testing.B) { + b.Run(def.name, func(b *testing.B) { for b.Loop() { store := def.factory() for _, k := range keys { @@ -503,7 +482,7 @@ func BenchmarkDense(b *testing.B) { {"/keyLen=2", keySets[1]}, {"/keyLen=3", keySets[2]}, } { - b.Run("impl="+def.name+tt.name, func(b *testing.B) { + b.Run(def.name+tt.name, func(b *testing.B) { for b.Loop() { store := def.factory() for _, k := range tt.keys { @@ -558,7 +537,6 @@ func BenchmarkSet(b *testing.B) { //nolint:gocognit func BenchmarkGet(b *testing.B) { for _, bench := range createTestStores(benchStoreConfigs) { - original := bench.store b.Run(bench.name, func(b *testing.B) { for keyLen := 8; keyLen < len(bench.config.present); keyLen += 4 { b.Run(fmt.Sprintf("keyLen=%d/existing=true", keyLen), func(b *testing.B) { @@ -566,10 +544,9 @@ func BenchmarkGet(b *testing.B) { if len(present) == 0 { b.Skipf("no present keys of length %d", keyLen) } - store := original.Clone() i := 0 for b.Loop() { - store.Get(present[i%len(present)]) + bench.store.Get(present[i%len(present)]) i++ } }) @@ -578,10 +555,9 @@ func BenchmarkGet(b *testing.B) { if len(absent) == 0 { b.Skipf("no absent keys of length %d", keyLen) } - store := original.Clone() i := 0 for b.Loop() { - store.Get(absent[i%len(absent)]) + bench.store.Get(absent[i%len(absent)]) i++ } }) diff --git a/fuzz_test.go b/fuzz_test.go index 4de0fb0..e1afc8b 100644 --- a/fuzz_test.go +++ b/fuzz_test.go @@ -65,21 +65,6 @@ func createFuzzStoreConfigs(size int) []*storeConfig { return []*storeConfig{&config} } -func TestBaseline(t *testing.T) { - t.Parallel() - fuzzStores := createTestStores(fuzzStoreConfigs) - ref := createReferenceStore(fuzzStoreConfigs[0]) - refForward := collect(ref.Range(forwardAll)) - refReverse := collect(ref.Range(reverseAll)) - for _, fuzz := range fuzzStores { - t.Run(fuzz.name, func(t *testing.T) { - t.Parallel() - assert.Equal(t, refForward, collect(fuzz.store.Range(forwardAll)), "forward") - assert.Equal(t, refReverse, collect(fuzz.store.Range(reverseAll)), "reverse") - }) - } -} - func FuzzGet(f *testing.F) { fuzzStores := createTestStores(fuzzStoreConfigs) ref := createReferenceStore(fuzzStoreConfigs[0]) @@ -142,11 +127,11 @@ func FuzzRange(f *testing.F) { } forward := From(begin).To(end) reverse := From(end).DownTo(begin) - refForward := collect(ref.Range(forward)) - refReverse := collect(ref.Range(reverse)) + refForward := ref.Range(forward) + refReverse := ref.Range(reverse) for _, fuzz := range fuzzStores { - assert.Equal(t, refForward, collect(fuzz.store.Range(forward)), "%s: %s", fuzz.def.name, forward) - assert.Equal(t, refReverse, collect(fuzz.store.Range(reverse)), "%s: %s", fuzz.def.name, reverse) + assertItersEqual(t, refForward, fuzz.store.Range(forward), "%s: %s", fuzz.def.name, forward) + assertItersEqual(t, refReverse, fuzz.store.Range(reverse), "%s: %s", fuzz.def.name, reverse) } }) } diff --git a/kv_test.go b/kv_test.go index b37ac35..e80f296 100644 --- a/kv_test.go +++ b/kv_test.go @@ -2,13 +2,9 @@ package kv_test import ( "bytes" - "encoding/binary" "fmt" "iter" - "math" "math/bits" - rand "math/rand/v2" - "reflect" "slices" "strings" "testing" @@ -59,12 +55,6 @@ type ( def *implDef config *storeConfig } - - // Used to test Range result sets. - entry struct { - key []byte - value byte - } ) const ( @@ -76,9 +66,9 @@ const ( var ( implDefs = []*implDef{ - {"reference", newReference}, - {"pointer-trie", asCloneable(kv.NewPointerTrie[byte])}, - {"array-trie", asCloneable(kv.NewArrayTrie[byte])}, + {"impl=reference", newReference}, + {"impl=pointer-trie", asCloneable(kv.NewPointerTrie[byte])}, + {"impl=array-trie", asCloneable(kv.NewArrayTrie[byte])}, } From = kv.From @@ -236,71 +226,6 @@ func asCloneable(factory func() kv.Store[byte]) func() TestStore { } } -func emptySeqInt(_ func(int) bool) {} - -func emptyAdjInt(_ int) iter.Seq[int] { - return emptySeqInt -} - -func emptyPathAdjInt(_ []int) iter.Seq[int] { - return emptySeqInt -} - -func cmpEntryForward(a, b entry) int { - return bytes.Compare(a.key, b.key) -} - -func cmpEntryReverse(a, b entry) int { - return bytes.Compare(b.key, a.key) -} - -func collect(itr iter.Seq2[[]byte, byte]) []entry { - entries := []entry{} - for k, v := range itr { - entries = append(entries, entry{k, v}) - } - return entries -} - -func randomBytes(n int, random *rand.Rand) []byte { - if n == 0 { - return []byte{} - } - k := (n-1)/8 + 1 - b := make([]byte, k*8) - for i := range k { - binary.BigEndian.PutUint64(b[i*8:], random.Uint64()) - } - return b[:n] -} - -func randomByte(random *rand.Rand) byte { - return byte(random.UintN(256)) -} - -// Returns a random key of with length chosen from a roughly normal distribution -// with the given mean. Lengths will range from 0 to 2*mean. -func randomKey(meanLen int, random *rand.Rand) []byte { - const bound = 4.0 // chosen experimentally - val := random.NormFloat64() - for val < -bound || val > +bound { - val = random.NormFloat64() - } - // val is in [-bound, +bound], translate that to [0, 2*mean] - val = (val + bound) * float64(meanLen) / bound - return randomBytes(int(math.Round(val)), random) -} - -func randomFixedLengthKey(keyLen int, random *rand.Rand) []byte { - return randomBytes(keyLen, random) -} - -func shuffle[S ~[]E, E any](slice S, random *rand.Rand) { - random.Shuffle(len(slice), func(i, j int) { - slice[i], slice[j] = slice[j], slice[i] - }) -} - // storeConfigs for all possible subsequences of presentKeys. func createTestStoreConfigs() []*storeConfig { result := []*storeConfig{} @@ -345,13 +270,6 @@ func createTestStoreConfigs() []*storeConfig { return result } -func TestTestStoreConfigRepeatability(t *testing.T) { - t.Parallel() - for i, config := range createTestStoreConfigs() { - assert.True(t, reflect.DeepEqual(testStoreConfigs[i], config)) - } -} - func createReferenceStore(config *storeConfig) TestStore { store := newReference() for k, v := range config.entries { @@ -362,13 +280,13 @@ func createReferenceStore(config *storeConfig) TestStore { func createTestStores(storeConfigs []*storeConfig) []*testStore { result := []*testStore{} - for _, def := range implDefs { - for _, config := range storeConfigs { + for _, config := range storeConfigs { + for _, def := range implDefs { store := def.factory() for k, v := range config.entries { store.Set([]byte(k), v) } - name := fmt.Sprintf("impl=%s/%s", def.name, config.name) + name := config.name + "/" + def.name result = append(result, &testStore{name, store, def, config}) } } @@ -410,20 +328,35 @@ func assertAbsent(t *testing.T, key []byte, store TestStore) { } } +// Temporary. +func entryIter(entries map[string]byte, keys iter.Seq[[]byte]) iter.Seq2[[]byte, byte] { + return func(yield func([]byte, byte) bool) { + for k := range keys { + v, ok := entries[string(k)] + if !ok { + panic(fmt.Sprintf("key %s not found", kv.KeyName(k))) + } + if !yield(k, v) { + return + } + } + } +} + // Test that store contains only the key/value pairs in entries, // and that Range(forward/reverse) returns them in the correct order. func assertSame(t *testing.T, entries map[string]byte, store TestStore) { - sliceEntries := []entry{} - for k, expected := range entries { + keys := [][]byte{} + for k, v := range entries { actual, ok := store.Get([]byte(k)) assert.True(t, ok) - assert.Equal(t, expected, actual) - sliceEntries = append(sliceEntries, entry{[]byte(k), expected}) + assert.Equal(t, v, actual) + keys = append(keys, []byte(k)) } - slices.SortFunc(sliceEntries, cmpEntryForward) - assert.Equal(t, sliceEntries, collect(store.Range(forwardAll))) - slices.SortFunc(sliceEntries, cmpEntryReverse) - assert.Equal(t, sliceEntries, collect(store.Range(reverseAll))) + slices.SortFunc(keys, bytes.Compare) + assertItersEqual(t, entryIter(entries, slices.Values(keys)), store.Range(forwardAll)) + slices.Reverse(keys) + assertItersEqual(t, entryIter(entries, slices.Values(keys)), store.Range(reverseAll)) } func TestNilArgPanics(t *testing.T) { @@ -519,6 +452,7 @@ func TestStoreString(t *testing.T) { } func checkFprint[V any](t *testing.T, expected string, seq iter.Seq2[[]byte, V]) { + t.Helper() var s strings.Builder n, err := kv.Fprint(&s, seq) require.NoError(t, err) @@ -568,7 +502,54 @@ func TestFprint(t *testing.T) { } } -//nolint:gocognit +// Functions to test iterators. + +func msgFrom(msgAndArgs ...any) string { + if len(msgAndArgs) == 0 { + return "" + } + msg, ok := msgAndArgs[0].(string) + if !ok { + return fmt.Sprintf("%+v of type %T must be a format string", msgAndArgs[0], msgAndArgs[0]) + } + return fmt.Sprintf(msg, msgAndArgs[1:]...) +} + +func assertItersEqual[K, V any](t *testing.T, expected, actual iter.Seq2[K, V], msgAndArgs ...any) { + t.Helper() + msg := msgFrom(msgAndArgs...) + i := 0 + next, stop := iter.Pull2(actual) + defer stop() + for expectedKey, expectedValue := range expected { + actualKey, actualValue, ok := next() + require.True(t, ok, "%s: too short, len == %d", msg, i) + assert.Equal(t, expectedKey, actualKey, "%s: keys at index %d differ", msg, i) + assert.Equal(t, expectedValue, actualValue, "%s: values at index %d differ", msg, i) + i++ + } + _, _, ok := next() + assert.False(t, ok, "%s: too long, len > %d", msg, i) +} + +// This tests that an early yield does not fail, +// and ensures those code paths get test coverage. +func assertEarlyYield(t *testing.T, size int, itr iter.Seq2[[]byte, byte]) { + t.Helper() + expectedCount := size + if expectedCount > 4 { + expectedCount = 4 + } + count := 0 + for range itr { + if count > 3 { + break + } + count++ + } + assert.Equal(t, expectedCount, count) +} + func TestStores(t *testing.T) { t.Parallel() for _, test := range createTestStores(testStoreConfigs) { @@ -602,28 +583,13 @@ func TestStores(t *testing.T) { t.Run("op=range", func(t *testing.T) { ref := createReferenceStore(test.config) for _, bounds := range test.config.forward { - assert.Equal(t, collect(ref.Range(&bounds)), collect(store.Range(&bounds)), - "%s", &bounds) + assertItersEqual(t, ref.Range(&bounds), store.Range(&bounds), "%s", &bounds) } for _, bounds := range test.config.reverse { - assert.Equal(t, collect(ref.Range(&bounds)), collect(store.Range(&bounds)), - "%s", &bounds) - } - // need an early yield for test coverage - count := 0 - for range store.Range(forwardAll) { - if count > 3 { - break - } - count++ - } - count = 0 - for range store.Range(reverseAll) { - if count > 3 { - break - } - count++ + assertItersEqual(t, ref.Range(&bounds), store.Range(&bounds), "%s", &bounds) } + assertEarlyYield(t, test.config.size, store.Range(forwardAll)) + assertEarlyYield(t, test.config.size, store.Range(reverseAll)) }) }) } @@ -672,152 +638,171 @@ func TestClone(t *testing.T) { // Things that failed at one point or another during testing. -func testFail1(t *testing.T, factory func() TestStore) { - t.Run("fail 1", func(t *testing.T) { - t.Parallel() - store := factory() - store.Set([]byte{5}, 0) - assert.Equal(t, - []entry{}, - collect(store.Range(From([]byte{5, 0}).To([]byte{6})))) - assert.Equal(t, - []entry{{[]byte{5}, 0}}, - collect(store.Range(From([]byte{4}).To([]byte{5, 0})))) - }) +func assertIterEmpty[V any](t *testing.T, actual iter.Seq2[[]byte, V]) bool { + t.Helper() + for key, value := range actual { + return assert.Fail(t, fmt.Sprintf("should be empty, contains {%s:%v}", kv.KeyName(key), value)) + } + return true } -func testFail2(t *testing.T, factory func() TestStore) { - t.Run("fail 2", func(t *testing.T) { - t.Parallel() - store := factory() - store.Set([]byte{0xB3, 0x9C}, 184) +func assertIterSingleton[V any](t *testing.T, expectedKey []byte, expectedValue V, actual iter.Seq2[[]byte, V]) bool { + t.Helper() + next, stop := iter.Pull2(actual) + defer stop() + key, value, ok := next() + assert.True(t, ok, "should not be empty") + assert.Equal(t, expectedKey, key) + assert.Equal(t, expectedValue, value) + _, _, ok = next() + assert.False(t, ok, "should have exactly one entry") + return true +} - // forgot to check isTerminal - actual, actualOk := store.Get([]byte{0xB3}) - assert.False(t, actualOk) - assert.Equal(t, byte(0), actual) +func TestFail1(t *testing.T) { + t.Parallel() + for _, def := range implDefs { + t.Run(def.name, func(t *testing.T) { + t.Parallel() + store := def.factory() + store.Set([]byte{5}, 0) + assertIterEmpty(t, store.Range(From([]byte{5, 0}).To([]byte{6}))) + assertIterSingleton(t, []byte{5}, 0, store.Range(From([]byte{4}).To([]byte{5, 0}))) + }) + } +} - actual, actualOk = store.Get([]byte{0xB3, 0x9C}) - assert.True(t, actualOk) - assert.Equal(t, byte(184), actual) - }) +func TestFail2(t *testing.T) { + t.Parallel() + for _, def := range implDefs { + t.Run(def.name, func(t *testing.T) { + t.Parallel() + store := def.factory() + store.Set([]byte{0xB3, 0x9C}, 184) + + // forgot to check isTerminal + actual, actualOk := store.Get([]byte{0xB3}) + assert.False(t, actualOk) + assert.Equal(t, zero, actual) + + actual, actualOk = store.Get([]byte{0xB3, 0x9C}) + assert.True(t, actualOk) + assert.Equal(t, byte(184), actual) + }) + } } -func testFail3(t *testing.T, factory func() TestStore) { - t.Run("fail 3", func(t *testing.T) { - t.Parallel() - store := factory() - store.Set([]byte{0xB3, 0x9C}, 184) +func TestFail3(t *testing.T) { + t.Parallel() + for _, def := range implDefs { + t.Run(def.name, func(t *testing.T) { + t.Parallel() + store := def.factory() + store.Set([]byte{0xB3, 0x9C}, 184) - actual, actualOk := store.Delete([]byte{0xB3}) - assert.False(t, actualOk) - assert.Equal(t, byte(0), actual) + actual, actualOk := store.Delete([]byte{0xB3}) + assert.False(t, actualOk) + assert.Equal(t, zero, actual) - // Make sure the subtree wasn't deleted. - actual, actualOk = store.Get([]byte{0xB3, 0x9C}) - assert.True(t, actualOk) - assert.Equal(t, byte(184), actual) - }) + // Make sure the subtree wasn't deleted. + actual, actualOk = store.Get([]byte{0xB3, 0x9C}) + assert.True(t, actualOk) + assert.Equal(t, byte(184), actual) + }) + } } -func testFail4(t *testing.T, factory func() TestStore) { - t.Run("fail 4", func(t *testing.T) { - t.Parallel() - store := factory() - store.Set([]byte{0x50, 0xEF}, 45) - assert.Equal(t, - []entry{}, - collect(store.Range(From([]byte{0x50}).DownTo([]byte{0x15})))) - }) +func TestFail4(t *testing.T) { + t.Parallel() + for _, def := range implDefs { + t.Run(def.name, func(t *testing.T) { + t.Parallel() + store := def.factory() + store.Set([]byte{0x50, 0xEF}, 45) + assertIterEmpty(t, store.Range(From([]byte{0x50}).DownTo([]byte{0x15}))) + }) + } } -func testFail5(t *testing.T, factory func() TestStore) { - t.Run("fail 5", func(t *testing.T) { - t.Parallel() - store := factory() - store.Set([]byte{0x50, 0xEF}, 45) - assert.Equal(t, - []entry{{[]byte{0x50, 0xEF}, 45}}, - collect(store.Range(From([]byte{0xFD}).DownTo([]byte{0x3D})))) - }) +func TestFail5(t *testing.T) { + t.Parallel() + for _, def := range implDefs { + t.Run(def.name, func(t *testing.T) { + t.Parallel() + store := def.factory() + store.Set([]byte{0x50, 0xEF}, 45) + assertIterSingleton(t, []byte{0x50, 0xEF}, 45, store.Range(From([]byte{0xFD}).DownTo([]byte{0x3D}))) + }) + } } -func testFail6(t *testing.T, factory func() TestStore) { - t.Run("fail 6", func(t *testing.T) { - t.Parallel() - store := factory() - store.Set([]byte{0x50, 0xEF}, 45) - assert.Equal(t, - []entry{{[]byte{0x50, 0xEF}, 45}}, - collect(store.Range(From([]byte{0x51}).DownTo([]byte{0x50})))) - }) +func TestFail6(t *testing.T) { + t.Parallel() + for _, def := range implDefs { + t.Run(def.name, func(t *testing.T) { + t.Parallel() + store := def.factory() + store.Set([]byte{0x50, 0xEF}, 45) + assertIterSingleton(t, []byte{0x50, 0xEF}, 45, store.Range(From([]byte{0x51}).DownTo([]byte{0x50}))) + }) + } } -func testFail7(t *testing.T, factory func() TestStore) { +func TestFail7(t *testing.T) { // Failure is due to continuing iteration past false yield(). // Failure requires the second Set. - t.Run("fail 7", func(t *testing.T) { - t.Parallel() - store := factory() - store.Set([]byte{3}, 0) - store.Set([]byte{4}, 0) - assert.Equal(t, - []entry{}, - collect(store.Range(From([]byte{1}).To([]byte{2})))) - }) + t.Parallel() + for _, def := range implDefs { + t.Run(def.name, func(t *testing.T) { + t.Parallel() + store := def.factory() + store.Set([]byte{3}, 0) + store.Set([]byte{4}, 0) + assertIterEmpty(t, store.Range(From([]byte{2}).DownTo([]byte{1}))) + }) + } } -func testFail8(t *testing.T, factory func() TestStore) { - t.Run("fail 8", func(t *testing.T) { - t.Parallel() - store := factory() - store.Set([]byte{}, 1) - store.Set([]byte{0}, 3) - store.Set([]byte{0x23}, 4) - store.Set([]byte{0x23, 0}, 5) - store.Set([]byte{0x23, 0xA5}, 6) - assert.Equal(t, - []entry{{[]byte{0x23, 0}, 5}}, - collect(store.Range(From([]byte{0x23, 0}).To([]byte{0x23, 0, 0})))) - }) +func TestFail8(t *testing.T) { + t.Parallel() + for _, def := range implDefs { + t.Run(def.name, func(t *testing.T) { + t.Parallel() + store := def.factory() + store.Set([]byte{}, 1) + store.Set([]byte{0}, 3) + store.Set([]byte{0x23}, 4) + store.Set([]byte{0x23, 0}, 5) + store.Set([]byte{0x23, 0xA5}, 6) + assertIterSingleton(t, []byte{0x23, 0}, 5, store.Range(From([]byte{0x23, 0}).To([]byte{0x23, 0, 0}))) + }) + } } -func testFail9(t *testing.T, factory func() TestStore) { +func TestFail9(t *testing.T) { // Test that removing the last value on a path removes the path. // Definite hack to detect this one, // but there's no good way to test this using the public API. // The alternative would be to have implementation-specific tests, // which is probably a better approach, but this works for now. - t.Run("fail 9", func(t *testing.T) { - t.Parallel() - store := factory() - sStore, ok := store.(fmt.Stringer) - if !ok { - t.Skipf("%T does not implement Stringer", store) - } - expected := sStore.String() - key := []byte{0x23} - store.Set(key, 6) - store.Delete(key) - assert.Equal(t, expected, sStore.String()) - }) -} - -func TestPastFailures(t *testing.T) { t.Parallel() for _, def := range implDefs { - factory := def.factory t.Run(def.name, func(t *testing.T) { t.Parallel() - testFail1(t, factory) - testFail2(t, factory) - testFail3(t, factory) - testFail4(t, factory) - testFail5(t, factory) - testFail6(t, factory) - testFail7(t, factory) - testFail8(t, factory) - testFail9(t, factory) + store := def.factory() + sStore, ok := store.(fmt.Stringer) + if !ok { + t.Skipf("%T does not implement Stringer", store) + } + expected := sStore.String() + key := []byte{0x23} + store.Set(key, 6) + store.Delete(key) + // make sure reference.dirty is false + for range store.Range(forwardAll) { + break + } + assert.Equal(t, expected, sStore.String()) }) } } diff --git a/reference_test.go b/reference_test.go index 5066c4a..ad327bb 100644 --- a/reference_test.go +++ b/reference_test.go @@ -1,7 +1,7 @@ package kv_test import ( - "cmp" + "bytes" "fmt" "iter" "maps" @@ -11,88 +11,155 @@ import ( "github.com/phiryll/kv" ) +// reference serves as an expected value to compare against while testing, +// and a source of entries from which to create a new Store. +// This implementation is meant to be as trivially correct as possible. +type reference struct { + entries map[string]byte + ascKeys [][]byte + dirty bool +} + func newReference() TestStore { - return reference{} + return &reference{ + entries: map[string]byte{}, + } } -// reference implements the TestStore interface, but it is not a trie. -// This serves as an expected value to compare against a Store[byte] implementation while testing. -type reference map[string]byte +func (r *reference) Clone() TestStore { + return &reference{ + maps.Clone(r.entries), + slices.Clone(r.ascKeys), + r.dirty, + } +} -func (r reference) Clone() TestStore { - return maps.Clone(r) +func (r *reference) refresh() { + if !r.dirty { + return + } + var keys [][]byte + for key := range r.entries { + keys = append(keys, []byte(key)) + } + slices.SortFunc(keys, bytes.Compare) + r.ascKeys = keys + r.dirty = false } -func (r reference) Set(key []byte, value byte) (byte, bool) { +func (r *reference) Get(key []byte) (byte, bool) { if key == nil { panic("key must be non-nil") } - index := string(key) - prev, ok := r[index] - r[index] = value - if ok { - return prev, true - } - return 0, false + value, ok := r.entries[string(key)] + return value, ok } -func (r reference) Get(key []byte) (byte, bool) { +func (r *reference) Set(key []byte, value byte) (byte, bool) { if key == nil { panic("key must be non-nil") } - value, ok := r[string(key)] - return value, ok + prev, ok := r.entries[string(key)] + r.entries[string(key)] = value + r.ascKeys = nil + r.dirty = true + return prev, ok } -func (r reference) Delete(key []byte) (byte, bool) { +func (r *reference) Delete(key []byte) (byte, bool) { if key == nil { panic("key must be non-nil") } - index := string(key) - value, ok := r[index] - delete(r, index) + value, ok := r.entries[string(key)] + if ok { + delete(r.entries, string(key)) + r.ascKeys = nil + r.dirty = true + } return value, ok } -// Does not work with NaN. -func negCompare[T cmp.Ordered](x, y T) int { - if x < y { - return +1 +func (r *reference) Range(bounds *Bounds) iter.Seq2[[]byte, byte] { + bounds = bounds.Clone() + if bounds.IsReverse { + return r.Desc(bounds.End, bounds.Begin) } - if x > y { - return -1 + return r.Asc(bounds.Begin, bounds.End) +} + +// Future methods on Store, replacing Range. +// Desc will behave differently, with low inclusive and high exclusive. + +func (r *reference) All() iter.Seq2[[]byte, byte] { + return func(yield func([]byte, byte) bool) { + for k, v := range r.entries { + if !yield([]byte(k), v) { + return + } + } } - return 0 } -func (r reference) Range(bounds *Bounds) iter.Seq2[[]byte, byte] { - bounds = bounds.Clone() +func (r *reference) between(low, high []byte) [][]byte { + lowIndex, highIndex := 0, len(r.ascKeys) + if low != nil { + lowIndex, _ = slices.BinarySearchFunc(r.ascKeys, low, bytes.Compare) + } + if high != nil { + highIndex, _ = slices.BinarySearchFunc(r.ascKeys, high, bytes.Compare) + } + return r.ascKeys[lowIndex:highIndex] +} + +func (r *reference) Asc(low, high []byte) iter.Seq2[[]byte, byte] { + if low != nil && high != nil && bytes.Compare(low, high) >= 0 { + panic("low >= high") + } return func(yield func([]byte, byte) bool) { - var keys []string - for k := range r { - if bounds.CompareKey([]byte(k)) == 0 { - keys = append(keys, k) + r.refresh() + for _, key := range r.between(low, high) { + if !yield(key, r.entries[string(key)]) { + return } } - if bounds.IsReverse { - slices.SortFunc(keys, negCompare) - } else { - slices.Sort(keys) + } +} + +func (r *reference) Desc(low, high []byte) iter.Seq2[[]byte, byte] { + if low != nil && high != nil && bytes.Compare(low, high) >= 0 { + panic("low >= high") + } + return func(yield func([]byte, byte) bool) { + r.refresh() + a, b := low, high + if a != nil { + a = nextKey(a) + } + if b != nil { + b = nextKey(b) } - for _, k := range keys { - if !yield([]byte(k), r[k]) { + for _, key := range slices.Backward(r.between(a, b)) { + if !yield(key, r.entries[string(key)]) { return } } } } -func (r reference) String() string { +func (r *reference) String() string { var s strings.Builder - s.WriteString("{") - for _, k := range slices.Sorted(maps.Keys(r)) { - fmt.Fprintf(&s, "%s:%v, ", kv.KeyName([]byte(k)), r[k]) + s.WriteString("{\n") + fmt.Fprintf(&s, " dirty: %v\n", r.dirty) + s.WriteString(" entries: {") + for key, value := range r.entries { + fmt.Fprintf(&s, "%s:%02X, ", kv.KeyName([]byte(key)), value) + } + s.WriteString(" }\n") + s.WriteString(" ascKeys: {") + for _, key := range r.ascKeys { + fmt.Fprintf(&s, "%s, ", kv.KeyName(key)) } + s.WriteString(" }\n") s.WriteString("}") return s.String() } diff --git a/traversers_test.go b/traversers_test.go index 0aa1eee..8fc7c76 100644 --- a/traversers_test.go +++ b/traversers_test.go @@ -9,6 +9,16 @@ import ( "github.com/stretchr/testify/assert" ) +func emptySeqInt(_ func(int) bool) {} + +func emptyAdjInt(_ int) iter.Seq[int] { + return emptySeqInt +} + +func emptyPathAdjInt(_ []int) iter.Seq[int] { + return emptySeqInt +} + // adjInt returns a simple adjFunction[int] for testing traversals. // If k <= limit, children(k) == [4*k+1, 4*k+2, 4*k+3]. // If k > limit, children(k) == [].