diff --git a/collection.go b/collection.go index ed0a445..2fd6ef6 100644 --- a/collection.go +++ b/collection.go @@ -103,6 +103,10 @@ type Collection struct { querySnapshotMu sync.Mutex querySnapshot atomic.Pointer[collectionQuerySnapshot] querySnapshotBuildCount atomic.Uint64 + queryLeases sync.WaitGroup + + retiredRuntimeMu sync.Mutex + retiredRuntimes map[*collectionSegmentRuntime]error } type collectionRuntimeKey struct { @@ -121,6 +125,9 @@ type collectionRuntimeIndexes struct { sparseExact map[string]*core.SparseFlatIndex fts map[string]*collectionFTSRuntime scalar dbsql.IndexSet + + closeMu sync.Mutex + closedDenseIndex map[string]struct{} } type collectionSegmentDocuments struct { @@ -133,14 +140,99 @@ type collectionSegmentRuntime struct { segmentID uint64 key collectionRuntimeKey indexes *collectionRuntimeIndexes + refs atomic.Int64 } type collectionQuerySnapshot struct { + schema CollectionSchema documents []Document segments []collectionSegmentDocuments runtimes []*collectionSegmentRuntime } +func (r *collectionSegmentRuntime) retain() { + if r == nil { + return + } + if refs := r.refs.Add(1); refs <= 1 { + panic("xvec: retain released collection segment runtime") + } +} + +func (r *collectionSegmentRuntime) release() error { + if r == nil { + return nil + } + refs := r.refs.Add(-1) + if refs < 0 { + panic("xvec: collection segment runtime reference count underflow") + } + if refs == 0 && r.indexes != nil { + if err := r.indexes.Close(); err != nil { + r.refs.Store(1) + return err + } + } + return nil +} + +func (s *collectionQuerySnapshot) retainRuntimes() { + if s == nil { + return + } + for _, runtime := range s.runtimes { + runtime.retain() + } +} + +func (c *Collection) releaseSnapshotRuntimes(s *collectionQuerySnapshot) error { + if s == nil { + return nil + } + errs := make([]error, 0, len(s.runtimes)) + for _, runtime := range s.runtimes { + errs = append(errs, c.releaseSegmentRuntime(runtime)) + } + return errors.Join(errs...) +} + +func (c *Collection) releaseSegmentRuntime(runtime *collectionSegmentRuntime) error { + err := runtime.release() + if err == nil { + return nil + } + c.retiredRuntimeMu.Lock() + if c.retiredRuntimes == nil { + c.retiredRuntimes = make(map[*collectionSegmentRuntime]error) + } + c.retiredRuntimes[runtime] = errors.Join(c.retiredRuntimes[runtime], err) + c.retiredRuntimeMu.Unlock() + return err +} + +func (c *Collection) closeRetiredSegmentRuntimes() error { + c.retiredRuntimeMu.Lock() + pending := make(map[*collectionSegmentRuntime]error, len(c.retiredRuntimes)) + for runtime, err := range c.retiredRuntimes { + pending[runtime] = err + } + c.retiredRuntimeMu.Unlock() + + errs := make([]error, 0, 2*len(pending)) + for runtime, previousErr := range pending { + err := runtime.release() + errs = append(errs, previousErr, err) + c.retiredRuntimeMu.Lock() + if err == nil { + delete(c.retiredRuntimes, runtime) + } else { + c.retiredRuntimes[runtime] = errors.Join(previousErr, err) + } + c.retiredRuntimeMu.Unlock() + } + return errors.Join(errs...) +} + func (c *Collection) querySnapshotLocked(ctx context.Context) (*collectionQuerySnapshot, error) { if snapshot := c.querySnapshot.Load(); snapshot != nil { return snapshot, nil @@ -162,14 +254,38 @@ func (c *Collection) querySnapshotLocked(ctx context.Context) (*collectionQueryS if err != nil { return nil, err } - snapshot := &collectionQuerySnapshot{documents: documents, segments: segments, runtimes: runtimes} + snapshot := &collectionQuerySnapshot{ + schema: c.schema.Clone(), documents: documents, segments: segments, runtimes: runtimes, + } + snapshot.retainRuntimes() c.querySnapshot.Store(snapshot) c.querySnapshotBuildCount.Add(1) return snapshot, nil } func (c *Collection) invalidateQuerySnapshotLocked() { - c.querySnapshot.Store(nil) + if snapshot := c.querySnapshot.Swap(nil); snapshot != nil { + _ = c.releaseSnapshotRuntimes(snapshot) + } +} + +// acquireQuerySnapshotLocked pins the immutable segment runtimes for one query. +// The caller must hold c.mu for reading so Close and invalidation cannot race +// the acquisition. The returned release function does not acquire c.mu. +func (c *Collection) acquireQuerySnapshotLocked(ctx context.Context) (*collectionQuerySnapshot, func(), error) { + snapshot, err := c.querySnapshotLocked(ctx) + if err != nil { + return nil, nil, err + } + snapshot.retainRuntimes() + c.queryLeases.Add(1) + var once sync.Once + return snapshot, func() { + once.Do(func() { + _ = c.releaseSnapshotRuntimes(snapshot) + c.queryLeases.Done() + }) + }, nil } func collectionRuntimeKeyFor(schema CollectionSchema, documents []Document) (collectionRuntimeKey, error) { @@ -215,7 +331,7 @@ func (c *Collection) segmentRuntimeIndexesLocked( // Match Alibaba zvec's query path: collect an owned list of shared segment // handles under a read lock, then search them without taking an exclusive // collection-index lock. Segment runtimes are immutable after publication; - // c.mu keeps writers and Close out for the lifetime of this query. + // query snapshots retain them until every in-flight lease is released. c.indexMu.RLock() if len(c.segmentIndexes) == len(requested) { ordered := make([]*collectionSegmentRuntime, len(requested)) @@ -245,7 +361,7 @@ func (c *Collection) segmentRuntimeIndexesLocked( created := make([]*collectionSegmentRuntime, 0) fail := func(err error) ([]*collectionSegmentRuntime, error) { for _, runtime := range created { - _ = runtime.indexes.Close() + _ = c.releaseSegmentRuntime(runtime) } return nil, err } @@ -274,6 +390,7 @@ func (c *Collection) segmentRuntimeIndexesLocked( runtime := &collectionSegmentRuntime{ segmentID: segment.metadata.ID, key: key, indexes: indexes, } + runtime.refs.Store(1) created = append(created, runtime) next[segment.metadata.ID] = runtime ordered = append(ordered, runtime) @@ -281,7 +398,7 @@ func (c *Collection) segmentRuntimeIndexesLocked( } for segmentID, runtime := range previous { if next[segmentID] != runtime { - _ = runtime.indexes.Close() + _ = c.releaseSegmentRuntime(runtime) } } c.segmentIndexes = next @@ -802,11 +919,20 @@ func (i *collectionRuntimeIndexes) Close() error { if i == nil { return nil } - seen := make(map[uintptr]struct{}) + i.closeMu.Lock() + defer i.closeMu.Unlock() + if i.closedDenseIndex == nil { + i.closedDenseIndex = make(map[string]struct{}) + } + seen := make(map[uintptr]error) var errs []error - for _, index := range i.denseNative { + for field, index := range i.denseNative { + if _, closed := i.closedDenseIndex[field]; closed { + continue + } closer, ok := index.(interface{ Close() error }) if !ok || isNilInterface(closer) { + i.closedDenseIndex[field] = struct{}{} continue } value := reflect.ValueOf(closer) @@ -815,12 +941,22 @@ func (i *collectionRuntimeIndexes) Close() error { pointer = value.Pointer() } if pointer != 0 { - if _, duplicate := seen[pointer]; duplicate { + if previousErr, duplicate := seen[pointer]; duplicate { + if previousErr == nil { + i.closedDenseIndex[field] = struct{}{} + } continue } - seen[pointer] = struct{}{} } - errs = append(errs, closer.Close()) + err := closer.Close() + if pointer != 0 { + seen[pointer] = err + } + if err != nil { + errs = append(errs, err) + continue + } + i.closedDenseIndex[field] = struct{}{} } return errors.Join(errs...) } @@ -996,17 +1132,21 @@ func (c *Collection) Close() error { c.mu.Lock() defer c.mu.Unlock() if c.closed { - return nil + return wrapCollectionError("close collection", c.path, c.closeRetiredSegmentRuntimes()) } c.closed = true + c.invalidateQuerySnapshotLocked() + c.queryLeases.Wait() c.indexMu.Lock() segmentIndexes := c.segmentIndexes c.segmentIndexes = nil c.indexMu.Unlock() - segmentErr := closeCollectionSegmentRuntimes(segmentIndexes) + _ = c.closeCollectionSegmentRuntimes(segmentIndexes) + runtimeErr := c.closeRetiredSegmentRuntimes() + storeErr := c.store.Close() return errors.Join( - wrapCollectionError("close collection", c.path, c.store.Close()), - segmentErr, + wrapCollectionError("close collection", c.path, storeErr), + wrapCollectionError("close collection", c.path, runtimeErr), ) } @@ -1034,21 +1174,24 @@ func (c *Collection) Destroy(ctx context.Context) error { return &Error{Code: ErrorCodeInvalidArgument, Op: "destroy collection", Path: c.path, Message: "refusing to remove an unsafe collection path"} } c.closed = true + c.invalidateQuerySnapshotLocked() + c.queryLeases.Wait() c.indexMu.Lock() segmentIndexes := c.segmentIndexes c.segmentIndexes = nil c.indexMu.Unlock() - indexErr := closeCollectionSegmentRuntimes(segmentIndexes) + _ = c.closeCollectionSegmentRuntimes(segmentIndexes) + indexErr := c.closeRetiredSegmentRuntimes() closeErr := c.store.Close() removeErr := os.RemoveAll(c.path) return wrapCollectionError("destroy collection", c.path, errors.Join(indexErr, closeErr, removeErr)) } -func closeCollectionSegmentRuntimes(runtimes map[uint64]*collectionSegmentRuntime) error { +func (c *Collection) closeCollectionSegmentRuntimes(runtimes map[uint64]*collectionSegmentRuntime) error { errs := make([]error, 0, len(runtimes)) for _, runtime := range runtimes { - if runtime != nil && runtime.indexes != nil { - errs = append(errs, runtime.indexes.Close()) + if runtime != nil { + errs = append(errs, c.releaseSegmentRuntime(runtime)) } } return errors.Join(errs...) @@ -3109,18 +3252,16 @@ func (c *Collection) MultiQuery(ctx context.Context, query MultiQuery) ([]Docume if err != nil { return nil, invalidArgument(op, "invalid filter: %v", err) } - documents, err := c.liveDocumentsLocked(ctx) - if err != nil { - return nil, wrapCollectionError(op, c.path, err) - } - segments, err := c.segmentDocumentsLocked(ctx) - if err != nil { - return nil, wrapCollectionError(op, c.path, err) - } - runtimes, err := c.segmentRuntimeIndexesLocked(ctx, segments) + snapshot, releaseSnapshot, err := c.acquireQuerySnapshotLocked(ctx) if err != nil { return nil, wrapCollectionError(op, c.path, err) } + schema := snapshot.schema + path := c.path + c.mu.RUnlock() + locked = false + defer releaseSnapshot() + documents, segments, runtimes := snapshot.documents, snapshot.segments, snapshot.runtimes runtimeConfig := c.runtimeConfig() filters, err := evaluateSegmentFilters(ctx, filterPlan, documents, segments, runtimes, runtimeConfig.InvertToForwardScanRatio) if err != nil { @@ -3139,7 +3280,7 @@ func (c *Collection) MultiQuery(ctx context.Context, query MultiQuery) ([]Docume if err != nil { return nil, err } - field, found := c.schema.Field(subQuery.Field) + field, found := schema.Field(subQuery.Field) if !found { return nil, invalidArgument(op, "sub-query %d field %q does not exist", index, subQuery.Field) } @@ -3173,7 +3314,7 @@ func (c *Collection) MultiQuery(ctx context.Context, query MultiQuery) ([]Docume if err != nil { return nil, wrapMultiQueryBranchError(op, c.path, index, err) } - materialized, err := c.materializeResults(documents, results, projection) + materialized, err := c.materializeResults(schema, documents, results, projection) if err != nil { return nil, err } @@ -3183,12 +3324,9 @@ func (c *Collection) MultiQuery(ctx context.Context, query MultiQuery) ([]Docume batches[index] = RerankBatch{Field: field.Clone(), Documents: materialized} } - // Release the collection lock before invoking caller code. The immutable - // snapshot, schema, and candidate batches remain owned by this call. - schema := c.schema.Clone() - path := c.path - c.mu.RUnlock() - locked = false + // Candidate generation no longer needs segment runtimes. Release the lease + // before invoking caller code so Close is not coupled to reranker latency. + releaseSnapshot() if err := ctx.Err(); err != nil { return nil, wrapCollectionError(op, path, err) } @@ -3827,7 +3965,12 @@ func (c *Collection) Query(ctx context.Context, query VectorQuery) ([]Document, } defer releaseRuntime() c.mu.RLock() - defer c.mu.RUnlock() + locked := true + defer func() { + if locked { + c.mu.RUnlock() + } + }() if err := c.requireOpenLocked(op); err != nil { return nil, err } @@ -3838,10 +3981,15 @@ func (c *Collection) Query(ctx context.Context, query VectorQuery) ([]Document, if err != nil { return nil, invalidArgument(op, "invalid filter: %v", err) } - snapshot, err := c.querySnapshotLocked(ctx) + snapshot, releaseSnapshot, err := c.acquireQuerySnapshotLocked(ctx) if err != nil { return nil, wrapCollectionError(op, c.path, err) } + schema := snapshot.schema + c.mu.RUnlock() + locked = false + defer releaseSnapshot() + documents, segments, runtimes := snapshot.documents, snapshot.segments, snapshot.runtimes runtimeConfig := c.runtimeConfig() filters, err := evaluateSegmentFilters(ctx, filterPlan, documents, segments, runtimes, runtimeConfig.InvertToForwardScanRatio) @@ -3861,13 +4009,13 @@ func (c *Collection) Query(ctx context.Context, query VectorQuery) ([]Document, } results = filterOnlyResults(documents, candidateFilter.predicate, query.TopK) case singleQueryTargetFTS: - field, found := c.schema.Field(query.Field) + field, found := schema.Field(query.Field) if !found { return nil, invalidArgument(op, "FTS field %q does not exist", query.Field) } results, err = c.searchFTSSegments(ctx, op, field, query.FTS, query.Params, query.TopK, documents, runtimes, filters) case singleQueryTargetDense, singleQueryTargetSparse, singleQueryTargetPrimaryKey: - field, found := c.schema.Field(query.Field) + field, found := schema.Field(query.Field) if !found || !field.DataType.IsVector() { return nil, invalidArgument(op, "vector field %q does not exist", query.Field) } @@ -3885,7 +4033,7 @@ func (c *Collection) Query(ctx context.Context, query VectorQuery) ([]Document, if err != nil { return nil, wrapCollectionError(op, c.path, err) } - return c.materializeResults(documents, results, query.Projection) + return c.materializeResults(schema, documents, results, query.Projection) } type singleQueryTarget uint8 @@ -4220,7 +4368,12 @@ func (c *Collection) GroupByQuery(ctx context.Context, query GroupByVectorQuery) } defer releaseRuntime() c.mu.RLock() - defer c.mu.RUnlock() + locked := true + defer func() { + if locked { + c.mu.RUnlock() + } + }() if err := c.requireOpenLocked(op); err != nil { return nil, err } @@ -4259,18 +4412,15 @@ func (c *Collection) GroupByQuery(ctx context.Context, query GroupByVectorQuery) if err != nil { return nil, invalidArgument(op, "invalid filter: %v", err) } - documents, err := c.liveDocumentsLocked(ctx) - if err != nil { - return nil, wrapCollectionError(op, c.path, err) - } - segments, err := c.segmentDocumentsLocked(ctx) - if err != nil { - return nil, wrapCollectionError(op, c.path, err) - } - runtimes, err := c.segmentRuntimeIndexesLocked(ctx, segments) + snapshot, releaseSnapshot, err := c.acquireQuerySnapshotLocked(ctx) if err != nil { return nil, wrapCollectionError(op, c.path, err) } + schema := snapshot.schema + c.mu.RUnlock() + locked = false + defer releaseSnapshot() + documents, segments, runtimes := snapshot.documents, snapshot.segments, snapshot.runtimes runtimeConfig := c.runtimeConfig() filters, err := evaluateSegmentFilters(ctx, filterPlan, documents, segments, runtimes, runtimeConfig.InvertToForwardScanRatio) if err != nil { @@ -4328,7 +4478,7 @@ func (c *Collection) GroupByQuery(ctx context.Context, query GroupByVectorQuery) metric = core.MetricIP } groups := core.MergeGroupResults(metric, query.GroupCount, query.TopKPerGroup, batches...) - return c.materializeGroups(documents, groups, query.Projection) + return c.materializeGroups(schema, documents, groups, query.Projection) } func (c *Collection) searchGroupSegment( @@ -4606,7 +4756,7 @@ func sparseValueToCore(value any) (core.SparseVector, error) { } } -func (c *Collection) materializeResults(documents []Document, results []core.Result, projection Projection) ([]Document, error) { +func (c *Collection) materializeResults(schema CollectionSchema, documents []Document, results []core.Result, projection Projection) ([]Document, error) { byID := make(map[uint64]Document, len(documents)) for _, document := range documents { byID[document.DocID] = document @@ -4618,7 +4768,7 @@ func (c *Collection) materializeResults(documents []Document, results []core.Res return nil, &Error{Code: ErrorCodeInternal, Op: "materialize query", Path: c.path, Message: fmt.Sprintf("document %d disappeared from snapshot", result.Key)} } document.Score = result.Score - projected, err := ProjectDocument(document, c.schema, projection) + projected, err := ProjectDocument(document, schema, projection) if err != nil { return nil, err } @@ -4627,10 +4777,10 @@ func (c *Collection) materializeResults(documents []Document, results []core.Res return output, nil } -func (c *Collection) materializeGroups(documents []Document, groups []core.GroupResult, projection Projection) ([]GroupResult, error) { +func (c *Collection) materializeGroups(schema CollectionSchema, documents []Document, groups []core.GroupResult, projection Projection) ([]GroupResult, error) { output := make([]GroupResult, len(groups)) for index, group := range groups { - materialized, err := c.materializeResults(documents, group.Results, projection) + materialized, err := c.materializeResults(schema, documents, group.Results, projection) if err != nil { return nil, err } diff --git a/collection_test.go b/collection_test.go index aea328a..1e59f8c 100644 --- a/collection_test.go +++ b/collection_test.go @@ -5243,6 +5243,280 @@ func TestCollectionQueryReusesSnapshotAndRebuildsAfterUpdate(t *testing.T) { require.Equal(t, uint64(5), collection.querySnapshotBuildCount.Load()) } +type failingCollectionDenseIndex struct { + collectionDenseIndex + closeErr error + closes int +} + +func (i *failingCollectionDenseIndex) Close() error { + i.closes++ + if i.closes == 1 { + return i.closeErr + } + return nil +} + +func TestCollectionRetainsFailedRuntimeCloseForRetry(t *testing.T) { + index, err := core.NewDenseFlatIndex(2, core.MetricIP) + require.NoError(t, err) + closeErr := errors.New("close runtime") + failing := &failingCollectionDenseIndex{collectionDenseIndex: index, closeErr: closeErr} + otherIndex, err := core.NewDenseFlatIndex(2, core.MetricIP) + require.NoError(t, err) + successful := &closingCollectionDenseIndex{collectionDenseIndex: otherIndex} + runtime := &collectionSegmentRuntime{indexes: &collectionRuntimeIndexes{ + denseNative: map[string]collectionDenseIndex{ + "embedding": failing, + "other": successful, + }, + }} + runtime.refs.Store(1) + collection := &Collection{} + + require.ErrorIs(t, collection.releaseSegmentRuntime(runtime), closeErr) + require.Equal(t, int64(1), runtime.refs.Load(), "failed close must preserve retry ownership") + require.ErrorIs(t, collection.closeRetiredSegmentRuntimes(), closeErr) + require.Equal(t, 2, failing.closes) + require.Equal(t, 1, successful.closes, "retry must not close an already closed index again") + require.Zero(t, runtime.refs.Load()) + require.NoError(t, collection.closeRetiredSegmentRuntimes()) +} + +func TestCollectionCloseReportsAndRetriesRetiredRuntimeError(t *testing.T) { + ctx := context.Background() + collection, err := CreateAndOpen(ctx, filepath.Join(t.TempDir(), "close-runtime-retry"), testMultiQuerySchema(), NewCollectionOptions()) + require.NoError(t, err) + _, err = collection.Insert(ctx, testMultiQueryDocuments()) + require.NoError(t, err) + _, err = collection.Query(ctx, VectorQuery{Field: "embedding", DenseVector: VectorFP32{1, 0}, TopK: 2}) + require.NoError(t, err) + + snapshot := collection.querySnapshot.Load() + require.NotNil(t, snapshot) + require.NotEmpty(t, snapshot.runtimes) + runtime := snapshot.runtimes[0] + closeErr := errors.New("close runtime") + failing := &failingCollectionDenseIndex{ + collectionDenseIndex: runtime.indexes.denseNative["embedding"], + closeErr: closeErr, + } + runtime.indexes.denseNative["embedding"] = failing + + require.ErrorIs(t, collection.Close(), closeErr) + require.Equal(t, 2, failing.closes, "Close should retry the retained runtime") + require.NoError(t, collection.Close()) +} + +type closingCollectionDenseIndex struct { + collectionDenseIndex + closes int +} + +func (i *closingCollectionDenseIndex) Close() error { + i.closes++ + return nil +} + +func TestCollectionSegmentRuntimeClosesAfterSnapshotAndQueryRelease(t *testing.T) { + index, err := core.NewDenseFlatIndex(2, core.MetricIP) + require.NoError(t, err) + closing := &closingCollectionDenseIndex{collectionDenseIndex: index} + runtime := &collectionSegmentRuntime{indexes: &collectionRuntimeIndexes{ + denseNative: map[string]collectionDenseIndex{"embedding": closing}, + }} + runtime.refs.Store(1) // segmentIndexes ownership + snapshot := &collectionQuerySnapshot{runtimes: []*collectionSegmentRuntime{runtime}} + snapshot.retainRuntimes() // published snapshot ownership + runtime.retain() // in-flight query ownership + + collection := &Collection{} + require.NoError(t, collection.releaseSnapshotRuntimes(snapshot)) + require.NoError(t, collection.releaseSegmentRuntime(runtime)) + require.Zero(t, closing.closes, "retired runtime closed while a query still held a lease") + require.NoError(t, collection.releaseSegmentRuntime(runtime)) + require.Equal(t, 1, closing.closes) +} + +type blockingCollectionDenseIndex struct { + collectionDenseIndex + started chan struct{} + release chan struct{} + closed chan struct{} + once sync.Once + closeOnce sync.Once +} + +func (i *blockingCollectionDenseIndex) Close() error { + if i.closed != nil { + i.closeOnce.Do(func() { close(i.closed) }) + } + if closer, ok := i.collectionDenseIndex.(interface{ Close() error }); ok { + return closer.Close() + } + return nil +} + +func (i *blockingCollectionDenseIndex) wait(ctx context.Context) error { + i.once.Do(func() { close(i.started) }) + select { + case <-i.release: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +func (i *blockingCollectionDenseIndex) SearchWithOptions(ctx context.Context, query []float32, options core.SearchOptions) ([]core.Result, error) { + if err := i.wait(ctx); err != nil { + return nil, err + } + return i.collectionDenseIndex.SearchWithOptions(ctx, query, options) +} + +func (i *blockingCollectionDenseIndex) SearchGroups(ctx context.Context, query []float32, options core.GroupByOptions) ([]core.GroupResult, error) { + if err := i.wait(ctx); err != nil { + return nil, err + } + return i.collectionDenseIndex.(core.DenseGroupSearcher).SearchGroups(ctx, query, options) +} + +func TestCollectionQueryDoesNotBlockConcurrentUpdate(t *testing.T) { + assertCollectionSearchDoesNotBlockConcurrentUpdate(t, func(ctx context.Context, collection *Collection) error { + _, err := collection.Query(ctx, VectorQuery{Field: "embedding", DenseVector: VectorFP32{1, 0}, TopK: 2}) + return err + }) +} + +func TestCollectionMultiQueryDoesNotBlockConcurrentUpdate(t *testing.T) { + assertCollectionSearchDoesNotBlockConcurrentUpdate(t, func(ctx context.Context, collection *Collection) error { + _, err := collection.MultiQuery(ctx, MultiQuery{Queries: []SubQuery{ + {Field: "embedding", DenseVector: VectorFP32{1, 0}, NumCandidates: 2}, + {Field: "embedding", DenseVector: VectorFP32{0, 1}, NumCandidates: 2}, + }, TopK: 2}) + return err + }) +} + +func TestCollectionGroupByQueryDoesNotBlockConcurrentUpdate(t *testing.T) { + assertCollectionSearchDoesNotBlockConcurrentUpdate(t, func(ctx context.Context, collection *Collection) error { + _, err := collection.GroupByQuery(ctx, GroupByVectorQuery{ + Field: "embedding", DenseVector: VectorFP32{1, 0}, + GroupByField: "category", GroupCount: 2, TopKPerGroup: 1, + }) + return err + }) +} + +func TestCollectionCloseWaitsForQueryLease(t *testing.T) { + ctx := context.Background() + collection, err := CreateAndOpen(ctx, filepath.Join(t.TempDir(), "close-query-lease"), testMultiQuerySchema(), NewCollectionOptions()) + require.NoError(t, err) + _, err = collection.Insert(ctx, testMultiQueryDocuments()) + require.NoError(t, err) + _, err = collection.Query(ctx, VectorQuery{Field: "embedding", DenseVector: VectorFP32{1, 0}, TopK: 2}) + require.NoError(t, err) + + snapshot := collection.querySnapshot.Load() + require.NotNil(t, snapshot) + require.NotEmpty(t, snapshot.runtimes) + runtime := snapshot.runtimes[0] + blocking := &blockingCollectionDenseIndex{ + collectionDenseIndex: runtime.indexes.denseFlat["embedding"], + started: make(chan struct{}), + release: make(chan struct{}), + closed: make(chan struct{}), + } + runtime.indexes.denseNative["embedding"] = blocking + runtime.indexes.denseFlat["embedding"] = blocking + + queryDone := make(chan error, 1) + go func() { + _, queryErr := collection.Query(ctx, VectorQuery{Field: "embedding", DenseVector: VectorFP32{1, 0}, TopK: 2}) + queryDone <- queryErr + }() + select { + case <-blocking.started: + case <-time.After(5 * time.Second): + t.Fatal("query did not reach the blocking index") + } + + closeDone := make(chan error, 1) + go func() { closeDone <- collection.Close() }() + select { + case <-blocking.closed: + t.Fatal("runtime closed before the query released its lease") + case <-time.After(100 * time.Millisecond): + } + close(blocking.release) + require.NoError(t, <-queryDone) + require.NoError(t, <-closeDone) + select { + case <-blocking.closed: + case <-time.After(5 * time.Second): + t.Fatal("runtime was not closed after the query released its lease") + } +} + +func assertCollectionSearchDoesNotBlockConcurrentUpdate(t *testing.T, search func(context.Context, *Collection) error) { + t.Helper() + ctx := context.Background() + collection, err := CreateAndOpen(ctx, filepath.Join(t.TempDir(), "query-write-concurrency"), testMultiQuerySchema(), NewCollectionOptions()) + require.NoError(t, err) + defer func() { require.NoError(t, collection.Close()) }() + _, err = collection.Insert(ctx, testMultiQueryDocuments()) + require.NoError(t, err) + _, err = collection.Query(ctx, VectorQuery{Field: "embedding", DenseVector: VectorFP32{1, 0}, TopK: 2}) + require.NoError(t, err) + + snapshot := collection.querySnapshot.Load() + require.NotNil(t, snapshot) + require.NotEmpty(t, snapshot.runtimes) + runtime := snapshot.runtimes[0] + originalNative := runtime.indexes.denseNative["embedding"] + originalFlat := runtime.indexes.denseFlat["embedding"] + blocking := &blockingCollectionDenseIndex{ + collectionDenseIndex: originalFlat, + started: make(chan struct{}), + release: make(chan struct{}), + } + runtime.indexes.denseNative["embedding"] = blocking + runtime.indexes.denseFlat["embedding"] = blocking + defer func() { + runtime.indexes.denseNative["embedding"] = originalNative + runtime.indexes.denseFlat["embedding"] = originalFlat + }() + + queryDone := make(chan error, 1) + go func() { queryDone <- search(ctx, collection) }() + select { + case <-blocking.started: + case <-time.After(5 * time.Second): + t.Fatal("query did not reach the blocking index") + } + + writeDone := make(chan error, 1) + go func() { + _, writeErr := collection.Update(ctx, []Document{{PrimaryKey: "a", Fields: map[string]any{"rating": int32(5)}}}) + writeDone <- writeErr + }() + + var writeErr error + writeCompletedBeforeQuery := false + select { + case writeErr = <-writeDone: + writeCompletedBeforeQuery = true + case <-time.After(5 * time.Second): + } + close(blocking.release) + require.NoError(t, <-queryDone) + if !writeCompletedBeforeQuery { + writeErr = <-writeDone + } + require.NoError(t, writeErr) + require.True(t, writeCompletedBeforeQuery, "concurrent update waited for the query to release the collection read lock") +} + func TestCollectionQuerySnapshotPublishesOnceUntilInvalidated(t *testing.T) { ctx := context.Background() collection, err := CreateAndOpen(ctx, filepath.Join(t.TempDir(), "query-snapshot"), testMultiQuerySchema(), NewCollectionOptions()) diff --git a/internal/ailego/reader.go b/internal/ailego/reader.go index 727b1bd..171af5d 100644 --- a/internal/ailego/reader.go +++ b/internal/ailego/reader.go @@ -87,6 +87,7 @@ type mmapReaderAt struct { data mmap.MMap size int64 closed bool + unmap func() error } func (r *mmapReaderAt) ReadAt(dst []byte, off int64) (int, error) { @@ -123,8 +124,15 @@ func (r *mmapReaderAt) Close() error { if r.closed { return nil } + unmap := r.unmap + if unmap == nil { + unmap = r.data.Unmap + } + if err := unmap(); err != nil { + return err + } r.closed = true - return r.data.Unmap() + return nil } func maxInt() int { return int(^uint(0) >> 1) } diff --git a/internal/ailego/reader_test.go b/internal/ailego/reader_test.go index 3879eb6..ccc6abd 100644 --- a/internal/ailego/reader_test.go +++ b/internal/ailego/reader_test.go @@ -16,6 +16,7 @@ package ailego import ( "bytes" + "errors" "io" "os" "path/filepath" @@ -26,6 +27,28 @@ import ( "github.com/stretchr/testify/require" ) +func TestMmapReaderCloseRetriesFailedUnmap(t *testing.T) { + closeErr := errors.New("unmap") + calls := 0 + reader := &mmapReaderAt{unmap: func() error { + calls++ + if calls == 1 { + return closeErr + } + return nil + }} + + require.ErrorIs(t, reader.Close(), closeErr) + _, err := reader.ReadAt(make([]byte, 1), 0) + require.ErrorIs(t, err, io.EOF) + require.NoError(t, reader.Close()) + require.Equal(t, 2, calls) + _, err = reader.ReadAt(make([]byte, 1), 0) + require.ErrorIs(t, err, os.ErrClosed) + require.NoError(t, reader.Close()) + require.Equal(t, 2, calls) +} + func TestOpenReaderAt(t *testing.T) { path := filepath.Join(t.TempDir(), "segment.dat") content := bytes.Repeat([]byte("zvec"), 1024) diff --git a/internal/core/diskann.go b/internal/core/diskann.go index 2b71a71..71b03c9 100644 --- a/internal/core/diskann.go +++ b/internal/core/diskann.go @@ -445,10 +445,12 @@ func (i *DiskANNIndex) Close() error { if i.closed { return nil } - i.closed = true if i.closer != nil { - return i.closer.Close() + if err := i.closer.Close(); err != nil { + return err + } } + i.closed = true return nil } diff --git a/internal/core/diskann_test.go b/internal/core/diskann_test.go index fb21cd5..0ed9030 100644 --- a/internal/core/diskann_test.go +++ b/internal/core/diskann_test.go @@ -33,6 +33,29 @@ import ( "github.com/stretchr/testify/require" ) +type failOnceDiskANNCloser struct { + err error + calls int +} + +func (c *failOnceDiskANNCloser) Close() error { + c.calls++ + if c.calls == 1 { + return c.err + } + return nil +} + +func TestDiskANNCloseRetriesFailedCloser(t *testing.T) { + closeErr := errors.New("close diskann") + closer := &failOnceDiskANNCloser{err: closeErr} + index := &DiskANNIndex{closer: closer} + + require.ErrorIs(t, index.Close(), closeErr) + require.NoError(t, index.Close()) + require.Equal(t, 2, closer.calls) +} + func TestDiskANNBuildSearchMetricsFilterRadiusAndRefiner(t *testing.T) { for _, metric := range []Metric{MetricL2, MetricIP, MetricCosine, MetricMIPSL2} { t.Run(diskANNMetricName(metric), func(t *testing.T) {