enhance: shard_interceptor check schema version for insert msg(#49455) (#49456)

related: #49455

Signed-off-by: MrPresent-Han <chun.han@gmail.com>
Co-authored-by: MrPresent-Han <chun.han@gmail.com>
This commit is contained in:
Chun Han
2026-04-30 15:57:51 +08:00
committed by GitHub
co-authored by MrPresent-Han
parent 0fa822fc0a
commit 72f83e01f1
10 changed files with 1123 additions and 837 deletions
@@ -450,9 +450,9 @@ func (_c *MockShardManager_CheckIfCollectionExists_Call) RunAndReturn(run func(i
return _c
}
// CheckIfCollectionSchemaVersionMatch provides a mock function with given fields: collectionID, schemaVersion
func (_m *MockShardManager) CheckIfCollectionSchemaVersionMatch(collectionID int64, schemaVersion int32) (int32, error) {
ret := _m.Called(collectionID, schemaVersion)
// CheckIfCollectionSchemaVersionMatch provides a mock function with given fields: header
func (_m *MockShardManager) CheckIfCollectionSchemaVersionMatch(header *message.InsertMessageHeader) (int32, error) {
ret := _m.Called(header)
if len(ret) == 0 {
panic("no return value specified for CheckIfCollectionSchemaVersionMatch")
@@ -460,17 +460,17 @@ func (_m *MockShardManager) CheckIfCollectionSchemaVersionMatch(collectionID int
var r0 int32
var r1 error
if rf, ok := ret.Get(0).(func(int64, int32) (int32, error)); ok {
return rf(collectionID, schemaVersion)
if rf, ok := ret.Get(0).(func(*message.InsertMessageHeader) (int32, error)); ok {
return rf(header)
}
if rf, ok := ret.Get(0).(func(int64, int32) int32); ok {
r0 = rf(collectionID, schemaVersion)
if rf, ok := ret.Get(0).(func(*message.InsertMessageHeader) int32); ok {
r0 = rf(header)
} else {
r0 = ret.Get(0).(int32)
}
if rf, ok := ret.Get(1).(func(int64, int32) error); ok {
r1 = rf(collectionID, schemaVersion)
if rf, ok := ret.Get(1).(func(*message.InsertMessageHeader) error); ok {
r1 = rf(header)
} else {
r1 = ret.Error(1)
}
@@ -484,15 +484,14 @@ type MockShardManager_CheckIfCollectionSchemaVersionMatch_Call struct {
}
// CheckIfCollectionSchemaVersionMatch is a helper method to define mock.On call
// - collectionID int64
// - schemaVersion int32
func (_e *MockShardManager_Expecter) CheckIfCollectionSchemaVersionMatch(collectionID interface{}, schemaVersion interface{}) *MockShardManager_CheckIfCollectionSchemaVersionMatch_Call {
return &MockShardManager_CheckIfCollectionSchemaVersionMatch_Call{Call: _e.mock.On("CheckIfCollectionSchemaVersionMatch", collectionID, schemaVersion)}
// - header *message.InsertMessageHeader
func (_e *MockShardManager_Expecter) CheckIfCollectionSchemaVersionMatch(header interface{}) *MockShardManager_CheckIfCollectionSchemaVersionMatch_Call {
return &MockShardManager_CheckIfCollectionSchemaVersionMatch_Call{Call: _e.mock.On("CheckIfCollectionSchemaVersionMatch", header)}
}
func (_c *MockShardManager_CheckIfCollectionSchemaVersionMatch_Call) Run(run func(collectionID int64, schemaVersion int32)) *MockShardManager_CheckIfCollectionSchemaVersionMatch_Call {
func (_c *MockShardManager_CheckIfCollectionSchemaVersionMatch_Call) Run(run func(header *message.InsertMessageHeader)) *MockShardManager_CheckIfCollectionSchemaVersionMatch_Call {
_c.Call.Run(func(args mock.Arguments) {
run(args[0].(int64), args[1].(int32))
run(args[0].(*message.InsertMessageHeader))
})
return _c
}
@@ -502,7 +501,7 @@ func (_c *MockShardManager_CheckIfCollectionSchemaVersionMatch_Call) Return(_a0
return _c
}
func (_c *MockShardManager_CheckIfCollectionSchemaVersionMatch_Call) RunAndReturn(run func(int64, int32) (int32, error)) *MockShardManager_CheckIfCollectionSchemaVersionMatch_Call {
func (_c *MockShardManager_CheckIfCollectionSchemaVersionMatch_Call) RunAndReturn(run func(*message.InsertMessageHeader) (int32, error)) *MockShardManager_CheckIfCollectionSchemaVersionMatch_Call {
_c.Call.Return(run)
return _c
}
+2 -2
View File
@@ -124,7 +124,7 @@ func repackInsertDataForStreamingService(
BinarySize: 0, // TODO: current not used, message estimate size is used.
},
},
SchemaVersion: schemaVersion,
SchemaVersion: &schemaVersion,
}).
WithBody(insertRequest).
WithCipher(ez).
@@ -207,7 +207,7 @@ func repackInsertDataWithPartitionKeyForStreamingService(
BinarySize: 0, // TODO: current not used, message estimate size is used.
},
},
SchemaVersion: schemaVersion,
SchemaVersion: &schemaVersion,
}).
WithBody(insertRequest).
WithCipher(ez).
+56
View File
@@ -8,6 +8,7 @@ import (
"github.com/stretchr/testify/mock"
"github.com/milvus-io/milvus-proto/go-api/v2/commonpb"
"github.com/milvus-io/milvus-proto/go-api/v2/milvuspb"
"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/allocator"
@@ -16,11 +17,66 @@ import (
"github.com/milvus-io/milvus/pkg/v2/common"
"github.com/milvus-io/milvus/pkg/v2/mq/msgstream"
"github.com/milvus-io/milvus/pkg/v2/proto/rootcoordpb"
"github.com/milvus-io/milvus/pkg/v2/streaming/util/message"
"github.com/milvus-io/milvus/pkg/v2/util/merr"
"github.com/milvus-io/milvus/pkg/v2/util/paramtable"
"github.com/milvus-io/milvus/pkg/v2/util/testutils"
)
func TestRepackInsertDataForStreamingServicePreservesExplicitZeroSchemaVersion(t *testing.T) {
paramtable.Init()
oldCache := globalMetaCache
cache := NewMockCache(t)
cache.On("GetPartitionID", mock.Anything, "db", "coll", "_default").Return(int64(200), nil)
globalMetaCache = cache
defer func() { globalMetaCache = oldCache }()
insertMsg := &msgstream.InsertMsg{
InsertRequest: &msgpb.InsertRequest{
Base: &commonpb.MsgBase{
MsgType: commonpb.MsgType_Insert,
SourceID: 1,
},
CollectionID: 100,
DbName: "db",
CollectionName: "coll",
PartitionName: "_default",
NumRows: 1,
FieldsData: []*schemapb.FieldData{
{
FieldName: "pk",
FieldId: 1,
Type: schemapb.DataType_Int64,
Field: &schemapb.FieldData_Scalars{
Scalars: &schemapb.ScalarField{
Data: &schemapb.ScalarField_LongData{
LongData: &schemapb.LongArray{Data: []int64{1}},
},
},
},
},
},
RowIDs: []int64{1},
Timestamps: []uint64{1},
},
}
result := &milvuspb.MutationResult{
IDs: &schemapb.IDs{
IdField: &schemapb.IDs_IntId{IntId: &schemapb.LongArray{Data: []int64{1}}},
},
}
msgs, err := repackInsertDataForStreamingService(context.Background(), []string{"ch"}, insertMsg, result, nil, 0)
assert.NoError(t, err)
assert.Len(t, msgs, 1)
msg := message.MustAsMutableInsertMessageV1(msgs[0])
header := msg.Header()
assert.NotNil(t, header.SchemaVersion)
assert.Equal(t, int32(0), header.GetSchemaVersion())
}
func TestInsertTask_CheckAligned(t *testing.T) {
var err error
@@ -145,17 +145,19 @@ func (impl *shardInterceptor) handleInsertMessage(ctx context.Context, msg messa
// Assign segment for insert message.
// !!! Current implementation a insert message only has one parition, but we need to merge the message for partition-key in future.
header := insertMsg.Header()
collectionID := header.GetCollectionId()
schemaVersion := header.GetSchemaVersion()
if correctSchemaVersion, err := impl.shardManager.CheckIfCollectionSchemaVersionMatch(header.GetCollectionId(), schemaVersion); err != nil {
if correctSchemaVersion, err := impl.shardManager.CheckIfCollectionSchemaVersionMatch(header); err != nil {
if errors.Is(err, shards.ErrCollectionNotFound) {
return nil, status.NewUnrecoverableError("collection %d not found", header.GetCollectionId())
return nil, status.NewUnrecoverableError("collection %d not found", collectionID)
}
if errors.Is(err, shards.ErrCollectionSchemaNotFound) {
return nil, status.NewUnrecoverableError("collection %d schema not provided by create collection message", header.GetCollectionId())
return nil, status.NewUnrecoverableError("collection %d schema not provided by create collection message", collectionID)
}
if errors.Is(err, shards.ErrCollectionSchemaVersionNotMatch) {
impl.shardManager.Logger().Warn("insertMessage schema version mismatch",
zap.Int64("collectionID", header.GetCollectionId()),
zap.Int64("collectionID", collectionID),
zap.Bool("schemaVersionProvided", header.SchemaVersion != nil),
zap.Int32("schemaVersion", schemaVersion),
zap.Int32("collectionSchemaVersion", correctSchemaVersion),
zap.Error(err))
@@ -163,7 +165,8 @@ func (impl *shardInterceptor) handleInsertMessage(ctx context.Context, msg messa
schemaVersion, correctSchemaVersion)
}
impl.shardManager.Logger().Error("unexpected error from CheckIfCollectionSchemaVersionMatch",
zap.Int64("collectionID", header.GetCollectionId()),
zap.Int64("collectionID", collectionID),
zap.Bool("schemaVersionProvided", header.SchemaVersion != nil),
zap.Int32("schemaVersion", schemaVersion),
zap.Error(err))
return nil, errors.Wrap(err, "CheckIfCollectionSchemaVersionMatch")
@@ -8,6 +8,10 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"go.uber.org/atomic"
"go.uber.org/zap"
"go.uber.org/zap/zapcore"
"go.uber.org/zap/zaptest/observer"
"google.golang.org/protobuf/proto"
"github.com/milvus-io/milvus-proto/go-api/v2/msgpb"
"github.com/milvus-io/milvus/internal/mocks/streamingnode/server/wal/interceptors/shard/mock_shards"
@@ -20,6 +24,164 @@ import (
"github.com/milvus-io/milvus/pkg/v2/streaming/walimpls/impls/rmq"
)
func TestShardInterceptorLogsOmittedSchemaVersionAsNotProvided(t *testing.T) {
core, logs := observer.New(zapcore.WarnLevel)
logger := &log.MLogger{Logger: zap.New(core)}
b := NewInterceptorBuilder()
shardManager := mock_shards.NewMockShardManager(t)
shardManager.EXPECT().Logger().Return(logger).Maybe()
i := b.Build(&interceptors.InterceptorBuildParam{
ShardManager: shardManager,
})
defer i.Close()
msg := message.NewInsertMessageBuilderV1().
WithVChannel("v1").
WithHeader(&messagespb.InsertMessageHeader{
CollectionId: 1,
Partitions: []*messagespb.PartitionSegmentAssignment{
{
PartitionId: 1,
Rows: 1,
BinarySize: 100,
},
},
}).
WithBody(&msgpb.InsertRequest{}).
MustBuildMutable().WithTimeTick(1)
insertHdrMatcher := mock.MatchedBy(func(h *message.InsertMessageHeader) bool {
return h != nil && h.GetCollectionId() == int64(1) && h.SchemaVersion == nil
})
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(insertHdrMatcher).Return(int32(5), shards.ErrCollectionSchemaVersionNotMatch)
msgID, err := i.DoAppend(context.Background(), msg, func(ctx context.Context, msg message.MutableMessage) (message.MessageID, error) {
return rmq.NewRmqID(1), nil
})
assert.Error(t, err)
assert.Nil(t, msgID)
entries := logs.FilterMessage("insertMessage schema version mismatch").All()
assert.Len(t, entries, 1)
assert.Equal(t, false, entries[0].ContextMap()["schemaVersionProvided"])
}
func TestShardInterceptorReportsExplicitZeroSchemaVersionInMismatchError(t *testing.T) {
b := NewInterceptorBuilder()
shardManager := mock_shards.NewMockShardManager(t)
shardManager.EXPECT().Logger().Return(log.With()).Maybe()
i := b.Build(&interceptors.InterceptorBuildParam{
ShardManager: shardManager,
})
defer i.Close()
zero := proto.Int32(0)
msg := message.NewInsertMessageBuilderV1().
WithVChannel("v1").
WithHeader(&messagespb.InsertMessageHeader{
CollectionId: 1,
Partitions: []*messagespb.PartitionSegmentAssignment{
{
PartitionId: 1,
Rows: 1,
BinarySize: 100,
},
},
SchemaVersion: zero,
}).
WithBody(&msgpb.InsertRequest{}).
MustBuildMutable().WithTimeTick(1)
insertHdrMatcher := mock.MatchedBy(func(h *message.InsertMessageHeader) bool {
return h != nil && h.GetCollectionId() == int64(1) && h.SchemaVersion != nil && h.GetSchemaVersion() == 0
})
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(insertHdrMatcher).Return(int32(5), shards.ErrCollectionSchemaVersionNotMatch)
msgID, err := i.DoAppend(context.Background(), msg, func(ctx context.Context, msg message.MutableMessage) (message.MessageID, error) {
return rmq.NewRmqID(1), nil
})
assert.Error(t, err)
assert.ErrorContains(t, err, "input schema version: 0")
assert.Nil(t, msgID)
}
func TestShardInterceptorPassesExplicitNonZeroSchemaVersion(t *testing.T) {
b := NewInterceptorBuilder()
shardManager := mock_shards.NewMockShardManager(t)
shardManager.EXPECT().Logger().Return(log.With()).Maybe()
i := b.Build(&interceptors.InterceptorBuildParam{
ShardManager: shardManager,
})
defer i.Close()
msg := message.NewInsertMessageBuilderV1().
WithVChannel("v1").
WithHeader(&messagespb.InsertMessageHeader{
CollectionId: 1,
Partitions: []*messagespb.PartitionSegmentAssignment{
{
PartitionId: 1,
Rows: 1,
BinarySize: 100,
},
},
SchemaVersion: proto.Int32(3),
}).
WithBody(&msgpb.InsertRequest{}).
MustBuildMutable().WithTimeTick(1)
insertHdrMatcher := mock.MatchedBy(func(h *message.InsertMessageHeader) bool {
return h != nil && h.GetCollectionId() == int64(1) && h.SchemaVersion != nil && h.GetSchemaVersion() == 3
})
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(insertHdrMatcher).Return(int32(3), nil)
shardManager.EXPECT().AssignSegment(mock.Anything).Return(&shards.AssignSegmentResult{SegmentID: 1, Acknowledge: atomic.NewInt32(1)}, nil)
msgID, err := i.DoAppend(context.Background(), msg, func(ctx context.Context, msg message.MutableMessage) (message.MessageID, error) {
return rmq.NewRmqID(1), nil
})
assert.NoError(t, err)
assert.NotNil(t, msgID)
}
func TestShardInterceptorPassesExplicitZeroSchemaVersion(t *testing.T) {
b := NewInterceptorBuilder()
shardManager := mock_shards.NewMockShardManager(t)
shardManager.EXPECT().Logger().Return(log.With()).Maybe()
i := b.Build(&interceptors.InterceptorBuildParam{
ShardManager: shardManager,
})
defer i.Close()
zero := proto.Int32(0)
msg := message.NewInsertMessageBuilderV1().
WithVChannel("v1").
WithHeader(&messagespb.InsertMessageHeader{
CollectionId: 1,
Partitions: []*messagespb.PartitionSegmentAssignment{
{
PartitionId: 1,
Rows: 1,
BinarySize: 100,
},
},
SchemaVersion: zero,
}).
WithBody(&msgpb.InsertRequest{}).
MustBuildMutable().WithTimeTick(1)
insertHdrMatcher := mock.MatchedBy(func(h *message.InsertMessageHeader) bool {
return h != nil && h.GetCollectionId() == int64(1) && h.SchemaVersion != nil && h.GetSchemaVersion() == 0
})
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(insertHdrMatcher).Return(int32(0), nil)
shardManager.EXPECT().AssignSegment(mock.Anything).Return(&shards.AssignSegmentResult{SegmentID: 1, Acknowledge: atomic.NewInt32(1)}, nil)
msgID, err := i.DoAppend(context.Background(), msg, func(ctx context.Context, msg message.MutableMessage) (message.MessageID, error) {
return rmq.NewRmqID(1), nil
})
assert.NoError(t, err)
assert.NotNil(t, msgID)
}
func TestShardInterceptor(t *testing.T) {
mockErr := errors.New("mock error")
@@ -202,44 +364,43 @@ func TestShardInterceptor(t *testing.T) {
WithBody(&msgpb.InsertRequest{}).
MustBuildMutable().WithTimeTick(1)
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(int64(1), int32(0)).Return(int32(0), nil)
shardManager.EXPECT().AssignSegment(mock.Anything).Return(&shards.AssignSegmentResult{SegmentID: 1, Acknowledge: atomic.NewInt32(1)}, nil)
insertHdrMatcher := mock.MatchedBy(func(h *message.InsertMessageHeader) bool {
return h != nil && h.GetCollectionId() == int64(1) && h.SchemaVersion == nil
})
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(insertHdrMatcher).Return(int32(0), nil).Once()
shardManager.EXPECT().AssignSegment(mock.Anything).Return(&shards.AssignSegmentResult{SegmentID: 1, Acknowledge: atomic.NewInt32(1)}, nil).Once()
msgID, err = i.DoAppend(ctx, msg, appender)
assert.NoError(t, err)
assert.NotNil(t, msgID)
shardManager.EXPECT().AssignSegment(mock.Anything).Unset()
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(int64(1), int32(0)).Return(int32(0), nil)
shardManager.EXPECT().AssignSegment(mock.Anything).Return(nil, mockErr)
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(insertHdrMatcher).Return(int32(0), nil).Once()
shardManager.EXPECT().AssignSegment(mock.Anything).Return(nil, mockErr).Once()
msgID, err = i.DoAppend(ctx, msg, appender)
assert.Error(t, err)
assert.Nil(t, msgID)
// ErrCollectionNotFound from schema version check must surface as an unrecoverable insert error.
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(int64(1), int32(0)).Unset()
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(int64(1), int32(0)).Return(int32(-1), shards.ErrCollectionNotFound)
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(insertHdrMatcher).Return(int32(-1), shards.ErrCollectionNotFound).Once()
msgID, err = i.DoAppend(ctx, msg, appender)
assert.Error(t, err)
assert.Nil(t, msgID)
// ErrCollectionSchemaNotFound must also become an unrecoverable insert error.
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(int64(1), int32(0)).Unset()
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(int64(1), int32(0)).Return(int32(-1), shards.ErrCollectionSchemaNotFound)
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(insertHdrMatcher).Return(int32(-1), shards.ErrCollectionSchemaNotFound).Once()
msgID, err = i.DoAppend(ctx, msg, appender)
assert.Error(t, err)
assert.Nil(t, msgID)
// ErrCollectionSchemaVersionNotMatch must surface as a schema-version-mismatch error
// so the proxy can refresh its cache and retry.
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(int64(1), int32(0)).Unset()
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(int64(1), int32(0)).Return(int32(5), shards.ErrCollectionSchemaVersionNotMatch)
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(insertHdrMatcher).Return(int32(5), shards.ErrCollectionSchemaVersionNotMatch).Once()
msgID, err = i.DoAppend(ctx, msg, appender)
assert.Error(t, err)
assert.Nil(t, msgID)
// Unexpected error from the schema version check must be propagated as-is.
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(int64(1), int32(0)).Unset()
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(int64(1), int32(0)).Return(int32(-1), mockErr)
shardManager.EXPECT().CheckIfCollectionSchemaVersionMatch(insertHdrMatcher).Return(int32(-1), mockErr).Once()
msgID, err = i.DoAppend(ctx, msg, appender)
assert.Error(t, err)
assert.Nil(t, msgID)
@@ -186,14 +186,15 @@ func (m *shardManagerImpl) AlterCollection(msg message.MutableAlterCollectionMes
return segmentIDs, nil
}
func (m *shardManagerImpl) CheckIfCollectionSchemaVersionMatch(collectionID int64, schemaVersion int32) (int32, error) {
func (m *shardManagerImpl) CheckIfCollectionSchemaVersionMatch(header *message.InsertMessageHeader) (int32, error) {
m.mu.Lock()
defer m.mu.Unlock()
return m.checkIfCollectionSchemaVersionMatch(collectionID, schemaVersion)
return m.checkIfCollectionSchemaVersionMatch(header)
}
func (m *shardManagerImpl) checkIfCollectionSchemaVersionMatch(collectionID int64, schemaVersion int32) (int32, error) {
func (m *shardManagerImpl) checkIfCollectionSchemaVersionMatch(header *message.InsertMessageHeader) (int32, error) {
collectionID := header.GetCollectionId()
collectionInfo, ok := m.collections[collectionID]
if !ok {
m.Logger().Warn("collection not found", zap.Int64("collectionID", collectionID))
@@ -202,19 +203,22 @@ func (m *shardManagerImpl) checkIfCollectionSchemaVersionMatch(collectionID int6
// Input schemaVersion 0 means the proxy did not set it (old proxy or old SDK).
// Skip the schema presence and version checks for backward compatibility during rolling
// upgrades, where a legacy collection may still have Schema == nil when an old proxy writes.
if schemaVersion == 0 {
if header.SchemaVersion == nil {
return collectionInfo.SchemaVersion(), nil
}
if collectionInfo.Schema == nil || collectionInfo.Schema.GetSchema() == nil {
m.Logger().Warn("collection schema not found", zap.Int64("collectionID", collectionID))
return -1, ErrCollectionSchemaNotFound
}
collectionSchemaVersion := collectionInfo.SchemaVersion()
if collectionSchemaVersion != schemaVersion {
if collectionSchemaVersion != header.GetSchemaVersion() {
m.Logger().Warn("collection schema version not match", zap.Int64("collectionID", collectionID),
zap.Int32("schemaVersion", schemaVersion),
zap.Int32("schemaVersion", header.GetSchemaVersion()),
zap.Int32("collectionSchemaVersion", collectionSchemaVersion))
return collectionSchemaVersion, ErrCollectionSchemaVersionNotMatch
}
return collectionSchemaVersion, nil
}
@@ -53,7 +53,8 @@ type ShardManager interface {
// Returns the IDs of flushed segments (non-empty only for schema changes).
AlterCollection(msg message.MutableAlterCollectionMessageV2) ([]int64, error)
CheckIfCollectionSchemaVersionMatch(collectionID int64, schemaVersion int32) (int32, error)
// CheckIfCollectionSchemaVersionMatch validates insert header schema version against in-memory collection state.
CheckIfCollectionSchemaVersionMatch(header *message.InsertMessageHeader) (int32, error)
Close()
}
@@ -6,6 +6,7 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/fieldmaskpb"
"github.com/milvus-io/milvus-proto/go-api/v2/msgpb"
@@ -321,7 +322,10 @@ func TestShardManagerSchemaVersionCheck(t *testing.T) {
}).(*shardManagerImpl)
// Test 1: CheckIfCollectionSchemaVersionMatch on non-existent collection
_, err := m.CheckIfCollectionSchemaVersionMatch(999, 1)
_, err := m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{
CollectionId: 999,
SchemaVersion: proto.Int32(1),
})
assert.ErrorIs(t, err, ErrCollectionNotFound)
// Test 2: Create collection with schema, then check version match
@@ -344,18 +348,24 @@ func TestShardManagerSchemaVersionCheck(t *testing.T) {
IntoImmutableMessage(rmq.NewRmqID(10))
m.CreateCollection(message.MustAsImmutableCreateCollectionMessageV1(createMsg))
// version match should succeed
ver, err := m.CheckIfCollectionSchemaVersionMatch(100, 1)
// version match should succeed when header carries explicit schema version
ver, err := m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{
CollectionId: 100,
SchemaVersion: proto.Int32(1),
})
assert.NoError(t, err)
assert.Equal(t, int32(1), ver)
// version 0 (old proxy) should skip check and succeed
ver, err = m.CheckIfCollectionSchemaVersionMatch(100, 0)
// legacy insert (no schema_version field): skip schema validation and return current version
ver, err = m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{CollectionId: 100})
assert.NoError(t, err)
assert.Equal(t, int32(1), ver)
// version mismatch should fail
ver, err = m.CheckIfCollectionSchemaVersionMatch(100, 2)
ver, err = m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{
CollectionId: 100,
SchemaVersion: proto.Int32(2),
})
assert.ErrorIs(t, err, ErrCollectionSchemaVersionNotMatch)
assert.Equal(t, int32(1), ver)
@@ -373,13 +383,46 @@ func TestShardManagerSchemaVersionCheck(t *testing.T) {
IntoImmutableMessage(rmq.NewRmqID(11))
m.CreateCollection(message.MustAsImmutableCreateCollectionMessageV1(createMsgNoSchema))
// collection exists but has no schema
_, err = m.CheckIfCollectionSchemaVersionMatch(101, 1)
// collection exists but has no schema — explicit version in header requires schema metadata
_, err = m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{
CollectionId: 101,
SchemaVersion: proto.Int32(1),
})
assert.ErrorIs(t, err, ErrCollectionSchemaNotFound)
// backward-compat: old proxy sends schemaVersion=0 against a schema-less collection,
// must succeed and return the legacy version 0 instead of ErrCollectionSchemaNotFound.
ver, err = m.CheckIfCollectionSchemaVersionMatch(101, 0)
_, err = m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{
CollectionId: 101,
SchemaVersion: proto.Int32(0),
})
assert.ErrorIs(t, err, ErrCollectionSchemaNotFound)
// legacy insert while collection schema version is still 0: allow
ver, err = m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{CollectionId: 101})
assert.NoError(t, err)
assert.Equal(t, int32(0), ver)
createMsgSchemaVersionZero := message.NewCreateCollectionMessageBuilderV1().
WithVChannel("v_schema_zero").
WithHeader(&message.CreateCollectionMessageHeader{
CollectionId: 103,
PartitionIds: []int64{203},
}).
WithBody(&msgpb.CreateCollectionRequest{
CollectionSchema: &schemapb.CollectionSchema{
Name: "test_schema_zero_collection",
Version: 0,
},
}).
MustBuildMutable().
WithTimeTick(350).
WithLastConfirmedUseMessageID().
IntoImmutableMessage(rmq.NewRmqID(12))
m.CreateCollection(message.MustAsImmutableCreateCollectionMessageV1(createMsgSchemaVersionZero))
ver, err = m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{
CollectionId: 103,
SchemaVersion: proto.Int32(0),
})
assert.NoError(t, err)
assert.Equal(t, int32(0), ver)
@@ -407,11 +450,17 @@ func TestShardManagerSchemaVersionCheck(t *testing.T) {
assert.NoError(t, err)
// now version 2 matches, version 1 doesn't
ver, err = m.CheckIfCollectionSchemaVersionMatch(100, 2)
ver, err = m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{
CollectionId: 100,
SchemaVersion: proto.Int32(2),
})
assert.NoError(t, err)
assert.Equal(t, int32(2), ver)
ver, err = m.CheckIfCollectionSchemaVersionMatch(100, 1)
ver, err = m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{
CollectionId: 100,
SchemaVersion: proto.Int32(1),
})
assert.ErrorIs(t, err, ErrCollectionSchemaVersionNotMatch)
assert.Equal(t, int32(2), ver)
@@ -448,11 +497,14 @@ func TestShardManagerSchemaVersionCheck(t *testing.T) {
},
}
m.mu.Unlock()
_, err = m.CheckIfCollectionSchemaVersionMatch(102, 1)
_, err = m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{
CollectionId: 102,
SchemaVersion: proto.Int32(1),
})
assert.ErrorIs(t, err, ErrCollectionSchemaNotFound)
// same backward-compat path: schemaVersion=0 must not be rejected by the nil inner schema.
ver, err = m.CheckIfCollectionSchemaVersionMatch(102, 0)
// legacy insert while effective collection schema version is still 0 (nil inner schema)
ver, err = m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{CollectionId: 102})
assert.NoError(t, err)
assert.Equal(t, int32(0), ver)
@@ -565,7 +617,10 @@ func TestAlterCollectionSchemaChange(t *testing.T) {
// The growing segment must have been flushed and fenced.
assert.Equal(t, []int64{segID}, segmentIDs)
// Schema must be updated to v2.
ver, err := m.CheckIfCollectionSchemaVersionMatch(collID, 2)
ver, err := m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{
CollectionId: collID,
SchemaVersion: proto.Int32(2),
})
assert.NoError(t, err)
assert.Equal(t, int32(2), ver)
})
@@ -609,7 +664,10 @@ func TestAlterCollectionSchemaChange(t *testing.T) {
assert.NoError(t, err)
assert.Empty(t, segmentIDs)
// Schema must remain at v1, not updated to v2.
ver, err := m.CheckIfCollectionSchemaVersionMatch(collID, 1)
ver, err := m.CheckIfCollectionSchemaVersionMatch(&message.InsertMessageHeader{
CollectionId: collID,
SchemaVersion: proto.Int32(1),
})
assert.NoError(t, err)
assert.Equal(t, int32(1), ver)
})
+2 -1
View File
@@ -158,7 +158,8 @@ message TimeTickMessageHeader {}
message InsertMessageHeader {
int64 collection_id = 1;
repeated PartitionSegmentAssignment partitions = 2;
int32 schema_version = 3;
// optional so consumers can distinguish omitted (legacy producer) from explicit value.
optional int32 schema_version = 3;
}
// PartitionSegmentAssignment is the segment assignment of a partition.
File diff suppressed because it is too large Load Diff