Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 21 additions & 5 deletions internal/engine/engine_map.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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
}
}
}

Expand All @@ -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
}

Expand Down Expand Up @@ -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
}

Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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
}
Expand All @@ -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
Expand All @@ -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
}
Expand Down
8 changes: 8 additions & 0 deletions internal/engine/engine_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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'])
}
}
41 changes: 40 additions & 1 deletion pkg/acor/versioned_failure_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import (
"context"
"encoding/json"
"errors"
"fmt"
"net"
"strings"
"sync/atomic"
Expand All @@ -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 {
Expand Down Expand Up @@ -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)
Expand Down
11 changes: 7 additions & 4 deletions pkg/acor/versioned_search.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
19 changes: 19 additions & 0 deletions pkg/acor/versioned_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
68 changes: 49 additions & 19 deletions pkg/acor/versioned_write.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 + `
Expand Down Expand Up @@ -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)
Expand All @@ -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
}
}
Expand Down Expand Up @@ -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
Expand All @@ -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)
Expand All @@ -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
Expand Down