mirror of
https://github.com/milvus-io/milvus.git
synced 2026-07-21 10:15:43 +00:00
related: #47259 Signed-off-by: MrPresent-Han <chun.han@gmail.com> Co-authored-by: MrPresent-Han <chun.han@gmail.com>
This commit is contained in:
co-authored by
MrPresent-Han
parent
1c080484e9
commit
8ef6c2f092
@@ -1499,6 +1499,28 @@ func (h *HandlersV2) search(ctx context.Context, c *gin.Context, anyReq any, dbN
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Check if search by primary keys or by vectors
|
||||
hasIDs := len(httpReq.Ids) > 0
|
||||
hasData := len(httpReq.Data) > 0
|
||||
|
||||
// Primary keys and query vectors are mutually exclusive
|
||||
if hasIDs && hasData {
|
||||
HTTPAbortReturn(c, http.StatusOK, gin.H{
|
||||
HTTPReturnCode: merr.Code(merr.ErrParameterInvalid),
|
||||
HTTPReturnMessage: "primary keys (ids) and query vectors (data) are mutually exclusive. Please provide either 'ids' or 'data', not both",
|
||||
})
|
||||
return nil, merr.ErrParameterInvalid
|
||||
}
|
||||
|
||||
// At least one of ids or data must be provided
|
||||
if !hasIDs && !hasData {
|
||||
HTTPAbortReturn(c, http.StatusOK, gin.H{
|
||||
HTTPReturnCode: merr.Code(merr.ErrMissingRequiredParameters),
|
||||
HTTPReturnMessage: "either 'ids' (for primary key search) or 'data' (for vector search) must be provided",
|
||||
})
|
||||
return nil, merr.ErrMissingRequiredParameters
|
||||
}
|
||||
|
||||
searchParams, err := generateSearchParams(httpReq.SearchParams)
|
||||
if err != nil {
|
||||
log.Ctx(ctx).Warn("high level restful api, generate SearchParams failed", zap.Error(err))
|
||||
@@ -1524,21 +1546,55 @@ func (h *HandlersV2) search(ctx context.Context, c *gin.Context, anyReq any, dbN
|
||||
}
|
||||
}
|
||||
|
||||
searchParams = append(searchParams, &commonpb.KeyValuePair{Key: proxy.AnnsFieldKey, Value: httpReq.AnnsField})
|
||||
body, _ := c.Get(gin.BodyBytesKey)
|
||||
placeholderGroup, err := generatePlaceholderGroup(ctx, string(body.([]byte)), collSchema, httpReq.AnnsField)
|
||||
if err != nil {
|
||||
log.Ctx(ctx).Warn("high level restful api, search with vector invalid", zap.Error(err))
|
||||
HTTPAbortReturn(c, http.StatusOK, gin.H{
|
||||
HTTPReturnCode: merr.Code(merr.ErrIncorrectParameterFormat),
|
||||
HTTPReturnMessage: merr.ErrIncorrectParameterFormat.Error() + ", error: " + err.Error(),
|
||||
})
|
||||
return nil, err
|
||||
if hasIDs {
|
||||
// Search by primary keys
|
||||
primaryField, ok := getPrimaryField(collSchema)
|
||||
if !ok {
|
||||
HTTPAbortReturn(c, http.StatusOK, gin.H{
|
||||
HTTPReturnCode: merr.Code(merr.ErrParameterInvalid),
|
||||
HTTPReturnMessage: "collection has no primary key field",
|
||||
})
|
||||
return nil, merr.ErrParameterInvalid
|
||||
}
|
||||
|
||||
// Convert ids to schemapb.IDs
|
||||
ids, err := convertIDsToSchemapbIDs(httpReq.Ids, primaryField)
|
||||
if err != nil {
|
||||
log.Ctx(ctx).Warn("high level restful api, convert ids to schemapb.IDs failed", zap.Error(err))
|
||||
HTTPAbortReturn(c, http.StatusOK, gin.H{
|
||||
HTTPReturnCode: merr.Code(merr.ErrParameterInvalid),
|
||||
HTTPReturnMessage: merr.ErrParameterInvalid.Error() + ", error: " + err.Error(),
|
||||
})
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Set Ids field using the oneof SearchInput field
|
||||
req.SearchInput = &milvuspb.SearchRequest_Ids{
|
||||
Ids: ids,
|
||||
}
|
||||
// Set anns_field in search params if provided
|
||||
if httpReq.AnnsField != "" {
|
||||
searchParams = append(searchParams, &commonpb.KeyValuePair{Key: proxy.AnnsFieldKey, Value: httpReq.AnnsField})
|
||||
}
|
||||
} else {
|
||||
// Search by vectors (existing logic)
|
||||
searchParams = append(searchParams, &commonpb.KeyValuePair{Key: proxy.AnnsFieldKey, Value: httpReq.AnnsField})
|
||||
body, _ := c.Get(gin.BodyBytesKey)
|
||||
placeholderGroup, err := generatePlaceholderGroup(ctx, string(body.([]byte)), collSchema, httpReq.AnnsField)
|
||||
if err != nil {
|
||||
log.Ctx(ctx).Warn("high level restful api, search with vector invalid", zap.Error(err))
|
||||
HTTPAbortReturn(c, http.StatusOK, gin.H{
|
||||
HTTPReturnCode: merr.Code(merr.ErrIncorrectParameterFormat),
|
||||
HTTPReturnMessage: merr.ErrIncorrectParameterFormat.Error() + ", error: " + err.Error(),
|
||||
})
|
||||
return nil, err
|
||||
}
|
||||
req.SearchInput = &milvuspb.SearchRequest_PlaceholderGroup{
|
||||
PlaceholderGroup: placeholderGroup,
|
||||
}
|
||||
}
|
||||
|
||||
req.SearchParams = searchParams
|
||||
req.SearchInput = &milvuspb.SearchRequest_PlaceholderGroup{
|
||||
PlaceholderGroup: placeholderGroup,
|
||||
}
|
||||
req.ExprTemplateValues = generateExpressionTemplate(httpReq.ExprParams)
|
||||
resp, err := wrapperProxyWithLimit(ctx, c, req, h.checkAuth, false, "/milvus.proto.milvus.MilvusService/Search", true, h.proxy, func(reqCtx context.Context, req any) (interface{}, error) {
|
||||
return h.proxy.Search(reqCtx, req.(*milvuspb.SearchRequest))
|
||||
|
||||
@@ -3026,3 +3026,78 @@ func TestTruncateCollection(t *testing.T) {
|
||||
sendReqAndVerify(t, testEngine, testcase.path, http.MethodPost, testcase)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSearchByPK(t *testing.T) {
|
||||
paramtable.Init()
|
||||
// disable rate limit
|
||||
paramtable.Get().Save(paramtable.Get().QuotaConfig.QuotaAndLimitsEnabled.Key, "false")
|
||||
defer paramtable.Get().Reset(paramtable.Get().QuotaConfig.QuotaAndLimitsEnabled.Key)
|
||||
|
||||
outputFields := []string{FieldBookID, FieldWordCount}
|
||||
mp := mocks.NewMockProxy(t)
|
||||
|
||||
// Mock for successful search by PK with int64 IDs
|
||||
mp.EXPECT().DescribeCollection(mock.Anything, mock.Anything).Return(&milvuspb.DescribeCollectionResponse{
|
||||
CollectionName: DefaultCollectionName,
|
||||
Schema: generateCollectionSchema(schemapb.DataType_Int64, false, true),
|
||||
ShardsNum: ShardNumDefault,
|
||||
Status: &StatusSuccess,
|
||||
}, nil).Times(5)
|
||||
|
||||
mp.EXPECT().Search(mock.Anything, mock.MatchedBy(func(req *milvuspb.SearchRequest) bool {
|
||||
// Verify the SearchInput is set with IDs
|
||||
return req.GetIds() != nil
|
||||
})).Return(&milvuspb.SearchResults{
|
||||
Status: commonSuccessStatus,
|
||||
Results: &schemapb.SearchResultData{
|
||||
TopK: int64(3),
|
||||
OutputFields: outputFields,
|
||||
FieldsData: generateFieldData(),
|
||||
Ids: generateIDs(schemapb.DataType_Int64, 3),
|
||||
Scores: DefaultScores,
|
||||
},
|
||||
}, nil).Times(2)
|
||||
|
||||
testEngine := initHTTPServerV2(mp, false)
|
||||
|
||||
queryTestCases := []requestBodyTestCase{}
|
||||
|
||||
// Test case 1: Search by PK with int64 IDs (JSON numbers decoded as float64)
|
||||
queryTestCases = append(queryTestCases, requestBodyTestCase{
|
||||
path: SearchAction,
|
||||
requestBody: []byte(`{"collectionName": "book", "ids": [1, 2, 3], "limit": 10, "outputFields": ["word_count"]}`),
|
||||
})
|
||||
|
||||
// Test case 2: Search by PK with string IDs for int64 PK
|
||||
queryTestCases = append(queryTestCases, requestBodyTestCase{
|
||||
path: SearchAction,
|
||||
requestBody: []byte(`{"collectionName": "book", "ids": ["1", "2", "3"], "limit": 10, "outputFields": ["word_count"]}`),
|
||||
})
|
||||
|
||||
// Test case 3: Search by PK with fractional float64 should fail
|
||||
queryTestCases = append(queryTestCases, requestBodyTestCase{
|
||||
path: SearchAction,
|
||||
requestBody: []byte(`{"collectionName": "book", "ids": [1.5, 2.9], "limit": 10, "outputFields": ["word_count"]}`),
|
||||
errMsg: "has fractional part",
|
||||
errCode: 1100, // ErrParameterInvalid
|
||||
})
|
||||
|
||||
// Test case 4: Search by PK with empty ids array should fail
|
||||
// Empty array is treated as "no ids provided" at request validation level
|
||||
queryTestCases = append(queryTestCases, requestBodyTestCase{
|
||||
path: SearchAction,
|
||||
requestBody: []byte(`{"collectionName": "book", "ids": [], "limit": 10, "outputFields": ["word_count"]}`),
|
||||
errMsg: "either 'ids' (for primary key search) or 'data' (for vector search) must be provided",
|
||||
errCode: 1802, // ErrIncorrectParameterFormat
|
||||
})
|
||||
|
||||
// Test case 5: Search by PK with invalid string value should fail
|
||||
queryTestCases = append(queryTestCases, requestBodyTestCase{
|
||||
path: SearchAction,
|
||||
requestBody: []byte(`{"collectionName": "book", "ids": ["not_a_number"], "limit": 10, "outputFields": ["word_count"]}`),
|
||||
errMsg: "invalid int64 id",
|
||||
errCode: 1100, // ErrParameterInvalid
|
||||
})
|
||||
|
||||
validateTestCases(t, testEngine, queryTestCases, false)
|
||||
}
|
||||
|
||||
@@ -333,7 +333,8 @@ func (req *CollectionDataReq) GetDbName() string { return req.DbName }
|
||||
type SearchReqV2 struct {
|
||||
DbName string `json:"dbName"`
|
||||
CollectionName string `json:"collectionName" binding:"required"`
|
||||
Data []interface{} `json:"data" binding:"required"`
|
||||
Data []interface{} `json:"data"`
|
||||
Ids []interface{} `json:"ids"`
|
||||
AnnsField string `json:"annsField"`
|
||||
PartitionNames []string `json:"partitionNames"`
|
||||
Filter string `json:"filter"`
|
||||
|
||||
@@ -174,6 +174,81 @@ func checkGetPrimaryKey(coll *schemapb.CollectionSchema, idResult gjson.Result)
|
||||
return filter, nil
|
||||
}
|
||||
|
||||
// convertIDsToSchemapbIDs converts a slice of interface{} (JSON ids) to schemapb.IDs
|
||||
// based on the primary key field type
|
||||
func convertIDsToSchemapbIDs(ids []interface{}, pkField *schemapb.FieldSchema) (*schemapb.IDs, error) {
|
||||
if len(ids) == 0 {
|
||||
return nil, errors.New("ids array cannot be empty")
|
||||
}
|
||||
|
||||
switch pkField.DataType {
|
||||
case schemapb.DataType_Int64:
|
||||
int64IDs := make([]int64, 0, len(ids))
|
||||
for i, id := range ids {
|
||||
var int64ID int64
|
||||
switch v := id.(type) {
|
||||
case int64:
|
||||
int64ID = v
|
||||
case int:
|
||||
int64ID = int64(v)
|
||||
case float64:
|
||||
// JSON numbers are decoded as float64
|
||||
// Check if the float has a fractional part
|
||||
if v != math.Trunc(v) {
|
||||
return nil, fmt.Errorf("invalid int64 id at index %d: %v has fractional part", i, v)
|
||||
}
|
||||
int64ID = int64(v)
|
||||
case string:
|
||||
// Try to parse string as int64
|
||||
parsed, err := strconv.ParseInt(v, 10, 64)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid int64 id at index %d: %v, error: %v", i, id, err)
|
||||
}
|
||||
int64ID = parsed
|
||||
default:
|
||||
return nil, fmt.Errorf("invalid id type at index %d: expected int64, got %T", i, id)
|
||||
}
|
||||
int64IDs = append(int64IDs, int64ID)
|
||||
}
|
||||
return &schemapb.IDs{
|
||||
IdField: &schemapb.IDs_IntId{
|
||||
IntId: &schemapb.LongArray{
|
||||
Data: int64IDs,
|
||||
},
|
||||
},
|
||||
}, nil
|
||||
|
||||
case schemapb.DataType_VarChar:
|
||||
stringIDs := make([]string, 0, len(ids))
|
||||
for i, id := range ids {
|
||||
var stringID string
|
||||
switch v := id.(type) {
|
||||
case string:
|
||||
stringID = v
|
||||
case int64, int, float64:
|
||||
// Convert number to string
|
||||
stringID = fmt.Sprintf("%v", v)
|
||||
default:
|
||||
return nil, fmt.Errorf("invalid id type at index %d: expected string, got %T", i, id)
|
||||
}
|
||||
if stringID == "" {
|
||||
return nil, fmt.Errorf("empty string id at index %d", i)
|
||||
}
|
||||
stringIDs = append(stringIDs, stringID)
|
||||
}
|
||||
return &schemapb.IDs{
|
||||
IdField: &schemapb.IDs_StrId{
|
||||
StrId: &schemapb.StringArray{
|
||||
Data: stringIDs,
|
||||
},
|
||||
},
|
||||
}, nil
|
||||
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported primary key type: %s", pkField.DataType.String())
|
||||
}
|
||||
}
|
||||
|
||||
// --------------------- collection details --------------------- //
|
||||
|
||||
func printFields(fields []*schemapb.FieldSchema) []gin.H {
|
||||
|
||||
@@ -2793,3 +2793,142 @@ func TestParseUsernamePassword(t *testing.T) {
|
||||
assert.Equal(t, "", password)
|
||||
})
|
||||
}
|
||||
|
||||
func TestConvertIDsToSchemapbIDs(t *testing.T) {
|
||||
int64PkField := &schemapb.FieldSchema{
|
||||
FieldID: common.StartOfUserFieldID,
|
||||
Name: "id",
|
||||
IsPrimaryKey: true,
|
||||
DataType: schemapb.DataType_Int64,
|
||||
}
|
||||
|
||||
varcharPkField := &schemapb.FieldSchema{
|
||||
FieldID: common.StartOfUserFieldID,
|
||||
Name: "id",
|
||||
IsPrimaryKey: true,
|
||||
DataType: schemapb.DataType_VarChar,
|
||||
}
|
||||
|
||||
t.Run("empty ids array", func(t *testing.T) {
|
||||
_, err := convertIDsToSchemapbIDs([]interface{}{}, int64PkField)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "ids array cannot be empty")
|
||||
})
|
||||
|
||||
t.Run("int64 pk with float64 values (whole numbers)", func(t *testing.T) {
|
||||
// JSON numbers are decoded as float64
|
||||
ids := []interface{}{float64(1), float64(2), float64(3)}
|
||||
result, err := convertIDsToSchemapbIDs(ids, int64PkField)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
intIds := result.GetIntId()
|
||||
assert.NotNil(t, intIds)
|
||||
assert.Equal(t, []int64{1, 2, 3}, intIds.Data)
|
||||
})
|
||||
|
||||
t.Run("int64 pk with float64 values having fractional part", func(t *testing.T) {
|
||||
ids := []interface{}{float64(1.5)}
|
||||
_, err := convertIDsToSchemapbIDs(ids, int64PkField)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "has fractional part")
|
||||
})
|
||||
|
||||
t.Run("int64 pk with float64 values - second element has fractional part", func(t *testing.T) {
|
||||
ids := []interface{}{float64(1), float64(2.9)}
|
||||
_, err := convertIDsToSchemapbIDs(ids, int64PkField)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "index 1")
|
||||
assert.Contains(t, err.Error(), "has fractional part")
|
||||
})
|
||||
|
||||
t.Run("int64 pk with int64 values", func(t *testing.T) {
|
||||
ids := []interface{}{int64(100), int64(200)}
|
||||
result, err := convertIDsToSchemapbIDs(ids, int64PkField)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
intIds := result.GetIntId()
|
||||
assert.NotNil(t, intIds)
|
||||
assert.Equal(t, []int64{100, 200}, intIds.Data)
|
||||
})
|
||||
|
||||
t.Run("int64 pk with int values", func(t *testing.T) {
|
||||
ids := []interface{}{int(10), int(20)}
|
||||
result, err := convertIDsToSchemapbIDs(ids, int64PkField)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
intIds := result.GetIntId()
|
||||
assert.NotNil(t, intIds)
|
||||
assert.Equal(t, []int64{10, 20}, intIds.Data)
|
||||
})
|
||||
|
||||
t.Run("int64 pk with valid string values", func(t *testing.T) {
|
||||
ids := []interface{}{"123", "456"}
|
||||
result, err := convertIDsToSchemapbIDs(ids, int64PkField)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
intIds := result.GetIntId()
|
||||
assert.NotNil(t, intIds)
|
||||
assert.Equal(t, []int64{123, 456}, intIds.Data)
|
||||
})
|
||||
|
||||
t.Run("int64 pk with invalid string values", func(t *testing.T) {
|
||||
ids := []interface{}{"not_a_number"}
|
||||
_, err := convertIDsToSchemapbIDs(ids, int64PkField)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "invalid int64 id")
|
||||
})
|
||||
|
||||
t.Run("int64 pk with invalid type", func(t *testing.T) {
|
||||
ids := []interface{}{true}
|
||||
_, err := convertIDsToSchemapbIDs(ids, int64PkField)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "invalid id type")
|
||||
})
|
||||
|
||||
t.Run("varchar pk with string values", func(t *testing.T) {
|
||||
ids := []interface{}{"abc", "def", "ghi"}
|
||||
result, err := convertIDsToSchemapbIDs(ids, varcharPkField)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
strIds := result.GetStrId()
|
||||
assert.NotNil(t, strIds)
|
||||
assert.Equal(t, []string{"abc", "def", "ghi"}, strIds.Data)
|
||||
})
|
||||
|
||||
t.Run("varchar pk with empty string", func(t *testing.T) {
|
||||
ids := []interface{}{""}
|
||||
_, err := convertIDsToSchemapbIDs(ids, varcharPkField)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "empty string id")
|
||||
})
|
||||
|
||||
t.Run("varchar pk with number values", func(t *testing.T) {
|
||||
ids := []interface{}{float64(123), int64(456), int(789)}
|
||||
result, err := convertIDsToSchemapbIDs(ids, varcharPkField)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
strIds := result.GetStrId()
|
||||
assert.NotNil(t, strIds)
|
||||
assert.Equal(t, []string{"123", "456", "789"}, strIds.Data)
|
||||
})
|
||||
|
||||
t.Run("varchar pk with invalid type", func(t *testing.T) {
|
||||
ids := []interface{}{[]int{1, 2, 3}}
|
||||
_, err := convertIDsToSchemapbIDs(ids, varcharPkField)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "invalid id type")
|
||||
})
|
||||
|
||||
t.Run("unsupported pk type", func(t *testing.T) {
|
||||
boolPkField := &schemapb.FieldSchema{
|
||||
FieldID: common.StartOfUserFieldID,
|
||||
Name: "id",
|
||||
IsPrimaryKey: true,
|
||||
DataType: schemapb.DataType_Bool,
|
||||
}
|
||||
ids := []interface{}{float64(1)}
|
||||
_, err := convertIDsToSchemapbIDs(ids, boolPkField)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "unsupported primary key type")
|
||||
})
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user