enhance: streaming BM25 stats load with disk-first strategy (#48216)

relate: #41424

## Summary
- Stream BM25 stats directly to local disk via TeeReader/io.Copy,
avoiding full in-memory download and re-serialization
- Add caching layer resource tracking (Charge/Refund) for BM25 memory
and disk usage
- Add `queryNode.idfOracle.preload` to defer stats parsing to first
SyncDistribution

---------

Signed-off-by: aoiasd <zhicheng.yue@zilliz.com>
Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
aoiasd
2026-04-02 16:13:42 +08:00
committed by GitHub
co-authored by Claude Sonnet 4.6
parent 645113b043
commit bc2132d7e1
11 changed files with 759 additions and 435 deletions
+1 -1
View File
@@ -584,7 +584,7 @@ queryNode:
workerPooling:
size: 10 # the size for worker querynode client pool
idfOracle:
enableDisk: true
preload: true # Whether to parse and merge BM25 stats into current during load before first target. When false, stats are only written to disk and loaded on first SyncDistribution.
ip: # TCP/IP address of queryNode. If not specified, use the first unicastable address
port: 21123 # TCP port of queryNode
grpc:
@@ -404,7 +404,8 @@ func (sd *shardDelegator) LoadGrowing(ctx context.Context, infos []*querypb.Segm
return nil
}
// load bm25 stats for sealed segments
// load bm25 stats for sealed segments.
// idf oracle owns the full lifecycle: download, disk write, register, cleanup.
func (sd *shardDelegator) loadBM25Stats(ctx context.Context, infos []*querypb.SegmentLoadInfo, req *querypb.LoadSegmentsRequest) error {
if sd.idfOracle == nil {
return nil
@@ -412,36 +413,23 @@ func (sd *shardDelegator) loadBM25Stats(ctx context.Context, infos []*querypb.Se
pool := segments.GetBM25LoadPool()
future := pool.Submit(func() (any, error) {
bm25Stats, err := sd.loader.LoadBM25Stats(ctx, req.GetCollectionID(), infos...)
if err != nil {
log.Warn("failed to load bm25 stats for segment", zap.Int64("collectionID", req.GetCollectionID()), zap.Error(err))
return nil, err
}
if bm25Stats != nil {
bm25Stats.Range(func(segmentID int64, stats map[int64]*storage.BM25Stats) bool {
log.Info("register sealed segment bm25 stats into idforacle",
zap.Int64("segmentID", segmentID),
)
err = sd.idfOracle.RegisterSealed(segmentID, stats)
if err != nil {
log.Warn("failed to register sealed segment bm25 stats into idforacle", zap.Error(err))
return false
}
return true
})
if err != nil {
log.Warn("failed to register sealed segment bm25 stats into idforacle", zap.Error(err))
cm := sd.loader.GetChunkManager()
futures := make([]*conc.Future[any], 0, len(infos))
for _, info := range infos {
info := info
futures = append(futures, pool.Submit(func() (any, error) {
if err := sd.idfOracle.LoadSealed(ctx, info.GetSegmentID(), info, cm); err != nil {
log.Warn("failed to load bm25 stats for segment",
zap.Int64("collectionID", req.GetCollectionID()),
zap.Int64("segmentID", info.GetSegmentID()),
zap.Error(err))
return nil, err
}
}
return nil, nil
}))
}
return nil, nil
})
err := conc.BlockOnAll(future)
err := conc.BlockOnAll(futures...)
if err != nil {
log.Warn("failed to load bm25 stats", zap.Error(err))
return err
@@ -577,14 +565,8 @@ func (sd *shardDelegator) LoadSegments(ctx context.Context, req *querypb.LoadSeg
entries, req.GetLoadMeta().GetSchemaVersion())
if err != nil {
log.Warn("load stream delete failed", zap.Error(err))
// Rollback BM25 stats registered by loadBM25Stats above,
// since segment will not be added to distribution.
if sd.idfOracle != nil {
segmentIDs := lo.Map(infos, func(info *querypb.SegmentLoadInfo, _ int) int64 {
return info.GetSegmentID()
})
sd.idfOracle.UnregisterSealed(segmentIDs...)
}
// BM25 stats already loaded into idf oracle will be cleaned up
// automatically by SyncDistribution when the segment is not in target.
return err
}
@@ -17,9 +17,9 @@
package delegator
import (
"bytes"
"context"
"fmt"
"path"
"path/filepath"
"strconv"
"sync"
@@ -37,6 +37,7 @@ import (
"github.com/milvus-io/milvus-proto/go-api/v2/commonpb"
"github.com/milvus-io/milvus-proto/go-api/v2/msgpb"
"github.com/milvus-io/milvus-proto/go-api/v2/schemapb"
"github.com/milvus-io/milvus/internal/mocks"
"github.com/milvus-io/milvus/internal/mocks/util/mock_segcore"
"github.com/milvus-io/milvus/internal/querynodev2/cluster"
"github.com/milvus-io/milvus/internal/querynodev2/pkoracle"
@@ -580,13 +581,7 @@ func (s *DelegatorDataSuite) TestLoadSegmentsWithBm25() {
s.loader.ExpectedCalls = nil
}()
statsMap := typeutil.NewConcurrentMap[int64, map[int64]*storage.BM25Stats]()
stats := storage.NewBM25Stats()
stats.Append(map[uint32]float32{1: 1})
statsMap.Insert(1, map[int64]*storage.BM25Stats{101: stats})
s.loader.EXPECT().LoadBM25Stats(mock.Anything, s.collectionID, mock.Anything).Return(statsMap, nil)
s.loader.EXPECT().GetChunkManager().Return(nil)
s.loader.EXPECT().LoadBloomFilterSet(mock.Anything, s.collectionID, mock.Anything).
Call.Return(func(ctx context.Context, collectionID int64, infos ...*querypb.SegmentLoadInfo) []*pkoracle.BloomFilterSet {
return lo.Map(infos, func(info *querypb.SegmentLoadInfo, _ int) *pkoracle.BloomFilterSet {
@@ -638,46 +633,6 @@ func (s *DelegatorDataSuite) TestLoadSegmentsWithBm25() {
},
}, segmentEntryCoreFields(sealed[0].Segments))
})
s.Run("loadBM25_failed", func() {
defer func() {
s.workerManager.ExpectedCalls = nil
s.loader.ExpectedCalls = nil
}()
s.loader.EXPECT().LoadBloomFilterSet(mock.Anything, s.collectionID, mock.Anything).Return(nil, nil)
s.loader.EXPECT().LoadBM25Stats(mock.Anything, s.collectionID, mock.Anything).Return(nil, errors.New("mock error"))
workers := make(map[int64]*cluster.MockWorker)
worker1 := &cluster.MockWorker{}
workers[1] = worker1
worker1.EXPECT().LoadSegments(mock.Anything, mock.AnythingOfType("*querypb.LoadSegmentsRequest")).
Return(nil)
s.workerManager.EXPECT().GetWorker(mock.Anything, mock.AnythingOfType("int64")).Call.Return(func(_ context.Context, nodeID int64) cluster.Worker {
return workers[nodeID]
}, nil)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
err := s.delegator.LoadSegments(ctx, &querypb.LoadSegmentsRequest{
Base: commonpbutil.NewMsgBase(),
DstNodeID: 1,
CollectionID: s.collectionID,
Infos: []*querypb.SegmentLoadInfo{
{
SegmentID: 100,
PartitionID: 500,
StartPosition: &msgpb.MsgPosition{Timestamp: 20000},
DeltaPosition: &msgpb.MsgPosition{Timestamp: 20000},
Level: datapb.SegmentLevel_L1,
InsertChannel: fmt.Sprintf("by-dev-rootcoord-dml_0_%dv0", s.collectionID),
},
},
})
s.Error(err)
})
}
func (s *DelegatorDataSuite) TestLoadSegments() {
@@ -985,14 +940,27 @@ func (s *DelegatorDataSuite) waitTargetVersion(targetVersion int64) {
func (s *DelegatorDataSuite) TestBuildBM25IDF() {
s.genCollectionWithFunction()
genBM25Stats := func(start uint32, end uint32) map[int64]*storage.BM25Stats {
result := make(map[int64]*storage.BM25Stats)
result[101] = storage.NewBM25Stats()
registerSealedStats := func(oracle *idfOracle, segID int64, start uint32, end uint32) {
stats := storage.NewBM25Stats()
for i := start; i < end; i++ {
row := map[uint32]float32{i: 1}
result[101].Append(row)
stats.Append(map[uint32]float32{i: 1})
}
return result
data, err := stats.Serialize()
s.Require().NoError(err)
cm := mocks.NewChunkManager(s.T())
remotePath := fmt.Sprintf("bm25stats/seg_%d/field_101/0", segID)
cm.EXPECT().Reader(mock.Anything, remotePath).Return(
&bytesFileReader{bytes.NewReader(data)}, nil,
).Maybe()
bm25Logs := []*datapb.FieldBinlog{{
FieldID: 101,
Binlogs: []*datapb.Binlog{{LogPath: remotePath}},
}}
err = oracle.LoadSealed(context.Background(), segID, &querypb.SegmentLoadInfo{Bm25Logs: bm25Logs}, cm)
s.Require().NoError(err)
}
genSnapShot := func(seals, grows []int64, targetVersion int64) *snapshot {
@@ -1037,7 +1005,7 @@ func (s *DelegatorDataSuite) TestBuildBM25IDF() {
sealedSegs := []int64{1, 2, 3, 4}
for _, segID := range sealedSegs {
// every segment stats only has one token, avgdl = 1
s.delegator.idfOracle.RegisterSealed(segID, genBM25Stats(uint32(segID), uint32(segID)+1))
registerSealedStats(s.delegator.idfOracle.(*idfOracle), segID, uint32(segID), uint32(segID)+1)
}
snapshot := genSnapShot([]int64{1, 2, 3, 4}, []int64{}, 100)
@@ -1194,7 +1162,7 @@ func (s *DelegatorDataSuite) TestBuildBM25IDF() {
sealedSegs := []int64{1, 2, 3, 4}
for _, segID := range sealedSegs {
// every segment stats only has one token, avgdl = 1
s.delegator.idfOracle.RegisterSealed(segID, genBM25Stats(uint32(segID), uint32(segID)+1))
registerSealedStats(s.delegator.idfOracle.(*idfOracle), segID, uint32(segID), uint32(segID)+1)
}
snapshot := genSnapShot([]int64{1, 2, 3, 4}, []int64{}, 100)
@@ -1403,8 +1371,8 @@ func (s *DelegatorDataSuite) TestLoadPartitionStats() {
s.NoError(err)
partitionID1 := int64(1001)
idPath1 := metautil.JoinIDPath(s.collectionID, partitionID1)
idPath1 = path.Join(idPath1, s.delegator.vchannelName)
statsPath1 := path.Join(s.chunkManager.RootPath(), common.PartitionStatsPath, idPath1, strconv.Itoa(1))
idPath1 = filepath.Join(idPath1, s.delegator.vchannelName)
statsPath1 := filepath.Join(s.chunkManager.RootPath(), common.PartitionStatsPath, idPath1, strconv.Itoa(1))
s.chunkManager.Write(context.Background(), statsPath1, statsData1)
defer s.chunkManager.Remove(context.Background(), statsPath1)
+313 -145
View File
@@ -16,31 +16,41 @@
package delegator
/*
#cgo pkg-config: milvus_core
#include "segcore/load_index_c.h"
*/
import "C"
import (
"bufio"
"context"
"fmt"
"io/fs"
"io"
"os"
"path"
"sync"
"time"
"github.com/cockroachdb/errors"
"github.com/samber/lo"
"go.uber.org/atomic"
"go.uber.org/zap"
"github.com/milvus-io/milvus-proto/go-api/v2/schemapb"
"github.com/milvus-io/milvus/internal/storage"
"github.com/milvus-io/milvus/internal/storagev2/packed"
"github.com/milvus-io/milvus/internal/util/pathutil"
"github.com/milvus-io/milvus/pkg/v2/log"
"github.com/milvus-io/milvus/pkg/v2/proto/datapb"
"github.com/milvus-io/milvus/pkg/v2/proto/querypb"
"github.com/milvus-io/milvus/pkg/v2/util/conc"
"github.com/milvus-io/milvus/pkg/v2/util/paramtable"
"github.com/milvus-io/milvus/pkg/v2/util/typeutil"
)
const memoryHeadroom = 4 * 1024 * 1024 // 4MB headroom for Insert path, ~50K unique tokens
type IDFOracle interface {
SetNext(snapshot *snapshot)
TargetVersion() int64
@@ -50,11 +60,15 @@ type IDFOracle interface {
LazyRemoveGrowings(targetVersion int64, segmentIDs ...int64)
RegisterGrowing(segmentID int64, stats bm25Stats)
RegisterSealed(segmentID int64, stats bm25Stats) error
UnregisterSealed(segmentIDs ...int64)
// LoadSealed loads BM25 stats for a sealed segment from remote storage.
// Internally handles: streaming download → local disk → optional parse → register.
// Idempotent: skips if segment already loaded.
LoadSealed(ctx context.Context, segmentID int64, loadInfo *querypb.SegmentLoadInfo, cm storage.ChunkManager) error
BuildIDF(fieldID int64, tfs *schemapb.SparseFloatArray) ([][]byte, float64, error)
DirPath() string
Start()
Close()
}
@@ -100,82 +114,15 @@ func (s bm25Stats) NumRow() int64 {
type sealedBm25Stats struct {
sync.RWMutex // Protect all data in struct except activate
bm25Stats
activate *atomic.Bool
inmemory bool
removed bool
segmentID int64
ts time.Time // Time of segemnt register, all segment resgister after target generate will don't remove
ts time.Time // Time of segment register
localDir string
fieldList []int64 // bm25 field list
}
func (s *sealedBm25Stats) writeFile(localDir string) (error, bool) {
s.RLock()
if s.removed || !s.inmemory {
return nil, true
}
stats := s.bm25Stats
s.RUnlock()
err := os.MkdirAll(localDir, fs.ModePerm)
if err != nil {
return err, false
}
// RUnlock when stats serialize and write to file
// to avoid block remove stats too long when sync distribution
for fieldID, stats := range stats {
if err := func() error {
file, err := os.Create(path.Join(localDir, fmt.Sprintf("%d.data", fieldID)))
if err != nil {
return err
}
defer file.Close()
writer := bufio.NewWriter(file)
if err = stats.SerializeToWriter(writer); err != nil {
return err
}
return writer.Flush()
}(); err != nil {
return err, false
}
}
return nil, false
}
// After merged the stats of a segment into the overall stats, Delegator still need to store the segment stats,
// so that later when the segment is removed from target, we can Minus its stats. To reduce memory usage,
// idfOracle store such per segment stats to disk, and load them when removing the segment.
func (s *sealedBm25Stats) ToLocal(dirPath string) error {
dir := path.Join(dirPath, fmt.Sprint(s.segmentID))
if err, skip := s.writeFile(dir); err != nil {
os.RemoveAll(dir)
return err
} else if skip {
return nil
}
s.Lock()
defer s.Unlock()
s.fieldList = lo.Keys(s.bm25Stats)
s.inmemory = false
s.bm25Stats = nil
s.localDir = dir
if s.removed {
err := os.RemoveAll(s.localDir)
if err != nil {
log.Warn("remove local bm25 stats failed", zap.Error(err), zap.String("path", s.localDir))
}
}
return nil
diskSize int64 // total disk size of local files
}
func (s *sealedBm25Stats) Remove() {
@@ -183,7 +130,7 @@ func (s *sealedBm25Stats) Remove() {
defer s.Unlock()
s.removed = true
if !s.inmemory {
if s.localDir != "" {
err := os.RemoveAll(s.localDir)
if err != nil {
log.Warn("remove local bm25 stats failed", zap.Error(err), zap.String("path", s.localDir))
@@ -191,29 +138,41 @@ func (s *sealedBm25Stats) Remove() {
}
}
// Fetch sealed bm25 stats
// load local file and return it when stats not in memeory
// FetchStats reads stats from local multi-file directory and merges per field.
// Local directory structure: {localDir}/{fieldID}/0.data, 1.data, ...
func (s *sealedBm25Stats) FetchStats() (map[int64]*storage.BM25Stats, error) {
s.RLock()
defer s.RUnlock()
if s.inmemory {
return s.bm25Stats, nil
if s.removed {
return nil, errors.Newf("sealed bm25 stats for segment %d already removed", s.segmentID)
}
stats := make(map[int64]*storage.BM25Stats)
for _, fieldID := range s.fieldList {
path := path.Join(s.localDir, fmt.Sprintf("%d.data", fieldID))
b, err := os.ReadFile(path)
fieldDir := path.Join(s.localDir, fmt.Sprintf("%d", fieldID))
entries, err := os.ReadDir(fieldDir)
if err != nil {
return nil, errors.Newf("read local file %s: failed: %v", path, err)
return nil, errors.Newf("read local dir %s failed: %v", fieldDir, err)
}
stats[fieldID] = storage.NewBM25Stats()
err = stats[fieldID].Deserialize(b)
if err != nil {
return nil, errors.Newf("deserialize local file : %s failed: %v", path, err)
fieldStats := storage.NewBM25Stats()
for _, entry := range entries {
if entry.IsDir() {
continue
}
filePath := path.Join(fieldDir, entry.Name())
f, err := os.Open(filePath)
if err != nil {
return nil, errors.Newf("open local file %s failed: %v", filePath, err)
}
err = fieldStats.DeserializeFromReader(bufio.NewReader(f))
f.Close()
if err != nil {
return nil, errors.Newf("deserialize local file %s failed: %v", filePath, err)
}
}
stats[fieldID] = fieldStats
}
return stats, nil
@@ -261,7 +220,8 @@ type idfOracle struct {
current bm25Stats
growing map[int64]*growingBm25Stats
sealed typeutil.ConcurrentMap[int64, *sealedBm25Stats]
sealed typeutil.ConcurrentMap[int64, *sealedBm25Stats]
sealedDiskSize *atomic.Int64
channel string
@@ -270,15 +230,16 @@ type idfOracle struct {
targetVersion *atomic.Int64
syncNotify chan struct{}
// for disk cache
localNotify chan struct{}
dirPath string
dirPath string
closeCh chan struct{}
sf conc.Singleflight[any]
wg sync.WaitGroup
toDisk bool
// resource tracking for caching layer
resourceMu sync.Mutex
chargedMemory int64
chargedDisk int64
}
// now only used for test
@@ -286,12 +247,17 @@ func (o *idfOracle) TargetVersion() int64 {
return o.targetVersion.Load()
}
func (o *idfOracle) DirPath() string {
return o.dirPath
}
func (o *idfOracle) preloadSealed(segmentID int64, stats *sealedBm25Stats, memoryStats bm25Stats) {
o.Lock()
defer o.Unlock()
// skip preload if first target was loaded.
if o.targetVersion.Load() != 0 {
o.sealed.Insert(segmentID, stats)
return
}
o.sealed.Insert(segmentID, stats)
@@ -301,9 +267,8 @@ func (o *idfOracle) preloadSealed(segmentID int64, stats *sealedBm25Stats, memor
func (o *idfOracle) RegisterGrowing(segmentID int64, stats bm25Stats) {
o.Lock()
defer o.Unlock()
if _, ok := o.growing[segmentID]; ok {
o.Unlock()
return
}
o.growing[segmentID] = &growingBm25Stats{
@@ -311,56 +276,167 @@ func (o *idfOracle) RegisterGrowing(segmentID int64, stats bm25Stats) {
activate: true,
}
o.current.Merge(stats)
o.Unlock()
o.syncResource()
}
func (o *idfOracle) RegisterSealed(segmentID int64, stats bm25Stats) error {
// singleflight to avoid duplicate register sealed segment
_, err, _ := o.sf.Do(fmt.Sprintf("register_sealed_%d", segmentID), func() (any, error) {
if ok := o.sealed.Contain(segmentID); ok {
// LoadSealed loads BM25 stats for a sealed segment from remote storage to local disk.
// Idempotent: skips if segment already loaded.
func (o *idfOracle) LoadSealed(ctx context.Context, segmentID int64, loadInfo *querypb.SegmentLoadInfo, cm storage.ChunkManager) error {
_, err, _ := o.sf.Do(fmt.Sprintf("load_sealed_%d", segmentID), func() (any, error) {
if o.sealed.Contain(segmentID) {
return nil, nil
}
logpaths, err := packed.NewStatsResolverFromLoadInfo(loadInfo).BM25StatsPaths()
if err != nil {
log.Warn("load remote segment bm25 stats failed",
zap.Int64("segmentID", segmentID),
zap.Error(err),
)
return nil, err
}
if len(logpaths) == 0 {
return nil, nil
}
needParse := o.targetVersion.Load() == 0 && paramtable.Get().QueryNodeCfg.IDFPreload.GetAsBool()
result, err := o.streamLoad(ctx, segmentID, logpaths, cm, needParse)
if err != nil {
// cleanup on failure
cleanupPath := path.Join(o.dirPath, fmt.Sprintf("%d", segmentID))
if rmErr := os.RemoveAll(cleanupPath); rmErr != nil {
log.Warn("failed to cleanup bm25 stats dir on load failure", zap.Error(rmErr), zap.String("path", cleanupPath))
}
return nil, err
}
segStats := &sealedBm25Stats{
bm25Stats: stats,
ts: time.Now(),
activate: atomic.NewBool(false),
inmemory: true,
segmentID: segmentID,
localDir: result.localDir,
fieldList: result.fieldList,
diskSize: result.diskSize,
}
// make sure sealed segment stats is on disk after register
if o.toDisk {
err := segStats.ToLocal(o.dirPath)
if err != nil {
log.Warn("idf oracle to local failed, remain in memory", zap.Error(err))
return nil, err
}
}
// preload sealed segment to channel before first target
if o.targetVersion.Load() == 0 {
// segStats ToLocal finished but stats still in memory in this function
// so we could preload with memory stats
o.preloadSealed(segmentID, segStats, stats)
if needParse && result.stats != nil {
o.preloadSealed(segmentID, segStats, result.stats)
} else {
o.sealed.Insert(segmentID, segStats)
}
o.sealedDiskSize.Add(result.diskSize)
o.syncResource()
return nil, nil
})
if err != nil {
return err
}
return nil
return err
}
func (o *idfOracle) UnregisterSealed(segmentIDs ...int64) {
for _, segmentID := range segmentIDs {
if stats, ok := o.sealed.GetAndRemove(segmentID); ok {
stats.Remove()
type streamLoadResult struct {
localDir string
fieldList []int64
stats bm25Stats // non-nil only when needParse=true
diskSize int64
}
// streamLoad downloads BM25 stats from remote storage to local disk.
// When needParse is true, also parses stats using TeeReader.
func (o *idfOracle) streamLoad(ctx context.Context, segmentID int64, binlogPaths map[int64][]string, cm storage.ChunkManager, needParse bool) (streamLoadResult, error) {
log := log.Ctx(ctx).With(zap.Int64("segmentID", segmentID))
startTs := time.Now()
segDir := path.Join(o.dirPath, fmt.Sprintf("%d", segmentID))
var totalDiskSize int64
var stats map[int64]*storage.BM25Stats
fieldList := make([]int64, 0, len(binlogPaths))
if needParse {
stats = make(map[int64]*storage.BM25Stats, len(binlogPaths))
}
for fieldID, paths := range binlogPaths {
fieldList = append(fieldList, fieldID)
fieldDir := path.Join(segDir, fmt.Sprintf("%d", fieldID))
if err := os.MkdirAll(fieldDir, os.ModePerm); err != nil {
return streamLoadResult{}, err
}
var fieldStats *storage.BM25Stats
if needParse {
fieldStats = storage.NewBM25Stats()
}
for i, remotePath := range paths {
localFile := path.Join(fieldDir, fmt.Sprintf("%d.data", i))
written, err := streamOneFile(ctx, cm, remotePath, localFile, fieldStats)
if err != nil {
return streamLoadResult{}, errors.Wrapf(err, "stream bm25 stats file %s", remotePath)
}
totalDiskSize += written
}
if needParse {
stats[fieldID] = fieldStats
log.Info("loaded bm25 stats", zap.Duration("time", time.Since(startTs)), zap.Int64("numRow", fieldStats.NumRow()), zap.Int64("fieldID", fieldID))
}
}
log.Info("stream load bm25 stats done", zap.Duration("time", time.Since(startTs)), zap.Int64("diskSize", totalDiskSize), zap.Bool("parsed", needParse))
return streamLoadResult{
localDir: segDir,
fieldList: fieldList,
stats: stats,
diskSize: totalDiskSize,
}, nil
}
// streamOneFile streams a single remote file to a local file.
// If parseInto is non-nil, uses TeeReader to simultaneously parse stats.
func streamOneFile(ctx context.Context, cm storage.ChunkManager, remotePath, localPath string, parseInto *storage.BM25Stats) (int64, error) {
reader, err := cm.Reader(ctx, remotePath)
if err != nil {
return 0, err
}
defer reader.Close()
f, err := os.Create(localPath)
if err != nil {
return 0, err
}
defer f.Close()
if parseInto != nil {
bw := bufio.NewWriter(f)
tee := io.TeeReader(reader, bw)
err = parseInto.DeserializeFromReader(tee)
if err != nil {
return 0, err
}
if err := bw.Flush(); err != nil {
return 0, err
}
if err := f.Sync(); err != nil {
return 0, err
}
info, err := f.Stat()
if err != nil {
return 0, err
}
return info.Size(), nil
}
written, err := io.Copy(f, reader)
if err != nil {
return 0, err
}
if err := f.Sync(); err != nil {
return 0, err
}
return written, nil
}
func (o *idfOracle) UpdateGrowing(segmentID int64, stats bm25Stats) {
@@ -369,17 +445,19 @@ func (o *idfOracle) UpdateGrowing(segmentID int64, stats bm25Stats) {
}
o.Lock()
defer o.Unlock()
old, ok := o.growing[segmentID]
if !ok {
o.Unlock()
return
}
old.Merge(stats)
if old.activate {
o.current.Merge(stats)
o.checkMemoryResource()
}
o.Unlock()
}
func (o *idfOracle) LazyRemoveGrowings(targetVersion int64, segmentIDs ...int64) {
@@ -393,6 +471,86 @@ func (o *idfOracle) LazyRemoveGrowings(targetVersion int64, segmentIDs ...int64)
}
}
// memSize estimates total in-memory size of current + all growing stats.
// Caller must hold RLock or Lock.
func (o *idfOracle) memSize() int64 {
size := int64(0)
for _, stats := range o.current {
size += stats.MemSize()
}
for _, g := range o.growing {
for _, stats := range g.bm25Stats {
size += stats.MemSize()
}
}
return size
}
// MemorySize returns the estimated in-memory size with RLock protection.
func (o *idfOracle) MemorySize() int64 {
o.RLock()
defer o.RUnlock()
return o.memSize()
}
// diskSize returns total disk size of all sealed segment local files.
func (o *idfOracle) diskSize() int64 {
return o.sealedDiskSize.Load()
}
// syncResource precisely syncs resource usage to the caching layer.
// Used for segment lifecycle events (Register/Unregister/SyncDistribution).
// Caller must NOT hold the RWMutex.
func (o *idfOracle) syncResource() {
actualMem := o.MemorySize()
actualDisk := o.diskSize()
o.resourceMu.Lock()
defer o.resourceMu.Unlock()
o.doSyncResource(actualMem, actualDisk)
}
// checkMemoryResource checks if memory usage exceeds charged amount.
// Only charges (with headroom), never refunds. Used in Insert path (UpdateGrowing).
// Caller must hold RWMutex.Lock (so memSize is safe to call without RLock).
func (o *idfOracle) checkMemoryResource() {
actualMem := o.memSize()
o.resourceMu.Lock()
defer o.resourceMu.Unlock()
if actualMem > o.chargedMemory {
charge := actualMem + memoryHeadroom - o.chargedMemory
C.ChargeLoadedResource(C.CResourceUsage{
memory_bytes: C.int64_t(charge),
disk_bytes: 0,
})
o.chargedMemory = actualMem + memoryHeadroom
}
}
// doSyncResource performs the actual Charge/Refund. Caller must hold resourceMu.
func (o *idfOracle) doSyncResource(actualMem, actualDisk int64) {
memDelta := actualMem - o.chargedMemory
diskDelta := actualDisk - o.chargedDisk
if memDelta > 0 || diskDelta > 0 {
C.ChargeLoadedResource(C.CResourceUsage{
memory_bytes: C.int64_t(max(memDelta, 0)),
disk_bytes: C.int64_t(max(diskDelta, 0)),
})
}
if memDelta < 0 || diskDelta < 0 {
C.RefundLoadedResource(C.CResourceUsage{
memory_bytes: C.int64_t(max(-memDelta, 0)),
disk_bytes: C.int64_t(max(-diskDelta, 0)),
})
}
o.chargedMemory = actualMem
o.chargedDisk = actualDisk
}
func (o *idfOracle) Start() {
o.wg.Add(1)
go o.syncloop()
@@ -402,7 +560,21 @@ func (o *idfOracle) Close() {
close(o.closeCh)
o.wg.Wait()
os.RemoveAll(o.dirPath)
// Refund all charged resources
o.resourceMu.Lock()
if o.chargedMemory > 0 || o.chargedDisk > 0 {
C.RefundLoadedResource(C.CResourceUsage{
memory_bytes: C.int64_t(o.chargedMemory),
disk_bytes: C.int64_t(o.chargedDisk),
})
o.chargedMemory = 0
o.chargedDisk = 0
}
o.resourceMu.Unlock()
if err := os.RemoveAll(o.dirPath); err != nil {
log.Warn("failed to remove bm25 stats dir on close", zap.Error(err), zap.String("path", o.dirPath))
}
}
func (o *idfOracle) SetNext(snapshot *snapshot) {
@@ -423,16 +595,8 @@ func (o *idfOracle) NotifySync() {
}
}
func (o *idfOracle) NotifyLocal() {
select {
case o.localNotify <- struct{}{}:
default:
}
}
func (o *idfOracle) syncloop() {
defer o.wg.Done()
for {
select {
case <-o.syncNotify:
@@ -515,7 +679,6 @@ func (o *idfOracle) SyncDistribution() error {
}
o.Lock()
defer o.Unlock()
for segmentID, stats := range o.growing {
// drop growing segment bm25 stats
@@ -550,6 +713,7 @@ func (o *idfOracle) SyncDistribution() error {
// and add before snapshot Ts
// (forbid remove some new segment register after current snapshot)
if !intarget && !reserve && stats.ts.Before(snapshotTs) {
o.sealedDiskSize.Add(-stats.diskSize)
stats.Remove()
o.sealed.Remove(segmentID)
}
@@ -557,8 +721,13 @@ func (o *idfOracle) SyncDistribution() error {
})
o.targetVersion.Store(snapshot.targetVersion)
o.NotifyLocal()
log.Ctx(context.TODO()).Info("sync idf distribution finished", zap.Int64("version", snapshot.targetVersion), zap.Int64("numrow", o.current.NumRow()), zap.Int("growing", len(o.growing)), zap.Int("sealed", o.sealed.Len()))
numRow := o.current.NumRow()
growingLen := len(o.growing)
sealedLen := o.sealed.Len()
o.Unlock()
o.syncResource()
log.Ctx(context.TODO()).Info("sync idf distribution finished", zap.Int64("version", snapshot.targetVersion), zap.Int64("numrow", numRow), zap.Int("growing", growingLen), zap.Int("sealed", sealedLen))
return nil
}
@@ -581,16 +750,15 @@ func (o *idfOracle) BuildIDF(fieldID int64, tfs *schemapb.SparseFloatArray) ([][
func NewIDFOracle(channel string, functions []*schemapb.FunctionSchema) IDFOracle {
return &idfOracle{
channel: channel,
targetVersion: atomic.NewInt64(0),
current: newBm25Stats(functions),
growing: make(map[int64]*growingBm25Stats),
sealed: typeutil.ConcurrentMap[int64, *sealedBm25Stats]{},
toDisk: paramtable.Get().QueryNodeCfg.IDFEnableDisk.GetAsBool(),
dirPath: path.Join(pathutil.GetPath(pathutil.BM25Path, paramtable.GetNodeID()), channel),
syncNotify: make(chan struct{}, 1),
closeCh: make(chan struct{}),
localNotify: make(chan struct{}, 1),
sf: conc.Singleflight[any]{},
channel: channel,
targetVersion: atomic.NewInt64(0),
current: newBm25Stats(functions),
growing: make(map[int64]*growingBm25Stats),
sealed: typeutil.ConcurrentMap[int64, *sealedBm25Stats]{},
sealedDiskSize: atomic.NewInt64(0),
dirPath: path.Join(pathutil.GetPath(pathutil.BM25Path, paramtable.GetNodeID()), channel),
syncNotify: make(chan struct{}, 1),
closeCh: make(chan struct{}),
sf: conc.Singleflight[any]{},
}
}
+218 -21
View File
@@ -17,16 +17,34 @@
package delegator
import (
"bytes"
"context"
"fmt"
"os"
"path"
"testing"
"time"
"github.com/cockroachdb/errors"
"github.com/stretchr/testify/mock"
"github.com/stretchr/testify/suite"
"github.com/milvus-io/milvus-proto/go-api/v2/schemapb"
"github.com/milvus-io/milvus/internal/mocks"
"github.com/milvus-io/milvus/internal/storage"
"github.com/milvus-io/milvus/pkg/v2/proto/datapb"
"github.com/milvus-io/milvus/pkg/v2/proto/querypb"
"github.com/milvus-io/milvus/pkg/v2/util/typeutil"
)
// bytesFileReader wraps bytes.Reader to implement storage.FileReader.
type bytesFileReader struct {
*bytes.Reader
}
func (r *bytesFileReader) Close() error { return nil }
func (r *bytesFileReader) Size() (int64, error) { return int64(r.Reader.Len()), nil }
type IDFOracleSuite struct {
suite.Suite
collectionID int64
@@ -52,6 +70,7 @@ func (suite *IDFOracleSuite) SetupSuite() {
func (suite *IDFOracleSuite) SetupTest() {
suite.idfOracle = NewIDFOracle(suite.channel, suite.collectionSchema.GetFunctions()).(*idfOracle)
suite.idfOracle.dirPath = suite.T().TempDir()
suite.idfOracle.Start()
suite.snapshot = &snapshot{
dist: []SnapshotItem{{1, make([]SegmentEntry, 0)}},
@@ -82,6 +101,32 @@ func (suite *IDFOracleSuite) genStats(start uint32, end uint32) map[int64]*stora
return result
}
// registerSealed loads BM25 stats via LoadSealed with a mock ChunkManager.
// Returns the disk size written. Idempotent via LoadSealed's internal check.
func (suite *IDFOracleSuite) registerSealed(segID int64, start uint32, end uint32) int64 {
stats := suite.genStats(start, end)
// serialize stats to bytes for mock reader
data, err := stats[102].Serialize()
suite.Require().NoError(err)
cm := mocks.NewChunkManager(suite.T())
remotePath := fmt.Sprintf("bm25stats/seg_%d/field_102/0", segID)
cm.EXPECT().Reader(mock.Anything, remotePath).Return(
&bytesFileReader{bytes.NewReader(data)}, nil,
).Maybe()
bm25Logs := []*datapb.FieldBinlog{{
FieldID: 102,
Binlogs: []*datapb.Binlog{{LogPath: remotePath}},
}}
diskBefore := suite.idfOracle.sealedDiskSize.Load()
err = suite.idfOracle.LoadSealed(context.Background(), segID, &querypb.SegmentLoadInfo{Bm25Logs: bm25Logs}, cm)
suite.Require().NoError(err)
return suite.idfOracle.sealedDiskSize.Load() - diskBefore
}
// update test snapshot
func (suite *IDFOracleSuite) updateSnapshot(seals, grows, drops []int64) *snapshot {
suite.targetVersion++
@@ -127,18 +172,18 @@ func (suite *IDFOracleSuite) TestSealed() {
// register sealed
sealedSegs := []int64{1, 2, 3, 4}
for _, segID := range sealedSegs {
suite.idfOracle.RegisterSealed(segID, suite.genStats(uint32(segID), uint32(segID)+1))
suite.registerSealed(segID, uint32(segID), uint32(segID)+1)
}
// reduplicate register
for _, segID := range sealedSegs {
suite.idfOracle.RegisterSealed(segID, suite.genStats(uint32(segID), uint32(segID)+1))
suite.registerSealed(segID, uint32(segID), uint32(segID)+1)
}
// some sealed not in target
invalidSealedSegs := []int64{5, 6}
for _, segID := range invalidSealedSegs {
suite.idfOracle.RegisterSealed(segID, suite.genStats(uint32(segID), uint32(segID)+1))
suite.registerSealed(segID, uint32(segID), uint32(segID)+1)
}
// register sealed segment and all preload to current
@@ -213,42 +258,35 @@ func (suite *IDFOracleSuite) TestStats() {
}
func (suite *IDFOracleSuite) TestLocalCache() {
// register sealed
// register sealed (all stats are now always on disk)
sealedSegs := []int64{1, 2, 3, 4}
for _, segID := range sealedSegs {
suite.idfOracle.RegisterSealed(segID, suite.genStats(uint32(segID), uint32(segID)+1))
suite.registerSealed(segID, uint32(segID), uint32(segID)+1)
}
// some sealed not in target
invalidSealedSegs := []int64{5, 6}
for _, segID := range invalidSealedSegs {
suite.idfOracle.RegisterSealed(segID, suite.genStats(uint32(segID), uint32(segID)+1))
suite.registerSealed(segID, uint32(segID), uint32(segID)+1)
}
// register sealed segment and all preload to current
suite.Equal(int64(len(sealedSegs)+len(invalidSealedSegs)), suite.idfOracle.current.NumRow())
// verify all sealed stats have local dir set
suite.idfOracle.sealed.Range(func(id int64, stats *sealedBm25Stats) bool {
stats.RLock()
defer stats.RUnlock()
suite.NotEmpty(stats.localDir)
return true
})
// update and sync snapshot make all sealed in target activate
suite.updateSnapshot(sealedSegs, []int64{}, []int64{})
suite.idfOracle.SetNext(suite.snapshot)
suite.waitTargetVersion(suite.targetVersion)
suite.Equal(int64(len(sealedSegs)), suite.idfOracle.current.NumRow())
suite.Require().Eventually(func() bool {
allInLocal := true
suite.idfOracle.sealed.Range(func(id int64, stats *sealedBm25Stats) bool {
stats.RLock()
defer stats.RUnlock()
if stats.inmemory == true {
allInLocal = false
return false
}
return true
})
return allInLocal
}, time.Minute, time.Millisecond*100)
// release some segments
releasedSeg := []int64{1, 2, 3}
suite.updateSnapshot([]int64{}, []int64{}, releasedSeg)
@@ -257,6 +295,165 @@ func (suite *IDFOracleSuite) TestLocalCache() {
suite.Equal(int64(1), suite.idfOracle.current.NumRow())
}
func (suite *IDFOracleSuite) TestFetchStatsRemoved() {
segID := int64(1)
suite.registerSealed(segID, 1, 5)
stats, ok := suite.idfOracle.sealed.Get(segID)
suite.True(ok)
// remove then fetch — should return error
stats.Remove()
_, err := stats.FetchStats()
suite.Error(err)
suite.Contains(err.Error(), "already removed")
}
func (suite *IDFOracleSuite) TestDiskSizeTracking() {
disk1 := suite.registerSealed(1, 1, 2)
disk2 := suite.registerSealed(2, 2, 3)
suite.Equal(disk1+disk2, suite.idfOracle.sealedDiskSize.Load())
// SyncDistribution with only seg 1 in target — seg 2 gets removed
suite.updateSnapshot([]int64{1}, []int64{}, []int64{})
suite.idfOracle.SetNext(suite.snapshot)
suite.waitTargetVersion(suite.targetVersion)
suite.Equal(disk1, suite.idfOracle.sealedDiskSize.Load())
// release seg 1
suite.updateSnapshot([]int64{}, []int64{}, []int64{1})
suite.idfOracle.SetNext(suite.snapshot)
suite.waitTargetVersion(suite.targetVersion)
suite.Equal(int64(0), suite.idfOracle.sealedDiskSize.Load())
}
func (suite *IDFOracleSuite) TestDiskSizeTrackingSyncDistribution() {
sealedSegs := []int64{1, 2, 3}
var totalDisk int64
for _, segID := range sealedSegs {
totalDisk += suite.registerSealed(segID, uint32(segID), uint32(segID)+1)
}
suite.Equal(totalDisk, suite.idfOracle.sealedDiskSize.Load())
// activate only seg 1,2 via SyncDistribution — seg 3 gets removed
suite.updateSnapshot([]int64{1, 2}, []int64{}, []int64{})
suite.idfOracle.SetNext(suite.snapshot)
suite.waitTargetVersion(suite.targetVersion)
suite.Equal(2, suite.idfOracle.sealed.Len())
suite.Less(suite.idfOracle.sealedDiskSize.Load(), totalDisk)
}
func (suite *IDFOracleSuite) TestMemorySize() {
// initial state — empty current stats
suite.Greater(suite.idfOracle.MemorySize(), int64(0)) // current has fixed overhead
// add growing segments — memory should increase
sizeBefore := suite.idfOracle.MemorySize()
suite.idfOracle.RegisterGrowing(1, suite.genStats(1, 100))
sizeAfter := suite.idfOracle.MemorySize()
suite.Greater(sizeAfter, sizeBefore)
}
func (suite *IDFOracleSuite) TestUpdateGrowingCheckMemory() {
suite.idfOracle.RegisterGrowing(1, suite.genStats(1, 2))
// repeated updates grow the stats
for i := uint32(2); i < 200; i++ {
suite.idfOracle.UpdateGrowing(1, suite.genStats(i, i+1))
}
suite.Equal(int64(199), suite.idfOracle.current.NumRow())
}
func (suite *IDFOracleSuite) TestLoadSealedIdempotent() {
suite.registerSealed(1, 1, 5)
suite.Equal(int64(4), suite.idfOracle.current.NumRow())
diskSize := suite.idfOracle.sealedDiskSize.Load()
// duplicate load — should be skipped
suite.registerSealed(1, 1, 5)
suite.Equal(int64(4), suite.idfOracle.current.NumRow())
suite.Equal(diskSize, suite.idfOracle.sealedDiskSize.Load())
}
func (suite *IDFOracleSuite) TestLoadSealedEmptyBm25Logs() {
cm := mocks.NewChunkManager(suite.T())
// nil bm25Logs
err := suite.idfOracle.LoadSealed(context.Background(), 1, &querypb.SegmentLoadInfo{}, cm)
suite.NoError(err)
suite.False(suite.idfOracle.sealed.Contain(1))
// empty bm25Logs
err = suite.idfOracle.LoadSealed(context.Background(), 2, &querypb.SegmentLoadInfo{Bm25Logs: []*datapb.FieldBinlog{}}, cm)
suite.NoError(err)
suite.False(suite.idfOracle.sealed.Contain(2))
}
func (suite *IDFOracleSuite) TestLoadSealedNoParse() {
// set targetVersion > 0 so needParse = false
suite.idfOracle.targetVersion.Store(1)
stats := suite.genStats(1, 5)
data, err := stats[102].Serialize()
suite.Require().NoError(err)
cm := mocks.NewChunkManager(suite.T())
remotePath := "bm25stats/seg_1/field_102/0"
cm.EXPECT().Reader(mock.Anything, remotePath).Return(
&bytesFileReader{bytes.NewReader(data)}, nil,
)
bm25Logs := []*datapb.FieldBinlog{{
FieldID: 102,
Binlogs: []*datapb.Binlog{{LogPath: remotePath}},
}}
err = suite.idfOracle.LoadSealed(context.Background(), 1, &querypb.SegmentLoadInfo{Bm25Logs: bm25Logs}, cm)
suite.NoError(err)
// segment registered but NOT preloaded (current stays 0)
suite.True(suite.idfOracle.sealed.Contain(1))
suite.Equal(int64(0), suite.idfOracle.current.NumRow())
// disk file should exist
segDir := path.Join(suite.idfOracle.dirPath, "1", "102")
entries, err := os.ReadDir(segDir)
suite.NoError(err)
suite.NotEmpty(entries)
// FetchStats should work (reads from disk)
sealedStats, ok := suite.idfOracle.sealed.Get(1)
suite.True(ok)
fetched, err := sealedStats.FetchStats()
suite.NoError(err)
suite.Equal(int64(4), fetched[102].NumRow())
}
func (suite *IDFOracleSuite) TestLoadSealedFailureCleanup() {
cm := mocks.NewChunkManager(suite.T())
remotePath := "bm25stats/seg_1/field_102/0"
cm.EXPECT().Reader(mock.Anything, remotePath).Return(
nil, errors.New("remote read failed"),
)
bm25Logs := []*datapb.FieldBinlog{{
FieldID: 102,
Binlogs: []*datapb.Binlog{{LogPath: remotePath}},
}}
err := suite.idfOracle.LoadSealed(context.Background(), 1, &querypb.SegmentLoadInfo{Bm25Logs: bm25Logs}, cm)
suite.Error(err)
// segment should NOT be registered
suite.False(suite.idfOracle.sealed.Contain(1))
suite.Equal(int64(0), suite.idfOracle.sealedDiskSize.Load())
// disk directory should be cleaned up
segDir := path.Join(suite.idfOracle.dirPath, "1")
_, statErr := os.Stat(segDir)
suite.True(os.IsNotExist(statErr))
}
func TestIDFOracle(t *testing.T) {
suite.Run(t, new(IDFOracleSuite))
}
+20 -50
View File
@@ -10,12 +10,9 @@ import (
mock "github.com/stretchr/testify/mock"
pkoracle "github.com/milvus-io/milvus/internal/querynodev2/pkoracle"
querypb "github.com/milvus-io/milvus/pkg/v2/proto/querypb"
storage "github.com/milvus-io/milvus/internal/storage"
typeutil "github.com/milvus-io/milvus/pkg/v2/util/typeutil"
querypb "github.com/milvus-io/milvus/pkg/v2/proto/querypb"
)
// MockLoader is an autogenerated mock type for the Loader type
@@ -107,76 +104,49 @@ func (_c *MockLoader_Load_Call) RunAndReturn(run func(context.Context, int64, co
return _c
}
// LoadBM25Stats provides a mock function with given fields: ctx, collectionID, infos
func (_m *MockLoader) LoadBM25Stats(ctx context.Context, collectionID int64, infos ...*querypb.SegmentLoadInfo) (*typeutil.ConcurrentMap[int64, map[int64]*storage.BM25Stats], error) {
_va := make([]interface{}, len(infos))
for _i := range infos {
_va[_i] = infos[_i]
}
var _ca []interface{}
_ca = append(_ca, ctx, collectionID)
_ca = append(_ca, _va...)
ret := _m.Called(_ca...)
// GetChunkManager provides a mock function with no fields
func (_m *MockLoader) GetChunkManager() storage.ChunkManager {
ret := _m.Called()
if len(ret) == 0 {
panic("no return value specified for LoadBM25Stats")
panic("no return value specified for GetChunkManager")
}
var r0 *typeutil.ConcurrentMap[int64, map[int64]*storage.BM25Stats]
var r1 error
if rf, ok := ret.Get(0).(func(context.Context, int64, ...*querypb.SegmentLoadInfo) (*typeutil.ConcurrentMap[int64, map[int64]*storage.BM25Stats], error)); ok {
return rf(ctx, collectionID, infos...)
}
if rf, ok := ret.Get(0).(func(context.Context, int64, ...*querypb.SegmentLoadInfo) *typeutil.ConcurrentMap[int64, map[int64]*storage.BM25Stats]); ok {
r0 = rf(ctx, collectionID, infos...)
var r0 storage.ChunkManager
if rf, ok := ret.Get(0).(func() storage.ChunkManager); ok {
r0 = rf()
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*typeutil.ConcurrentMap[int64, map[int64]*storage.BM25Stats])
r0 = ret.Get(0).(storage.ChunkManager)
}
}
if rf, ok := ret.Get(1).(func(context.Context, int64, ...*querypb.SegmentLoadInfo) error); ok {
r1 = rf(ctx, collectionID, infos...)
} else {
r1 = ret.Error(1)
}
return r0, r1
return r0
}
// MockLoader_LoadBM25Stats_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'LoadBM25Stats'
type MockLoader_LoadBM25Stats_Call struct {
// MockLoader_GetChunkManager_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'GetChunkManager'
type MockLoader_GetChunkManager_Call struct {
*mock.Call
}
// LoadBM25Stats is a helper method to define mock.On call
// - ctx context.Context
// - collectionID int64
// - infos ...*querypb.SegmentLoadInfo
func (_e *MockLoader_Expecter) LoadBM25Stats(ctx interface{}, collectionID interface{}, infos ...interface{}) *MockLoader_LoadBM25Stats_Call {
return &MockLoader_LoadBM25Stats_Call{Call: _e.mock.On("LoadBM25Stats",
append([]interface{}{ctx, collectionID}, infos...)...)}
// GetChunkManager is a helper method to define mock.On call
func (_e *MockLoader_Expecter) GetChunkManager() *MockLoader_GetChunkManager_Call {
return &MockLoader_GetChunkManager_Call{Call: _e.mock.On("GetChunkManager")}
}
func (_c *MockLoader_LoadBM25Stats_Call) Run(run func(ctx context.Context, collectionID int64, infos ...*querypb.SegmentLoadInfo)) *MockLoader_LoadBM25Stats_Call {
func (_c *MockLoader_GetChunkManager_Call) Run(run func()) *MockLoader_GetChunkManager_Call {
_c.Call.Run(func(args mock.Arguments) {
variadicArgs := make([]*querypb.SegmentLoadInfo, len(args)-2)
for i, a := range args[2:] {
if a != nil {
variadicArgs[i] = a.(*querypb.SegmentLoadInfo)
}
}
run(args[0].(context.Context), args[1].(int64), variadicArgs...)
run()
})
return _c
}
func (_c *MockLoader_LoadBM25Stats_Call) Return(_a0 *typeutil.ConcurrentMap[int64, map[int64]*storage.BM25Stats], _a1 error) *MockLoader_LoadBM25Stats_Call {
_c.Call.Return(_a0, _a1)
func (_c *MockLoader_GetChunkManager_Call) Return(_a0 storage.ChunkManager) *MockLoader_GetChunkManager_Call {
_c.Call.Return(_a0)
return _c
}
func (_c *MockLoader_LoadBM25Stats_Call) RunAndReturn(run func(context.Context, int64, ...*querypb.SegmentLoadInfo) (*typeutil.ConcurrentMap[int64, map[int64]*storage.BM25Stats], error)) *MockLoader_LoadBM25Stats_Call {
func (_c *MockLoader_GetChunkManager_Call) RunAndReturn(run func() storage.ChunkManager) *MockLoader_GetChunkManager_Call {
_c.Call.Return(run)
return _c
}
+27 -68
View File
@@ -87,8 +87,8 @@ type Loader interface {
// LoadBloomFilterSet loads needed statslog for RemoteSegment.
LoadBloomFilterSet(ctx context.Context, collectionID int64, infos ...*querypb.SegmentLoadInfo) ([]*pkoracle.BloomFilterSet, error)
// LoadBM25Stats loads BM25 statslog for RemoteSegment
LoadBM25Stats(ctx context.Context, collectionID int64, infos ...*querypb.SegmentLoadInfo) (*typeutil.ConcurrentMap[int64, map[int64]*storage.BM25Stats], error)
// GetChunkManager returns the chunk manager for remote storage access.
GetChunkManager() storage.ChunkManager
// LoadIndex append index for segment and remove vector binlogs.
LoadIndex(ctx context.Context,
@@ -633,49 +633,8 @@ func (loader *segmentLoader) waitSegmentLoadDone(ctx context.Context, segmentTyp
return nil
}
func (loader *segmentLoader) LoadBM25Stats(ctx context.Context, collectionID int64, infos ...*querypb.SegmentLoadInfo) (*typeutil.ConcurrentMap[int64, map[int64]*storage.BM25Stats], error) {
segmentNum := len(infos)
if segmentNum == 0 {
return nil, nil
}
log.Info("start loading bm25 stats for remote...", zap.Int64("collectionID", collectionID), zap.Int("segmentNum", segmentNum))
loadedStats := typeutil.NewConcurrentMap[int64, map[int64]*storage.BM25Stats]()
loadRemoteBM25Func := func(idx int) error {
loadInfo := infos[idx]
segmentID := loadInfo.SegmentID
stats := make(map[int64]*storage.BM25Stats)
log.Info("loading bm25 stats for remote...", zap.Int64("collectionID", collectionID), zap.Int64("segment", segmentID))
logpaths, err := packed.NewStatsResolverFromLoadInfo(loadInfo).BM25StatsPaths()
if err != nil {
log.Warn("load remote segment bm25 stats paths failed",
zap.Int64("segmentID", segmentID),
zap.Error(err),
)
return err
}
err = loader.loadBm25Stats(ctx, segmentID, stats, logpaths)
if err != nil {
log.Warn("load remote segment bm25 stats failed",
zap.Int64("segmentID", segmentID),
zap.Error(err),
)
return err
}
loadedStats.Insert(segmentID, stats)
return nil
}
err := funcutil.ProcessFuncParallel(segmentNum, segmentNum, loadRemoteBM25Func, "loadRemoteBM25Func")
if err != nil {
// no partial success here
log.Warn("failed to load bm25 stats for remote segment", zap.Int64("collectionID", collectionID), zap.Error(err))
return nil, err
}
return loadedStats, nil
func (loader *segmentLoader) GetChunkManager() storage.ChunkManager {
return loader.cm
}
// load single bloom filter
@@ -1239,29 +1198,6 @@ func (loader *segmentLoader) loadFieldsIndex(ctx context.Context,
return nil
}
func (loader *segmentLoader) loadFieldIndex(ctx context.Context, segment *LocalSegment, indexInfo *querypb.FieldIndexInfo) error {
filteredPaths := make([]string, 0, len(indexInfo.IndexFilePaths))
for _, indexPath := range indexInfo.IndexFilePaths {
if path.Base(indexPath) != storage.IndexParamsKey {
filteredPaths = append(filteredPaths, indexPath)
}
}
indexInfo.IndexFilePaths = filteredPaths
fieldType, err := loader.getFieldType(segment.Collection(), indexInfo.FieldID)
if err != nil {
return err
}
collection := loader.manager.Collection.Get(segment.Collection())
if collection == nil {
return merr.WrapErrCollectionNotLoaded(segment.Collection(), "failed to load field index")
}
return segment.LoadIndex(ctx, indexInfo, fieldType)
}
func (loader *segmentLoader) loadBm25Stats(ctx context.Context, segmentID int64, stats map[int64]*storage.BM25Stats, binlogPaths map[int64][]string) error {
log := log.Ctx(ctx).With(
zap.Int64("segmentID", segmentID),
@@ -1307,6 +1243,29 @@ func (loader *segmentLoader) loadBm25Stats(ctx context.Context, segmentID int64,
return nil
}
func (loader *segmentLoader) loadFieldIndex(ctx context.Context, segment *LocalSegment, indexInfo *querypb.FieldIndexInfo) error {
filteredPaths := make([]string, 0, len(indexInfo.IndexFilePaths))
for _, indexPath := range indexInfo.IndexFilePaths {
if path.Base(indexPath) != storage.IndexParamsKey {
filteredPaths = append(filteredPaths, indexPath)
}
}
indexInfo.IndexFilePaths = filteredPaths
fieldType, err := loader.getFieldType(segment.Collection(), indexInfo.FieldID)
if err != nil {
return err
}
collection := loader.manager.Collection.Get(segment.Collection())
if collection == nil {
return merr.WrapErrCollectionNotLoaded(segment.Collection(), "failed to load field index")
}
return segment.LoadIndex(ctx, indexInfo, fieldType)
}
func (loader *segmentLoader) loadBloomFilter(ctx context.Context, segmentID int64, bfs *pkoracle.BloomFilterSet,
binlogPaths []string,
) error {
@@ -102,22 +102,6 @@ func (suite *SegmentLoaderSuite) SetupTest() {
suite.manager.Collection.PutOrRef(suite.collectionID, suite.schema, indexMeta, loadMeta)
}
func (suite *SegmentLoaderSuite) SetupBM25() {
// Dependencies
suite.manager = NewManager()
suite.loader = NewLoader(context.Background(), suite.manager, suite.chunkManager)
initcore.InitRemoteChunkManager(paramtable.Get())
suite.schema = mock_segcore.GenTestBM25CollectionSchema("test")
indexMeta := mock_segcore.GenTestIndexMeta(suite.collectionID, suite.schema)
loadMeta := &querypb.LoadMetaInfo{
LoadType: querypb.LoadType_LoadCollection,
CollectionID: suite.collectionID,
PartitionIDs: []int64{suite.partitionID},
}
suite.manager.Collection.PutOrRef(suite.collectionID, suite.schema, indexMeta, loadMeta)
}
func (suite *SegmentLoaderSuite) TearDownTest() {
ctx := context.Background()
for i := 0; i < suite.segmentNum; i++ {
@@ -437,41 +421,6 @@ func (suite *SegmentLoaderSuite) TestLoadDeltaLogs() {
}
}
func (suite *SegmentLoaderSuite) TestLoadBm25Stats() {
suite.SetupBM25()
msgLength := 1
sparseFieldID := mock_segcore.SimpleSparseFloatVectorField.ID
loadInfos := make([]*querypb.SegmentLoadInfo, 0, suite.segmentNum)
for i := 0; i < suite.segmentNum; i++ {
segmentID := suite.segmentID + int64(i)
bm25logs, err := mock_segcore.SaveBM25Log(suite.collectionID, suite.partitionID, segmentID, sparseFieldID, msgLength, suite.chunkManager)
suite.NoError(err)
loadInfos = append(loadInfos, &querypb.SegmentLoadInfo{
SegmentID: segmentID,
PartitionID: suite.partitionID,
CollectionID: suite.collectionID,
Bm25Logs: []*datapb.FieldBinlog{bm25logs},
NumOfRows: int64(msgLength),
InsertChannel: fmt.Sprintf("by-dev-rootcoord-dml_0_%dv0", suite.collectionID),
})
}
statsMap, err := suite.loader.LoadBM25Stats(context.Background(), suite.collectionID, loadInfos...)
suite.NoError(err)
for i := 0; i < suite.segmentNum; i++ {
segmentID := suite.segmentID + int64(i)
stats, ok := statsMap.Get(segmentID)
suite.True(ok)
fieldStats, ok := stats[sparseFieldID]
suite.True(ok)
suite.Equal(int64(msgLength), fieldStats.NumRow())
}
}
func (suite *SegmentLoaderSuite) TestLoadDupDeltaLogs() {
ctx := context.Background()
loadInfos := make([]*querypb.SegmentLoadInfo, 0, suite.segmentNum)
+48
View File
@@ -514,6 +514,54 @@ func DeserializeBloomFilterStats(paths []string, blobs []*Blob) ([]*PrimaryKeySt
return DeserializeStats(blobs)
}
// bm25StatsPerEntryBytes is the estimated memory cost per entry in the rowsWithToken map.
// Go map overhead per entry: key(4) + value(4) + bucket/pointer overhead (~72B) ≈ 80 bytes.
const bm25StatsPerEntryBytes = 80
// MemSize estimates the in-memory size of this BM25Stats in bytes.
// len(map) is O(1) in Go (reads hmap.count directly).
func (m *BM25Stats) MemSize() int64 {
// Fixed overhead: numRow(8) + numToken(8) + map header (~100B)
return 120 + int64(len(m.rowsWithToken))*bm25StatsPerEntryBytes
}
// DeserializeFromReader reads BM25 stats from an io.Reader and accumulates into self.
// Unlike Deserialize([]byte), this does not require knowing the total size upfront.
func (m *BM25Stats) DeserializeFromReader(r io.Reader) error {
var version int32
if err := binary.Read(r, common.Endian, &version); err != nil {
return err
}
var numRow, tokenNum int64
if err := binary.Read(r, common.Endian, &numRow); err != nil {
return err
}
if err := binary.Read(r, common.Endian, &tokenNum); err != nil {
return err
}
m.numRow += numRow
m.numToken += tokenNum
var key uint32
var value int32
for {
if err := binary.Read(r, common.Endian, &key); err != nil {
if err == io.EOF {
break
}
return err
}
if err := binary.Read(r, common.Endian, &value); err != nil {
return err
}
m.rowsWithToken[key] += value
}
return nil
}
// DeserializeStats deserializes @blobs as []*PrimaryKeyStats
func DeserializeStats(blobs []*Blob) ([]*PrimaryKeyStats, error) {
results := make([]*PrimaryKeyStats, 0, len(blobs))
+82
View File
@@ -17,6 +17,8 @@
package storage
import (
"bytes"
"encoding/binary"
"testing"
"github.com/stretchr/testify/assert"
@@ -263,3 +265,83 @@ func TestMarshalStats(t *testing.T) {
assert.True(t, stat1[0].BF.Test(b))
}
}
func TestBM25Stats_MemSize(t *testing.T) {
stats := NewBM25Stats()
baseSize := stats.MemSize()
assert.Equal(t, int64(120), baseSize)
// Add tokens and verify size grows
for i := uint32(0); i < 100; i++ {
stats.Append(map[uint32]float32{i: 1})
}
assert.Equal(t, int64(120+100*bm25StatsPerEntryBytes), stats.MemSize())
}
func TestBM25Stats_DeserializeFromReader(t *testing.T) {
t.Run("roundtrip", func(t *testing.T) {
original := NewBM25Stats()
for i := uint32(0); i < 50; i++ {
original.Append(map[uint32]float32{i: 1})
}
data, err := original.Serialize()
assert.NoError(t, err)
restored := NewBM25Stats()
err = restored.DeserializeFromReader(bytes.NewReader(data))
assert.NoError(t, err)
assert.Equal(t, original.NumRow(), restored.NumRow())
assert.Equal(t, original.GetAvgdl(), restored.GetAvgdl())
})
t.Run("accumulate_multiple", func(t *testing.T) {
s1 := NewBM25Stats()
s1.Append(map[uint32]float32{1: 1, 2: 1})
d1, _ := s1.Serialize()
s2 := NewBM25Stats()
s2.Append(map[uint32]float32{2: 1, 3: 1})
d2, _ := s2.Serialize()
merged := NewBM25Stats()
assert.NoError(t, merged.DeserializeFromReader(bytes.NewReader(d1)))
assert.NoError(t, merged.DeserializeFromReader(bytes.NewReader(d2)))
assert.Equal(t, int64(2), merged.NumRow())
})
t.Run("truncated_header", func(t *testing.T) {
// Only 10 bytes, header needs 20 (version + numRow + tokenNum)
data := make([]byte, 10)
restored := NewBM25Stats()
err := restored.DeserializeFromReader(bytes.NewReader(data))
assert.Error(t, err)
})
t.Run("truncated_value", func(t *testing.T) {
// Valid header + key but truncated value
buf := new(bytes.Buffer)
binary.Write(buf, common.Endian, int32(0)) // version
binary.Write(buf, common.Endian, int64(1)) // numRow
binary.Write(buf, common.Endian, int64(1)) // numToken
binary.Write(buf, common.Endian, uint32(42)) // key
binary.Write(buf, common.Endian, int16(1)) // truncated value (2 bytes instead of 4)
restored := NewBM25Stats()
err := restored.DeserializeFromReader(buf)
assert.Error(t, err)
})
t.Run("empty_tokens", func(t *testing.T) {
// Valid header, zero tokens
buf := new(bytes.Buffer)
binary.Write(buf, common.Endian, int32(0)) // version
binary.Write(buf, common.Endian, int64(5)) // numRow
binary.Write(buf, common.Endian, int64(0)) // numToken
restored := NewBM25Stats()
err := restored.DeserializeFromReader(buf)
assert.NoError(t, err)
assert.Equal(t, int64(5), restored.NumRow())
})
}
+6 -5
View File
@@ -3479,19 +3479,20 @@ type queryNodeConfig struct {
EnabledGrowingSegmentJSONKeyStats ParamItem `refreshable:"false"`
// Idf Oracle
IDFEnableDisk ParamItem `refreshable:"true"`
IDFPreload ParamItem `refreshable:"true"`
// partial search
PartialResultRequiredDataRatio ParamItem `refreshable:"true"`
}
func (p *queryNodeConfig) init(base *BaseTable) {
p.IDFEnableDisk = ParamItem{
Key: "queryNode.idfOracle.enableDisk",
Version: "2.6.0",
p.IDFPreload = ParamItem{
Key: "queryNode.idfOracle.preload",
Version: "2.6.8",
Export: true,
DefaultValue: "true",
Doc: "Whether to parse and merge BM25 stats into current during load before first target. When false, stats are only written to disk and loaded on first SyncDistribution.",
}
p.IDFEnableDisk.Init(base.mgr)
p.IDFPreload.Init(base.mgr)
p.SoPath = ParamItem{
Key: "queryNode.soPath",