feat: support restful search_by_pk(#47259) (#47260)

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:
Chun Han
2026-01-23 17:33:31 +08:00
committed by GitHub
co-authored by MrPresent-Han
parent 1c080484e9
commit 8ef6c2f092
5 changed files with 360 additions and 14 deletions
@@ -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")
})
}