fix: authorize REST import commit and abort (#50691)

## Summary
- Add authorization for REST v2 import commit and abort endpoints.
- Resolve the import job collection via GetImportProgress before
checking PrivilegeImport.
- Cover denied non-privileged users, allowed root users, and
authorization-disabled behavior.

Fixes #50458

## Test Plan
- [x] make build-cpp-with-unittest
- [x] go test -tags dynamic,test -gcflags=\"all=-N -l\" -count=1
./internal/distributed/proxy/httpserver/... -run
'TestCommitImportJob|TestAbortImportJob'
- [x] go test -tags dynamic,test -gcflags=\"all=-N -l\" -count=1
./internal/distributed/proxy/httpserver/... -skip '^TestSearchV2$'
- [x] GOTOOLCHAIN=auto make static-check

## Notes
- Full httpserver package test currently has an unrelated existing
TestSearchV2/search#09 expectation mismatch: expected `Mismatch type
uint8`, actual contains `Mismatch type []uint8`.

---------

Signed-off-by: Yihao Dai <yihao.dai@zilliz.com>
This commit is contained in:
yihao.dai
2026-06-26 06:46:26 +08:00
committed by GitHub
parent 66bf1f168b
commit 486c8fe18b
2 changed files with 289 additions and 99 deletions
@@ -3363,66 +3363,98 @@ func (h *HandlersV2) createImportJob(ctx context.Context, c *gin.Context, anyReq
func (h *HandlersV2) getImportJobProcess(ctx context.Context, c *gin.Context, anyReq any, dbName string) (interface{}, error) {
jobIDGetter := anyReq.(JobIDGetter)
response, err := h.getImportProgress(ctx, c, dbName, jobIDGetter.GetJobID())
if err != nil {
return response, err
}
if err := h.checkImportPrivilege(ctx, c, dbName, response.GetCollectionName()); err != nil {
return nil, err
}
returnData := make(map[string]interface{})
returnData["jobId"] = jobIDGetter.GetJobID()
returnData["createTime"] = response.GetCreateTime()
returnData["collectionName"] = response.GetCollectionName()
returnData["completeTime"] = response.GetCompleteTime()
returnData["state"] = response.GetState().String()
returnData["progress"] = response.GetProgress()
returnData["importedRows"] = response.GetImportedRows()
returnData["totalRows"] = response.GetTotalRows()
reason := response.GetReason()
if reason != "" {
returnData["reason"] = reason
}
details := make([]map[string]interface{}, 0)
totalFileSize := int64(0)
for _, taskProgress := range response.GetTaskProgresses() {
detail := make(map[string]interface{})
detail["fileName"] = taskProgress.GetFileName()
detail["fileSize"] = taskProgress.GetFileSize()
detail["progress"] = taskProgress.GetProgress()
detail["completeTime"] = taskProgress.GetCompleteTime()
detail["state"] = taskProgress.GetState()
detail["importedRows"] = taskProgress.GetImportedRows()
detail["totalRows"] = taskProgress.GetTotalRows()
reason = taskProgress.GetReason()
if reason != "" {
detail["reason"] = reason
}
details = append(details, detail)
totalFileSize += taskProgress.GetFileSize()
}
returnData["fileSize"] = totalFileSize
returnData["details"] = details
HTTPReturn(c, http.StatusOK, gin.H{HTTPReturnCode: merr.Code(nil), HTTPReturnData: returnData})
return response, nil
}
func (h *HandlersV2) getImportProgress(ctx context.Context, c *gin.Context, dbName string, jobID string) (*internalpb.GetImportProgressResponse, error) {
req := &internalpb.GetImportProgressRequest{
DbName: dbName,
JobID: jobIDGetter.GetJobID(),
JobID: jobID,
}
c.Set(ContextRequest, req)
resp, err := wrapperProxy(ctx, c, req, false, false, "/milvus.proto.milvus.MilvusService/GetImportProgress", func(reqCtx context.Context, req any) (interface{}, error) {
return h.proxy.GetImportProgress(reqCtx, req.(*internalpb.GetImportProgressRequest))
})
if err == nil {
response := resp.(*internalpb.GetImportProgressResponse)
if h.checkAuth {
err := checkAuthorizationHelper(ctx, c, &milvuspb.ImportAuthPlaceholder{
DbName: dbName,
CollectionName: response.GetCollectionName(),
})
if err != nil {
HTTPReturn(c, http.StatusForbidden, gin.H{
HTTPReturnCode: merr.Code(err),
HTTPReturnMessage: err.Error(),
})
return nil, err
}
}
returnData := make(map[string]interface{})
returnData["jobId"] = jobIDGetter.GetJobID()
returnData["createTime"] = response.GetCreateTime()
returnData["collectionName"] = response.GetCollectionName()
returnData["completeTime"] = response.GetCompleteTime()
returnData["state"] = response.GetState().String()
returnData["progress"] = response.GetProgress()
returnData["importedRows"] = response.GetImportedRows()
returnData["totalRows"] = response.GetTotalRows()
reason := response.GetReason()
if reason != "" {
returnData["reason"] = reason
}
details := make([]map[string]interface{}, 0)
totalFileSize := int64(0)
for _, taskProgress := range response.GetTaskProgresses() {
detail := make(map[string]interface{})
detail["fileName"] = taskProgress.GetFileName()
detail["fileSize"] = taskProgress.GetFileSize()
detail["progress"] = taskProgress.GetProgress()
detail["completeTime"] = taskProgress.GetCompleteTime()
detail["state"] = taskProgress.GetState()
detail["importedRows"] = taskProgress.GetImportedRows()
detail["totalRows"] = taskProgress.GetTotalRows()
reason = taskProgress.GetReason()
if reason != "" {
detail["reason"] = reason
}
details = append(details, detail)
totalFileSize += taskProgress.GetFileSize()
}
returnData["fileSize"] = totalFileSize
returnData["details"] = details
HTTPReturn(c, http.StatusOK, gin.H{HTTPReturnCode: merr.Code(nil), HTTPReturnData: returnData})
if err != nil {
return nil, err
}
return resp, err
return resp.(*internalpb.GetImportProgressResponse), nil
}
func (h *HandlersV2) checkImportPrivilege(ctx context.Context, c *gin.Context, dbName string, collectionName string) error {
if !h.checkAuth {
return nil
}
authErr := checkAuthorizationHelper(ctx, c, &milvuspb.ImportAuthPlaceholder{
DbName: dbName,
CollectionName: collectionName,
})
if authErr != nil {
HTTPReturn(c, http.StatusForbidden, gin.H{
HTTPReturnCode: merr.Code(authErr),
HTTPReturnMessage: authErr.Error(),
})
return authErr
}
return nil
}
func (h *HandlersV2) checkImportJobAuth(ctx context.Context, c *gin.Context, dbName string, jobID string) error {
if !h.checkAuth {
return nil
}
// Commit/abort requests only carry jobID. Import privilege is collection-scoped,
// so resolve the job's collection before checking PrivilegeImport.
response, err := h.getImportProgress(ctx, c, dbName, jobID)
if err != nil {
return err
}
return h.checkImportPrivilege(ctx, c, dbName, response.GetCollectionName())
}
func (h *HandlersV2) refreshExternalCollection(ctx context.Context, c *gin.Context, anyReq any, dbName string) (interface{}, error) {
@@ -3455,6 +3487,9 @@ func (h *HandlersV2) commitImportJob(ctx context.Context, c *gin.Context, anyReq
HTTPAbortReturn(c, http.StatusOK, gin.H{HTTPReturnCode: merr.Code(paramErr), HTTPReturnMessage: paramErr.Error()})
return nil, paramErr
}
if err := h.checkImportJobAuth(ctx, c, dbName, jobIDGetter.GetJobID()); err != nil {
return nil, err
}
req := &datapb.CommitImportRequest{
Base: commonpbutil.NewMsgBase(),
JobId: jobID,
@@ -3498,6 +3533,9 @@ func (h *HandlersV2) abortImportJob(ctx context.Context, c *gin.Context, anyReq
HTTPAbortReturn(c, http.StatusOK, gin.H{HTTPReturnCode: merr.Code(paramErr), HTTPReturnMessage: paramErr.Error()})
return nil, paramErr
}
if err := h.checkImportJobAuth(ctx, c, dbName, jobIDGetter.GetJobID()); err != nil {
return nil, err
}
req := &datapb.AbortImportRequest{
Base: commonpbutil.NewMsgBase(),
JobId: jobID,
@@ -4344,58 +4344,210 @@ func TestSearchByPK(t *testing.T) {
func TestCommitImportJob(t *testing.T) {
paramtable.Init()
mp := mocks.NewMockProxy(t)
mp.EXPECT().CommitImport(mock.Anything, mock.Anything).Return(commonSuccessStatus, nil).Once()
testEngine := initHTTPServerV2(mp, false)
queryTestCases := []requestBodyTestCase{}
queryTestCases = append(queryTestCases, requestBodyTestCase{
path: versionalV2(ImportJobCategory, CommitAction),
requestBody: []byte(`{"jobId": "123"}`),
})
queryTestCases = append(queryTestCases, requestBodyTestCase{
path: versionalV2(ImportJobCategory, CommitAction),
requestBody: []byte(`{"jobId": "not-a-number"}`),
errCode: 1100, // ErrParameterInvalid
})
for _, testcase := range queryTestCases {
t.Run(testcase.path, func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, testcase.path, bytes.NewReader(testcase.requestBody))
w := httptest.NewRecorder()
testEngine.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
returnBody := &ReturnErrMsg{}
err := json.Unmarshal(w.Body.Bytes(), returnBody)
assert.Nil(t, err)
assert.Equal(t, testcase.errCode, returnBody.Code)
t.Run("authorization disabled", func(t *testing.T) {
paramtable.Get().Save(proxy.Params.CommonCfg.AuthorizationEnabled.Key, "false")
defer paramtable.Get().Reset(proxy.Params.CommonCfg.AuthorizationEnabled.Key)
mp := mocks.NewMockProxy(t)
mp.EXPECT().CommitImport(mock.Anything, mock.Anything).Return(commonSuccessStatus, nil).Once()
testEngine := initHTTPServerV2(mp, false)
queryTestCases := []requestBodyTestCase{}
queryTestCases = append(queryTestCases, requestBodyTestCase{
path: versionalV2(ImportJobCategory, CommitAction),
requestBody: []byte(`{"jobId": "123"}`),
})
}
queryTestCases = append(queryTestCases, requestBodyTestCase{
path: versionalV2(ImportJobCategory, CommitAction),
requestBody: []byte(`{"jobId": "not-a-number"}`),
errCode: 1100, // ErrParameterInvalid
})
for _, testcase := range queryTestCases {
t.Run(testcase.path, func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, testcase.path, bytes.NewReader(testcase.requestBody))
w := httptest.NewRecorder()
testEngine.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
returnBody := &ReturnErrMsg{}
err := json.Unmarshal(w.Body.Bytes(), returnBody)
assert.Nil(t, err)
assert.Equal(t, testcase.errCode, returnBody.Code)
})
}
})
t.Run("authorization enabled denies user without import privilege", func(t *testing.T) {
paramtable.Get().Save(proxy.Params.CommonCfg.AuthorizationEnabled.Key, "true")
defer paramtable.Get().Reset(proxy.Params.CommonCfg.AuthorizationEnabled.Key)
mp := mocks.NewMockProxy(t)
mp.EXPECT().GetImportProgress(mock.Anything, mock.Anything).Return(&internalpb.GetImportProgressResponse{
Status: &StatusSuccess,
CollectionName: DefaultCollectionName,
}, nil).Once()
testEngine := initHTTPServerV2(mp, true)
req := httptest.NewRequest(http.MethodPost, versionalV2(ImportJobCategory, CommitAction), bytes.NewReader([]byte(`{"jobId": "123"}`)))
req.SetBasicAuth("test", "test")
w := httptest.NewRecorder()
testEngine.ServeHTTP(w, req)
assert.Equal(t, http.StatusForbidden, w.Code)
returnBody := &ReturnErrMsg{}
err := json.Unmarshal(w.Body.Bytes(), returnBody)
assert.Nil(t, err)
assert.Equal(t, int32(65535), returnBody.Code)
assert.Contains(t, returnBody.Message, "PrivilegeImport: permission deny to test")
mp.AssertNotCalled(t, "CommitImport", mock.Anything, mock.Anything)
})
t.Run("authorization enabled returns get progress error", func(t *testing.T) {
paramtable.Get().Save(proxy.Params.CommonCfg.AuthorizationEnabled.Key, "true")
defer paramtable.Get().Reset(proxy.Params.CommonCfg.AuthorizationEnabled.Key)
mp := mocks.NewMockProxy(t)
mp.EXPECT().GetImportProgress(mock.Anything, mock.Anything).Return(nil, merr.ErrImportFailed).Once()
testEngine := initHTTPServerV2(mp, true)
req := httptest.NewRequest(http.MethodPost, versionalV2(ImportJobCategory, CommitAction), bytes.NewReader([]byte(`{"jobId": "123"}`)))
req.SetBasicAuth(util.UserRoot, getDefaultRootPassword())
w := httptest.NewRecorder()
testEngine.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
returnBody := &ReturnErrMsg{}
err := json.Unmarshal(w.Body.Bytes(), returnBody)
assert.Nil(t, err)
assert.Equal(t, merr.Code(merr.ErrImportFailed), returnBody.Code)
mp.AssertNotCalled(t, "CommitImport", mock.Anything, mock.Anything)
})
t.Run("authorization enabled allows root", func(t *testing.T) {
paramtable.Get().Save(proxy.Params.CommonCfg.AuthorizationEnabled.Key, "true")
defer paramtable.Get().Reset(proxy.Params.CommonCfg.AuthorizationEnabled.Key)
mp := mocks.NewMockProxy(t)
mp.EXPECT().GetImportProgress(mock.Anything, mock.Anything).Return(&internalpb.GetImportProgressResponse{
Status: &StatusSuccess,
CollectionName: DefaultCollectionName,
}, nil).Once()
mp.EXPECT().CommitImport(mock.Anything, mock.Anything).Return(commonSuccessStatus, nil).Once()
testEngine := initHTTPServerV2(mp, true)
req := httptest.NewRequest(http.MethodPost, versionalV2(ImportJobCategory, CommitAction), bytes.NewReader([]byte(`{"jobId": "123"}`)))
req.SetBasicAuth(util.UserRoot, getDefaultRootPassword())
w := httptest.NewRecorder()
testEngine.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
returnBody := &ReturnErrMsg{}
err := json.Unmarshal(w.Body.Bytes(), returnBody)
assert.Nil(t, err)
assert.Equal(t, int32(0), returnBody.Code)
})
}
func TestAbortImportJob(t *testing.T) {
paramtable.Init()
mp := mocks.NewMockProxy(t)
mp.EXPECT().AbortImport(mock.Anything, mock.Anything).Return(commonSuccessStatus, nil).Once()
testEngine := initHTTPServerV2(mp, false)
queryTestCases := []requestBodyTestCase{}
queryTestCases = append(queryTestCases, requestBodyTestCase{
path: versionalV2(ImportJobCategory, AbortAction),
requestBody: []byte(`{"jobId": "123"}`),
})
queryTestCases = append(queryTestCases, requestBodyTestCase{
path: versionalV2(ImportJobCategory, AbortAction),
requestBody: []byte(`{"jobId": "not-a-number"}`),
errCode: 1100, // ErrParameterInvalid
})
for _, testcase := range queryTestCases {
t.Run(testcase.path, func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, testcase.path, bytes.NewReader(testcase.requestBody))
w := httptest.NewRecorder()
testEngine.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
returnBody := &ReturnErrMsg{}
err := json.Unmarshal(w.Body.Bytes(), returnBody)
assert.Nil(t, err)
assert.Equal(t, testcase.errCode, returnBody.Code)
t.Run("authorization disabled", func(t *testing.T) {
paramtable.Get().Save(proxy.Params.CommonCfg.AuthorizationEnabled.Key, "false")
defer paramtable.Get().Reset(proxy.Params.CommonCfg.AuthorizationEnabled.Key)
mp := mocks.NewMockProxy(t)
mp.EXPECT().AbortImport(mock.Anything, mock.Anything).Return(commonSuccessStatus, nil).Once()
testEngine := initHTTPServerV2(mp, false)
queryTestCases := []requestBodyTestCase{}
queryTestCases = append(queryTestCases, requestBodyTestCase{
path: versionalV2(ImportJobCategory, AbortAction),
requestBody: []byte(`{"jobId": "123"}`),
})
}
queryTestCases = append(queryTestCases, requestBodyTestCase{
path: versionalV2(ImportJobCategory, AbortAction),
requestBody: []byte(`{"jobId": "not-a-number"}`),
errCode: 1100, // ErrParameterInvalid
})
for _, testcase := range queryTestCases {
t.Run(testcase.path, func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, testcase.path, bytes.NewReader(testcase.requestBody))
w := httptest.NewRecorder()
testEngine.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
returnBody := &ReturnErrMsg{}
err := json.Unmarshal(w.Body.Bytes(), returnBody)
assert.Nil(t, err)
assert.Equal(t, testcase.errCode, returnBody.Code)
})
}
})
t.Run("authorization enabled denies user without import privilege", func(t *testing.T) {
paramtable.Get().Save(proxy.Params.CommonCfg.AuthorizationEnabled.Key, "true")
defer paramtable.Get().Reset(proxy.Params.CommonCfg.AuthorizationEnabled.Key)
mp := mocks.NewMockProxy(t)
mp.EXPECT().GetImportProgress(mock.Anything, mock.Anything).Return(&internalpb.GetImportProgressResponse{
Status: &StatusSuccess,
CollectionName: DefaultCollectionName,
}, nil).Once()
testEngine := initHTTPServerV2(mp, true)
req := httptest.NewRequest(http.MethodPost, versionalV2(ImportJobCategory, AbortAction), bytes.NewReader([]byte(`{"jobId": "123"}`)))
req.SetBasicAuth("test", "test")
w := httptest.NewRecorder()
testEngine.ServeHTTP(w, req)
assert.Equal(t, http.StatusForbidden, w.Code)
returnBody := &ReturnErrMsg{}
err := json.Unmarshal(w.Body.Bytes(), returnBody)
assert.Nil(t, err)
assert.Equal(t, int32(65535), returnBody.Code)
assert.Contains(t, returnBody.Message, "PrivilegeImport: permission deny to test")
mp.AssertNotCalled(t, "AbortImport", mock.Anything, mock.Anything)
})
t.Run("authorization enabled returns get progress error", func(t *testing.T) {
paramtable.Get().Save(proxy.Params.CommonCfg.AuthorizationEnabled.Key, "true")
defer paramtable.Get().Reset(proxy.Params.CommonCfg.AuthorizationEnabled.Key)
mp := mocks.NewMockProxy(t)
mp.EXPECT().GetImportProgress(mock.Anything, mock.Anything).Return(nil, merr.ErrImportFailed).Once()
testEngine := initHTTPServerV2(mp, true)
req := httptest.NewRequest(http.MethodPost, versionalV2(ImportJobCategory, AbortAction), bytes.NewReader([]byte(`{"jobId": "123"}`)))
req.SetBasicAuth(util.UserRoot, getDefaultRootPassword())
w := httptest.NewRecorder()
testEngine.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
returnBody := &ReturnErrMsg{}
err := json.Unmarshal(w.Body.Bytes(), returnBody)
assert.Nil(t, err)
assert.Equal(t, merr.Code(merr.ErrImportFailed), returnBody.Code)
mp.AssertNotCalled(t, "AbortImport", mock.Anything, mock.Anything)
})
t.Run("authorization enabled allows root", func(t *testing.T) {
paramtable.Get().Save(proxy.Params.CommonCfg.AuthorizationEnabled.Key, "true")
defer paramtable.Get().Reset(proxy.Params.CommonCfg.AuthorizationEnabled.Key)
mp := mocks.NewMockProxy(t)
mp.EXPECT().GetImportProgress(mock.Anything, mock.Anything).Return(&internalpb.GetImportProgressResponse{
Status: &StatusSuccess,
CollectionName: DefaultCollectionName,
}, nil).Once()
mp.EXPECT().AbortImport(mock.Anything, mock.Anything).Return(commonSuccessStatus, nil).Once()
testEngine := initHTTPServerV2(mp, true)
req := httptest.NewRequest(http.MethodPost, versionalV2(ImportJobCategory, AbortAction), bytes.NewReader([]byte(`{"jobId": "123"}`)))
req.SetBasicAuth(util.UserRoot, getDefaultRootPassword())
w := httptest.NewRecorder()
testEngine.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
returnBody := &ReturnErrMsg{}
err := json.Unmarshal(w.Body.Bytes(), returnBody)
assert.Nil(t, err)
assert.Equal(t, int32(0), returnBody.Code)
})
}