diff --git a/client/bulkwriter/bulk_import.go b/client/bulkwriter/bulk_import.go index fb0226a688..d90605a00b 100644 --- a/client/bulkwriter/bulk_import.go +++ b/client/bulkwriter/bulk_import.go @@ -307,6 +307,128 @@ func GetImportProgress(ctx context.Context, option *GetImportProgressOption) (*G return result, result.CheckStatus() } +type CommitImportOption struct { + URL string `json:"-"` + JobID string `json:"jobId"` + ClusterID string `json:"clusterId,omitempty"` + APIKey string `json:"-"` +} + +func (opt *CommitImportOption) GetRequest() ([]byte, error) { + return json.Marshal(opt) +} + +func (opt *CommitImportOption) WithAPIKey(key string) *CommitImportOption { + opt.APIKey = key + return opt +} + +func NewCommitImportOption(uri string, jobID string) *CommitImportOption { + return &CommitImportOption{ + URL: uri, + JobID: jobID, + } +} + +func NewCloudCommitImportOption(uri string, jobID string, apiKey string, clusterID string) *CommitImportOption { + return &CommitImportOption{ + URL: uri, + JobID: jobID, + APIKey: apiKey, + ClusterID: clusterID, + } +} + +type CommitImportResponse struct { + ResponseBase +} + +// CommitImport is the API wrapper for the restful import commit API. +func CommitImport(ctx context.Context, option *CommitImportOption) (*CommitImportResponse, error) { + url := option.URL + "/v2/vectordb/jobs/import/commit" + + bs, err := option.GetRequest() + if err != nil { + return nil, err + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(bs)) + if err != nil { + return nil, err + } + req.Header.Set("Content-Type", "application/json") + if option.APIKey != "" { + req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", option.APIKey)) + } + + result := &CommitImportResponse{} + if err := doPostRequest(req, result); err != nil { + return nil, err + } + return result, result.CheckStatus() +} + +type AbortImportOption struct { + URL string `json:"-"` + JobID string `json:"jobId"` + ClusterID string `json:"clusterId,omitempty"` + APIKey string `json:"-"` +} + +func (opt *AbortImportOption) GetRequest() ([]byte, error) { + return json.Marshal(opt) +} + +func (opt *AbortImportOption) WithAPIKey(key string) *AbortImportOption { + opt.APIKey = key + return opt +} + +func NewAbortImportOption(uri string, jobID string) *AbortImportOption { + return &AbortImportOption{ + URL: uri, + JobID: jobID, + } +} + +func NewCloudAbortImportOption(uri string, jobID string, apiKey string, clusterID string) *AbortImportOption { + return &AbortImportOption{ + URL: uri, + JobID: jobID, + APIKey: apiKey, + ClusterID: clusterID, + } +} + +type AbortImportResponse struct { + ResponseBase +} + +// AbortImport is the API wrapper for the restful import abort API. +func AbortImport(ctx context.Context, option *AbortImportOption) (*AbortImportResponse, error) { + url := option.URL + "/v2/vectordb/jobs/import/abort" + + bs, err := option.GetRequest() + if err != nil { + return nil, err + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(bs)) + if err != nil { + return nil, err + } + req.Header.Set("Content-Type", "application/json") + if option.APIKey != "" { + req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", option.APIKey)) + } + + result := &AbortImportResponse{} + if err := doPostRequest(req, result); err != nil { + return nil, err + } + return result, result.CheckStatus() +} + func doPostRequest(req *http.Request, response any) error { client := &http.Client{} resp, err := client.Do(req) diff --git a/client/bulkwriter/bulk_import_test.go b/client/bulkwriter/bulk_import_test.go index 7f3fb26d4f..b3cc855752 100644 --- a/client/bulkwriter/bulk_import_test.go +++ b/client/bulkwriter/bulk_import_test.go @@ -163,6 +163,76 @@ func (s *BulkImportSuite) TestGetImportProgress() { }) } +func (s *BulkImportSuite) TestCommitImport() { + s.Run("normal_case", func() { + svr := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, req *http.Request) { + authHeader := req.Header.Get("Authorization") + s.Equal("Bearer root:Milvus", authHeader) + s.True(strings.Contains(req.URL.Path, "/v2/vectordb/jobs/import/commit")) + rw.Write([]byte(`{"status":0, "data":{}}`)) + })) + defer svr.Close() + + resp, err := CommitImport(context.Background(), + NewCommitImportOption(svr.URL, "123").WithAPIKey("root:Milvus")) + s.NoError(err) + s.EqualValues(0, resp.Status) + }) + + s.Run("status_error", func() { + svr := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, req *http.Request) { + s.True(strings.Contains(req.URL.Path, "/v2/vectordb/jobs/import/commit")) + rw.Write([]byte(`{"status":1100, "message": "commit failed"}`)) + })) + defer svr.Close() + + _, err := CommitImport(context.Background(), NewCommitImportOption(svr.URL, "123")) + s.Error(err) + }) + + s.Run("server_closed", func() { + svr := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, req *http.Request) {})) + svr.Close() + _, err := CommitImport(context.Background(), NewCommitImportOption(svr.URL, "123")) + s.Error(err) + }) +} + +func (s *BulkImportSuite) TestAbortImport() { + s.Run("normal_case", func() { + svr := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, req *http.Request) { + authHeader := req.Header.Get("Authorization") + s.Equal("Bearer root:Milvus", authHeader) + s.True(strings.Contains(req.URL.Path, "/v2/vectordb/jobs/import/abort")) + rw.Write([]byte(`{"status":0, "data":{}}`)) + })) + defer svr.Close() + + resp, err := AbortImport(context.Background(), + NewAbortImportOption(svr.URL, "123").WithAPIKey("root:Milvus")) + s.NoError(err) + s.EqualValues(0, resp.Status) + }) + + s.Run("status_error", func() { + svr := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, req *http.Request) { + s.True(strings.Contains(req.URL.Path, "/v2/vectordb/jobs/import/abort")) + rw.Write([]byte(`{"status":1100, "message": "abort failed"}`)) + })) + defer svr.Close() + + _, err := AbortImport(context.Background(), NewAbortImportOption(svr.URL, "123")) + s.Error(err) + }) + + s.Run("server_closed", func() { + svr := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, req *http.Request) {})) + svr.Close() + _, err := AbortImport(context.Background(), NewAbortImportOption(svr.URL, "123")) + s.Error(err) + }) +} + func TestBulkImportAPIs(t *testing.T) { suite.Run(t, new(BulkImportSuite)) }