From 32a4de9c6da1572781d6c2a5c5fad10444c09285 Mon Sep 17 00:00:00 2001 From: Sung-Kyu Yoo Date: Tue, 15 Sep 2026 07:53:01 +0900 Subject: [PATCH] perf: optimize V3 search and staging --- internal/engine/engine_map.go | 26 +++++++++--- internal/engine/engine_test.go | 8 ++++ pkg/acor/versioned_failure_test.go | 41 +++++++++++++++++- pkg/acor/versioned_search.go | 11 +++-- pkg/acor/versioned_test.go | 19 +++++++++ pkg/acor/versioned_write.go | 68 +++++++++++++++++++++--------- 6 files changed, 144 insertions(+), 29 deletions(-) diff --git a/internal/engine/engine_map.go b/internal/engine/engine_map.go index 59e8403..710731d 100644 --- a/internal/engine/engine_map.go +++ b/internal/engine/engine_map.go @@ -34,6 +34,8 @@ type memEfficientEngine struct { guard *buildGuard trie mapTrie bloom *bloomFilter + // asciiRoot skips Bloom hashing for ASCII characters that cannot start a keyword. + asciiRoot [128]bool } func newMemEfficientEngine() *memEfficientEngine { @@ -123,9 +125,13 @@ func (e *memEfficientEngine) buildFromSequence(keywords iter.Seq[string], count e.trie = trie e.bloom = newBloomFilter(len(firstRunes), 0.01) + e.asciiRoot = [128]bool{} for r := range firstRunes { e.guard.check() e.bloom.add(r) + if r >= 0 && r < utf8.RuneSelf { + e.asciiRoot[r] = true + } } } @@ -138,7 +144,7 @@ func (e *memEfficientEngine) find(text string) []string { state := 0 for _, ch := range text { - if e.bloom.skipAtRoot(state == 0, ch) { + if e.skipAtRoot(state == 0, ch) { continue } @@ -170,7 +176,7 @@ func (e *memEfficientEngine) findSet(text string) []string { state := 0 for _, ch := range text { - if e.bloom.skipAtRoot(state == 0, ch) { + if e.skipAtRoot(state == 0, ch) { continue } @@ -201,7 +207,7 @@ func (e *memEfficientEngine) findIndex(text string) map[string][]int { runeIndex := 0 for _, ch := range text { - if e.bloom.skipAtRoot(state == 0, ch) { + if e.skipAtRoot(state == 0, ch) { runeIndex++ continue } @@ -235,7 +241,7 @@ func (e *memEfficientEngine) matchString(text string, emit func(keyword string, runeIndex := 0 for _, ch := range text { - if e.bloom.skipAtRoot(state == 0, ch) { + if e.skipAtRoot(state == 0, ch) { runeIndex++ continue } @@ -258,6 +264,16 @@ func (e *memEfficientEngine) matchString(text string, emit func(keyword string, } } +func (e *memEfficientEngine) skipAtRoot(atRoot bool, ch rune) bool { + if !atRoot { + return false + } + if ch >= 0 && ch < utf8.RuneSelf { + return !e.asciiRoot[ch] + } + return e.bloom != nil && !e.bloom.mightContain(ch) +} + func (e *memEfficientEngine) matchStream(next func() (rune, bool), emit func(keyword string, start, end int) bool) { if len(e.trie.nodes) <= 1 { return @@ -271,7 +287,7 @@ func (e *memEfficientEngine) matchStream(next func() (rune, bool), emit func(key if !ok { return } - if e.bloom.skipAtRoot(state == 0, ch) { + if e.skipAtRoot(state == 0, ch) { runeIndex++ continue } diff --git a/internal/engine/engine_test.go b/internal/engine/engine_test.go index f367127..2a69fa7 100644 --- a/internal/engine/engine_test.go +++ b/internal/engine/engine_test.go @@ -210,3 +210,11 @@ func TestEngineEmptyKeywords(t *testing.T) { }) } } + +func TestMemoryEfficientBuildsASCIIRootIndex(t *testing.T) { + e := newMemEfficientEngine() + e.buildFromKeywords(keywordSet("alpha", "한국")) + if !e.asciiRoot['a'] || e.asciiRoot['z'] { + t.Fatalf("ascii root index = a:%v z:%v, want a:true z:false", e.asciiRoot['a'], e.asciiRoot['z']) + } +} diff --git a/pkg/acor/versioned_failure_test.go b/pkg/acor/versioned_failure_test.go index 0711520..e530caa 100644 --- a/pkg/acor/versioned_failure_test.go +++ b/pkg/acor/versioned_failure_test.go @@ -6,6 +6,7 @@ import ( "context" "encoding/json" "errors" + "fmt" "net" "strings" "sync/atomic" @@ -21,13 +22,29 @@ type v3FaultHook struct { failChunks atomic.Bool suppressPublish bool chunksWritten atomic.Int64 + stagePipelines atomic.Int64 } func (h *v3FaultHook) DialHook(next redis.DialHook) redis.DialHook { return func(ctx context.Context, network, addr string) (net.Conn, error) { return next(ctx, network, addr) } } func (h *v3FaultHook) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook { - return next + return func(ctx context.Context, cmds []redis.Cmder) error { + staged := false + for _, cmd := range cmds { + args := cmd.Args() + if cmd.Name() == "eval" && args[1] == v3StageScript { + staged = true + if strings.Contains(args[5].(string), ":chunk:") { + h.chunksWritten.Add(1) + } + } + } + if staged { + h.stagePipelines.Add(1) + } + return next(ctx, cmds) + } } func (h *v3FaultHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook { return func(ctx context.Context, cmd redis.Cmder) error { @@ -93,6 +110,28 @@ func TestVersionedLostCommitReceiptAndReuse(t *testing.T) { t.Fatal("single addition rewrote unaffected buckets") } } + +func TestVersionedBatchStagesChangedBucketsInOnePipeline(t *testing.T) { + ctx := context.Background() + server := miniredis.RunT(t) + v := openV3Test(t, server, "batch-stage") + hook := &v3FaultHook{} + v.client.AddHook(hook) + words := []string{"first"} + for i := 0; ; i++ { + word := fmt.Sprintf("word-%d", i) + if v3BucketNumber(words[0]) != v3BucketNumber(word) { + words = append(words, word) + break + } + } + if _, err := v.AddMany(ctx, v.Status().ServingVersion, words); err != nil { + t.Fatal(err) + } + if got := hook.stagePipelines.Load(); got != 1 { + t.Fatalf("stage pipelines = %d, want 1", got) + } +} func TestVersionedPollingAndBuildFailure(t *testing.T) { ctx := context.Background() server := miniredis.RunT(t) diff --git a/pkg/acor/versioned_search.go b/pkg/acor/versioned_search.go index 0ad2ba4..2697936 100644 --- a/pkg/acor/versioned_search.go +++ b/pkg/acor/versioned_search.go @@ -117,16 +117,19 @@ func (v *VersionedCollection) FindStream(ctx context.Context, r io.Reader, onMat // FindBatch scans every input against one serving engine, preserving input order. func (v *VersionedCollection) FindBatch(ctx context.Context, texts []string) ([][]string, error) { - ac, err := v.search(ctx) - if err != nil { + if err := v.check(ctx); err != nil { return nil, err } + e := v.current.Load() + if e == nil { + return nil, ErrVersionedClosed + } out := make([][]string, len(texts)) for i, text := range texts { - out[i], err = ac.FindContext(ctx, text) - if err != nil { + if err := ctx.Err(); err != nil { return nil, err } + out[i] = e.engine.Find(normalizeText(text, v.opts.CaseSensitive)) } return out, nil } diff --git a/pkg/acor/versioned_test.go b/pkg/acor/versioned_test.go index 65f04bf..44b9142 100644 --- a/pkg/acor/versioned_test.go +++ b/pkg/acor/versioned_test.go @@ -164,6 +164,25 @@ func TestVersionedSearchParity(t *testing.T) { t.Fatal(ap, bp) } } + +func TestVersionedFindBatchAvoidsPerTextAdapterAllocations(t *testing.T) { + ctx := context.Background() + v := openV3Test(t, miniredis.RunT(t), "batch-search") + r, err := v.Replace(ctx, v.Status().ServingVersion, []string{"needle"}) + if err != nil { + t.Fatal(err) + } + waitV3(t, v, r.Version) + allocs := testing.AllocsPerRun(100, func() { + found, err := v.FindBatch(ctx, []string{"needle"}) + if err != nil || !reflect.DeepEqual(found, [][]string{{"needle"}}) { + t.Fatal(found, err) + } + }) + if allocs > 6 { + t.Fatalf("FindBatch allocations = %.0f, want at most 6", allocs) + } +} func TestVersionedConcurrentWriters(t *testing.T) { ctx := context.Background() server := miniredis.RunT(t) diff --git a/pkg/acor/versioned_write.go b/pkg/acor/versioned_write.go index db2d7f1..16f74de 100644 --- a/pkg/acor/versioned_write.go +++ b/pkg/acor/versioned_write.go @@ -16,6 +16,11 @@ import ( const v3ChunkBytes = 1 << 20 const v3Add = "add" +type v3Stage struct { + key, registry, id string + data []byte +} + // Every preparation write is fenced, including writes by an expired process // resuming after Prune. Registration and content creation are one atomic step. const v3StageScript = v3Now + ` @@ -89,6 +94,7 @@ func (v *VersionedCollection) change(ctx context.Context, expected Version, word } next := *old result := &WriteResult{PreviousVersion: expected, Version: expected, OperationID: v3ID()} + stages := make([]v3Stage, 0) var buckets [v3BucketCount][]string for _, w := range normalized { b := v3BucketNumber(w) @@ -108,22 +114,21 @@ func (v *VersionedCollection) change(ctx context.Context, expected Version, word } result.Added += added result.Removed += removed - b, stageErr := v.stageBucket(ctx, l, after) - if stageErr != nil { - return nil, stageErr - } + b, bucketStages := v3Stages(after) + stages = append(stages, bucketStages...) next.Buckets[i] = b } - return v.prepareCommit(ctx, l, &next, result) + return v.prepareCommit(ctx, l, &next, result, stages) } -func (v *VersionedCollection) prepareCommit(ctx context.Context, l *v3Lease, next *v3Manifest, result *WriteResult) (*WriteResult, error) { +func (v *VersionedCollection) prepareCommit(ctx context.Context, l *v3Lease, next *v3Manifest, result *WriteResult, stages []v3Stage) (*WriteResult, error) { if result.Added != 0 || result.Removed != 0 { next.Version = Version(v.id + "." + v3ID()) next.Sequence++ next.Count += result.Added - result.Removed result.Version = next.Version data, _ := json.Marshal(next) - if err := v.stage(ctx, l, "gen:"+string(next.Version), "generations", string(next.Version), data); err != nil { + stages = append(stages, v3Stage{key: "gen:" + string(next.Version), registry: "generations", id: string(next.Version), data: data}) + if err := v.stageAll(ctx, l, stages); err != nil { return nil, err } } @@ -198,7 +203,38 @@ func (v *VersionedCollection) stage(ctx context.Context, l *v3Lease, key, regist } return nil } +func (v *VersionedCollection) stageAll(ctx context.Context, l *v3Lease, stages []v3Stage) error { + if len(stages) == 0 { + return nil + } + cmds := make([]*redis.Cmd, 0, len(stages)) + _, err := v.client.Pipelined(ctx, func(pipe redis.Pipeliner) error { + for _, stage := range stages { + cmds = append(cmds, pipe.Eval(ctx, v3StageScript, + []string{v.key("maintenance"), v.key("writers"), v.key(stage.key), v.key(stage.registry)}, + l.member, stage.data, stage.id)) + } + return nil + }) + if err != nil { + return err + } + for _, cmd := range cmds { + ok, err := cmd.Int() + if err != nil { + return err + } + if ok != 1 { + return ErrLeaseExpired + } + } + return nil +} func (v *VersionedCollection) stageBucket(ctx context.Context, l *v3Lease, words []string) (v3Bucket, error) { + b, stages := v3Stages(words) + return b, v.stageAll(ctx, l, stages) +} +func v3Stages(words []string) (v3Bucket, []v3Stage) { b := v3Bucket{Count: len(words)} if len(words) == 0 { return b, nil @@ -208,16 +244,14 @@ func (v *VersionedCollection) stageBucket(ctx context.Context, l *v3Lease, words // Size includes JSON quotes, escaping, commas and brackets. Oversize single // keywords form independent chunks and cannot make neighboring chunks exceed 1 MiB. start, size := 0, 2 - flush := func(end int) error { + stages := make([]v3Stage, 0, 1) + flush := func(end int) { part, _ := json.Marshal(words[start:end]) h := v3Hash(part) - if err := v.stage(ctx, l, "chunk:"+h, "chunks", h, part); err != nil { - return err - } + stages = append(stages, v3Stage{key: "chunk:" + h, registry: "chunks", id: h, data: part}) b.Chunks = append(b.Chunks, h) start = end size = 2 - return nil } for i, w := range words { encoded, _ := json.Marshal(w) @@ -226,17 +260,13 @@ func (v *VersionedCollection) stageBucket(ctx context.Context, l *v3Lease, words n++ } if size+n > v3ChunkBytes && i > start { - if err := flush(i); err != nil { - return b, err - } + flush(i) n = len(encoded) } size += n } - if err := flush(len(words)); err != nil { - return b, err - } - return b, nil + flush(len(words)) + return b, stages } // ResolveOperation retrieves a durable successful commit receipt. redis.Nil