diff --git a/cmd/ragflow_server.go b/cmd/ragflow_server.go index 086bddd503..cf0ec173b2 100644 --- a/cmd/ragflow_server.go +++ b/cmd/ragflow_server.go @@ -367,7 +367,7 @@ func main() { } defer redis.Close() - if err = storage.Init(); err != nil { + if err = storage.Init(ctx); err != nil { common.Error("Failed to initialize storage factory", err) } defer storage.CloseStorage() diff --git a/internal/admin/handler.go b/internal/admin/handler.go index 053286f452..ed3ff8d87b 100644 --- a/internal/admin/handler.go +++ b/internal/admin/handler.go @@ -1226,7 +1226,8 @@ func (h *Handler) PingStore(c *gin.Context) { return } - if storageImpl.Health() { + ctx := c.Request.Context() + if storageImpl.Health(ctx) { common.SuccessNoMessage(c, "SUCCESS") } else { common.ErrorWithCode(c, common.CodeServerError, "storage health check failed") diff --git a/internal/admin/service.go b/internal/admin/service.go index c5d6a00152..2179814d73 100644 --- a/internal/admin/service.go +++ b/internal/admin/service.go @@ -1006,7 +1006,7 @@ func (s *Service) ListServices(ctx context.Context) ([]ServiceStatus, error) { // storage engine storageImpl := storage.GetStorageFactory().GetStorage() - storageHealth := storageImpl.Health() + storageHealth := storageImpl.Health(ctx) if storageHealth { results = append(results, newServiceStatus("storage_engine", storageImpl.Type(), "alive", time.Now(), "")) } else { diff --git a/internal/agent/component/docs_generator.go b/internal/agent/component/docs_generator.go index 5ad6d9762a..ae75bd39bb 100644 --- a/internal/agent/component/docs_generator.go +++ b/internal/agent/component/docs_generator.go @@ -388,7 +388,7 @@ func storeAgentAttachment(ctx context.Context, docID string, payload []byte) (bo return false, fmt.Errorf("storage not initialized") } bucket := fmt.Sprintf("%s-downloads", tenantID) - if err := storageImpl.Put(bucket, docID, payload); err != nil { + if err = storageImpl.Put(ctx, bucket, docID, payload); err != nil { return false, err } return true, nil diff --git a/internal/agent/component/docs_generator_test.go b/internal/agent/component/docs_generator_test.go index f1320e886a..6e2443be40 100644 --- a/internal/agent/component/docs_generator_test.go +++ b/internal/agent/component/docs_generator_test.go @@ -192,7 +192,7 @@ func TestDocsGenerator_Invoke_StoresAgentAttachment(t *testing.T) { } state := canvas.NewCanvasState("run-1", "task-1") state.Sys["tenant_id"] = "tenant-1" - ctx := canvas.WithState(context.Background(), state) + ctx := canvas.WithState(t.Context(), state) out, err := c.Invoke(ctx, nil, map[string]any{}) if err != nil { @@ -205,7 +205,7 @@ func TestDocsGenerator_Invoke_StoresAgentAttachment(t *testing.T) { if !ok || docID == "" { t.Fatalf("doc_id = %v, want non-empty string", out["doc_id"]) } - blob, err := memStorage.Get("tenant-1-downloads", docID) + blob, err := memStorage.Get(ctx, "tenant-1-downloads", docID) if err != nil { t.Fatalf("stored blob missing: %v", err) } diff --git a/internal/handler/document.go b/internal/handler/document.go index 98d2a1af3d..b94797da60 100644 --- a/internal/handler/document.go +++ b/internal/handler/document.go @@ -75,7 +75,7 @@ type documentServiceIface interface { UploadEmptyDocument(ctx context.Context, kb *entity.Knowledgebase, tenantID, name string) (map[string]interface{}, common.ErrorCode, error) DownloadDocument(ctx context.Context, datasetID, docID string) (*document.DownloadDocumentResp, error) UpdateDatasetDocument(ctx context.Context, userID, datasetID, documentID string, req *document.UpdateDatasetDocumentRequest, present map[string]bool) (*document.UpdateDatasetDocumentResponse, common.ErrorCode, error) - BatchUpdateDocumentMetadatas(ctx context.Context, datasetID string, selector *document.DocumentMetadataSelector, updates []document.DocumentMetadataUpdate, deletes []document.DocumentMetadataDelete) (*document.BatchUpdateDocumentMetadatasResponse, common.ErrorCode, error) + BatchUpdateDocumentMetadatas(ctx context.Context, datasetID string, selector *document.MetadataSelector, updates []document.MetadataUpdate, deletes []document.MetadataDelete) (*document.BatchUpdateMetadatasResponse, common.ErrorCode, error) ListIngestionTasks(ctx context.Context, userID string, datasetID *string, page, pageSize int) ([]*entity.IngestionTask, error) IngestDocuments(ctx context.Context, datasetID, userID string, docIDs []string) ([]*service.ParseDocumentResponse, error) StopIngestionTasks(ctx context.Context, tasks []string, userID string) ([]*entity.IngestionTask, error) @@ -1803,9 +1803,9 @@ func (h *DocumentHandler) UploadInfo(c *gin.Context) { } type documentMetadataBatchRequest struct { - Selector *document.DocumentMetadataSelector `json:"selector"` - Updates []document.DocumentMetadataUpdate `json:"updates"` - Deletes []document.DocumentMetadataDelete `json:"deletes"` + Selector *document.MetadataSelector `json:"selector"` + Updates []document.MetadataUpdate `json:"updates"` + Deletes []document.MetadataDelete `json:"deletes"` } func (h *DocumentHandler) MetadataBatchUpdate(c *gin.Context) { @@ -1887,15 +1887,15 @@ func inferJSONType(err error) string { return "unknown" } -func parseMetadataSelector(raw interface{}) (*document.DocumentMetadataSelector, string) { +func parseMetadataSelector(raw interface{}) (*document.MetadataSelector, string) { if raw == nil { - return &document.DocumentMetadataSelector{}, "" + return &document.MetadataSelector{}, "" } m, ok := raw.(map[string]interface{}) if !ok { return nil, "selector must be an object." } - selector := &document.DocumentMetadataSelector{} + selector := &document.MetadataSelector{} if v, ok := m["document_ids"]; ok && v != nil { ids, ok := v.([]interface{}) if !ok { @@ -1915,15 +1915,15 @@ func parseMetadataSelector(raw interface{}) (*document.DocumentMetadataSelector, return selector, "" } -func parseMetadataUpdates(raw interface{}) ([]document.DocumentMetadataUpdate, string) { +func parseMetadataUpdates(raw interface{}) ([]document.MetadataUpdate, string) { if raw == nil { - return []document.DocumentMetadataUpdate{}, "" + return []document.MetadataUpdate{}, "" } arr, ok := raw.([]interface{}) if !ok { return nil, "updates and deletes must be lists." } - updates := make([]document.DocumentMetadataUpdate, 0, len(arr)) + updates := make([]document.MetadataUpdate, 0, len(arr)) for _, item := range arr { m, ok := item.(map[string]interface{}) if !ok { @@ -1934,20 +1934,20 @@ func parseMetadataUpdates(raw interface{}) ([]document.DocumentMetadataUpdate, s return nil, "Each update requires key and value." } value := m["value"] - updates = append(updates, document.DocumentMetadataUpdate{Key: key, Value: value}) + updates = append(updates, document.MetadataUpdate{Key: key, Value: value}) } return updates, "" } -func parseMetadataDeletes(raw interface{}) ([]document.DocumentMetadataDelete, string) { +func parseMetadataDeletes(raw interface{}) ([]document.MetadataDelete, string) { if raw == nil { - return []document.DocumentMetadataDelete{}, "" + return []document.MetadataDelete{}, "" } arr, ok := raw.([]interface{}) if !ok { return nil, "updates and deletes must be lists." } - deletes := make([]document.DocumentMetadataDelete, 0, len(arr)) + deletes := make([]document.MetadataDelete, 0, len(arr)) for _, item := range arr { m, ok := item.(map[string]interface{}) if !ok { @@ -1957,7 +1957,7 @@ func parseMetadataDeletes(raw interface{}) ([]document.DocumentMetadataDelete, s if key == "" { return nil, "Each delete requires key." } - deletes = append(deletes, document.DocumentMetadataDelete{Key: key}) + deletes = append(deletes, document.MetadataDelete{Key: key}) } return deletes, "" } diff --git a/internal/handler/document_test.go b/internal/handler/document_test.go index fe602412d0..f6d06ca243 100644 --- a/internal/handler/document_test.go +++ b/internal/handler/document_test.go @@ -97,7 +97,7 @@ const uploadTestDatasetID = "123e4567-e89b-12d3-a456-426614174000" func (f *fakeDocumentService) UpdateDatasetDocument(ctx context.Context, userID, datasetID, documentID string, req *document.UpdateDatasetDocumentRequest, present map[string]bool) (*document.UpdateDatasetDocumentResponse, common.ErrorCode, error) { return nil, common.CodeSuccess, nil } -func (f *fakeDocumentService) BatchUpdateDocumentMetadatas(ctx context.Context, datasetID string, selector *document.DocumentMetadataSelector, updates []document.DocumentMetadataUpdate, deletes []document.DocumentMetadataDelete) (*document.BatchUpdateDocumentMetadatasResponse, common.ErrorCode, error) { +func (f *fakeDocumentService) BatchUpdateDocumentMetadatas(ctx context.Context, datasetID string, selector *document.MetadataSelector, updates []document.MetadataUpdate, deletes []document.MetadataDelete) (*document.BatchUpdateMetadatasResponse, common.ErrorCode, error) { return nil, common.CodeSuccess, nil } diff --git a/internal/handler/file.go b/internal/handler/file.go index c16618517e..71757d7859 100644 --- a/internal/handler/file.go +++ b/internal/handler/file.go @@ -486,7 +486,7 @@ func (h *FileHandler) Download(c *gin.Context) { var blob []byte var getErr error if file.Location != nil && *file.Location != "" { - blob, getErr = storageImpl.Get(file.ParentID, *file.Location) + blob, getErr = storageImpl.Get(ctx, file.ParentID, *file.Location) } // If blob is empty, try fallback via file2document @@ -497,7 +497,7 @@ func (h *FileHandler) Download(c *gin.Context) { common.ResponseWithCodeData(c, common.CodeServerError, nil, "Failed to get file storage address: "+err.Error()) return } - blob, getErr = storageImpl.Get(storageAddr.Bucket, storageAddr.Name) + blob, getErr = storageImpl.Get(ctx, storageAddr.Bucket, storageAddr.Name) } // Check if we got valid data diff --git a/internal/ingestion/component/document_storage.go b/internal/ingestion/component/document_storage.go index 2b9789aa14..844acddcd4 100644 --- a/internal/ingestion/component/document_storage.go +++ b/internal/ingestion/component/document_storage.go @@ -45,8 +45,8 @@ var ResolveDocumentStorageOverride func(docID string) (*DocumentStorageRef, erro // the source PDF to crop section images on demand) can reuse the same // storage resolution the Parser uses. func FetchBinary(ctx context.Context, bucket, path string) ([]byte, error) { - stg := resolveStorage() - if stg == nil { + storageImpl := resolveStorage() + if storageImpl == nil { return nil, fmt.Errorf("no storage backend registered") } @@ -56,7 +56,7 @@ func FetchBinary(ctx context.Context, bucket, path string) ([]byte, error) { } done := make(chan result, 1) go func() { - data, err := stg.Get(bucket, path) + data, err := storageImpl.Get(ctx, bucket, path) done <- result{data: data, err: err} }() select { diff --git a/internal/ingestion/component/extractor_tag.go b/internal/ingestion/component/extractor_tag.go index 69276dc16a..9a2b0590a7 100644 --- a/internal/ingestion/component/extractor_tag.go +++ b/internal/ingestion/component/extractor_tag.go @@ -245,7 +245,7 @@ func (c *ExtractorComponent) loadTagFileIndexed(ctx context.Context) (*indexedTa return nil, false } tenantID := globals.GlobalOrInput(ctx, nil, "tenant_id", "") - data, err := stg.Get(f.ParentID, *f.Location, tenantID) + data, err := stg.Get(ctx, f.ParentID, *f.Location, tenantID) if err != nil { common.Warn(fmt.Sprintf("extractor tags: load tag source %q/%q: %v", f.ParentID, *f.Location, err)) return nil, false diff --git a/internal/ingestion/component/file_test.go b/internal/ingestion/component/file_test.go index 2de65cba67..043fa18063 100644 --- a/internal/ingestion/component/file_test.go +++ b/internal/ingestion/component/file_test.go @@ -83,7 +83,8 @@ func TestFileComponent_Registered(t *testing.T) { // even when explicit storage overrides are supplied. func TestFileComponent_Invoke_HappyPath(t *testing.T) { ms := withMemoryStorage(t) - if err := ms.Put("bucketA", "path/to/file.txt", []byte("hello, ragflow")); err != nil { + ctx := t.Context() + if err := ms.Put(ctx, "bucketA", "path/to/file.txt", []byte("hello, ragflow")); err != nil { t.Fatalf("seed: %v", err) } @@ -113,9 +114,10 @@ func TestFileComponent_Invoke_HappyPath(t *testing.T) { func TestFileComponent_Invoke_ResolvesDocIDViaDocumentLocation(t *testing.T) { ms := withMemoryStorage(t) + ctx := t.Context() db := withFileComponentTestDB(t) location := "docs/from-document.bin" - if err := ms.Put("kb-doc", location, []byte("doc-location")); err != nil { + if err := ms.Put(ctx, "kb-doc", location, []byte("doc-location")); err != nil { t.Fatalf("seed storage: %v", err) } docName := "report.pdf" @@ -151,9 +153,10 @@ func TestFileComponent_Invoke_ResolvesDocIDViaDocumentLocation(t *testing.T) { func TestFileComponent_Invoke_ResolvesDocIDViaFileMapping(t *testing.T) { ms := withMemoryStorage(t) + ctx := t.Context() db := withFileComponentTestDB(t) location := "tenant-root/from-file.bin" - if err := ms.Put("folder-1", location, []byte("file-mapping")); err != nil { + if err := ms.Put(ctx, "folder-1", location, []byte("file-mapping")); err != nil { t.Fatalf("seed storage: %v", err) } docName := "deck.pptx" @@ -249,8 +252,9 @@ func TestFileComponent_Invoke_MissingDoc(t *testing.T) { // bucket/path overrides still flow through for downstream Parser use. func TestFileComponent_Invoke_IncludesCheckpointPath(t *testing.T) { ms := withMemoryStorage(t) + ctx := t.Context() const wantPath = "checkpoint/expected/path.bin" - if err := ms.Put("b", wantPath, []byte("x")); err != nil { + if err := ms.Put(ctx, "b", wantPath, []byte("x")); err != nil { t.Fatalf("seed: %v", err) } c := &FileComponent{} diff --git a/internal/ingestion/component/image_uploader.go b/internal/ingestion/component/image_uploader.go index ef9cb7c97a..b1be30fca3 100644 --- a/internal/ingestion/component/image_uploader.go +++ b/internal/ingestion/component/image_uploader.go @@ -32,12 +32,12 @@ type ImageUploader func(ctx context.Context, kbID, chunkID string, data []byte) // returns img_id "-". Mirrors Python image2id's storage_put + // f"{bucket}-{objname}" (rag/utils/base64_image.py:80-82), minus re-encoding: // the bytes are stored in whatever format the chunker produced. -func DefaultImageUploader(_ context.Context, kbID, chunkID string, data []byte) (string, error) { - stg := resolveStorage() - if stg == nil { +func DefaultImageUploader(ctx context.Context, kbID, chunkID string, data []byte) (string, error) { + storageImpl := resolveStorage() + if storageImpl == nil { return "", fmt.Errorf("no storage backend registered") } - if err := stg.Put(kbID, chunkID, data); err != nil { + if err := storageImpl.Put(ctx, kbID, chunkID, data); err != nil { return "", fmt.Errorf("store chunk image (%q,%q): %w", kbID, chunkID, err) } return kbID + "-" + chunkID, nil diff --git a/internal/ingestion/component/parser_test.go b/internal/ingestion/component/parser_test.go index 0924b14674..99769388a4 100644 --- a/internal/ingestion/component/parser_test.go +++ b/internal/ingestion/component/parser_test.go @@ -244,7 +244,8 @@ func TestParserComponent_Invoke_ResolvesBinaryFromDocID(t *testing.T) { ms := withMemoryStorage(t) db := withFileComponentTestDB(t) location := "docs/from-parser.txt" - if err := ms.Put("kb-parser", location, []byte("alpha\fbeta")); err != nil { + ctx := t.Context() + if err := ms.Put(ctx, "kb-parser", location, []byte("alpha\fbeta")); err != nil { t.Fatalf("seed storage: %v", err) } docName := "parser.txt" @@ -263,7 +264,7 @@ func TestParserComponent_Invoke_ResolvesBinaryFromDocID(t *testing.T) { } c := &ParserComponent{Param: schema.ParserParam{}.Defaults()} - out, err := c.Invoke(context.Background(), db, map[string]any{"doc_id": "doc-parser"}) + out, err := c.Invoke(ctx, db, map[string]any{"doc_id": "doc-parser"}) if err != nil { t.Fatalf("Invoke: %v", err) } @@ -281,12 +282,13 @@ func TestParserComponent_Invoke_ResolvesBinaryFromDocID(t *testing.T) { func TestParserComponent_Invoke_ResolvesBinaryFromBucketPath(t *testing.T) { ms := withMemoryStorage(t) - if err := ms.Put("bucket-1", "docs/explicit.txt", []byte("bucket content")); err != nil { + ctx := t.Context() + if err := ms.Put(ctx, "bucket-1", "docs/explicit.txt", []byte("bucket content")); err != nil { t.Fatalf("seed storage: %v", err) } c := &ParserComponent{Param: schema.ParserParam{}.Defaults()} - out, err := c.Invoke(context.Background(), nil, map[string]any{ + out, err := c.Invoke(ctx, nil, map[string]any{ "bucket": "bucket-1", "path": "docs/explicit.txt", }) diff --git a/internal/ingestion/task/pipeline_real_integration_test.go b/internal/ingestion/task/pipeline_real_integration_test.go index 0042ff90f4..6e68d68e20 100644 --- a/internal/ingestion/task/pipeline_real_integration_test.go +++ b/internal/ingestion/task/pipeline_real_integration_test.go @@ -32,6 +32,7 @@ import ( func TestPipelineExecutor_Run_RealCanvasDSL_UsesGeneralPipeline(t *testing.T) { requireTokenizerPool(t) + ctx := t.Context() mustLoadTaskTestConfig(t) origDB := dao.DB @@ -55,7 +56,7 @@ func TestPipelineExecutor_Run_RealCanvasDSL_UsesGeneralPipeline(t *testing.T) { } templateBytes = disableTokenizerEmbeddingForTaskTemplate(t, templateBytes) var templateDSL entity.JSONMap - if err := json.Unmarshal(templateBytes, &templateDSL); err != nil { + if err = json.Unmarshal(templateBytes, &templateDSL); err != nil { t.Fatalf("unmarshal template dsl: %v", err) } @@ -71,10 +72,10 @@ func TestPipelineExecutor_Run_RealCanvasDSL_UsesGeneralPipeline(t *testing.T) { content := "Alpha paragraph\n\nBeta paragraph." mustSeedTaskRealPipelineDocument(t, realDB, realStorage, tenantID, kbID, docID, fileID, bucket, objectPath, docName, content) - if err := realDB.Model(&entity.Document{}).Where("id = ?", docID).Update("pipeline_id", canvasID).Error; err != nil { + if err = realDB.Model(&entity.Document{}).Where("id = ?", docID).Update("pipeline_id", canvasID).Error; err != nil { t.Fatalf("set document pipeline_id: %v", err) } - if err := realDB.Create(&entity.UserCanvas{ + if err = realDB.Create(&entity.UserCanvas{ ID: canvasID, UserID: tenantID, Permission: "me", @@ -84,8 +85,9 @@ func TestPipelineExecutor_Run_RealCanvasDSL_UsesGeneralPipeline(t *testing.T) { t.Fatalf("create user canvas: %v", err) } t.Cleanup(func() { - _ = realDB.Where("id = ?", canvasID).Delete(&entity.UserCanvas{}).Error - cleanupTaskRealPipelineDocument(realDB, realStorage, tenantID, kbID, docID, fileID, bucket, objectPath) + cleanUpCtx := context.Background() + _ = realDB.WithContext(cleanUpCtx).Where("id = ?", canvasID).Delete(&entity.UserCanvas{}).Error + cleanupTaskRealPipelineDocument(cleanUpCtx, realDB, realStorage, tenantID, kbID, docID, fileID, bucket, objectPath) }) taskCtx := &TaskContext{ @@ -115,7 +117,7 @@ func TestPipelineExecutor_Run_RealCanvasDSL_UsesGeneralPipeline(t *testing.T) { return nil, nil }) - if _, err := svc.Execute(context.Background()); err != nil { + if _, err = svc.Execute(ctx); err != nil { t.Fatalf("Run: %v", err) } if len(inserted) != 1 { @@ -136,6 +138,7 @@ func TestPipelineExecutor_Run_RealCanvasDSL_UsesGeneralPipeline(t *testing.T) { func TestPipelineExecutor_Run_RealPDF_ProducesIndexedChunks(t *testing.T) { requireTokenizerPool(t) + ctx := t.Context() // Loads service config (server.Init side effect) without requiring any // external MySQL/MinIO/ES. The pipeline runs against an in-memory sqlite @@ -198,7 +201,7 @@ func TestPipelineExecutor_Run_RealPDF_ProducesIndexedChunks(t *testing.T) { } t.Cleanup(func() { _ = realDB.Where("id = ?", canvasID).Delete(&entity.UserCanvas{}).Error - cleanupTaskRealPipelineDocument(realDB, realStorage, tenantID, kbID, docID, fileID, bucket, objectPath) + cleanupTaskRealPipelineDocument(ctx, realDB, realStorage, tenantID, kbID, docID, fileID, bucket, objectPath) }) taskCtx := &TaskContext{ @@ -229,7 +232,6 @@ func TestPipelineExecutor_Run_RealPDF_ProducesIndexedChunks(t *testing.T) { }). WithLogCreateFunc(func(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error { return nil }) - ctx := t.Context() if _, err = svc.Execute(ctx); err != nil { t.Fatalf("Run: %v", err) } @@ -269,6 +271,7 @@ func TestPipelineExecutor_Run_RealPDF_ProducesIndexedChunks(t *testing.T) { func TestRunPipeline_RealPipelineOutput_ProducesIndexFields(t *testing.T) { requireTokenizerPool(t) + ctx := t.Context() mustLoadTaskTestConfig(t) origDB := dao.DB @@ -309,10 +312,9 @@ func TestRunPipeline_RealPipelineOutput_ProducesIndexFields(t *testing.T) { mustSeedTaskRealPipelineDocument(t, realDB, realStorage, tenantID, kbID, docID, fileID, bucket, objectPath, docName, content) t.Cleanup(func() { - cleanupTaskRealPipelineDocument(realDB, realStorage, tenantID, kbID, docID, fileID, bucket, objectPath) + cleanupTaskRealPipelineDocument(ctx, realDB, realStorage, tenantID, kbID, docID, fileID, bucket, objectPath) }) - ctx := t.Context() pipelineOut, err := pipe.Run(ctx, map[string]any{ "doc_id": docID, }, nil) @@ -618,7 +620,8 @@ func mustSeedTaskRealPipelineDocumentBytes( }).Error; err != nil { t.Fatalf("create kb: %v", err) } - if err := stg.Put(bucket, objectPath, content); err != nil { + ctx := t.Context() + if err := stg.Put(ctx, bucket, objectPath, content); err != nil { t.Fatalf("put real minio object: %v", err) } if err := db.Create(&entity.File{ @@ -657,13 +660,13 @@ func mustSeedTaskRealPipelineDocumentBytes( } } -func cleanupTaskRealPipelineDocument(db *gorm.DB, stg storage.Storage, tenantID, kbID, docID, fileID, bucket, objectPath string) { - _ = db.Where("document_id = ?", docID).Delete(&entity.File2Document{}).Error - _ = db.Where("id = ?", docID).Delete(&entity.Document{}).Error - _ = db.Where("id = ?", fileID).Delete(&entity.File{}).Error - _ = db.Where("id = ?", kbID).Delete(&entity.Knowledgebase{}).Error - _ = db.Where("id = ?", tenantID).Delete(&entity.Tenant{}).Error - _ = stg.Remove(bucket, objectPath) +func cleanupTaskRealPipelineDocument(ctx context.Context, db *gorm.DB, storageImpl storage.Storage, tenantID, kbID, docID, fileID, bucket, objectPath string) { + _ = db.WithContext(ctx).Where("document_id = ?", docID).Delete(&entity.File2Document{}).Error + _ = db.WithContext(ctx).Where("id = ?", docID).Delete(&entity.Document{}).Error + _ = db.WithContext(ctx).Where("id = ?", fileID).Delete(&entity.File{}).Error + _ = db.WithContext(ctx).Where("id = ?", kbID).Delete(&entity.Knowledgebase{}).Error + _ = db.WithContext(ctx).Where("id = ?", tenantID).Delete(&entity.Tenant{}).Error + _ = storageImpl.Remove(ctx, bucket, objectPath) } func deepCopyTaskChunks(in []map[string]any) []map[string]any { diff --git a/internal/service/agent_run_e2e_test.go b/internal/service/agent_run_e2e_test.go index 2bb04e42c9..50cc67b737 100644 --- a/internal/service/agent_run_e2e_test.go +++ b/internal/service/agent_run_e2e_test.go @@ -1937,6 +1937,7 @@ func TestRunAgent_RunTracker_AttachCheckpoint_CallSequence(t *testing.T) { // TestRunAgent_FilesPopulateIteration verifies the full upload object -> // sys.files -> Parallel/Iteration item path. func TestRunAgent_FilesPopulateIteration(t *testing.T) { + ctx := t.Context() testDB := setupServiceTestDB(t) if err := testDB.AutoMigrate( &entity.UserCanvas{}, @@ -1958,7 +1959,7 @@ func TestRunAgent_FilesPopulateIteration(t *testing.T) { sessionID := "session-files-e2e" versionID := "v-files-e2e" memory := storage.NewMemoryStorage() - if err := memory.Put("user-1-downloads", "upload-1", []byte("iteration payload")); err != nil { + if err := memory.Put(ctx, "user-1-downloads", "upload-1", []byte("iteration payload")); err != nil { t.Fatalf("put upload: %v", err) } factory := storage.GetStorageFactory() @@ -2042,7 +2043,7 @@ func TestRunAgent_FilesPopulateIteration(t *testing.T) { svc := NewAgentService() events, err := svc.RunAgent( - context.Background(), + ctx, "user-1", canvasID, sessionID, diff --git a/internal/service/chat_session.go b/internal/service/chat_session.go index b386a5f726..f331659371 100644 --- a/internal/service/chat_session.go +++ b/internal/service/chat_session.go @@ -432,7 +432,7 @@ func (s *ChatSessionService) DeleteSessions(ctx context.Context, userID, chatID continue } - s.removeSessionUploadFiles(userID, session) + s.removeSessionUploadFiles(ctx, userID, session) if err = s.chatSessionDAO.DeleteByID(ctx, dao.DB, sid); err != nil { return nil, "", common.CodeServerError, err @@ -485,7 +485,7 @@ func stringSliceFromValue(value interface{}) ([]string, bool) { return ids, true } -func (s *ChatSessionService) removeSessionUploadFiles(userID string, session *entity.ChatSession) { +func (s *ChatSessionService) removeSessionUploadFiles(ctx context.Context, userID string, session *entity.ChatSession) { messages := parseMessages(session.Message) bucket := fmt.Sprintf("%s-downloads", userID) storageImpl := storage.GetStorageFactory().GetStorage() @@ -511,7 +511,7 @@ func (s *ChatSessionService) removeSessionUploadFiles(userID string, session *en continue } - if err := storageImpl.Remove(bucket, fileID); err != nil { + if err := storageImpl.Remove(ctx, bucket, fileID); err != nil { common.Warn("Failed to delete chat upload blob", zap.String("bucket", bucket), zap.String("file_id", fileID), diff --git a/internal/service/chunk/chunk.go b/internal/service/chunk/chunk.go index 8edd37ae2f..76bd75a71d 100644 --- a/internal/service/chunk/chunk.go +++ b/internal/service/chunk/chunk.go @@ -1366,7 +1366,7 @@ func (s *ChunkService) AddChunk(ctx context.Context, req *service.AddChunkReques if err != nil { return nil, addChunkError{code: common.CodeDataError, message: err.Error()} } - if err := s.storeChunkImage(req.DatasetID, chunkID, imageBinary); err != nil { + if err = s.storeChunkImage(ctx, req.DatasetID, chunkID, imageBinary); err != nil { return nil, addChunkError{code: common.CodeDataError, message: "Failed to store chunk image"} } chunkData["img_id"] = fmt.Sprintf("%s-%s", req.DatasetID, chunkID) @@ -1396,12 +1396,12 @@ func (s *ChunkService) AddChunk(ctx context.Context, req *service.AddChunkReques ctx, cancel := context.WithTimeout(context.Background(), 600*time.Second) defer cancel() - if _, err := s.docEngine.InsertChunks(ctx, []map[string]interface{}{chunkData}, indexName, req.DatasetID); err != nil { + if _, err = s.docEngine.InsertChunks(ctx, []map[string]interface{}{chunkData}, indexName, req.DatasetID); err != nil { return nil, addChunkError{code: common.CodeServerError, message: fmt.Sprintf("insert chunk: %v", err)} } tokenNum := int64(s.numTokens(req.Content)) - if err := s.incrementChunkStats(req.DocumentID, req.DatasetID, tokenNum, 1, 0); err != nil { + if err = s.incrementChunkStats(req.DocumentID, req.DatasetID, tokenNum, 1, 0); err != nil { return nil, addChunkError{code: common.CodeServerError, message: fmt.Sprintf("increment chunk stats: %v", err)} } @@ -1640,7 +1640,7 @@ func (s *ChunkService) decrementChunkStats(docID, kbID string, tokenNum, chunkNu }) } -func (s *ChunkService) storeChunkImage(bucket, chunkID string, imageBinary []byte) error { +func (s *ChunkService) storeChunkImage(ctx context.Context, bucket, chunkID string, imageBinary []byte) error { if s.storeChunkImageFunc != nil { return s.storeChunkImageFunc(bucket, chunkID, imageBinary) } @@ -1656,11 +1656,11 @@ func (s *ChunkService) storeChunkImage(bucket, chunkID string, imageBinary []byt releaseChunkImageMergeLock(lockKey) }() - if !storageImpl.ObjExist(bucket, chunkID) { - return storageImpl.Put(bucket, chunkID, imageBinary) + if !storageImpl.ObjExist(ctx, bucket, chunkID) { + return storageImpl.Put(ctx, bucket, chunkID, imageBinary) } - oldBinary, err := storageImpl.Get(bucket, chunkID) + oldBinary, err := storageImpl.Get(ctx, bucket, chunkID) if err != nil { return err } @@ -1684,10 +1684,10 @@ func (s *ChunkService) storeChunkImage(bucket, chunkID string, imageBinary []byt draw.Draw(combined, image.Rect(0, oldBounds.Dy(), newBounds.Dx(), oldBounds.Dy()+newBounds.Dy()), newImage, newBounds.Min, draw.Src) var buf bytes.Buffer - if err := jpeg.Encode(&buf, combined, nil); err != nil { + if err = jpeg.Encode(&buf, combined, nil); err != nil { return err } - return storageImpl.Put(bucket, chunkID, buf.Bytes()) + return storageImpl.Put(ctx, bucket, chunkID, buf.Bytes()) } func acquireChunkImageMergeLock(key string) *chunkImageMergeLock { diff --git a/internal/service/chunk/chunk_test.go b/internal/service/chunk/chunk_test.go index b2555a6f44..2a54adac26 100644 --- a/internal/service/chunk/chunk_test.go +++ b/internal/service/chunk/chunk_test.go @@ -690,6 +690,7 @@ func TestStoreChunkImageMergesExistingImage(t *testing.T) { exists: true, oldBinary: oldImage, } + ctx := t.Context() factory := storage.GetStorageFactory() originalStorage := factory.GetStorage() @@ -699,7 +700,7 @@ func TestStoreChunkImageMergesExistingImage(t *testing.T) { }) svc := &ChunkService{} - if err := svc.storeChunkImage("kb-1", "chunk-1", newImage); err != nil { + if err := svc.storeChunkImage(ctx, "kb-1", "chunk-1", newImage); err != nil { t.Fatalf("storeChunkImage() error = %v", err) } if mockStorage.putCalls != 1 { @@ -1184,30 +1185,38 @@ type chunkImageStorage struct { putCalls int } -func (s *chunkImageStorage) Type() string { return "chunk_image_storage" } -func (s *chunkImageStorage) Health() bool { return true } -func (s *chunkImageStorage) Put(bucket, fnm string, binary []byte, tenantID ...string) error { +func (s *chunkImageStorage) Type() string { return "chunk_image_storage" } +func (s *chunkImageStorage) Health(_ context.Context) bool { return true } +func (s *chunkImageStorage) Put(ctx context.Context, bucket, fnm string, binary []byte, tenantID ...string) error { s.putCalls++ return nil } -func (s *chunkImageStorage) Get(bucket, fnm string, tenantID ...string) ([]byte, error) { +func (s *chunkImageStorage) Get(ctx context.Context, bucket, fnm string, tenantID ...string) ([]byte, error) { return s.oldBinary, nil } -func (s *chunkImageStorage) Remove(bucket, fnm string, tenantID ...string) error { return nil } -func (s *chunkImageStorage) ObjExist(bucket, fnm string, tenantID ...string) bool { return s.exists } +func (s *chunkImageStorage) Remove(ctx context.Context, bucket, fnm string, tenantID ...string) error { + return nil +} +func (s *chunkImageStorage) ObjExist(ctx context.Context, bucket, fnm string, tenantID ...string) bool { + return s.exists +} // ListObjects lists all objects in a bucket -func (s *chunkImageStorage) ListObjects(bucket string, tenantID ...string) ([]string, error) { +func (s *chunkImageStorage) ListObjects(ctx context.Context, bucket string, tenantID ...string) ([]string, error) { return []string{}, nil } -func (s *chunkImageStorage) GetPresignedURL(bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { +func (s *chunkImageStorage) GetPresignedURL(ctx context.Context, bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { return "", nil } -func (s *chunkImageStorage) BucketExists(bucket string) bool { return true } -func (s *chunkImageStorage) RemoveBucket(bucket string) error { return nil } -func (s *chunkImageStorage) Copy(srcBucket, srcPath, destBucket, destPath string) bool { return false } -func (s *chunkImageStorage) Move(srcBucket, srcPath, destBucket, destPath string) bool { return false } -func (s *chunkImageStorage) Close() error { return nil } +func (s *chunkImageStorage) BucketExists(ctx context.Context, bucket string) bool { return true } +func (s *chunkImageStorage) RemoveBucket(ctx context.Context, bucket string) error { return nil } +func (s *chunkImageStorage) Copy(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { + return false +} +func (s *chunkImageStorage) Move(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { + return false +} +func (s *chunkImageStorage) Close() error { return nil } func mustEncodePNG(t *testing.T, rect image.Rectangle) []byte { t.Helper() diff --git a/internal/service/document/document.go b/internal/service/document/document.go index 206fdab8a3..4df6d19b97 100644 --- a/internal/service/document/document.go +++ b/internal/service/document/document.go @@ -17,8 +17,10 @@ package document import ( + "context" "errors" "ragflow/internal/service" + "ragflow/internal/storage" "regexp" "time" @@ -219,10 +221,10 @@ type IngestDocumentRequest struct { // StartParseOptions controls StartParseDocuments behavior. type StartParseOptions struct { - // ApplyKB merges the knowledgebase's parser_config (llm_id, metadata) + // ApplyKB merges the knowledge base's parser_config (llm_id, metadata) // into the document before parsing. ApplyKB bool - // RerunWithDelete clears prior chunks/tasks/counters before re-parsing. + // RerunWithDelete clears prior chunks/tasks/counters before reparsing. RerunWithDelete bool } @@ -252,7 +254,7 @@ const knowledgebaseFolderName = ".knowledgebase" const maxUploadDocSize = 128 * 1024 * 1024 // MetadataUpdate is one update item: set key to value. -type DocumentMetadataUpdate struct { +type MetadataUpdate struct { Key string `json:"key"` Value interface{} `json:"value"` Match interface{} `json:"match,omitempty"` @@ -260,19 +262,44 @@ type DocumentMetadataUpdate struct { } // MetadataDelete removes a whole key, or a specific value from a list field. -type DocumentMetadataDelete struct { +type MetadataDelete struct { Key string `json:"key"` Value interface{} `json:"value,omitempty"` } // MetadataSelector selects which documents to target. -type DocumentMetadataSelector struct { +type MetadataSelector struct { DocumentIDs []string `json:"document_ids"` MetadataCondition map[string]interface{} `json:"metadata_condition"` } -// BatchUpdateDocumentMetadatasResponse summarises the operation. -type BatchUpdateDocumentMetadatasResponse struct { +// BatchUpdateMetadatasResponse summarises the operation. +type BatchUpdateMetadatasResponse struct { Updated int `json:"updated"` MatchedDocs int `json:"matched_docs"` } + +// removeObjectBestEffort retries blob deletion on a context that survives the +// originating request. Bounded by a timeout so a wedged storage SDK cannot +// block the caller forever. It always uses the parent request's storage impl, +// but deliberately NOT the request context, because a cancelled request must +// not leak the blob it already wrote (or orphan a blob whose row was deleted). +func removeObjectBestEffort(storageImpl storage.Storage, bucket, object string) error { + ctx, cancel := context.WithTimeout(context.WithoutCancel(context.Background()), 30*time.Second) + defer cancel() + + var lastErr error + for attempt := 0; attempt < 3; attempt++ { + if err := storageImpl.Remove(ctx, bucket, object); err != nil { + lastErr = err + // Treat cancellation of the *new* cleanup ctx as terminal. + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + break + } + time.Sleep(time.Duration(attempt+1) * 100 * time.Millisecond) + continue + } + return nil + } + return lastErr +} diff --git a/internal/service/document/document_artifact.go b/internal/service/document/document_artifact.go index 7a320e5f4e..99a5dfa817 100644 --- a/internal/service/document/document_artifact.go +++ b/internal/service/document/document_artifact.go @@ -25,7 +25,7 @@ func (s *DocumentService) GetDocumentImage(ctx context.Context, imageID string) return nil, fmt.Errorf("storage not initialized") } - return storageImpl.Get(parts[0], parts[1]) + return storageImpl.Get(ctx, parts[0], parts[1]) } // GetDocumentArtifact retrieves a sandbox artifact from object storage. @@ -61,11 +61,11 @@ func (s *DocumentService) GetDocumentArtifact(ctx context.Context, filename, use } bucket := sandboxArtifactBucket() - if !storageImpl.ObjExist(bucket, basename) { + if !storageImpl.ObjExist(ctx, bucket, basename) { return nil, ErrArtifactNotFound } - data, err := storageImpl.Get(bucket, basename) + data, err := storageImpl.Get(ctx, bucket, basename) if err != nil { return nil, err } @@ -209,7 +209,7 @@ func (s *DocumentService) GetDocumentPreview(ctx context.Context, docID string) return nil, fmt.Errorf("storage not initialized") } - data, err := storageImpl.Get(bucket, name) + data, err := storageImpl.Get(ctx, bucket, name) if err != nil { return nil, err } diff --git a/internal/service/document/document_crud.go b/internal/service/document/document_crud.go index f5b0f5e10d..48fa9fb38b 100644 --- a/internal/service/document/document_crud.go +++ b/internal/service/document/document_crud.go @@ -83,7 +83,7 @@ func (s *DocumentService) DownloadDocument(ctx context.Context, datasetID, docID return nil, fmt.Errorf("storage not initialized") } - data, err := storageImpl.Get(bucket, name) + data, err := storageImpl.Get(ctx, bucket, name) if err != nil { return nil, err } @@ -279,14 +279,11 @@ func (s *DocumentService) RemoveDocumentKeepFile(ctx context.Context, docID stri if err != nil { return err } - if _, delErr := s.taskDAO.DeleteByDocIDs(ctx, dao.DB, []string{docID}); delErr != nil { - common.Logger.Warn(fmt.Sprintf("RemoveDocumentKeepFile: failed to delete tasks for %s: %v", docID, delErr)) - } if _, delErr := s.taskDAO.DeleteByDocIDs(ctx, dao.DB, []string{docID}); delErr != nil { if errors.Is(delErr, context.Canceled) || errors.Is(delErr, context.DeadlineExceeded) { return fmt.Errorf("RemoveDocumentKeepFile: failed to delete tasks for %s: %w", docID, delErr) } - common.Logger.Warn(fmt.Sprintf("RemoveDocumentKeepFile: failed to delete tasks for %s: %v", docID, delErr)) + common.Warn(fmt.Sprintf("RemoveDocumentKeepFile: failed to delete tasks for %s: %v", docID, delErr)) } return s.deleteDocRecordWithCounters(ctx, doc, kb.ID) } @@ -357,7 +354,7 @@ func (s *DocumentService) deleteDocEngineData(docID, tenantID, kbID string) { ctx := context.Background() indexName := fmt.Sprintf("ragflow_%s", tenantID) if _, delErr := s.docEngine.DeleteChunks(ctx, map[string]interface{}{"doc_id": docID}, indexName, kbID); delErr != nil { - common.Logger.Warn(fmt.Sprintf("deleteDocEngineData: failed to delete chunks for %s: %v", docID, delErr)) + common.Warn(fmt.Sprintf("deleteDocEngineData: failed to delete chunks for %s: %v", docID, delErr)) } // Notify the dataset-level post-processing consumer (§11) that this document's // source + per-doc compiled chunks are gone. The consumer removes the @@ -369,7 +366,7 @@ func (s *DocumentService) deleteDocEngineData(docID, tenantID, kbID string) { pubCtx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() if err := knowledge_compile.PublishDeleted(pubCtx, tenantID, kbID, docID, 0); err != nil { - common.Logger.Warn(fmt.Sprintf("deleteDocEngineData: publish doc_deleted for %s failed: %v", docID, err)) + common.Warn(fmt.Sprintf("deleteDocEngineData: publish doc_deleted for %s failed: %v", docID, err)) } if s.metadataSvc != nil { _ = s.DeleteDocumentAllMetadata(ctx, docID) // logs internally @@ -422,7 +419,7 @@ func (s *DocumentService) rollbackAddFileFromKBError(ctx context.Context, doc *e func (s *DocumentService) cleanupFileReferences(ctx context.Context, docID string) error { mappings, mapErr := s.file2DocumentDAO.GetByDocumentID(ctx, dao.DB, docID) if mapErr != nil { - common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: failed to get f2d mappings for %s: %v", docID, mapErr)) + common.Warn(fmt.Sprintf("cleanupFileReferences: failed to get f2d mappings for %s: %v", docID, mapErr)) return mapErr } if len(mappings) == 0 { @@ -442,7 +439,7 @@ func (s *DocumentService) cleanupFileReferences(ctx context.Context, docID strin // Delete all file2document rows for this document if delErr := s.file2DocumentDAO.DeleteByDocumentID(ctx, dao.DB, docID); delErr != nil { - common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: failed to delete f2d for %s: %v", docID, delErr)) + common.Warn(fmt.Sprintf("cleanupFileReferences: failed to delete f2d for %s: %v", docID, delErr)) return delErr } @@ -451,7 +448,7 @@ func (s *DocumentService) cleanupFileReferences(ctx context.Context, docID strin for _, fileID := range fileIDs { remaining, remErr := s.file2DocumentDAO.GetByFileID(ctx, dao.DB, fileID) if remErr != nil { - common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: failed to check remaining f2d for %s: %v", fileID, remErr)) + common.Warn(fmt.Sprintf("cleanupFileReferences: failed to check remaining f2d for %s: %v", fileID, remErr)) continue } if len(remaining) > 0 { @@ -461,21 +458,22 @@ func (s *DocumentService) cleanupFileReferences(ctx context.Context, docID strin fileDAO := dao.NewFileDAO() file, fErr := fileDAO.GetByID(ctx, dao.DB, fileID) if fErr != nil || file == nil { - common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: file not found %s: %v", fileID, fErr)) + common.Warn(fmt.Sprintf("cleanupFileReferences: file not found %s: %v", fileID, fErr)) continue } if entity.FileSource(file.SourceType) != entity.FileSourceKnowledgebase { continue // linked from file management — unlink only, keep the file } if _, delErr := fileDAO.DeleteByIDs(ctx, dao.DB, []string{fileID}); delErr != nil { - common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: failed to delete file %s: %v", fileID, delErr)) + common.Warn(fmt.Sprintf("cleanupFileReferences: failed to delete file %s: %v", fileID, delErr)) continue // keep the blob so the live file row still has its object } if file.Location != nil && *file.Location != "" { storageImpl := storage.GetStorageFactory().GetStorage() if storageImpl != nil { - if rmErr := storageImpl.Remove(file.ParentID, *file.Location); rmErr != nil { - common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: failed to remove blob %s/%s: %v", file.ParentID, *file.Location, rmErr)) + rmErr := removeObjectBestEffort(storageImpl, file.ParentID, *file.Location) + if rmErr != nil { + common.Warn(fmt.Sprintf("cleanupFileReferences: failed to remove blob %s/%s: %v", file.ParentID, *file.Location, rmErr)) } } } diff --git a/internal/service/document/document_metadata.go b/internal/service/document/document_metadata.go index ea5d500124..a8757f35db 100644 --- a/internal/service/document/document_metadata.go +++ b/internal/service/document/document_metadata.go @@ -563,12 +563,12 @@ func (s *DocumentService) patchDocumentMetadata(ctx context.Context, docID strin func (s *DocumentService) BatchUpdateDocumentMetadatas( ctx context.Context, datasetID string, - selector *DocumentMetadataSelector, - updates []DocumentMetadataUpdate, - deletes []DocumentMetadataDelete, -) (*BatchUpdateDocumentMetadatasResponse, common.ErrorCode, error) { + selector *MetadataSelector, + updates []MetadataUpdate, + deletes []MetadataDelete, +) (*BatchUpdateMetadatasResponse, common.ErrorCode, error) { if selector == nil { - selector = &DocumentMetadataSelector{} + selector = &MetadataSelector{} } if code, err := validateBatchUpdateDocumentMetadatasRequest(selector, updates, deletes); err != nil { return nil, code, err @@ -635,7 +635,7 @@ func (s *DocumentService) BatchUpdateDocumentMetadatas( // Early-exit when conditions given but nothing matched. rawConds, _ := selector.MetadataCondition["conditions"] if rawConds != nil && len(targetDocIDs) == 0 { - return &BatchUpdateDocumentMetadatasResponse{Updated: 0, MatchedDocs: 0}, common.CodeSuccess, nil + return &BatchUpdateMetadatasResponse{Updated: 0, MatchedDocs: 0}, common.CodeSuccess, nil } } @@ -675,13 +675,13 @@ func (s *DocumentService) BatchUpdateDocumentMetadatas( updated++ } - return &BatchUpdateDocumentMetadatasResponse{Updated: updated, MatchedDocs: len(ids)}, common.CodeSuccess, nil + return &BatchUpdateMetadatasResponse{Updated: updated, MatchedDocs: len(ids)}, common.CodeSuccess, nil } func validateBatchUpdateDocumentMetadatasRequest( - selector *DocumentMetadataSelector, - updates []DocumentMetadataUpdate, - deletes []DocumentMetadataDelete, + selector *MetadataSelector, + updates []MetadataUpdate, + deletes []MetadataDelete, ) (common.ErrorCode, error) { for _, upd := range updates { if strings.TrimSpace(upd.Key) == "" || upd.Value == nil { @@ -778,7 +778,7 @@ func cloneDocumentMetadataValue(v interface{}) interface{} { } } -func applyDocumentMetadataUpdates(meta map[string]interface{}, updates []DocumentMetadataUpdate) bool { +func applyDocumentMetadataUpdates(meta map[string]interface{}, updates []MetadataUpdate) bool { changed := false for _, upd := range updates { key := strings.TrimSpace(upd.Key) @@ -854,7 +854,7 @@ func applyDocumentMetadataUpdates(meta map[string]interface{}, updates []Documen return changed } -func applyDocumentMetadataDeletes(meta map[string]interface{}, deletes []DocumentMetadataDelete) bool { +func applyDocumentMetadataDeletes(meta map[string]interface{}, deletes []MetadataDelete) bool { changed := false for _, del := range deletes { key := strings.TrimSpace(del.Key) diff --git a/internal/service/document/document_parse.go b/internal/service/document/document_parse.go index a3fbfce8d4..53959a8809 100644 --- a/internal/service/document/document_parse.go +++ b/internal/service/document/document_parse.go @@ -104,14 +104,14 @@ func (s *DocumentService) clearDocumentParseResults(ctx context.Context, doc *en } indexName := fmt.Sprintf("ragflow_%s", tenantID) - exists, err := s.docEngine.ChunkStoreExists(context.Background(), indexName, doc.KbID) + exists, err := s.docEngine.ChunkStoreExists(ctx, indexName, doc.KbID) if err != nil { return err } if !exists { return nil } - if _, err := s.docEngine.DeleteChunks(context.Background(), map[string]interface{}{"doc_id": doc.ID}, indexName, doc.KbID); err != nil { + if _, err = s.docEngine.DeleteChunks(ctx, map[string]interface{}{"doc_id": doc.ID}, indexName, doc.KbID); err != nil { return err } return nil @@ -380,8 +380,8 @@ func (s *DocumentService) resetDocumentForReparse(ctx context.Context, doc *enti } if s.docEngine != nil { indexName := fmt.Sprintf("ragflow_%s", tenantID) - s.deleteChunkImages(doc, indexName) - if _, err = s.docEngine.DeleteChunks(context.Background(), map[string]interface{}{"doc_id": doc.ID}, indexName, doc.KbID); err != nil { + s.deleteChunkImages(ctx, doc, indexName) + if _, err = s.docEngine.DeleteChunks(ctx, map[string]interface{}{"doc_id": doc.ID}, indexName, doc.KbID); err != nil { return err } } @@ -390,7 +390,7 @@ func (s *DocumentService) resetDocumentForReparse(ctx context.Context, doc *enti return nil } -func (s *DocumentService) deleteChunkImages(doc *entity.Document, indexName string) { +func (s *DocumentService) deleteChunkImages(ctx context.Context, doc *entity.Document, indexName string) { if s.docEngine == nil { return } @@ -401,7 +401,7 @@ func (s *DocumentService) deleteChunkImages(doc *entity.Document, indexName stri const pageSize = 1000 for offset := 0; ; offset += pageSize { - result, err := s.docEngine.Search(context.Background(), &enginetypes.SearchRequest{ + result, err := s.docEngine.Search(ctx, &enginetypes.SearchRequest{ IndexNames: []string{indexName}, KbIDs: []string{doc.KbID}, Offset: offset, @@ -420,8 +420,8 @@ func (s *DocumentService) deleteChunkImages(doc *entity.Document, indexName stri if !ok { continue } - if storageImpl.ObjExist(doc.KbID, imageKey) { - _ = storageImpl.Remove(doc.KbID, imageKey) + if storageImpl.ObjExist(ctx, doc.KbID, imageKey) { + _ = storageImpl.Remove(ctx, doc.KbID, imageKey) } } } @@ -499,7 +499,7 @@ func (s *DocumentService) updateDocumentStatusOnly(ctx context.Context, doc *ent indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) return s.docEngine.UpdateChunks( - context.Background(), + ctx, map[string]interface{}{"doc_id": doc.ID}, map[string]interface{}{"available_int": status}, indexName, diff --git a/internal/service/document/document_test.go b/internal/service/document/document_test.go index 1ab2613baf..29941d0a9f 100644 --- a/internal/service/document/document_test.go +++ b/internal/service/document/document_test.go @@ -63,36 +63,36 @@ func newFakeUploadStorage() *fakeUploadStorage { } func (f *fakeUploadStorage) Type() string { return "fake_upload_storage" } -func (f *fakeUploadStorage) Health() bool { return true } +func (f *fakeUploadStorage) Health(_ context.Context) bool { return true } func (f *fakeUploadStorage) key(bucket, fnm string) string { return bucket + "/" + fnm } -func (f *fakeUploadStorage) Put(bucket, fnm string, binary []byte, tenantID ...string) error { +func (f *fakeUploadStorage) Put(ctx context.Context, bucket, fnm string, binary []byte, tenantID ...string) error { f.objects[f.key(bucket, fnm)] = append([]byte(nil), binary...) return nil } -func (f *fakeUploadStorage) Get(bucket, fnm string, tenantID ...string) ([]byte, error) { +func (f *fakeUploadStorage) Get(ctx context.Context, bucket, fnm string, tenantID ...string) ([]byte, error) { v, ok := f.objects[f.key(bucket, fnm)] if !ok { return nil, errors.New("not found") } return append([]byte(nil), v...), nil } -func (f *fakeUploadStorage) Remove(bucket, fnm string, tenantID ...string) error { +func (f *fakeUploadStorage) Remove(ctx context.Context, bucket, fnm string, tenantID ...string) error { delete(f.objects, f.key(bucket, fnm)) return nil } -func (f *fakeUploadStorage) ObjExist(bucket, fnm string, tenantID ...string) bool { +func (f *fakeUploadStorage) ObjExist(ctx context.Context, bucket, fnm string, tenantID ...string) bool { _, ok := f.objects[f.key(bucket, fnm)] return ok } -func (f *fakeUploadStorage) ListObjects(bucket string, tenantID ...string) ([]string, error) { +func (f *fakeUploadStorage) ListObjects(ctx context.Context, bucket string, tenantID ...string) ([]string, error) { return []string{}, nil } -func (f *fakeUploadStorage) GetPresignedURL(bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { +func (f *fakeUploadStorage) GetPresignedURL(ctx context.Context, bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { return "", nil } -func (f *fakeUploadStorage) BucketExists(bucket string) bool { return true } -func (f *fakeUploadStorage) RemoveBucket(bucket string) error { return nil } -func (f *fakeUploadStorage) Copy(srcBucket, srcPath, destBucket, destPath string) bool { +func (f *fakeUploadStorage) BucketExists(ctx context.Context, bucket string) bool { return true } +func (f *fakeUploadStorage) RemoveBucket(ctx context.Context, bucket string) error { return nil } +func (f *fakeUploadStorage) Copy(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { v, ok := f.objects[f.key(srcBucket, srcPath)] if !ok { return false @@ -100,8 +100,8 @@ func (f *fakeUploadStorage) Copy(srcBucket, srcPath, destBucket, destPath string f.objects[f.key(destBucket, destPath)] = append([]byte(nil), v...) return true } -func (f *fakeUploadStorage) Move(srcBucket, srcPath, destBucket, destPath string) bool { - if !f.Copy(srcBucket, srcPath, destBucket, destPath) { +func (f *fakeUploadStorage) Move(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { + if !f.Copy(ctx, srcBucket, srcPath, destBucket, destPath) { return false } delete(f.objects, f.key(srcBucket, srcPath)) @@ -796,7 +796,7 @@ func TestUploadLocalDocuments_MirrorsPythonCoreFields(t *testing.T) { t.Fatalf("parser_config=%v", cfg) } - storedBlob, err := mockStorage.Get(kb.ID, "nested/path/deck(1).pptx") + storedBlob, err := mockStorage.Get(ctx, kb.ID, "nested/path/deck(1).pptx") if err != nil { t.Fatalf("blob not stored: %v", err) } @@ -2119,12 +2119,12 @@ func TestBatchUpdateDocumentMetadatasMatchesPythonSemantics(t *testing.T) { svc.docEngine = engine svc.metadataSvc = service.NewMetadataServiceForTest(dao.NewKnowledgebaseDAO(), engine) ctx := t.Context() - resp, code, err := svc.BatchUpdateDocumentMetadatas(ctx, "kb-1", &DocumentMetadataSelector{ + resp, code, err := svc.BatchUpdateDocumentMetadatas(ctx, "kb-1", &MetadataSelector{ DocumentIDs: []string{"doc-1", "doc-2", "doc-3"}, - }, []DocumentMetadataUpdate{ + }, []MetadataUpdate{ {Key: "tags", Value: "new", Match: "old"}, {Key: "category", Value: "paper"}, - }, []DocumentMetadataDelete{ + }, []MetadataDelete{ {Key: "author", Value: "alice"}, }) if err != nil { @@ -2180,9 +2180,9 @@ func TestBatchUpdateDocumentMetadatasDoesNotReplaceWhenCurrentSearchIsStale(t *t svc.docEngine = engine svc.metadataSvc = service.NewMetadataServiceForTest(dao.NewKnowledgebaseDAO(), engine) ctx := t.Context() - resp, code, err := svc.BatchUpdateDocumentMetadatas(ctx, "kb-1", &DocumentMetadataSelector{ + resp, code, err := svc.BatchUpdateDocumentMetadatas(ctx, "kb-1", &MetadataSelector{ DocumentIDs: []string{"doc-1"}, - }, []DocumentMetadataUpdate{ + }, []MetadataUpdate{ {Key: "category", Value: "paper"}, }, nil) if err != nil || code != common.CodeSuccess { @@ -2217,9 +2217,9 @@ func TestBatchUpdateDocumentMetadatasDeletesEmptyMetadataAndNoOps(t *testing.T) svc.docEngine = engine svc.metadataSvc = service.NewMetadataServiceForTest(dao.NewKnowledgebaseDAO(), engine) ctx := t.Context() - resp, code, err := svc.BatchUpdateDocumentMetadatas(ctx, "kb-1", &DocumentMetadataSelector{ + resp, code, err := svc.BatchUpdateDocumentMetadatas(ctx, "kb-1", &MetadataSelector{ DocumentIDs: []string{"doc-1", "doc-2"}, - }, nil, []DocumentMetadataDelete{{Key: "status", Value: "draft"}}) + }, nil, []MetadataDelete{{Key: "status", Value: "draft"}}) if err != nil || code != common.CodeSuccess { t.Fatalf("delete batch failed: code=%v err=%v", code, err) } @@ -2246,9 +2246,9 @@ func TestBatchUpdateDocumentMetadatasNormalizesNumberValues(t *testing.T) { svc.docEngine = engine svc.metadataSvc = service.NewMetadataServiceForTest(dao.NewKnowledgebaseDAO(), engine) ctx := t.Context() - resp, code, err := svc.BatchUpdateDocumentMetadatas(ctx, "kb-1", &DocumentMetadataSelector{ + resp, code, err := svc.BatchUpdateDocumentMetadatas(ctx, "kb-1", &MetadataSelector{ DocumentIDs: []string{"doc-1"}, - }, []DocumentMetadataUpdate{ + }, []MetadataUpdate{ {Key: "score", Value: "42", ValueType: "number"}, }, nil) if err != nil || code != common.CodeSuccess { @@ -2276,7 +2276,7 @@ func TestBatchUpdateDocumentMetadatasNormalizesNumberValues(t *testing.T) { func TestBatchUpdateDocumentMetadatasRejectsMissingValue(t *testing.T) { svc := testDocumentService(t) ctx := t.Context() - resp, code, err := svc.BatchUpdateDocumentMetadatas(ctx, "kb-1", &DocumentMetadataSelector{}, []DocumentMetadataUpdate{ + resp, code, err := svc.BatchUpdateDocumentMetadatas(ctx, "kb-1", &MetadataSelector{}, []MetadataUpdate{ {Key: "status"}, }, nil) if err == nil { diff --git a/internal/service/document/document_upload.go b/internal/service/document/document_upload.go index 035c6205a4..7d7752141b 100644 --- a/internal/service/document/document_upload.go +++ b/internal/service/document/document_upload.go @@ -7,13 +7,12 @@ import ( "mime/multipart" "net/http" "path/filepath" - "ragflow/internal/dao" - "strings" - "ragflow/internal/common" + "ragflow/internal/dao" "ragflow/internal/entity" "ragflow/internal/storage" "ragflow/internal/utility" + "strings" ) // UploadLocalDocuments stores each uploaded file in object storage and inserts a @@ -86,10 +85,10 @@ func (s *DocumentService) UploadLocalDocuments(ctx context.Context, kb *entity.K if safeParent != "" { location = safeParent + "/" + filename } - for storageImpl.ObjExist(kb.ID, location) { + for storageImpl.ObjExist(ctx, kb.ID, location) { location += "_" } - if err = storageImpl.Put(kb.ID, location, blob); err != nil { + if err = storageImpl.Put(ctx, kb.ID, location, blob); err != nil { errMsgs = append(errMsgs, fh.Filename+": "+err.Error()) continue } @@ -97,7 +96,10 @@ func (s *DocumentService) UploadLocalDocuments(ctx context.Context, kb *entity.K doc := s.newDatasetDocument(kb, tenantID, filename, location, string(filetype), merged, "local", int64(len(blob)), blob) if err = s.InsertDocument(doc); err != nil { // Roll back the orphaned blob so a failed insert doesn't leak storage. - _ = storageImpl.Remove(kb.ID, location) + rmErr := removeObjectBestEffort(storageImpl, kb.ID, location) + if rmErr != nil { + common.Warn(fmt.Sprintf("upload rollback: failed to remove orphaned blob %s/%s: %v", kb.ID, location, rmErr)) + } errMsgs = append(errMsgs, fh.Filename+": "+err.Error()) continue } @@ -105,11 +107,14 @@ func (s *DocumentService) UploadLocalDocuments(ctx context.Context, kb *entity.K // Linkage failed: roll back the document row and blob so the partial // state doesn't leave an invisible (unlisted) document behind. err = s.rollbackAddFileFromKBError(ctx, doc, kb.ID, err) - _ = storageImpl.Remove(kb.ID, location) + rmErr := removeObjectBestEffort(storageImpl, kb.ID, location) + if rmErr != nil { + common.Warn(fmt.Sprintf("UploadLocalDocuments: failed to remove blob %s/%s: %v", kb.ID, location, rmErr)) + } errMsgs = append(errMsgs, fh.Filename+": "+err.Error()) continue } - // Only reserve the name once the write fully succeeds. + // Only reserve the name once write fully succeeds. taken[filename] = true results = append(results, docToRawMap(doc)) } @@ -271,21 +276,27 @@ func (s *DocumentService) UploadWebDocument(ctx context.Context, kb *entity.Know } location := filename - for storageImpl.ObjExist(kb.ID, location) { + for storageImpl.ObjExist(ctx, kb.ID, location) { location += "_" } - if err = storageImpl.Put(kb.ID, location, blob); err != nil { + if err = storageImpl.Put(ctx, kb.ID, location, blob); err != nil { return nil, common.CodeServerError, err } doc := s.newDatasetDocument(kb, tenantID, filename, location, string(filetype), kb.ParserConfig, "web", int64(len(blob)), blob) if err = s.InsertDocument(doc); err != nil { - _ = storageImpl.Remove(kb.ID, location) + rmErr := removeObjectBestEffort(storageImpl, kb.ID, location) + if rmErr != nil { + common.Warn(fmt.Sprintf("UploadWebDocument: failed to insert document, remove blob %s/%s: %v", kb.ID, location, rmErr)) + } return nil, common.CodeServerError, err } if err = s.addFileFromKB(ctx, doc, kbFolder.ID, kb.TenantID); err != nil { err = s.rollbackAddFileFromKBError(ctx, doc, kb.ID, err) - _ = storageImpl.Remove(kb.ID, location) + rmErr := removeObjectBestEffort(storageImpl, kb.ID, location) + if rmErr != nil { + common.Warn(fmt.Sprintf("UploadWebDocument: failed to add file from knowledge base, remove blob %s/%s: %v", kb.ID, location, rmErr)) + } return nil, common.CodeServerError, err } return docToRawMap(doc), common.CodeSuccess, nil diff --git a/internal/service/file/file_commit.go b/internal/service/file/file_commit.go index 85097220e6..09096b0cb6 100644 --- a/internal/service/file/file_commit.go +++ b/internal/service/file/file_commit.go @@ -122,7 +122,7 @@ func (s *FileCommitService) CreateCommit(ctx context.Context, folderID, authorID objKey := ".objects/" + hashHex if storageImpl != nil { - if err := storageImpl.Put(folderID, objKey, contentBytes); err != nil { + if err := storageImpl.Put(ctx, folderID, objKey, contentBytes); err != nil { return fmt.Errorf("failed to store object: %w", err) } } @@ -588,7 +588,7 @@ func (s *FileCommitService) GetUncommittedChanges(ctx context.Context, folderID } if liveFile, ok := liveMap[fid]; ok { - liveHash := computeLiveFileHash(folderID, fid, liveFile) + liveHash := computeLiveFileHash(ctx, folderID, fid, liveFile) committedHash := "" if h, ok := committedEntry["hash"].(string); ok { committedHash = h @@ -779,7 +779,7 @@ func (s *FileCommitService) GetCommitFileContent(ctx context.Context, folderID, return nil, fmt.Errorf("storage not initialized") } - blob, err := storageImpl.Get(folderID, objKey) + blob, err := storageImpl.Get(ctx, folderID, objKey) if err != nil { return nil, fmt.Errorf("failed to read file content from storage: %w", err) } @@ -822,7 +822,7 @@ func (s *FileCommitService) GetFileVersionHistory(ctx context.Context, fileID st } // computeLiveFileHash computes the SHA256 hash of current file content from storage -func computeLiveFileHash(folderID, fileID string, file *entity.File) string { +func computeLiveFileHash(ctx context.Context, folderID, fileID string, file *entity.File) string { if file.Location == nil || *file.Location == "" { return "" } @@ -832,7 +832,7 @@ func computeLiveFileHash(folderID, fileID string, file *entity.File) string { return "" } - data, err := storageImpl.Get(folderID, *file.Location) + data, err := storageImpl.Get(ctx, folderID, *file.Location) if err != nil { return "" } diff --git a/internal/service/file/file_content.go b/internal/service/file/file_content.go index 56af088322..09a8f548a3 100644 --- a/internal/service/file/file_content.go +++ b/internal/service/file/file_content.go @@ -85,7 +85,7 @@ func (s *FileService) DownloadAgentFile(ctx context.Context, tenantID, location bucketName := fmt.Sprintf("%s-downloads", tenantID) - blob, err := storageImpl.Get(bucketName, location) + blob, err := storageImpl.Get(ctx, bucketName, location) if err != nil { return nil, fmt.Errorf("failed to read file from storage: %w", err) } @@ -132,7 +132,7 @@ func (s *FileService) GetFileContents(ctx context.Context, uid string, fileDicts return nil, nil, fmt.Errorf("no authorization") } - data, derr := storageImpl.Get(createdBy+"-downloads", id) + data, derr := storageImpl.Get(ctx, createdBy+"-downloads", id) if derr != nil || len(data) == 0 { continue } @@ -156,7 +156,7 @@ func (s *FileService) GetFileContents(ctx context.Context, uid string, fileDicts return texts, images, nil } -// parseAgentUploads resolves descriptors returned by upload_info from the +// ParseAgentUploads resolves descriptors returned by upload_info from the // caller's downloads bucket and converts them to sys.files values. func (s *FileService) ParseAgentUploads(ctx context.Context, userID string, fileDicts []map[string]interface{}, layoutRecognize string) ([]string, error) { storageImpl := storage.GetStorageFactory().GetStorage() @@ -177,7 +177,7 @@ func (s *FileService) ParseAgentUploads(ctx context.Context, userID string, file return nil, fmt.Errorf("file %q: created_by does not match the current user", name) } - data, err := storageImpl.Get(createdBy+"-downloads", id) + data, err := storageImpl.Get(ctx, createdBy+"-downloads", id) if err != nil { return nil, fmt.Errorf("file %q: read upload: %w", name, err) } diff --git a/internal/service/file/file_delete.go b/internal/service/file/file_delete.go index 0fef740e1d..443935faf2 100644 --- a/internal/service/file/file_delete.go +++ b/internal/service/file/file_delete.go @@ -62,7 +62,7 @@ func (s *FileService) deleteSingleFile(ctx context.Context, file *entity.File) e if file.Location != nil && *file.Location != "" { storageImpl := storage.GetStorageFactory().GetStorage() if storageImpl != nil { - if err := storageImpl.Remove(file.ParentID, *file.Location); err != nil { + if err := storageImpl.Remove(ctx, file.ParentID, *file.Location); err != nil { common.Logger.Error(fmt.Sprintf("Fail to remove object: %s/%s, error: %v", file.ParentID, *file.Location, err)) } } diff --git a/internal/service/file/file_folder.go b/internal/service/file/file_folder.go index 1811eeb19b..a8c3f8aba2 100644 --- a/internal/service/file/file_folder.go +++ b/internal/service/file/file_folder.go @@ -560,7 +560,7 @@ func (s *FileService) moveEntryRecursive(ctx context.Context, sourceFile *entity // Calculate new location newLocation := effectiveName - for storageImpl.ObjExist(destFolder.ID, newLocation) { + for storageImpl.ObjExist(ctx, destFolder.ID, newLocation) { newLocation += "_" } @@ -569,7 +569,7 @@ func (s *FileService) moveEntryRecursive(ctx context.Context, sourceFile *entity return fmt.Errorf("file location is empty") } - if !storageImpl.Move(sourceFile.ParentID, *sourceFile.Location, destFolder.ID, newLocation) { + if !storageImpl.Move(ctx, sourceFile.ParentID, *sourceFile.Location, destFolder.ID, newLocation) { return fmt.Errorf("move file failed at storage layer") } diff --git a/internal/service/file/file_test.go b/internal/service/file/file_test.go index 65b9af04d7..97626a06fc 100644 --- a/internal/service/file/file_test.go +++ b/internal/service/file/file_test.go @@ -47,11 +47,11 @@ func testFileService() *FileService { func (f *fakeStorage) Type() string { return "fake_storage" } -func (f *fakeStorage) Health() bool { +func (f *fakeStorage) Health(ctx context.Context) bool { return true } -func (f *fakeStorage) Put(bucket, fnm string, binary []byte, tenantID ...string) error { +func (f *fakeStorage) Put(ctx context.Context, bucket, fnm string, binary []byte, tenantID ...string) error { f.lastBucket = bucket f.lastFnm = fnm f.blob = binary @@ -59,42 +59,42 @@ func (f *fakeStorage) Put(bucket, fnm string, binary []byte, tenantID ...string) return f.err } -func (f *fakeStorage) Get(bucket, fnm string, tenantID ...string) ([]byte, error) { +func (f *fakeStorage) Get(ctx context.Context, bucket, fnm string, tenantID ...string) ([]byte, error) { f.getCalls++ f.lastBucket = bucket f.lastFnm = fnm return f.blob, f.err } -func (f *fakeStorage) Remove(bucket, fnm string, tenantID ...string) error { +func (f *fakeStorage) Remove(ctx context.Context, bucket, fnm string, tenantID ...string) error { panic("not implemented in fakeStorage") } -func (f *fakeStorage) ObjExist(bucket, fnm string, tenantID ...string) bool { +func (f *fakeStorage) ObjExist(ctx context.Context, bucket, fnm string, tenantID ...string) bool { return f.exists && f.lastBucket == bucket && f.lastFnm == fnm } -func (f *fakeStorage) GetPresignedURL(bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { +func (f *fakeStorage) GetPresignedURL(ctx context.Context, bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { panic("not implemented in fakeStorage") } -func (f *fakeStorage) BucketExists(bucket string) bool { +func (f *fakeStorage) BucketExists(ctx context.Context, bucket string) bool { panic("not implemented in fakeStorage") } -func (f *fakeStorage) ListObjects(bucket string, tenantID ...string) ([]string, error) { +func (f *fakeStorage) ListObjects(ctx context.Context, bucket string, tenantID ...string) ([]string, error) { panic("not implemented in fakeStorage") } -func (f *fakeStorage) RemoveBucket(bucket string) error { +func (f *fakeStorage) RemoveBucket(ctx context.Context, bucket string) error { panic("not implemented in fakeStorage") } -func (f *fakeStorage) Copy(srcBucket, srcPath, destBucket, destPath string) bool { +func (f *fakeStorage) Copy(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { panic("not implemented in fakeStorage") } -func (f *fakeStorage) Move(srcBucket, srcPath, destBucket, destPath string) bool { +func (f *fakeStorage) Move(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { panic("not implemented in fakeStorage") } @@ -102,7 +102,8 @@ func (f *fakeStorage) Close() error { return nil } func TestFileService_GetFileContents_NotAccessible(t *testing.T) { memory := storage.NewMemoryStorage() - if err := memory.Put("other-user-downloads", "loc-1", []byte("secret")); err != nil { + ctx := t.Context() + if err := memory.Put(ctx, "other-user-downloads", "loc-1", []byte("secret")); err != nil { t.Fatalf("put: %v", err) } factory := storage.GetStorageFactory() @@ -130,7 +131,8 @@ func TestFileService_GetFileContents_NotAccessible(t *testing.T) { func TestFileService_GetFileContents_Accessible(t *testing.T) { memory := storage.NewMemoryStorage() - if err := memory.Put("user-1-downloads", "loc-1", []byte("allowed content")); err != nil { + ctx := t.Context() + if err := memory.Put(ctx, "user-1-downloads", "loc-1", []byte("allowed content")); err != nil { t.Fatalf("put: %v", err) } factory := storage.GetStorageFactory() @@ -159,10 +161,10 @@ func TestFileService_GetFileContents_Accessible(t *testing.T) { func TestFileService_ParseAgentUploads_TextAndImageInRequestOrder(t *testing.T) { ctx := t.Context() memory := storage.NewMemoryStorage() - if err := memory.Put("user-1-downloads", "text-id", []byte("uploaded text")); err != nil { + if err := memory.Put(ctx, "user-1-downloads", "text-id", []byte("uploaded text")); err != nil { t.Fatalf("put text: %v", err) } - if err := memory.Put("user-1-downloads", "image-id", []byte("png")); err != nil { + if err := memory.Put(ctx, "user-1-downloads", "image-id", []byte("png")); err != nil { t.Fatalf("put image: %v", err) } factory := storage.GetStorageFactory() diff --git a/internal/service/file/file_upload.go b/internal/service/file/file_upload.go index 7d2a34998b..beca433599 100644 --- a/internal/service/file/file_upload.go +++ b/internal/service/file/file_upload.go @@ -90,7 +90,7 @@ func (s *FileService) UploadFile(ctx context.Context, tenantID, parentID string, } location := fileObjNames[len(fileObjNames)-1] - for storageImpl.ObjExist(lastFolder.ID, location) { + for storageImpl.ObjExist(ctx, lastFolder.ID, location) { location += "_" } @@ -107,7 +107,7 @@ func (s *FileService) UploadFile(ctx context.Context, tenantID, parentID string, return nil, fmt.Errorf("failed to read file data: %w", err) } - if err = storageImpl.Put(lastFolder.ID, location, data); err != nil { + if err = storageImpl.Put(ctx, lastFolder.ID, location, data); err != nil { return nil, fmt.Errorf("failed to store file: %w", err) } @@ -169,7 +169,7 @@ func (s *FileService) UploadInfos(ctx context.Context, userID string, files []*m contentType = http.DetectContentType(data) } filename, contentType, data = utility.NormalizeUploadInfoContent(filename, contentType, data) - resp, err := s.storeUploadInfoBlob(storageImpl, userID, filename, contentType, data) + resp, err := s.storeUploadInfoBlob(ctx, storageImpl, userID, filename, contentType, data) if err != nil { return nil, err } @@ -245,10 +245,10 @@ func (s *FileService) checkUploadInfoHealth(ctx context.Context, userID, filenam return nil } -func (s *FileService) storeUploadInfoBlob(storageImpl storage.Storage, userID, filename, contentType string, data []byte) (map[string]interface{}, error) { +func (s *FileService) storeUploadInfoBlob(ctx context.Context, storageImpl storage.Storage, userID, filename, contentType string, data []byte) (map[string]interface{}, error) { location := utility.GenerateUUID() bucket := fmt.Sprintf("%s-downloads", userID) - if err := storageImpl.Put(bucket, location, data); err != nil { + if err := storageImpl.Put(ctx, bucket, location, data); err != nil { return nil, fmt.Errorf("failed to store file: %w", err) } ext := "" diff --git a/internal/service/file/file_url.go b/internal/service/file/file_url.go index 20b2aecacd..071048101b 100644 --- a/internal/service/file/file_url.go +++ b/internal/service/file/file_url.go @@ -46,7 +46,7 @@ func (s *FileService) UploadFromURL(ctx context.Context, tenantID, rawURL string return nil, err } filename, contentType, data = utility.NormalizeUploadInfoContent(filename, contentType, data) - return s.storeUploadInfoBlob(storageImpl, tenantID, filename, contentType, data) + return s.storeUploadInfoBlob(ctx, storageImpl, tenantID, filename, contentType, data) } func normalizeRemoteUploadFilename(rawURL, contentType string, data []byte) string { diff --git a/internal/service/skill_indexer.go b/internal/service/skill_indexer.go index 8e26427a20..6ab11ef7c6 100644 --- a/internal/service/skill_indexer.go +++ b/internal/service/skill_indexer.go @@ -786,7 +786,7 @@ func (s *SkillIndexerService) getFileContent(ctx context.Context, tenantID strin // Fallback to tenantID if ParentID is empty (should not happen) bucket = tenantID } - content, err := storageImpl.Get(bucket, *file.Location) + content, err := storageImpl.Get(ctx, bucket, *file.Location) if err != nil { return nil, fmt.Errorf("failed to get file from storage (bucket=%s, location=%s): %w", bucket, *file.Location, err) } diff --git a/internal/service/system.go b/internal/service/system.go index ee76a25e4a..aec429f0ca 100644 --- a/internal/service/system.go +++ b/internal/service/system.go @@ -104,15 +104,15 @@ type StatusResponse struct { // GetStatus gets health status for core system dependencies. func (s *SystemService) GetStatus(ctx context.Context) (*StatusResponse, error) { return &StatusResponse{ - DocEngine: s.getDocEngineStatus(), - Storage: s.getStorageStatus(), - Database: s.getDatabaseStatus(), + DocEngine: s.getDocEngineStatus(ctx), + Storage: s.getStorageStatus(ctx), + Database: s.getDatabaseStatus(ctx), Redis: s.getRedisStatus(ctx), TaskExecutorHeartbeats: s.getTaskExecutorHeartbeats(ctx), }, nil } -func (s *SystemService) getDocEngineStatus() ComponentStatus { +func (s *SystemService) getDocEngineStatus(ctx context.Context) ComponentStatus { cfg := server.GetConfig() docEngineType := "" if cfg != nil { @@ -130,9 +130,9 @@ func (s *SystemService) getDocEngineStatus() ComponentStatus { } } - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + timeOutCtx, cancel := context.WithTimeout(ctx, 5*time.Second) defer cancel() - if err := docEngine.Ping(ctx); err != nil { + if err := docEngine.Ping(timeOutCtx); err != nil { return ComponentStatus{ "type": docEngine.GetType(), "status": "red", @@ -148,7 +148,7 @@ func (s *SystemService) getDocEngineStatus() ComponentStatus { } } -func (s *SystemService) getStorageStatus() ComponentStatus { +func (s *SystemService) getStorageStatus(ctx context.Context) ComponentStatus { cfg := server.GetConfig() storageType := "" if cfg != nil { @@ -166,10 +166,10 @@ func (s *SystemService) getStorageStatus() ComponentStatus { } } - _, cancel := context.WithTimeout(context.Background(), 5*time.Second) + timeOutCtx, cancel := context.WithTimeout(ctx, 5*time.Second) defer cancel() - if !factory.Health() { + if !factory.Health(timeOutCtx) { return ComponentStatus{ "type": storageType, "status": "red", @@ -185,7 +185,7 @@ func (s *SystemService) getStorageStatus() ComponentStatus { } } -func (s *SystemService) getDatabaseStatus() ComponentStatus { +func (s *SystemService) getDatabaseStatus(ctx context.Context) ComponentStatus { cfg := server.GetConfig() databaseType := "" if cfg != nil { @@ -212,7 +212,7 @@ func (s *SystemService) getDatabaseStatus() ComponentStatus { } } - if err = sqlDB.Ping(); err != nil { + if err = sqlDB.PingContext(ctx); err != nil { return ComponentStatus{ "type": databaseType, "status": "red", @@ -351,7 +351,7 @@ func GetComponentsHealthz(ctx context.Context) (*HealthzResponse, bool) { storageOK, storageMeta := timedHealthCheck(func() error { store := storage.GetStorageFactory().GetStorage() - if store == nil || !store.Health() { + if store == nil || !store.Health(ctx) { return fmt.Errorf("storage is not healthy") } return nil diff --git a/internal/storage/gcs.go b/internal/storage/gcs.go index 1310eac49d..fafd1c66ca 100644 --- a/internal/storage/gcs.go +++ b/internal/storage/gcs.go @@ -37,21 +37,19 @@ type GCSStorage struct { } // NewGCSStorage creates a new GCS storage instance -func NewGCSStorage(config config.GCSConfig) (*GCSStorage, error) { +func NewGCSStorage(ctx context.Context, config config.GCSConfig) (*GCSStorage, error) { gcsStorage := &GCSStorage{ config: config, } - if err := gcsStorage.connect(); err != nil { + if err := gcsStorage.connect(ctx); err != nil { return nil, err } return gcsStorage, nil } -func (m *GCSStorage) connect() error { - - ctx := context.Background() +func (m *GCSStorage) connect(ctx context.Context) error { client, err := storage.NewClient(ctx) if err != nil { @@ -62,8 +60,8 @@ func (m *GCSStorage) connect() error { return nil } -func (m *GCSStorage) reconnect() { - if err := m.connect(); err != nil { +func (m *GCSStorage) reconnect(ctx context.Context) { + if err := m.connect(ctx); err != nil { common.Fatal(fmt.Sprintf("Failed to reconnect to GCS, %s", err.Error())) } } @@ -71,13 +69,12 @@ func (m *GCSStorage) reconnect() { func (m *GCSStorage) Type() string { return "gcs" } // Health checks GCS service availability -func (m *GCSStorage) Health() bool { - return m.BucketExists(m.config.Bucket) +func (m *GCSStorage) Health(ctx context.Context) bool { + return m.BucketExists(ctx, m.config.Bucket) } // Put uploads an object to GCS -func (m *GCSStorage) Put(bucket, fnm string, binary []byte, tenantID ...string) error { - ctx := context.Background() +func (m *GCSStorage) Put(ctx context.Context, bucket, fnm string, binary []byte, tenantID ...string) error { obj := m.client.Bucket(bucket).Object(fnm) w := obj.NewWriter(ctx) @@ -93,8 +90,7 @@ func (m *GCSStorage) Put(bucket, fnm string, binary []byte, tenantID ...string) } // Get retrieves an object from GCS -func (m *GCSStorage) Get(bucket, fnm string, tenantID ...string) ([]byte, error) { - ctx := context.Background() +func (m *GCSStorage) Get(ctx context.Context, bucket, fnm string, tenantID ...string) ([]byte, error) { r, err := m.client.Bucket(bucket).Object(fnm).NewReader(ctx) if err != nil { @@ -111,8 +107,7 @@ func (m *GCSStorage) Get(bucket, fnm string, tenantID ...string) ([]byte, error) } // Remove removes an object from GCS -func (m *GCSStorage) Remove(bucketName, objectName string, tenantID ...string) error { - ctx := context.Background() +func (m *GCSStorage) Remove(ctx context.Context, bucketName, objectName string, tenantID ...string) error { obj := m.client.Bucket(bucketName).Object(objectName) if err := obj.Delete(ctx); err != nil { @@ -123,8 +118,7 @@ func (m *GCSStorage) Remove(bucketName, objectName string, tenantID ...string) e } // ObjExist checks if an object exists in GCS -func (m *GCSStorage) ObjExist(bucketName, objectName string, tenantID ...string) bool { - ctx := context.Background() +func (m *GCSStorage) ObjExist(ctx context.Context, bucketName, objectName string, tenantID ...string) bool { obj := m.client.Bucket(bucketName).Object(objectName) @@ -136,8 +130,7 @@ func (m *GCSStorage) ObjExist(bucketName, objectName string, tenantID ...string) return true } -func (m *GCSStorage) ListObjects(bucket string, tenantID ...string) ([]string, error) { - ctx := context.Background() +func (m *GCSStorage) ListObjects(ctx context.Context, bucket string, tenantID ...string) ([]string, error) { bucketObject := m.client.Bucket(bucket) it := bucketObject.Objects(ctx, nil) @@ -159,7 +152,7 @@ func (m *GCSStorage) ListObjects(bucket string, tenantID ...string) ([]string, e } // GetPresignedURL generates a presigned URL for accessing an object -func (m *GCSStorage) GetPresignedURL(bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { +func (m *GCSStorage) GetPresignedURL(ctx context.Context, bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { bucketObject := m.client.Bucket(bucket) objectPath := fmt.Sprintf("%s/%s", bucket, fnm) @@ -175,14 +168,12 @@ func (m *GCSStorage) GetPresignedURL(bucket, fnm string, expires time.Duration, } // BucketExists checks if a bucket exists -func (m *GCSStorage) BucketExists(bucket string) bool { +func (m *GCSStorage) BucketExists(ctx context.Context, bucket string) bool { actualBucket := bucket if m.config.Bucket != "" { actualBucket = m.config.Bucket } - ctx := context.Background() - _, err := m.client.Bucket(actualBucket).Attrs(ctx) if err != nil { return false @@ -192,13 +183,11 @@ func (m *GCSStorage) BucketExists(bucket string) bool { } // RemoveBucket removes a bucket and all its objects -func (m *GCSStorage) RemoveBucket(bucketName string) error { +func (m *GCSStorage) RemoveBucket(ctx context.Context, bucketName string) error { if bucketName == "" { return fmt.Errorf("attempt to delete bucket without name") } - ctx := context.Background() - bucket := m.client.Bucket(bucketName) it := bucket.Objects(ctx, nil) @@ -224,8 +213,8 @@ func (m *GCSStorage) RemoveBucket(bucketName string) error { } // Copy copies an object from source to destination -func (m *GCSStorage) Copy(srcBucket, srcObject, destBucket, destObject string) bool { - ctx := context.Background() +func (m *GCSStorage) Copy(ctx context.Context, srcBucket, srcObject, destBucket, destObject string) bool { + src := m.client.Bucket(srcBucket).Object(srcObject) dst := m.client.Bucket(destBucket).Object(destObject) copier := dst.CopierFrom(src) @@ -238,10 +227,16 @@ func (m *GCSStorage) Copy(srcBucket, srcObject, destBucket, destObject string) b } // Move moves an object from source to destination -func (m *GCSStorage) Move(srcBucket, srcPath, destBucket, destPath string) bool { - if m.Copy(srcBucket, srcPath, destBucket, destPath) { - if err := m.Remove(srcBucket, srcPath); err != nil { +func (m *GCSStorage) Move(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { + if m.Copy(ctx, srcBucket, srcPath, destBucket, destPath) { + if err := m.Remove(ctx, srcBucket, srcPath); err != nil { common.Warn("Failed to remove source object after copy", zap.String("bucket", srcBucket), zap.String("key", srcPath), zap.Error(err)) + rollbackCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + defer cancel() + err = m.Remove(rollbackCtx, destBucket, destPath) + if err != nil { + common.Warn("Failed to roll back copied destination object", zap.String("bucket", destBucket), zap.String("key", destPath), zap.Error(err)) + } return false } return true diff --git a/internal/storage/memory.go b/internal/storage/memory.go index 145a324f2e..a67f6143a2 100644 --- a/internal/storage/memory.go +++ b/internal/storage/memory.go @@ -17,10 +17,14 @@ package storage import ( + "context" "errors" "fmt" + "ragflow/internal/common" "sync" "time" + + "go.uber.org/zap" ) // ErrMemoryNotFound is returned when a key does not exist in the in-memory backend. @@ -49,13 +53,13 @@ func NewMemoryStorage() Storage { func (m *MemoryStorage) Type() string { return "memory_storage" } // Health always reports healthy for the in-memory backend. -func (m *MemoryStorage) Health() bool { +func (m *MemoryStorage) Health(ctx context.Context) bool { return true } // Put uploads an object to the in-memory backend, creating the bucket // on demand if it does not yet exist. The stored bytes are a defensive copy. -func (m *MemoryStorage) Put(bucket, fnm string, binary []byte, tenantID ...string) error { +func (m *MemoryStorage) Put(ctx context.Context, bucket, fnm string, binary []byte, tenantID ...string) error { if bucket == "" { return fmt.Errorf("memory storage: bucket is required") } @@ -63,6 +67,10 @@ func (m *MemoryStorage) Put(bucket, fnm string, binary []byte, tenantID ...strin return fmt.Errorf("memory storage: key is required") } + if err := ctx.Err(); err != nil { + return err + } + m.mu.Lock() defer m.mu.Unlock() @@ -80,7 +88,7 @@ func (m *MemoryStorage) Put(bucket, fnm string, binary []byte, tenantID ...strin // Get retrieves an object from the in-memory backend. Returns // ErrMemoryNotFound when the bucket or key is missing. -func (m *MemoryStorage) Get(bucket, fnm string, tenantID ...string) ([]byte, error) { +func (m *MemoryStorage) Get(ctx context.Context, bucket, fnm string, tenantID ...string) ([]byte, error) { m.mu.RLock() defer m.mu.RUnlock() @@ -100,7 +108,11 @@ func (m *MemoryStorage) Get(bucket, fnm string, tenantID ...string) ([]byte, err // Remove deletes an object from the in-memory backend. Removing a // non-existent key is a no-op and returns nil. -func (m *MemoryStorage) Remove(bucket, fnm string, tenantID ...string) error { +func (m *MemoryStorage) Remove(ctx context.Context, bucket, fnm string, tenantID ...string) error { + if err := ctx.Err(); err != nil { + return err + } + m.mu.Lock() defer m.mu.Unlock() @@ -113,7 +125,7 @@ func (m *MemoryStorage) Remove(bucket, fnm string, tenantID ...string) error { } // ObjExist reports whether the given bucket and key are present. -func (m *MemoryStorage) ObjExist(bucket, fnm string, tenantID ...string) bool { +func (m *MemoryStorage) ObjExist(ctx context.Context, bucket, fnm string, tenantID ...string) bool { m.mu.RLock() defer m.mu.RUnlock() @@ -125,7 +137,7 @@ func (m *MemoryStorage) ObjExist(bucket, fnm string, tenantID ...string) bool { return ok } -func (m *MemoryStorage) ListObjects(bucket string, tenantID ...string) ([]string, error) { +func (m *MemoryStorage) ListObjects(ctx context.Context, bucket string, tenantID ...string) ([]string, error) { m.mu.RLock() defer m.mu.RUnlock() @@ -146,7 +158,7 @@ func (m *MemoryStorage) ListObjects(bucket string, tenantID ...string) ([]string // GetPresignedURL returns a deterministic, non-network URL string for tests. // Format: memory:///?exp= -func (m *MemoryStorage) GetPresignedURL(bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { +func (m *MemoryStorage) GetPresignedURL(ctx context.Context, bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { m.mu.RLock() defer m.mu.RUnlock() @@ -163,7 +175,7 @@ func (m *MemoryStorage) GetPresignedURL(bucket, fnm string, expires time.Duratio } // BucketExists reports whether the named bucket has been created. -func (m *MemoryStorage) BucketExists(bucket string) bool { +func (m *MemoryStorage) BucketExists(ctx context.Context, bucket string) bool { m.mu.RLock() defer m.mu.RUnlock() @@ -173,7 +185,11 @@ func (m *MemoryStorage) BucketExists(bucket string) bool { // RemoveBucket deletes a bucket and all of its keys. Removing a // non-existent bucket is a no-op and returns nil. -func (m *MemoryStorage) RemoveBucket(bucket string) error { +func (m *MemoryStorage) RemoveBucket(ctx context.Context, bucket string) error { + if err := ctx.Err(); err != nil { + return err + } + m.mu.Lock() defer m.mu.Unlock() @@ -184,7 +200,11 @@ func (m *MemoryStorage) RemoveBucket(bucket string) error { // Copy duplicates an object from srcBucket/srcKey to destBucket/destKey. // The source is left untouched. Returns false if the source does not exist // or if the destination bucket creation fails. -func (m *MemoryStorage) Copy(srcBucket, srcPath, destBucket, destPath string) bool { +func (m *MemoryStorage) Copy(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { + if err := ctx.Err(); err != nil { + return false + } + m.mu.RLock() srcBucketMap, ok := m.objects[srcBucket] if !ok { @@ -213,11 +233,22 @@ func (m *MemoryStorage) Copy(srcBucket, srcPath, destBucket, destPath string) bo // Move transfers an object to a new location, deleting the source on success. // Returns false if the source does not exist or the copy step fails. -func (m *MemoryStorage) Move(srcBucket, srcPath, destBucket, destPath string) bool { - if !m.Copy(srcBucket, srcPath, destBucket, destPath) { +func (m *MemoryStorage) Move(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { + + if err := ctx.Err(); err != nil { return false } - if err := m.Remove(srcBucket, srcPath); err != nil { + + if !m.Copy(ctx, srcBucket, srcPath, destBucket, destPath) { + return false + } + if err := m.Remove(ctx, srcBucket, srcPath); err != nil { + rollbackCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + defer cancel() + err = m.Remove(rollbackCtx, destBucket, destPath) + if err != nil { + common.Warn("Failed to roll back copied destination object", zap.String("bucket", destBucket), zap.String("key", destPath), zap.Error(err)) + } return false } return true diff --git a/internal/storage/memory_test.go b/internal/storage/memory_test.go index 86ec4c94a4..2c20ad09e0 100644 --- a/internal/storage/memory_test.go +++ b/internal/storage/memory_test.go @@ -38,13 +38,14 @@ func newTestMemory(t *testing.T) *MemoryStorage { func TestMemoryStorage_PutGet(t *testing.T) { ms := newTestMemory(t) + ctx := t.Context() payload := []byte("hello, world") - if err := ms.Put("b1", "k1", payload); err != nil { + if err := ms.Put(ctx, "b1", "k1", payload); err != nil { t.Fatalf("Put returned error: %v", err) } - got, err := ms.Get("b1", "k1") + got, err := ms.Get(ctx, "b1", "k1") if err != nil { t.Fatalf("Get returned error: %v", err) } @@ -54,7 +55,7 @@ func TestMemoryStorage_PutGet(t *testing.T) { // Mutating the caller's slice after Put must not affect stored data. payload[0] = 'X' - got2, err := ms.Get("b1", "k1") + got2, err := ms.Get(ctx, "b1", "k1") if err != nil { t.Fatalf("Get returned error: %v", err) } @@ -65,110 +66,115 @@ func TestMemoryStorage_PutGet(t *testing.T) { func TestMemoryStorage_GetMissing(t *testing.T) { ms := newTestMemory(t) + ctx := t.Context() - if _, err := ms.Get("missing-bucket", "k"); !errors.Is(err, ErrMemoryNotFound) { + if _, err := ms.Get(ctx, "missing-bucket", "k"); !errors.Is(err, ErrMemoryNotFound) { t.Fatalf("Get on missing bucket: expected ErrMemoryNotFound, got %v", err) } - if err := ms.Put("b1", "exists", []byte("data")); err != nil { + if err := ms.Put(ctx, "b1", "exists", []byte("data")); err != nil { t.Fatalf("Put failed: %v", err) } - if _, err := ms.Get("b1", "missing-key"); !errors.Is(err, ErrMemoryNotFound) { + if _, err := ms.Get(ctx, "b1", "missing-key"); !errors.Is(err, ErrMemoryNotFound) { t.Fatalf("Get on missing key: expected ErrMemoryNotFound, got %v", err) } } func TestMemoryStorage_ObjExist(t *testing.T) { ms := newTestMemory(t) + ctx := t.Context() - if ms.ObjExist("b1", "k1") { + if ms.ObjExist(ctx, "b1", "k1") { t.Fatalf("ObjExist on empty bucket returned true") } - if err := ms.Put("b1", "k1", []byte("v")); err != nil { + if err := ms.Put(ctx, "b1", "k1", []byte("v")); err != nil { t.Fatalf("Put failed: %v", err) } - if !ms.ObjExist("b1", "k1") { + if !ms.ObjExist(ctx, "b1", "k1") { t.Fatalf("ObjExist after Put returned false") } - if ms.ObjExist("b1", "other") { + if ms.ObjExist(ctx, "b1", "other") { t.Fatalf("ObjExist for sibling key returned true") } - if ms.ObjExist("other-bucket", "k1") { + if ms.ObjExist(ctx, "other-bucket", "k1") { t.Fatalf("ObjExist for sibling bucket returned true") } } func TestMemoryStorage_Remove(t *testing.T) { ms := newTestMemory(t) + ctx := t.Context() // Idempotent: removing a key from a missing bucket is a no-op. - if err := ms.Remove("ghost", "k"); err != nil { + if err := ms.Remove(ctx, "ghost", "k"); err != nil { t.Fatalf("Remove on missing bucket returned error: %v", err) } - if err := ms.Put("b1", "k1", []byte("v")); err != nil { + if err := ms.Put(ctx, "b1", "k1", []byte("v")); err != nil { t.Fatalf("Put failed: %v", err) } - if err := ms.Remove("b1", "k1"); err != nil { + if err := ms.Remove(ctx, "b1", "k1"); err != nil { t.Fatalf("Remove failed: %v", err) } - if ms.ObjExist("b1", "k1") { + if ms.ObjExist(ctx, "b1", "k1") { t.Fatalf("ObjExist after Remove returned true") } // Removing the same key again must not error. - if err := ms.Remove("b1", "k1"); err != nil { + if err := ms.Remove(ctx, "b1", "k1"); err != nil { t.Fatalf("Remove on already-removed key returned error: %v", err) } } func TestMemoryStorage_RemoveBucket(t *testing.T) { ms := newTestMemory(t) + ctx := t.Context() for _, k := range []string{"a", "b", "c"} { - if err := ms.Put("b1", k, []byte(k)); err != nil { + if err := ms.Put(ctx, "b1", k, []byte(k)); err != nil { t.Fatalf("Put failed: %v", err) } } - if err := ms.Put("b2", "x", []byte("x")); err != nil { + if err := ms.Put(ctx, "b2", "x", []byte("x")); err != nil { t.Fatalf("Put failed: %v", err) } - if err := ms.RemoveBucket("b1"); err != nil { + if err := ms.RemoveBucket(ctx, "b1"); err != nil { t.Fatalf("RemoveBucket failed: %v", err) } - if ms.BucketExists("b1") { + if ms.BucketExists(ctx, "b1") { t.Fatalf("BucketExists returned true after RemoveBucket") } - if !ms.BucketExists("b2") { + if !ms.BucketExists(ctx, "b2") { t.Fatalf("sibling bucket was removed unexpectedly") } // Idempotent: removing a missing bucket is a no-op. - if err := ms.RemoveBucket("b1"); err != nil { + if err := ms.RemoveBucket(ctx, "b1"); err != nil { t.Fatalf("RemoveBucket on missing bucket returned error: %v", err) } } func TestMemoryStorage_CopyMove(t *testing.T) { ms := newTestMemory(t) + ctx := t.Context() - if err := ms.Put("src", "k", []byte("payload")); err != nil { + if err := ms.Put(ctx, "src", "k", []byte("payload")); err != nil { t.Fatalf("Put failed: %v", err) } // Copy preserves source. - if !ms.Copy("src", "k", "dst", "k2") { + if !ms.Copy(ctx, "src", "k", "dst", "k2") { t.Fatalf("Copy returned false on existing source") } - if !ms.ObjExist("src", "k") { + if !ms.ObjExist(ctx, "src", "k") { t.Fatalf("source missing after Copy") } - if !ms.ObjExist("dst", "k2") { + if !ms.ObjExist(ctx, "dst", "k2") { t.Fatalf("destination missing after Copy") } - got, err := ms.Get("dst", "k2") + got, err := ms.Get(ctx, "dst", "k2") if err != nil { t.Fatalf("Get copy failed: %v", err) } @@ -177,53 +183,55 @@ func TestMemoryStorage_CopyMove(t *testing.T) { } // Move deletes the source. - if !ms.Move("src", "k", "dst2", "k3") { + if !ms.Move(ctx, "src", "k", "dst2", "k3") { t.Fatalf("Move returned false on existing source") } - if ms.ObjExist("src", "k") { + if ms.ObjExist(ctx, "src", "k") { t.Fatalf("source still exists after Move") } - if !ms.ObjExist("dst2", "k3") { + if !ms.ObjExist(ctx, "dst2", "k3") { t.Fatalf("destination missing after Move") } // Copy/Move on missing source returns false. - if ms.Copy("src", "k", "dst", "k4") { + if ms.Copy(ctx, "src", "k", "dst", "k4") { t.Fatalf("Copy on missing source returned true") } - if ms.Move("src", "k", "dst", "k4") { + if ms.Move(ctx, "src", "k", "dst", "k4") { t.Fatalf("Move on missing source returned true") } } func TestMemoryStorage_BucketExists(t *testing.T) { ms := newTestMemory(t) + ctx := t.Context() - if ms.BucketExists("b1") { + if ms.BucketExists(ctx, "b1") { t.Fatalf("BucketExists returned true for empty backend") } - if err := ms.Put("b1", "k", []byte("v")); err != nil { + if err := ms.Put(ctx, "b1", "k", []byte("v")); err != nil { t.Fatalf("Put failed: %v", err) } - if !ms.BucketExists("b1") { + if !ms.BucketExists(ctx, "b1") { t.Fatalf("BucketExists returned false after Put") } - if err := ms.RemoveBucket("b1"); err != nil { + if err := ms.RemoveBucket(ctx, "b1"); err != nil { t.Fatalf("RemoveBucket failed: %v", err) } - if ms.BucketExists("b1") { + if ms.BucketExists(ctx, "b1") { t.Fatalf("BucketExists returned true after RemoveBucket") } } func TestMemoryStorage_PresignedURL(t *testing.T) { ms := newTestMemory(t) + ctx := t.Context() - if err := ms.Put("b1", "k1", []byte("v")); err != nil { + if err := ms.Put(ctx, "b1", "k1", []byte("v")); err != nil { t.Fatalf("Put failed: %v", err) } - url, err := ms.GetPresignedURL("b1", "k1", time.Minute) + url, err := ms.GetPresignedURL(ctx, "b1", "k1", time.Minute) if err != nil { t.Fatalf("GetPresignedURL failed: %v", err) } @@ -237,20 +245,23 @@ func TestMemoryStorage_PresignedURL(t *testing.T) { t.Fatalf("presigned URL has unexpected scheme: %s", url) } - if _, err := ms.GetPresignedURL("b1", "missing", time.Minute); !errors.Is(err, ErrMemoryNotFound) { + if _, err := ms.GetPresignedURL(ctx, "b1", "missing", time.Minute); !errors.Is(err, ErrMemoryNotFound) { t.Fatalf("GetPresignedURL on missing key: expected ErrMemoryNotFound, got %v", err) } } func TestMemoryStorage_Health(t *testing.T) { ms := newTestMemory(t) - if !ms.Health() { + ctx := t.Context() + + if !ms.Health(ctx) { t.Fatalf("Health returned false for in-memory backend") } } func TestMemoryStorage_Concurrent(t *testing.T) { ms := newTestMemory(t) + ctx := t.Context() const writers = 100 var wg sync.WaitGroup @@ -261,7 +272,7 @@ func TestMemoryStorage_Concurrent(t *testing.T) { defer wg.Done() key := fmt.Sprintf("k-%d", i) payload := []byte(fmt.Sprintf("payload-%d", i)) - if err := ms.Put("race", key, payload); err != nil { + if err := ms.Put(ctx, "race", key, payload); err != nil { t.Errorf("Put failed for %s: %v", key, err) return } @@ -272,7 +283,7 @@ func TestMemoryStorage_Concurrent(t *testing.T) { for i := 0; i < writers; i++ { key := fmt.Sprintf("k-%d", i) want := fmt.Sprintf("payload-%d", i) - got, err := ms.Get("race", key) + got, err := ms.Get(ctx, "race", key) if err != nil { t.Fatalf("Get %s failed: %v", key, err) } @@ -284,18 +295,19 @@ func TestMemoryStorage_Concurrent(t *testing.T) { func TestMemoryStorage_Inspect(t *testing.T) { ms := newTestMemory(t) + ctx := t.Context() if got := ms.Inspect(); len(got) != 0 { t.Fatalf("Inspect on empty backend returned %d entries", len(got)) } - if err := ms.Put("b1", "k1", []byte("12345")); err != nil { + if err := ms.Put(ctx, "b1", "k1", []byte("12345")); err != nil { t.Fatalf("Put failed: %v", err) } - if err := ms.Put("b1", "k2", []byte("hello")); err != nil { + if err := ms.Put(ctx, "b1", "k2", []byte("hello")); err != nil { t.Fatalf("Put failed: %v", err) } - if err := ms.Put("b2", "only", []byte("x")); err != nil { + if err := ms.Put(ctx, "b2", "only", []byte("x")); err != nil { t.Fatalf("Put failed: %v", err) } @@ -325,10 +337,10 @@ func TestMemoryStorage_Inspect(t *testing.T) { } // After cleanup, Inspect should be empty again. - if err := ms.RemoveBucket("b1"); err != nil { + if err := ms.RemoveBucket(ctx, "b1"); err != nil { t.Fatalf("RemoveBucket failed: %v", err) } - if err := ms.RemoveBucket("b2"); err != nil { + if err := ms.RemoveBucket(ctx, "b2"); err != nil { t.Fatalf("RemoveBucket failed: %v", err) } if got := ms.Inspect(); len(got) != 0 { diff --git a/internal/storage/minio.go b/internal/storage/minio.go index cccf13f450..bba8500ec4 100644 --- a/internal/storage/minio.go +++ b/internal/storage/minio.go @@ -110,7 +110,7 @@ func (m *MinioStorage) resolveBucketAndPath(bucket, fnm string) (string, string) func (m *MinioStorage) Type() string { return "minio" } // Health checks MinIO service availability -func (m *MinioStorage) Health() bool { +func (m *MinioStorage) Health(ctx context.Context) bool { cancelFunction, err := m.client.HealthCheck(time.Second * 5) if cancelFunction != nil { defer cancelFunction() @@ -125,11 +125,9 @@ func (m *MinioStorage) Health() bool { } // Put uploads an object to MinIO -func (m *MinioStorage) Put(bucket, fnm string, binary []byte, tenantID ...string) error { +func (m *MinioStorage) Put(ctx context.Context, bucket, fnm string, binary []byte, tenantID ...string) error { bucket, fnm = m.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - var err error for i := 0; i < 3; i++ { @@ -138,16 +136,26 @@ func (m *MinioStorage) Put(bucket, fnm string, binary []byte, tenantID ...string if m.bucket == "" { exists, err = m.client.BucketExists(ctx, bucket) if err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } common.Warn("Failed to check bucket existence", zap.String("bucket", bucket), zap.Error(err)) m.reconnect() - time.Sleep(time.Second) + if err = sleepOrAbort(ctx, time.Second); err != nil { + return err + } continue } if !exists { if err = m.client.MakeBucket(ctx, bucket, minio.MakeBucketOptions{}); err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } common.Warn("Failed to create bucket", zap.String("bucket", bucket), zap.Error(err)) m.reconnect() - time.Sleep(time.Second) + if err = sleepOrAbort(ctx, time.Second); err != nil { + return err + } continue } } @@ -156,9 +164,14 @@ func (m *MinioStorage) Put(bucket, fnm string, binary []byte, tenantID ...string reader := bytes.NewReader(binary) _, err = m.client.PutObject(ctx, bucket, fnm, reader, int64(len(binary)), minio.PutObjectOptions{}) if err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } common.Warn("Failed to put object", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) m.reconnect() - time.Sleep(time.Second) + if err = sleepOrAbort(ctx, time.Second); err != nil { + return err + } continue } @@ -169,26 +182,34 @@ func (m *MinioStorage) Put(bucket, fnm string, binary []byte, tenantID ...string } // Get retrieves an object from MinIO -func (m *MinioStorage) Get(bucket, fnm string, tenantID ...string) ([]byte, error) { +func (m *MinioStorage) Get(ctx context.Context, bucket, fnm string, tenantID ...string) ([]byte, error) { bucket, fnm = m.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - for i := 0; i < 2; i++ { obj, err := m.client.GetObject(ctx, bucket, fnm, minio.GetObjectOptions{}) if err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return nil, ctxErr + } common.Warn("failed to get object", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) m.reconnect() - time.Sleep(time.Second) + if err = sleepOrAbort(ctx, time.Second); err != nil { + return nil, err + } continue } defer obj.Close() buf := new(bytes.Buffer) if _, err = buf.ReadFrom(obj); err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return nil, ctxErr + } common.Warn("failed to read object data", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) m.reconnect() - time.Sleep(time.Second) + if err = sleepOrAbort(ctx, time.Second); err != nil { + return nil, err + } continue } @@ -199,11 +220,9 @@ func (m *MinioStorage) Get(bucket, fnm string, tenantID ...string) ([]byte, erro } // Remove removes an object from MinIO -func (m *MinioStorage) Remove(bucket, fnm string, tenantID ...string) error { +func (m *MinioStorage) Remove(ctx context.Context, bucket, fnm string, tenantID ...string) error { bucket, fnm = m.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - if err := m.client.RemoveObject(ctx, bucket, fnm, minio.RemoveObjectOptions{}); err != nil { common.Warn("Failed to remove object", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) return err @@ -213,11 +232,9 @@ func (m *MinioStorage) Remove(bucket, fnm string, tenantID ...string) error { } // ObjExist checks if an object exists in MinIO -func (m *MinioStorage) ObjExist(bucket, fnm string, tenantID ...string) bool { +func (m *MinioStorage) ObjExist(ctx context.Context, bucket, fnm string, tenantID ...string) bool { bucket, fnm = m.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - exists, err := m.client.BucketExists(ctx, bucket) if err != nil || !exists { return false @@ -237,17 +254,20 @@ func (m *MinioStorage) ObjExist(bucket, fnm string, tenantID ...string) bool { } // GetPresignedURL generates a presigned URL for accessing an object -func (m *MinioStorage) GetPresignedURL(bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { +func (m *MinioStorage) GetPresignedURL(ctx context.Context, bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { bucket, fnm = m.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - for i := 0; i < 10; i++ { url, err := m.client.PresignedGetObject(ctx, bucket, fnm, expires, nil) if err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return "", ctxErr + } common.Warn("Failed to get presigned URL", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) m.reconnect() - time.Sleep(time.Second) + if err = sleepOrAbort(ctx, time.Second); err != nil { + return "", err + } continue } @@ -258,14 +278,12 @@ func (m *MinioStorage) GetPresignedURL(bucket, fnm string, expires time.Duration } // BucketExists checks if a bucket exists -func (m *MinioStorage) BucketExists(bucket string) bool { +func (m *MinioStorage) BucketExists(ctx context.Context, bucket string) bool { actualBucket := bucket if m.bucket != "" { actualBucket = m.bucket } - ctx := context.Background() - exists, err := m.client.BucketExists(ctx, actualBucket) if err != nil { common.Warn("Failed to check bucket existence", zap.String("bucket", actualBucket), zap.Error(err)) @@ -275,8 +293,7 @@ func (m *MinioStorage) BucketExists(bucket string) bool { return exists } -func (m *MinioStorage) ListObjects(bucket string, tenantID ...string) ([]string, error) { - ctx := context.Background() +func (m *MinioStorage) ListObjects(ctx context.Context, bucket string, tenantID ...string) ([]string, error) { var objects []string for obj := range m.client.ListObjects(ctx, bucket, minio.ListObjectsOptions{ @@ -294,7 +311,7 @@ func (m *MinioStorage) ListObjects(bucket string, tenantID ...string) ([]string, } // RemoveBucket removes a bucket and all its objects -func (m *MinioStorage) RemoveBucket(bucket string) error { +func (m *MinioStorage) RemoveBucket(ctx context.Context, bucket string) error { actualBucket := bucket origBucket := bucket @@ -302,8 +319,6 @@ func (m *MinioStorage) RemoveBucket(bucket string) error { actualBucket = m.bucket } - ctx := context.Background() - // Build prefix for single-bucket mode prefix := "" if m.bucket != "" { @@ -346,12 +361,10 @@ func (m *MinioStorage) RemoveBucket(bucket string) error { } // Copy copies an object from source to destination -func (m *MinioStorage) Copy(srcBucket, srcPath, destBucket, destPath string) bool { +func (m *MinioStorage) Copy(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { srcBucket, srcPath = m.resolveBucketAndPath(srcBucket, srcPath) destBucket, destPath = m.resolveBucketAndPath(destBucket, destPath) - ctx := context.Background() - // Ensure destination bucket exists if m.bucket == "" { exists, err := m.client.BucketExists(ctx, destBucket) @@ -394,10 +407,16 @@ func (m *MinioStorage) Copy(srcBucket, srcPath, destBucket, destPath string) boo } // Move moves an object from source to destination -func (m *MinioStorage) Move(srcBucket, srcPath, destBucket, destPath string) bool { - if m.Copy(srcBucket, srcPath, destBucket, destPath) { - if err := m.Remove(srcBucket, srcPath); err != nil { +func (m *MinioStorage) Move(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { + if m.Copy(ctx, srcBucket, srcPath, destBucket, destPath) { + if err := m.Remove(ctx, srcBucket, srcPath); err != nil { common.Warn("Failed to remove source object after copy", zap.String("bucket", srcBucket), zap.String("key", srcPath), zap.Error(err)) + rollbackCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + defer cancel() + err = m.Remove(rollbackCtx, destBucket, destPath) + if err != nil { + common.Warn("Failed to roll back copied destination object", zap.String("bucket", destBucket), zap.String("key", destPath), zap.Error(err)) + } return false } return true diff --git a/internal/storage/oss.go b/internal/storage/oss.go index 3adba24b36..d5ebc869e0 100644 --- a/internal/storage/oss.go +++ b/internal/storage/oss.go @@ -21,6 +21,7 @@ import ( "context" "errors" "fmt" + "ragflow/internal/common" "ragflow/internal/server/config" "time" @@ -42,22 +43,21 @@ type OSSStorage struct { } // NewOSSStorage creates a new OSS storage instance -func NewOSSStorage(config config.OSSConfig) (*OSSStorage, error) { +func NewOSSStorage(ctx context.Context, config config.OSSConfig) (*OSSStorage, error) { storage := &OSSStorage{ bucket: config.Bucket, prefixPath: config.PrefixPath, config: config, } - if err := storage.connect(); err != nil { + if err := storage.connect(ctx); err != nil { return nil, err } return storage, nil } -func (o *OSSStorage) connect() error { - ctx := context.Background() +func (o *OSSStorage) connect(ctx context.Context) error { // Create static credentials creds := credentials.NewStaticCredentialsProvider( @@ -83,9 +83,9 @@ func (o *OSSStorage) connect() error { return nil } -func (o *OSSStorage) reconnect() { - if err := o.connect(); err != nil { - zap.L().Error("Failed to reconnect to OSS", zap.Error(err)) +func (o *OSSStorage) reconnect(ctx context.Context) { + if err := o.connect(ctx); err != nil { + common.Error("Failed to reconnect to OSS", err) } } @@ -106,7 +106,7 @@ func (o *OSSStorage) resolveBucketAndPath(bucket, fnm string) (string, string) { func (o *OSSStorage) Type() string { return "oss" } // Health checks OSS service availability -func (o *OSSStorage) Health() bool { +func (o *OSSStorage) Health(ctx context.Context) bool { bucket := o.bucket if bucket == "" { bucket = "health-check-bucket" @@ -118,15 +118,13 @@ func (o *OSSStorage) Health() bool { } binary := []byte("_t@@@1") - ctx := context.Background() - // Ensure bucket exists - if !o.BucketExists(bucket) { + if !o.BucketExists(ctx, bucket) { _, err := o.client.CreateBucket(ctx, &s3.CreateBucketInput{ Bucket: aws.String(bucket), }) if err != nil { - zap.L().Error("Failed to create bucket for health check", zap.String("bucket", bucket), zap.Error(err)) + common.Error("Failed to create bucket for health check", err, zap.String("bucket", bucket)) return false } } @@ -140,7 +138,7 @@ func (o *OSSStorage) Health() bool { }) if err != nil { - zap.L().Error("Health check failed", zap.Error(err)) + common.Error("Health check failed", err) return false } @@ -148,24 +146,27 @@ func (o *OSSStorage) Health() bool { } // Put uploads an object to OSS -func (o *OSSStorage) Put(bucket, fnm string, binary []byte, tenantID ...string) error { +func (o *OSSStorage) Put(ctx context.Context, bucket, fnm string, binary []byte, tenantID ...string) error { bucket, fnm = o.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - for i := 0; i < 2; i++ { // Ensure bucket exists - if !o.BucketExists(bucket) { + if !o.BucketExists(ctx, bucket) { _, err := o.client.CreateBucket(ctx, &s3.CreateBucketInput{ Bucket: aws.String(bucket), }) if err != nil { - zap.L().Error("Failed to create bucket", zap.String("bucket", bucket), zap.Error(err)) - o.reconnect() - time.Sleep(time.Second) + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } + common.Error("Failed to create bucket", err, zap.String("bucket", bucket)) + o.reconnect(ctx) + if err = sleepOrAbort(ctx, time.Second); err != nil { + return err + } continue } - zap.L().Info("Created bucket", zap.String("bucket", bucket)) + common.Info("Created bucket", zap.String("bucket", bucket)) } reader := bytes.NewReader(binary) @@ -175,9 +176,14 @@ func (o *OSSStorage) Put(bucket, fnm string, binary []byte, tenantID ...string) Body: reader, }) if err != nil { - zap.L().Error("Failed to put object", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) - o.reconnect() - time.Sleep(time.Second) + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } + common.Error("Failed to put object", err, zap.String("bucket", bucket), zap.String("key", fnm)) + o.reconnect(ctx) + if err = sleepOrAbort(ctx, time.Second); err != nil { + return err + } continue } @@ -188,29 +194,37 @@ func (o *OSSStorage) Put(bucket, fnm string, binary []byte, tenantID ...string) } // Get retrieves an object from OSS -func (o *OSSStorage) Get(bucket, fnm string, tenantID ...string) ([]byte, error) { +func (o *OSSStorage) Get(ctx context.Context, bucket, fnm string, tenantID ...string) ([]byte, error) { bucket, fnm = o.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - for i := 0; i < 2; i++ { result, err := o.client.GetObject(ctx, &s3.GetObjectInput{ Bucket: aws.String(bucket), Key: aws.String(fnm), }) if err != nil { - zap.L().Error("Failed to get object", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) - o.reconnect() - time.Sleep(time.Second) + if ctxErr := ctx.Err(); ctxErr != nil { + return nil, ctxErr + } + common.Error("Failed to get object", err, zap.String("bucket", bucket), zap.String("key", fnm)) + o.reconnect(ctx) + if err = sleepOrAbort(ctx, time.Second); err != nil { + return nil, err + } continue } defer result.Body.Close() buf := new(bytes.Buffer) - if _, err := buf.ReadFrom(result.Body); err != nil { - zap.L().Error("Failed to read object data", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) - o.reconnect() - time.Sleep(time.Second) + if _, err = buf.ReadFrom(result.Body); err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return nil, ctxErr + } + common.Error("Failed to read object data", err, zap.String("bucket", bucket), zap.String("key", fnm)) + o.reconnect(ctx) + if err = sleepOrAbort(ctx, time.Second); err != nil { + return nil, err + } continue } @@ -221,17 +235,15 @@ func (o *OSSStorage) Get(bucket, fnm string, tenantID ...string) ([]byte, error) } // Remove removes an object from OSS -func (o *OSSStorage) Remove(bucket, fnm string, tenantID ...string) error { +func (o *OSSStorage) Remove(ctx context.Context, bucket, fnm string, tenantID ...string) error { bucket, fnm = o.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - _, err := o.client.DeleteObject(ctx, &s3.DeleteObjectInput{ Bucket: aws.String(bucket), Key: aws.String(fnm), }) if err != nil { - zap.L().Error("Failed to remove object", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) + common.Error("Failed to remove object", err, zap.String("bucket", bucket), zap.String("key", fnm)) return err } @@ -239,11 +251,9 @@ func (o *OSSStorage) Remove(bucket, fnm string, tenantID ...string) error { } // ObjExist checks if an object exists in OSS -func (o *OSSStorage) ObjExist(bucket, fnm string, tenantID ...string) bool { +func (o *OSSStorage) ObjExist(ctx context.Context, bucket, fnm string, tenantID ...string) bool { bucket, fnm = o.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - _, err := o.client.HeadObject(ctx, &s3.HeadObjectInput{ Bucket: aws.String(bucket), Key: aws.String(fnm), @@ -258,8 +268,7 @@ func (o *OSSStorage) ObjExist(bucket, fnm string, tenantID ...string) bool { return true } -func (o *OSSStorage) ListObjects(bucket string, tenantID ...string) ([]string, error) { - ctx := context.Background() +func (o *OSSStorage) ListObjects(ctx context.Context, bucket string, tenantID ...string) ([]string, error) { listInput := &s3.ListObjectsV2Input{ Bucket: aws.String(bucket), @@ -279,11 +288,9 @@ func (o *OSSStorage) ListObjects(bucket string, tenantID ...string) ([]string, e return objects, nil } -func (o *OSSStorage) GetPresignedURL(bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { +func (o *OSSStorage) GetPresignedURL(ctx context.Context, bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { bucket, fnm = o.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - presignClient := s3.NewPresignClient(o.client) for i := 0; i < 10; i++ { @@ -292,8 +299,8 @@ func (o *OSSStorage) GetPresignedURL(bucket, fnm string, expires time.Duration, Key: aws.String(fnm), }, s3.WithPresignExpires(expires)) if err != nil { - zap.L().Error("Failed to generate presigned URL", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) - o.reconnect() + common.Error("Failed to generate presigned URL", err, zap.String("bucket", bucket), zap.String("key", fnm)) + o.reconnect(ctx) time.Sleep(time.Second) continue } @@ -305,19 +312,17 @@ func (o *OSSStorage) GetPresignedURL(bucket, fnm string, expires time.Duration, } // BucketExists checks if a bucket exists -func (o *OSSStorage) BucketExists(bucket string) bool { +func (o *OSSStorage) BucketExists(ctx context.Context, bucket string) bool { actualBucket := bucket if o.bucket != "" { actualBucket = o.bucket } - ctx := context.Background() - _, err := o.client.HeadBucket(ctx, &s3.HeadBucketInput{ Bucket: aws.String(actualBucket), }) if err != nil { - zap.L().Debug("Bucket does not exist or error", zap.String("bucket", actualBucket), zap.Error(err)) + common.Error("Bucket does not exist or error", err, zap.String("bucket", actualBucket)) return false } @@ -325,16 +330,14 @@ func (o *OSSStorage) BucketExists(bucket string) bool { } // RemoveBucket removes a bucket and all its objects -func (o *OSSStorage) RemoveBucket(bucket string) error { +func (o *OSSStorage) RemoveBucket(ctx context.Context, bucket string) error { actualBucket := bucket if o.bucket != "" { actualBucket = o.bucket } - ctx := context.Background() - // Check if bucket exists - if !o.BucketExists(actualBucket) { + if !o.BucketExists(ctx, actualBucket) { return nil } @@ -346,17 +349,17 @@ func (o *OSSStorage) RemoveBucket(bucket string) error { for { result, err := o.client.ListObjectsV2(ctx, listInput) if err != nil { - zap.L().Error("Failed to list objects", zap.String("bucket", actualBucket), zap.Error(err)) + common.Error("Failed to list objects", err, zap.String("bucket", actualBucket)) return err } for _, obj := range result.Contents { - _, err := o.client.DeleteObject(ctx, &s3.DeleteObjectInput{ + _, err = o.client.DeleteObject(ctx, &s3.DeleteObjectInput{ Bucket: aws.String(actualBucket), Key: obj.Key, }) if err != nil { - zap.L().Error("Failed to delete object", zap.String("bucket", actualBucket), zap.Error(err)) + common.Error("Failed to delete object", err, zap.String("bucket", actualBucket)) } } @@ -371,7 +374,7 @@ func (o *OSSStorage) RemoveBucket(bucket string) error { Bucket: aws.String(actualBucket), }) if err != nil { - zap.L().Error("Failed to delete bucket", zap.String("bucket", actualBucket), zap.Error(err)) + common.Error("Failed to delete bucket", err, zap.String("bucket", actualBucket)) return err } @@ -379,12 +382,10 @@ func (o *OSSStorage) RemoveBucket(bucket string) error { } // Copy copies an object from source to destination -func (o *OSSStorage) Copy(srcBucket, srcPath, destBucket, destPath string) bool { +func (o *OSSStorage) Copy(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { srcBucket, srcPath = o.resolveBucketAndPath(srcBucket, srcPath) destBucket, destPath = o.resolveBucketAndPath(destBucket, destPath) - ctx := context.Background() - copySource := fmt.Sprintf("%s/%s", srcBucket, srcPath) _, err := o.client.CopyObject(ctx, &s3.CopyObjectInput{ @@ -393,7 +394,7 @@ func (o *OSSStorage) Copy(srcBucket, srcPath, destBucket, destPath string) bool CopySource: aws.String(copySource), }) if err != nil { - zap.L().Error("Failed to copy object", zap.String("src", copySource), zap.String("dest", fmt.Sprintf("%s/%s", destBucket, destPath)), zap.Error(err)) + common.Error("Failed to copy object", err, zap.String("src", copySource), zap.String("dest", fmt.Sprintf("%s/%s", destBucket, destPath))) return false } @@ -401,10 +402,16 @@ func (o *OSSStorage) Copy(srcBucket, srcPath, destBucket, destPath string) bool } // Move moves an object from source to destination -func (o *OSSStorage) Move(srcBucket, srcPath, destBucket, destPath string) bool { - if o.Copy(srcBucket, srcPath, destBucket, destPath) { - if err := o.Remove(srcBucket, srcPath); err != nil { - zap.L().Error("Failed to remove source object after copy", zap.String("bucket", srcBucket), zap.String("key", srcPath), zap.Error(err)) +func (o *OSSStorage) Move(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { + if o.Copy(ctx, srcBucket, srcPath, destBucket, destPath) { + if err := o.Remove(ctx, srcBucket, srcPath); err != nil { + common.Error("Failed to remove source object after copy", err, zap.String("bucket", srcBucket), zap.String("key", srcPath)) + rollbackCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + defer cancel() + err = o.Remove(rollbackCtx, destBucket, destPath) + if err != nil { + common.Warn("Failed to roll back copied destination object", zap.String("bucket", destBucket), zap.String("key", destPath), zap.Error(err)) + } return false } return true diff --git a/internal/storage/s3.go b/internal/storage/s3.go index 82e381abe0..c146e0f0cc 100644 --- a/internal/storage/s3.go +++ b/internal/storage/s3.go @@ -21,6 +21,7 @@ import ( "context" "errors" "fmt" + "ragflow/internal/common" "ragflow/internal/server/config" "time" @@ -41,20 +42,19 @@ type S3Storage struct { } // NewS3Storage creates a new S3 storage instance -func NewS3Storage(config config.S3Config) (*S3Storage, error) { +func NewS3Storage(ctx context.Context, config config.S3Config) (*S3Storage, error) { storage := &S3Storage{ config: config, } - if err := storage.connect(); err != nil { + if err := storage.connect(ctx); err != nil { return nil, err } return storage, nil } -func (s *S3Storage) connect() error { - ctx := context.Background() +func (s *S3Storage) connect(ctx context.Context) error { var opts []func(*s3Config.LoadOptions) error @@ -91,9 +91,9 @@ func (s *S3Storage) connect() error { return nil } -func (s *S3Storage) reconnect() { - if err := s.connect(); err != nil { - zap.L().Error("Failed to reconnect to S3", zap.Error(err)) +func (s *S3Storage) reconnect(ctx context.Context) { + if err := s.connect(ctx); err != nil { + common.Error("Failed to reconnect to S3", err, zap.Error(err)) } } @@ -114,7 +114,7 @@ func (s *S3Storage) resolveBucketAndPath(bucket, fnm string) (string, string) { func (s *S3Storage) Type() string { return "s3" } // Health checks S3 service availability -func (s *S3Storage) Health() bool { +func (s *S3Storage) Health(ctx context.Context) bool { bucket := s.bucket if bucket == "" { bucket = "health-check-bucket" @@ -126,15 +126,13 @@ func (s *S3Storage) Health() bool { } binary := []byte("_t@@@1") - ctx := context.Background() - // Ensure bucket exists - if !s.BucketExists(bucket) { + if !s.BucketExists(ctx, bucket) { _, err := s.client.CreateBucket(ctx, &s3.CreateBucketInput{ Bucket: aws.String(bucket), }) if err != nil { - zap.L().Error("Failed to create bucket for health check", zap.String("bucket", bucket), zap.Error(err)) + common.Error("Failed to create bucket for health check", err, zap.String("bucket", bucket), zap.Error(err)) return false } } @@ -148,7 +146,7 @@ func (s *S3Storage) Health() bool { }) if err != nil { - zap.L().Error("Health check failed", zap.Error(err)) + common.Error("Health check failed", err, zap.Error(err)) return false } @@ -156,24 +154,27 @@ func (s *S3Storage) Health() bool { } // Put uploads an object to S3 -func (s *S3Storage) Put(bucket, fnm string, binary []byte, tenantID ...string) error { +func (s *S3Storage) Put(ctx context.Context, bucket, fnm string, binary []byte, tenantID ...string) error { bucket, fnm = s.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - for i := 0; i < 2; i++ { // Ensure bucket exists - if !s.BucketExists(bucket) { + if !s.BucketExists(ctx, bucket) { _, err := s.client.CreateBucket(ctx, &s3.CreateBucketInput{ Bucket: aws.String(bucket), }) if err != nil { - zap.L().Error("Failed to create bucket", zap.String("bucket", bucket), zap.Error(err)) - s.reconnect() - time.Sleep(time.Second) + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } + common.Error("Failed to create bucket", err, zap.String("bucket", bucket), zap.Error(err)) + s.reconnect(ctx) + if err = sleepOrAbort(ctx, time.Second); err != nil { + return err + } continue } - zap.L().Info("Created bucket", zap.String("bucket", bucket)) + common.Info("Created bucket", zap.String("bucket", bucket)) } reader := bytes.NewReader(binary) @@ -183,9 +184,14 @@ func (s *S3Storage) Put(bucket, fnm string, binary []byte, tenantID ...string) e Body: reader, }) if err != nil { - zap.L().Error("Failed to put object", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) - s.reconnect() - time.Sleep(time.Second) + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } + common.Error("Failed to put object", err, zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) + s.reconnect(ctx) + if err = sleepOrAbort(ctx, time.Second); err != nil { + return err + } continue } @@ -196,28 +202,26 @@ func (s *S3Storage) Put(bucket, fnm string, binary []byte, tenantID ...string) e } // Get retrieves an object from S3 -func (s *S3Storage) Get(bucket, fnm string, tenantID ...string) ([]byte, error) { +func (s *S3Storage) Get(ctx context.Context, bucket, fnm string, tenantID ...string) ([]byte, error) { bucket, fnm = s.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - for i := 0; i < 2; i++ { result, err := s.client.GetObject(ctx, &s3.GetObjectInput{ Bucket: aws.String(bucket), Key: aws.String(fnm), }) if err != nil { - zap.L().Error("Failed to get object", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) - s.reconnect() + common.Error("Failed to get object", err, zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) + s.reconnect(ctx) time.Sleep(time.Second) continue } defer result.Body.Close() buf := new(bytes.Buffer) - if _, err := buf.ReadFrom(result.Body); err != nil { - zap.L().Error("Failed to read object data", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) - s.reconnect() + if _, err = buf.ReadFrom(result.Body); err != nil { + common.Error("Failed to read object data", err, zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) + s.reconnect(ctx) time.Sleep(time.Second) continue } @@ -229,17 +233,15 @@ func (s *S3Storage) Get(bucket, fnm string, tenantID ...string) ([]byte, error) } // Remove removes an object from S3 -func (s *S3Storage) Remove(bucket, fnm string, tenantID ...string) error { +func (s *S3Storage) Remove(ctx context.Context, bucket, fnm string, tenantID ...string) error { bucket, fnm = s.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - _, err := s.client.DeleteObject(ctx, &s3.DeleteObjectInput{ Bucket: aws.String(bucket), Key: aws.String(fnm), }) if err != nil { - zap.L().Error("Failed to remove object", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) + common.Error("Failed to remove object", err, zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) return err } @@ -247,11 +249,9 @@ func (s *S3Storage) Remove(bucket, fnm string, tenantID ...string) error { } // ObjExist checks if an object exists in S3 -func (s *S3Storage) ObjExist(bucket, fnm string, tenantID ...string) bool { +func (s *S3Storage) ObjExist(ctx context.Context, bucket, fnm string, tenantID ...string) bool { bucket, fnm = s.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - _, err := s.client.HeadObject(ctx, &s3.HeadObjectInput{ Bucket: aws.String(bucket), Key: aws.String(fnm), @@ -266,8 +266,7 @@ func (s *S3Storage) ObjExist(bucket, fnm string, tenantID ...string) bool { return true } -func (s *S3Storage) ListObjects(bucket string, tenantID ...string) ([]string, error) { - ctx := context.Background() +func (s *S3Storage) ListObjects(ctx context.Context, bucket string, tenantID ...string) ([]string, error) { listInput := &s3.ListObjectsV2Input{ Bucket: aws.String(bucket), @@ -288,11 +287,9 @@ func (s *S3Storage) ListObjects(bucket string, tenantID ...string) ([]string, er } // GetPresignedURL generates a presigned URL for accessing an object -func (s *S3Storage) GetPresignedURL(bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { +func (s *S3Storage) GetPresignedURL(ctx context.Context, bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) { bucket, fnm = s.resolveBucketAndPath(bucket, fnm) - ctx := context.Background() - presignClient := s3.NewPresignClient(s.client) for i := 0; i < 10; i++ { @@ -301,8 +298,8 @@ func (s *S3Storage) GetPresignedURL(bucket, fnm string, expires time.Duration, t Key: aws.String(fnm), }, s3.WithPresignExpires(expires)) if err != nil { - zap.L().Error("Failed to generate presigned URL", zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) - s.reconnect() + common.Error("Failed to generate presigned URL", err, zap.String("bucket", bucket), zap.String("key", fnm), zap.Error(err)) + s.reconnect(ctx) time.Sleep(time.Second) continue } @@ -314,19 +311,17 @@ func (s *S3Storage) GetPresignedURL(bucket, fnm string, expires time.Duration, t } // BucketExists checks if a bucket exists -func (s *S3Storage) BucketExists(bucket string) bool { +func (s *S3Storage) BucketExists(ctx context.Context, bucket string) bool { actualBucket := bucket if s.bucket != "" { actualBucket = s.bucket } - ctx := context.Background() - _, err := s.client.HeadBucket(ctx, &s3.HeadBucketInput{ Bucket: aws.String(actualBucket), }) if err != nil { - zap.L().Debug("Bucket does not exist or error", zap.String("bucket", actualBucket), zap.Error(err)) + common.Debug("Bucket does not exist or error", zap.String("bucket", actualBucket), zap.Error(err)) return false } @@ -334,16 +329,14 @@ func (s *S3Storage) BucketExists(bucket string) bool { } // RemoveBucket removes a bucket and all its objects -func (s *S3Storage) RemoveBucket(bucket string) error { +func (s *S3Storage) RemoveBucket(ctx context.Context, bucket string) error { actualBucket := bucket if s.bucket != "" { actualBucket = s.bucket } - ctx := context.Background() - // Check if bucket exists - if !s.BucketExists(actualBucket) { + if !s.BucketExists(ctx, actualBucket) { return nil } @@ -355,17 +348,17 @@ func (s *S3Storage) RemoveBucket(bucket string) error { for { result, err := s.client.ListObjectsV2(ctx, listInput) if err != nil { - zap.L().Error("Failed to list objects", zap.String("bucket", actualBucket), zap.Error(err)) + common.Error("Failed to list objects", err, zap.String("bucket", actualBucket), zap.Error(err)) return err } for _, obj := range result.Contents { - _, err := s.client.DeleteObject(ctx, &s3.DeleteObjectInput{ + _, err = s.client.DeleteObject(ctx, &s3.DeleteObjectInput{ Bucket: aws.String(actualBucket), Key: obj.Key, }) if err != nil { - zap.L().Error("Failed to delete object", zap.String("bucket", actualBucket), zap.Error(err)) + common.Error("Failed to delete object", err, zap.String("bucket", actualBucket), zap.Error(err)) } } @@ -380,7 +373,7 @@ func (s *S3Storage) RemoveBucket(bucket string) error { Bucket: aws.String(actualBucket), }) if err != nil { - zap.L().Error("Failed to delete bucket", zap.String("bucket", actualBucket), zap.Error(err)) + common.Error("Failed to delete bucket", err, zap.String("bucket", actualBucket), zap.Error(err)) return err } @@ -388,12 +381,10 @@ func (s *S3Storage) RemoveBucket(bucket string) error { } // Copy copies an object from source to destination -func (s *S3Storage) Copy(srcBucket, srcPath, destBucket, destPath string) bool { +func (s *S3Storage) Copy(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { srcBucket, srcPath = s.resolveBucketAndPath(srcBucket, srcPath) destBucket, destPath = s.resolveBucketAndPath(destBucket, destPath) - ctx := context.Background() - copySource := fmt.Sprintf("%s/%s", srcBucket, srcPath) _, err := s.client.CopyObject(ctx, &s3.CopyObjectInput{ @@ -402,7 +393,7 @@ func (s *S3Storage) Copy(srcBucket, srcPath, destBucket, destPath string) bool { CopySource: aws.String(copySource), }) if err != nil { - zap.L().Error("Failed to copy object", zap.String("src", copySource), zap.String("dest", fmt.Sprintf("%s/%s", destBucket, destPath)), zap.Error(err)) + common.Error("Failed to copy object", err, zap.String("src", copySource), zap.String("dest", fmt.Sprintf("%s/%s", destBucket, destPath)), zap.Error(err)) return false } @@ -410,10 +401,16 @@ func (s *S3Storage) Copy(srcBucket, srcPath, destBucket, destPath string) bool { } // Move moves an object from source to destination -func (s *S3Storage) Move(srcBucket, srcPath, destBucket, destPath string) bool { - if s.Copy(srcBucket, srcPath, destBucket, destPath) { - if err := s.Remove(srcBucket, srcPath); err != nil { - zap.L().Error("Failed to remove source object after copy", zap.String("bucket", srcBucket), zap.String("key", srcPath), zap.Error(err)) +func (s *S3Storage) Move(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool { + if s.Copy(ctx, srcBucket, srcPath, destBucket, destPath) { + if err := s.Remove(ctx, srcBucket, srcPath); err != nil { + common.Error("Failed to remove source object after copy", err, zap.String("bucket", srcBucket), zap.String("key", srcPath), zap.Error(err)) + rollbackCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + defer cancel() + err = s.Remove(rollbackCtx, destBucket, destPath) + if err != nil { + common.Warn("Failed to roll back copied destination object", zap.String("bucket", destBucket), zap.String("key", destPath), zap.Error(err)) + } return false } return true diff --git a/internal/storage/storage_factory.go b/internal/storage/storage_factory.go index e47f09583f..4fc54a29a6 100644 --- a/internal/storage/storage_factory.go +++ b/internal/storage/storage_factory.go @@ -17,10 +17,12 @@ package storage import ( + "context" "fmt" "ragflow/internal/common" "ragflow/internal/server" "sync" + "time" ) var ( @@ -43,12 +45,12 @@ func GetStorageFactory() *StorageFactory { } // Init initializes the storage factory with configuration -func Init() error { +func Init(ctx context.Context) error { factory := GetStorageFactory() globalConfig := server.GetConfig() // Initialize storage based on type - if err := factory.initStorage(); err != nil { + if err := factory.initStorage(ctx); err != nil { return err } @@ -65,17 +67,17 @@ func CloseStorage() error { return factory.storage.Close() } -func (f *StorageFactory) initStorage() error { +func (f *StorageFactory) initStorage(ctx context.Context) error { globalConfig := server.GetConfig() switch globalConfig.StorageEngineType() { case "minio": return f.initMinio() case "s3": - return f.initS3() + return f.initS3(ctx) case "oss": - return f.initOSS() + return f.initOSS(ctx) case "gcs": - return f.initGCS() + return f.initGCS(ctx) default: return fmt.Errorf("unsupported storage type: %s", globalConfig.StorageEngineType()) } @@ -95,9 +97,9 @@ func (f *StorageFactory) initMinio() error { return nil } -func (f *StorageFactory) initS3() error { +func (f *StorageFactory) initS3(ctx context.Context) error { globalConfig := server.GetConfig() - storage, err := NewS3Storage(globalConfig.GetS3Config()) + storage, err := NewS3Storage(ctx, globalConfig.GetS3Config()) if err != nil { return fmt.Errorf("failed to create S3 storage: %w", err) } @@ -109,9 +111,9 @@ func (f *StorageFactory) initS3() error { return nil } -func (f *StorageFactory) initOSS() error { +func (f *StorageFactory) initOSS(ctx context.Context) error { globalConfig := server.GetConfig() - storage, err := NewOSSStorage(globalConfig.GetOSSConfig()) + storage, err := NewOSSStorage(ctx, globalConfig.GetOSSConfig()) if err != nil { return fmt.Errorf("failed to create OSS storage: %w", err) } @@ -123,9 +125,9 @@ func (f *StorageFactory) initOSS() error { return nil } -func (f *StorageFactory) initGCS() error { +func (f *StorageFactory) initGCS(ctx context.Context) error { globalConfig := server.GetConfig() - storage, err := NewGCSStorage(globalConfig.GetGCSConfig()) + storage, err := NewGCSStorage(ctx, globalConfig.GetGCSConfig()) if err != nil { return fmt.Errorf("failed to create GCS storage: %w", err) } @@ -144,38 +146,23 @@ func (f *StorageFactory) GetStorage() Storage { return f.storage } -// Create creates a new storage instance based on the storage type -// This is the factory method equivalent to Python's StorageFactory.create() -//func (f *StorageFactory) Create(storageType StorageType) (Storage, error) { -// var storage Storage -// var err error -// -// switch storageType { -// case StorageMinio: -// storage, err = NewMinioStorage(f.config.Minio) -// if err != nil { -// return nil, fmt.Errorf("MinIO config not available: %w, %v", err, f.config.Minio) -// } -// case StorageAWSS3: -// storage, err = NewS3Storage(f.config.S3) -// if err != nil { -// return nil, fmt.Errorf("S3 config not available: %w, %v", err, f.config.S3) -// } -// case StorageOSS: -// storage, err = NewOSSStorage(f.config.OSS) -// if err != nil { -// return nil, fmt.Errorf("OSS config not available: %w, %v", err, f.config.OSS) -// } -// default: -// return nil, fmt.Errorf("unsupported storage type: %v", storageType) -// } -// -// return storage, nil -//} - // SetStorage sets the storage instance (useful for testing) func (f *StorageFactory) SetStorage(storage Storage) { f.mu.Lock() defer f.mu.Unlock() f.storage = storage } + +func sleepOrAbort(ctx context.Context, d time.Duration) error { + if err := ctx.Err(); err != nil { + return err + } + timer := time.NewTimer(d) + defer timer.Stop() + select { + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: + return nil + } +} diff --git a/internal/storage/types.go b/internal/storage/types.go index bf778db1fe..d9d6857ae0 100644 --- a/internal/storage/types.go +++ b/internal/storage/types.go @@ -17,6 +17,7 @@ package storage import ( + "context" "time" ) @@ -59,43 +60,43 @@ type Storage interface { Type() string // Health checks the storage service availability - Health() bool + Health(ctx context.Context) bool // Put uploads an object to storage // bucket: the bucket/container name // fnm: the file/object name (key) // binary: the data to upload // tenantID: optional tenant identifier - Put(bucket, fnm string, binary []byte, tenantID ...string) error + Put(ctx context.Context, bucket, fnm string, binary []byte, tenantID ...string) error // Get retrieves an object from storage // Returns the data or nil if not found - Get(bucket, fnm string, tenantID ...string) ([]byte, error) + Get(ctx context.Context, bucket, fnm string, tenantID ...string) ([]byte, error) // Remove removes an object from storage - Remove(bucket, fnm string, tenantID ...string) error + Remove(ctx context.Context, bucket, fnm string, tenantID ...string) error // ObjExist checks if an object exists - ObjExist(bucket, fnm string, tenantID ...string) bool + ObjExist(ctx context.Context, bucket, fnm string, tenantID ...string) bool // ListObjects list all objects of the bucket - ListObjects(bucket string, tenantID ...string) ([]string, error) + ListObjects(ctx context.Context, bucket string, tenantID ...string) ([]string, error) // GetPresignedURL generates a presigned URL for accessing an object // expires: duration until the URL expires - GetPresignedURL(bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) + GetPresignedURL(ctx context.Context, bucket, fnm string, expires time.Duration, tenantID ...string) (string, error) // BucketExists checks if a bucket exists - BucketExists(bucket string) bool + BucketExists(ctx context.Context, bucket string) bool // RemoveBucket removes a bucket and all its objects - RemoveBucket(bucket string) error + RemoveBucket(ctx context.Context, bucket string) error // Copy copies an object from source to destination - Copy(srcBucket, srcPath, destBucket, destPath string) bool + Copy(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool // Move moves an object from source to destination - Move(srcBucket, srcPath, destBucket, destPath string) bool + Move(ctx context.Context, srcBucket, srcPath, destBucket, destPath string) bool // Close closes the storage connection Close() error