diff --git a/internal/dao/tenant.go b/internal/dao/tenant.go index 8ce77f68c3..657b7f8987 100644 --- a/internal/dao/tenant.go +++ b/internal/dao/tenant.go @@ -17,7 +17,10 @@ package dao import ( + "context" "ragflow/internal/entity" + + "gorm.io/gorm" ) // TenantDAO tenant data access object @@ -29,10 +32,10 @@ func NewTenantDAO() *TenantDAO { } // GetJoinedTenantsByUserID get joined tenants by user ID -func (dao *TenantDAO) GetJoinedTenantsByUserID(userID string) ([]*TenantWithRole, error) { +func (dao *TenantDAO) GetJoinedTenantsByUserID(ctx context.Context, db *gorm.DB, userID string) ([]*TenantWithRole, error) { var results []*TenantWithRole - err := DB.Model(&entity.Tenant{}). + err := db.WithContext(ctx).Model(&entity.Tenant{}). Select("tenant.id as tenant_id, tenant.name, tenant.llm_id, tenant.embd_id, tenant.asr_id, tenant.img2txt_id, user_tenant.role"). Joins("INNER JOIN user_tenant ON user_tenant.tenant_id = tenant.id"). Where("user_tenant.user_id = ? AND user_tenant.status = ? AND user_tenant.role = ? AND tenant.status = ?", userID, "1", "normal", "1"). @@ -75,10 +78,10 @@ type TenantInfo struct { } // GetInfoByUserID get tenant information for the owner tenant of a user -func (dao *TenantDAO) GetInfoByUserID(userID string) ([]*TenantInfo, error) { +func (dao *TenantDAO) GetInfoByUserID(ctx context.Context, db *gorm.DB, userID string) ([]*TenantInfo, error) { var results []*TenantInfo - err := DB.Model(&entity.Tenant{}). + err := db.WithContext(ctx).Model(&entity.Tenant{}). Select("tenant.id as tenant_id, tenant.name, tenant.llm_id, tenant.tenant_llm_id, tenant.embd_id, tenant.tenant_embd_id, tenant.rerank_id, tenant.tenant_rerank_id, tenant.asr_id, tenant.tenant_asr_id, tenant.img2txt_id, tenant.tenant_img2txt_id, tenant.tts_id, tenant.tenant_tts_id, tenant.ocr_id, tenant.tenant_ocr_id, tenant.parser_ids, user_tenant.role"). Joins("INNER JOIN user_tenant ON user_tenant.tenant_id = tenant.id"). Where("user_tenant.user_id = ? AND user_tenant.status = ? AND user_tenant.role = ? AND tenant.status = ?", userID, "1", "owner", "1"). @@ -88,9 +91,9 @@ func (dao *TenantDAO) GetInfoByUserID(userID string) ([]*TenantInfo, error) { } // GetByID gets tenant by ID -func (dao *TenantDAO) GetByID(id string) (*entity.Tenant, error) { +func (dao *TenantDAO) GetByID(ctx context.Context, db *gorm.DB, id string) (*entity.Tenant, error) { var tenant entity.Tenant - err := DB.Where("id = ? AND status = ?", id, "1").First(&tenant).Error + err := db.WithContext(ctx).Where("id = ? AND status = ?", id, "1").First(&tenant).Error if err != nil { return nil, err } @@ -98,21 +101,21 @@ func (dao *TenantDAO) GetByID(id string) (*entity.Tenant, error) { } // Create creates a new tenant -func (dao *TenantDAO) Create(tenant *entity.Tenant) error { - return DB.Create(tenant).Error +func (dao *TenantDAO) Create(ctx context.Context, db *gorm.DB, tenant *entity.Tenant) error { + return db.WithContext(ctx).Create(tenant).Error } // Delete deletes a tenant by ID (soft delete) -func (dao *TenantDAO) Delete(id string) error { - return DB.Model(&entity.Tenant{}).Where("id = ?", id).Update("status", "0").Error +func (dao *TenantDAO) Delete(ctx context.Context, db *gorm.DB, id string) error { + return db.WithContext(ctx).Model(&entity.Tenant{}).Where("id = ?", id).Update("status", "0").Error } // Update updates a tenant by ID -func (dao *TenantDAO) Update(id string, updates map[string]interface{}) error { - return DB.Model(&entity.Tenant{}).Where("id = ?", id).Updates(updates).Error +func (dao *TenantDAO) Update(ctx context.Context, db *gorm.DB, id string, updates map[string]interface{}) error { + return db.WithContext(ctx).Model(&entity.Tenant{}).Where("id = ?", id).Updates(updates).Error } // HardDelete hard deletes a tenant by ID -func (dao *TenantDAO) HardDelete(id string) error { - return DB.Unscoped().Where("id = ?", id).Delete(&entity.Tenant{}).Error +func (dao *TenantDAO) HardDelete(ctx context.Context, db *gorm.DB, id string) error { + return db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&entity.Tenant{}).Error } diff --git a/internal/dao/tenant_test.go b/internal/dao/tenant_test.go index fe5ea3cd12..40a6f371d7 100644 --- a/internal/dao/tenant_test.go +++ b/internal/dao/tenant_test.go @@ -47,6 +47,7 @@ func useTenantDAOTestDB(t *testing.T, db *gorm.DB) { func TestTenantDAODeleteSoftDeletesTenant(t *testing.T) { db := setupTenantDAOTestDB(t) useTenantDAOTestDB(t, db) + ctx := t.Context() active := "1" tenant := &entity.Tenant{ @@ -59,22 +60,22 @@ func TestTenantDAODeleteSoftDeletesTenant(t *testing.T) { ParserIDs: "naive", Status: &active, } - if err := NewTenantDAO().Create(tenant); err != nil { + if err := NewTenantDAO().Create(ctx, db, tenant); err != nil { t.Fatalf("Create() error = %v", err) } - if err := NewTenantDAO().Delete(tenant.ID); err != nil { + if err := NewTenantDAO().Delete(ctx, db, tenant.ID); err != nil { t.Fatalf("Delete() error = %v", err) } var got entity.Tenant - if err := db.Where("id = ?", tenant.ID).First(&got).Error; err != nil { + if err := db.WithContext(ctx).Where("id = ?", tenant.ID).First(&got).Error; err != nil { t.Fatalf("failed to reload tenant: %v", err) } if got.Status == nil || *got.Status != "0" { t.Fatalf("status = %v, want 0", got.Status) } - if _, err := NewTenantDAO().GetByID(tenant.ID); err == nil { + if _, err := NewTenantDAO().GetByID(ctx, db, tenant.ID); err == nil { t.Fatalf("GetByID() after Delete() error = nil, want not found") } } @@ -82,6 +83,7 @@ func TestTenantDAODeleteSoftDeletesTenant(t *testing.T) { func TestTenantDAOUpdateStatus(t *testing.T) { db := setupTenantDAOTestDB(t) useTenantDAOTestDB(t, db) + ctx := t.Context() active := "1" tenant := &entity.Tenant{ @@ -94,16 +96,16 @@ func TestTenantDAOUpdateStatus(t *testing.T) { ParserIDs: "naive", Status: &active, } - if err := NewTenantDAO().Create(tenant); err != nil { + if err := NewTenantDAO().Create(ctx, db, tenant); err != nil { t.Fatalf("Create() error = %v", err) } - if err := NewTenantDAO().Update(tenant.ID, map[string]interface{}{"status": "0"}); err != nil { + if err := NewTenantDAO().Update(ctx, db, tenant.ID, map[string]interface{}{"status": "0"}); err != nil { t.Fatalf("Update() error = %v", err) } var got entity.Tenant - if err := db.Where("id = ?", tenant.ID).First(&got).Error; err != nil { + if err := db.WithContext(ctx).Where("id = ?", tenant.ID).First(&got).Error; err != nil { t.Fatalf("failed to reload tenant: %v", err) } if got.Status == nil || *got.Status != "0" { diff --git a/internal/handler/dataset.go b/internal/handler/dataset.go index 672e255cc0..7c949d1882 100644 --- a/internal/handler/dataset.go +++ b/internal/handler/dataset.go @@ -78,7 +78,8 @@ func (h *DatasetsHandler) ListDatasets(c *gin.Context) { } if c.Query("type") == "filter" { - data, code, err := h.datasetsService.ListDatasetFilters(user.ID) + ctx := c.Request.Context() + data, code, err := h.datasetsService.ListDatasetFilters(ctx, user.ID) if err != nil { common.ErrorWithCode(c, code, err.Error()) return diff --git a/internal/handler/document.go b/internal/handler/document.go index 5db4f9b30a..e29e569cea 100644 --- a/internal/handler/document.go +++ b/internal/handler/document.go @@ -836,7 +836,7 @@ func (h *DocumentHandler) UploadDocuments(c *gin.Context) { common.ResponseWithCodeData(c, common.CodeDataError, nil, fmt.Sprintf("Can't find the dataset with ID %s!", datasetID)) return } - if !h.datasetService.CheckKBTeamPermission(kb, tenantID) { + if !h.datasetService.CheckKBTeamPermission(ctx, kb, tenantID) { common.ResponseWithCodeData(c, common.CodeAuthenticationError, nil, "No authorization.") return } diff --git a/internal/handler/mindmap.go b/internal/handler/mindmap.go index 11c8c99dd4..0b71743e82 100644 --- a/internal/handler/mindmap.go +++ b/internal/handler/mindmap.go @@ -20,14 +20,13 @@ import ( "context" "encoding/json" "fmt" - "regexp" - "strings" - "time" - "ragflow/internal/common" "ragflow/internal/entity" modelModule "ragflow/internal/entity/models" "ragflow/internal/service" + "regexp" + "strings" + "time" ) type mindMapRunConfig struct { @@ -64,18 +63,19 @@ func runMindMap(ctx context.Context, config mindMapRunConfig) (mindMapNode, erro } modelID, _ := config.SearchConfig["chat_id"].(string) messages := []modelModule.Message{{Role: "system", Content: mindMapPrompt(strings.Join(sections, "\n"))}, {Role: "user", Content: "Output:"}} - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Minute) - defer cancel() + + streamCtx, streamCancel := context.WithTimeout(ctx, 10*time.Minute) + defer streamCancel() // search_config chat_id can be a stale tenant_model ID that no longer // exists. ResolveModelConfig tries ID lookup first, then falls back to // composite-name parsing which fails for bare IDs. If the configured // model can't be resolved, fall back to the tenant's default chat model // (mirrors Python's gen_mindmap get_tenant_default_model_by_type). - ch, streamErr := config.LLM.ChatStream(ctx, modelTenantID, modelID, messages, &modelModule.ChatConfig{}) + ch, streamErr := config.LLM.ChatStream(streamCtx, modelTenantID, modelID, messages, &modelModule.ChatConfig{}) if streamErr != nil && config.TenantSvc != nil { - if defaultModel, err := config.TenantSvc.GetDefaultModelName(modelTenantID, entity.ModelTypeChat); err == nil && defaultModel != "" && defaultModel != modelID { - ch, streamErr = config.LLM.ChatStream(ctx, modelTenantID, defaultModel, messages, &modelModule.ChatConfig{}) + if defaultModel, err := config.TenantSvc.GetDefaultModelName(streamCtx, modelTenantID, entity.ModelTypeChat); err == nil && defaultModel != "" && defaultModel != modelID { + ch, streamErr = config.LLM.ChatStream(streamCtx, modelTenantID, defaultModel, messages, &modelModule.ChatConfig{}) } } if streamErr != nil { diff --git a/internal/handler/providers.go b/internal/handler/providers.go index fafc82ea19..c11a4318ad 100644 --- a/internal/handler/providers.go +++ b/internal/handler/providers.go @@ -1425,8 +1425,9 @@ func (h *ProviderHandler) ListTenantAddedModels(c *gin.Context) { modelType := c.Query("type") ownerTenantID := c.Query("owner_tenant_id") + ctx := c.Request.Context() - addedModels, code, err := h.modelProviderService.ListTenantAddedModels(user.ID, ownerTenantID, modelType) + addedModels, code, err := h.modelProviderService.ListTenantAddedModels(ctx, user.ID, ownerTenantID, modelType) if err != nil { common.ErrorWithCode(c, code, err.Error()) return diff --git a/internal/handler/searchbot.go b/internal/handler/searchbot.go index bc0b95d41c..069fd36558 100644 --- a/internal/handler/searchbot.go +++ b/internal/handler/searchbot.go @@ -260,7 +260,8 @@ func (h *SearchBotHandler) Ask(c *gin.Context) { } } if modelID == "" && h.tenantSvc != nil { - defaultModel, err := h.tenantSvc.GetDefaultModelName(user.ID, entity.ModelTypeChat) + ctx := c.Request.Context() + defaultModel, err := h.tenantSvc.GetDefaultModelName(ctx, user.ID, entity.ModelTypeChat) if err == nil && defaultModel != "" { modelID = defaultModel } diff --git a/internal/handler/tenant.go b/internal/handler/tenant.go index 5849757da6..d2a8af09b2 100644 --- a/internal/handler/tenant.go +++ b/internal/handler/tenant.go @@ -75,7 +75,8 @@ func (h *TenantHandler) setDefaultModels(c *gin.Context, wrapModels bool) { return } - err := h.tenantService.SetTenantDefaultModels(user.ID, req.ModelProvider, req.ModelInstance, req.ModelName, req.ModelType, req.ModelID) + ctx := c.Request.Context() + err := h.tenantService.SetTenantDefaultModels(ctx, user.ID, req.ModelProvider, req.ModelInstance, req.ModelName, req.ModelType, req.ModelID) if err != nil { common.ResponseWithCodeData(c, common.CodeExceptionError, false, err.Error()) return @@ -100,8 +101,9 @@ func (h *TenantHandler) GetDefaultModels(c *gin.Context) { common.ErrorWithCode(c, errorCode, errorMessage) return } + ctx := c.Request.Context() - defaultModels, err := h.tenantService.ListTenantDefaultModels(user.ID) + defaultModels, err := h.tenantService.ListTenantDefaultModels(ctx, user.ID) if err != nil { common.ResponseWithCodeData(c, common.CodeExceptionError, false, err.Error()) return @@ -132,8 +134,9 @@ func (h *TenantHandler) TenantInfo(c *gin.Context) { common.ErrorWithCode(c, errorCode, errorMessage) return } + ctx := c.Request.Context() - tenantInfo, err := h.tenantService.GetTenantInfo(user.ID) + tenantInfo, err := h.tenantService.GetTenantInfo(ctx, user.ID) if err != nil { common.ResponseWithCodeData(c, common.CodeExceptionError, false, err.Error()) return diff --git a/internal/ingestion/component/dispatch_model.go b/internal/ingestion/component/dispatch_model.go index b88fca9334..73d3f5d2a1 100644 --- a/internal/ingestion/component/dispatch_model.go +++ b/internal/ingestion/component/dispatch_model.go @@ -20,6 +20,7 @@ package component import ( + "context" "encoding/json" "fmt" "strings" @@ -37,9 +38,9 @@ type tenantModelExtra struct { var resolveTenantModelByType = defaultResolveTenantModelByType -func defaultResolveTenantModelByType(tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { +func defaultResolveTenantModelByType(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { tenantDAO := dao.NewTenantDAO() - tenant, err := tenantDAO.GetByID(tenantID) + tenant, err := tenantDAO.GetByID(ctx, db, tenantID) if err != nil { return nil, "", nil, 0, err } diff --git a/internal/ingestion/component/docx_vision_dispatch.go b/internal/ingestion/component/docx_vision_dispatch.go index 5c0ac2ad70..24a71a0297 100644 --- a/internal/ingestion/component/docx_vision_dispatch.go +++ b/internal/ingestion/component/docx_vision_dispatch.go @@ -39,6 +39,8 @@ import ( modelModule "ragflow/internal/entity/models" "ragflow/internal/ingestion/component/schema" "ragflow/internal/utility" + + "gorm.io/gorm" ) var ( @@ -75,6 +77,7 @@ var ( // markdown branch only concatenates text and never calls the vision model. func maybeDispatchDOCXVision( ctx context.Context, + db *gorm.DB, fileType utility.FileType, dispatched parserDispatchResult, inputs map[string]any, @@ -95,7 +98,7 @@ func maybeDispatchDOCXVision( } // Resolve the tenant's IMAGE2TEXT model. - driver, modelName, apiConfig, _, err := resolveTenantModelByType(tenantID, entity.ModelTypeImage2Text) + driver, modelName, apiConfig, _, err := resolveTenantModelByType(ctx, db, tenantID, entity.ModelTypeImage2Text) if err != nil { // Model not available — skip vision enhancement silently, // matching Python's try/except pass behaviour. diff --git a/internal/ingestion/component/docx_vision_dispatch_test.go b/internal/ingestion/component/docx_vision_dispatch_test.go index 7554a2adac..11d76ab4f3 100644 --- a/internal/ingestion/component/docx_vision_dispatch_test.go +++ b/internal/ingestion/component/docx_vision_dispatch_test.go @@ -17,12 +17,15 @@ package component import ( "context" + "ragflow/internal/dao" "sync" "testing" "ragflow/internal/entity" modelModule "ragflow/internal/entity/models" "ragflow/internal/utility" + + "gorm.io/gorm" ) // docxVisionFakeDriver satisfies modelModule.ModelDriver but never reaches the @@ -66,7 +69,7 @@ func (c *docxVisionCaptureInvoker) invoke( // TestMaybeDispatchDOCXVision_EnhancesJSONImages verifies Diff 2.4: DOCX vision // enhancement must trigger on the JSON output path (like Python's // enhance_media_sections_with_vision in parser.py:_doc) and must NOT trigger on -// the markdown path. Image items with a non-empty `image` field get their VLM +// the Markdown path. Image items with a non-empty `image` field get their VLM // description appended to `text`; table items (no image) and text items are // left untouched. func TestMaybeDispatchDOCXVision_EnhancesJSONImages(t *testing.T) { @@ -79,7 +82,7 @@ func TestMaybeDispatchDOCXVision_EnhancesJSONImages(t *testing.T) { docxVisionPromptBuilder = origPrompt }() - resolveTenantModelByType = func(tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { + resolveTenantModelByType = func(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { return &docxVisionFakeDriver{}, "docx-vision-model", &modelModule.APIConfig{}, 0, nil } invoker := &docxVisionCaptureInvoker{} @@ -95,9 +98,11 @@ func TestMaybeDispatchDOCXVision_EnhancesJSONImages(t *testing.T) { {"text": "
", "image": nil, "doc_type_kwd": "table"}, }, } + ctx := t.Context() res, handled, err := maybeDispatchDOCXVision( - context.Background(), + ctx, + dao.DB, utility.FileTypeDOCX, dispatched, map[string]any{"tenant_id": "t1"}, @@ -142,7 +147,7 @@ func TestMaybeDispatchDOCXVision_JSONOnly(t *testing.T) { }() called := false - resolveTenantModelByType = func(tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { + resolveTenantModelByType = func(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { called = true return &docxVisionFakeDriver{}, "m", &modelModule.APIConfig{}, 0, nil } @@ -158,9 +163,11 @@ func TestMaybeDispatchDOCXVision_JSONOnly(t *testing.T) { Markdown: "![Image](data:image/png;base64,abc)", File: map[string]any{"figures": []map[string]any{{"image": "abc", "marker": "x"}}}, } + ctx := t.Context() res, handled, err := maybeDispatchDOCXVision( - context.Background(), + ctx, + dao.DB, utility.FileTypeDOCX, dispatched, map[string]any{"tenant_id": "t1"}, diff --git a/internal/ingestion/component/markdown_vision_dispatch.go b/internal/ingestion/component/markdown_vision_dispatch.go index 069900b7df..4eff26ec6f 100644 --- a/internal/ingestion/component/markdown_vision_dispatch.go +++ b/internal/ingestion/component/markdown_vision_dispatch.go @@ -35,6 +35,8 @@ import ( "ragflow/internal/entity" "ragflow/internal/utility" + + "gorm.io/gorm" ) var ( @@ -55,6 +57,7 @@ var ( // "image" base64 field. It returns (result, handled, error). func maybeDispatchMarkdownVision( ctx context.Context, + db *gorm.DB, fileType utility.FileType, dispatched parserDispatchResult, inputs map[string]any, @@ -101,7 +104,7 @@ func maybeDispatchMarkdownVision( } // Resolve the tenant's IMAGE2TEXT model. - driver, modelName, apiConfig, _, err := resolveTenantModelByType(tenantID, entity.ModelTypeImage2Text) + driver, modelName, apiConfig, _, err := resolveTenantModelByType(ctx, db, tenantID, entity.ModelTypeImage2Text) if err != nil { // Model not available — skip vision enhancement silently, // matching Python's try/except pass behaviour. diff --git a/internal/ingestion/component/media_dispatch.go b/internal/ingestion/component/media_dispatch.go index e627c125d0..8d184bbb27 100644 --- a/internal/ingestion/component/media_dispatch.go +++ b/internal/ingestion/component/media_dispatch.go @@ -47,12 +47,15 @@ import ( "ragflow/internal/ingestion/component/schema" "ragflow/internal/parser/parser" "ragflow/internal/utility" + + "gorm.io/gorm" ) // Video dispatch: IMAGE2TEXT vision chat --- func maybeDispatchVideo( ctx context.Context, + db *gorm.DB, fileType utility.FileType, filename string, binary []byte, @@ -69,14 +72,14 @@ func maybeDispatchVideo( tenantID := getStringOr(inputs, "tenant_id", "") if tenantID == "" { return parserDispatchResult{}, true, - fmt.Errorf("Parser: video requires tenant_id") + fmt.Errorf("parser: video requires tenant_id") } // Resolve the tenant's IMAGE2TEXT model. - driver, modelName, apiConfig, _, err := resolveTenantModelByType(tenantID, entity.ModelTypeImage2Text) + driver, modelName, apiConfig, _, err := resolveTenantModelByType(ctx, db, tenantID, entity.ModelTypeImage2Text) if err != nil { return parserDispatchResult{}, true, - fmt.Errorf("Parser: video image2text model: %w", err) + fmt.Errorf("parser: video image2text model: %w", err) } videoPrompt, _ := setup["prompt"].(string) @@ -98,7 +101,7 @@ func maybeDispatchVideo( resp, err := driver.ChatWithMessages(ctx, modelName, messages, apiConfig, &modelModule.ChatConfig{Vision: &vision}, nil) if err != nil { return parserDispatchResult{}, true, - fmt.Errorf("Parser: video describe: %w", err) + fmt.Errorf("parser: video describe: %w", err) } txt := "" if resp != nil && resp.Answer != nil { @@ -126,6 +129,7 @@ func maybeDispatchVideo( func maybeDispatchImage( ctx context.Context, + db *gorm.DB, fileType utility.FileType, filename string, binary []byte, @@ -142,7 +146,7 @@ func maybeDispatchImage( tenantID := getStringOr(inputs, "tenant_id", "") if tenantID == "" { return parserDispatchResult{}, true, - fmt.Errorf("Parser: image requires tenant_id") + fmt.Errorf("parser: image requires tenant_id") } // --- Phase 1: OCR --- @@ -196,14 +200,14 @@ func maybeDispatchImage( } // Short OCR text (or no text): supplement with VLM describe. - driver, modelName, apiConfig, _, err := resolveTenantModelByType(tenantID, entity.ModelTypeImage2Text) + driver, modelName, apiConfig, _, err := resolveTenantModelByType(ctx, db, tenantID, entity.ModelTypeImage2Text) if err != nil { - // If VLM is unavailable but we have OCR text, return it. + // If VLM is unavailable, but we have OCR text, return it. if ocrText != "" { return imageDispatchResult(ocrText, dataURI), true, nil } return parserDispatchResult{}, true, - fmt.Errorf("Parser: picture image2text model: %w", err) + fmt.Errorf("parser: picture image2text model: %w", err) } prompt := "Describe this image in detail." @@ -227,7 +231,7 @@ func maybeDispatchImage( return imageDispatchResult(ocrText, dataURI), true, nil } return parserDispatchResult{}, true, - fmt.Errorf("Parser: picture describe: %w", err) + fmt.Errorf("parser: picture describe: %w", err) } vlmText := "" if resp != nil && resp.Answer != nil { @@ -271,6 +275,7 @@ func imageDispatchResult(text, dataURI string) parserDispatchResult { func maybeDispatchAudio( ctx context.Context, + db *gorm.DB, fileType utility.FileType, filename string, binary []byte, @@ -287,19 +292,19 @@ func maybeDispatchAudio( tenantID := getStringOr(inputs, "tenant_id", "") if tenantID == "" { return parserDispatchResult{}, true, - fmt.Errorf("Parser: audio requires tenant_id") + fmt.Errorf("parser: audio requires tenant_id") } - driver, modelName, apiConfig, _, err := resolveTenantModelByType(tenantID, entity.ModelTypeSpeech2Text) + driver, modelName, apiConfig, _, err := resolveTenantModelByType(ctx, db, tenantID, entity.ModelTypeSpeech2Text) if err != nil { return parserDispatchResult{}, true, - fmt.Errorf("Parser: audio speech2text model: %w", err) + fmt.Errorf("parser: audio speech2text model: %w", err) } tmpFile, err := writeTempAudioFile(filename, binary) if err != nil { return parserDispatchResult{}, true, - fmt.Errorf("Parser: audio temp file: %w", err) + fmt.Errorf("parser: audio temp file: %w", err) } defer os.Remove(tmpFile) diff --git a/internal/ingestion/component/media_dispatch_test.go b/internal/ingestion/component/media_dispatch_test.go index 1d684e44de..010dc0888b 100644 --- a/internal/ingestion/component/media_dispatch_test.go +++ b/internal/ingestion/component/media_dispatch_test.go @@ -17,6 +17,7 @@ package component import ( "context" + "ragflow/internal/dao" "strings" "sync" "testing" @@ -25,6 +26,8 @@ import ( "ragflow/internal/entity" modelModule "ragflow/internal/entity/models" "ragflow/internal/utility" + + "gorm.io/gorm" ) // imagePromptCaptureDriver embeds ModelDriver so it satisfies the interface @@ -79,7 +82,7 @@ func TestMaybeDispatchImage_UsesSystemPrompt(t *testing.T) { defer func() { resolveTenantModelByType = origResolver }() drv := &imagePromptCaptureDriver{} - resolveTenantModelByType = func(tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { + resolveTenantModelByType = func(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { return drv, "img-model", &modelModule.APIConfig{}, 0, nil } @@ -91,8 +94,10 @@ func TestMaybeDispatchImage_UsesSystemPrompt(t *testing.T) { setups["image"]["prompt"] = "legacy prompt" setups["image"]["system_prompt"] = "自定义视觉提示" + ctx := t.Context() res, dispatched, err := maybeDispatchImage( - context.Background(), + ctx, + dao.DB, utility.FileTypeVISUAL, "test.png", []byte("not-a-real-image"), @@ -139,13 +144,15 @@ func TestMaybeDispatchImage_ReturnsJSONWithImage(t *testing.T) { defer func() { resolveTenantModelByType = origResolver }() drv := &imagePromptCaptureDriver{} - resolveTenantModelByType = func(tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { + resolveTenantModelByType = func(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { return drv, "img-model", &modelModule.APIConfig{}, 0, nil } setups := defaultSetups() + ctx := t.Context() res, dispatched, err := maybeDispatchImage( - context.Background(), + ctx, + dao.DB, utility.FileTypeVISUAL, "test.png", []byte("not-a-real-image"), @@ -187,14 +194,16 @@ func TestMaybeDispatchImage_HardcodesJSONOutput(t *testing.T) { defer func() { resolveTenantModelByType = origResolver }() drv := &imagePromptCaptureDriver{} - resolveTenantModelByType = func(tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { + resolveTenantModelByType = func(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { return drv, "img-model", &modelModule.APIConfig{}, 0, nil } setups := defaultSetups() setups["image"]["output_format"] = "text" // legacy/override; must be ignored + ctx := t.Context() res, _, err := maybeDispatchImage( - context.Background(), + ctx, + dao.DB, utility.FileTypeVISUAL, "test.png", []byte("not-a-real-image"), @@ -235,15 +244,17 @@ func TestMaybeDispatchAudio_JSONCarriesTranscription(t *testing.T) { const want = "hello world" drv := &audioTranscribeDriver{transcription: want} - resolveTenantModelByType = func(tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { + resolveTenantModelByType = func(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { return drv, "asr-model", &modelModule.APIConfig{}, 0, nil } setups := defaultSetups() setups["audio"]["output_format"] = "json" + ctx := t.Context() res, dispatched, err := maybeDispatchAudio( - context.Background(), + ctx, + dao.DB, utility.FileTypeAURAL, "test.mp3", []byte("fake-audio"), @@ -279,15 +290,17 @@ func TestMaybeDispatchAudio_TextCarriesTranscription(t *testing.T) { const want = "hello world" drv := &audioTranscribeDriver{transcription: want} - resolveTenantModelByType = func(tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { + resolveTenantModelByType = func(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { return drv, "asr-model", &modelModule.APIConfig{}, 0, nil } setups := defaultSetups() setups["audio"]["output_format"] = "text" + ctx := t.Context() res, dispatched, err := maybeDispatchAudio( - context.Background(), + ctx, + dao.DB, utility.FileTypeAURAL, "test.mp3", []byte("fake-audio"), @@ -320,7 +333,7 @@ func TestMaybeDispatchMarkdownVision_EnhancesTables(t *testing.T) { defer func() { resolveTenantModelByType = origResolver }() drv := &imagePromptCaptureDriver{} - resolveTenantModelByType = func(tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { + resolveTenantModelByType = func(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { return drv, "img-model", &modelModule.APIConfig{}, 0, nil } @@ -331,8 +344,10 @@ func TestMaybeDispatchMarkdownVision_EnhancesTables(t *testing.T) { }, } + ctx := t.Context() res, handled, err := maybeDispatchMarkdownVision( - context.Background(), + ctx, + dao.DB, utility.FileTypeMarkdown, dispatched, map[string]any{"tenant_id": "t1"}, diff --git a/internal/ingestion/component/parser.go b/internal/ingestion/component/parser.go index 8724424ab1..8659b9c092 100644 --- a/internal/ingestion/component/parser.go +++ b/internal/ingestion/component/parser.go @@ -443,7 +443,7 @@ func (c *ParserComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map[st } } - dispatched, handledVision, visionErr := maybeDispatchPDFVision(ctx, fileTypeExt, filename, binary, inputs, c.Setups) + dispatched, handledVision, visionErr := maybeDispatchPDFVision(ctx, db, fileTypeExt, filename, binary, inputs, c.Setups) if visionErr != nil { return nil, visionErr } @@ -452,7 +452,7 @@ func (c *ParserComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map[st if !handledVision { // Video dispatch: IMAGE2TEXT vision chat. // Mirrors Python's _video(). - dispatched, handledMedia, visionErr = maybeDispatchVideo(ctx, fileTypeExt, filename, binary, inputs, c.Setups) + dispatched, handledMedia, visionErr = maybeDispatchVideo(ctx, db, fileTypeExt, filename, binary, inputs, c.Setups) if visionErr != nil { return nil, visionErr } @@ -461,7 +461,7 @@ func (c *ParserComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map[st if !handledVision && !handledMedia { // Image/Picture dispatch: OCR + IMAGE2TEXT vision describe. // Mirrors Python's rag/app/picture.py:chunk() image branch. - dispatched, handledImage, visionErr = maybeDispatchImage(ctx, fileTypeExt, filename, binary, inputs, c.Setups) + dispatched, handledImage, visionErr = maybeDispatchImage(ctx, db, fileTypeExt, filename, binary, inputs, c.Setups) if visionErr != nil { return nil, visionErr } @@ -470,7 +470,7 @@ func (c *ParserComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map[st if !handledVision && !handledMedia && !handledImage { // Audio dispatch: SPEECH2TEXT transcription. // Mirrors Python's rag/app/audio.py:chunk(). - dispatched, handledAudio, visionErr = maybeDispatchAudio(ctx, fileTypeExt, filename, binary, inputs, c.Setups) + dispatched, handledAudio, visionErr = maybeDispatchAudio(ctx, db, fileTypeExt, filename, binary, inputs, c.Setups) if visionErr != nil { return nil, visionErr } @@ -483,13 +483,13 @@ func (c *ParserComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map[st // append vision-model descriptions to embedded image items // (doc_type_kwd "image"). Mirrors Python's // enhance_media_sections_with_vision in parser.py:_doc. - dispatched, _, _ = maybeDispatchDOCXVision(ctx, fileTypeExt, dispatched, inputs, c.Setups) + dispatched, _, _ = maybeDispatchDOCXVision(ctx, db, fileTypeExt, dispatched, inputs, c.Setups) // Markdown vision figure enhancement: enrich parsed - // markdown JSON items with LLM-generated descriptions of + // Markdown JSON items with LLM-generated descriptions of // referenced images (![alt](url)). Mirrors Python's - // enhance_media_sections_with_vision in _markdown. - dispatched, _, _ = maybeDispatchMarkdownVision(ctx, fileTypeExt, dispatched, inputs) + // enhance_media_sections_with_vision in _Markdown. + dispatched, _, _ = maybeDispatchMarkdownVision(ctx, db, fileTypeExt, dispatched, inputs) } // Known/supported families must fail loudly when dispatch or // parsing breaks. Only unknown families keep the raw-text fallback. @@ -536,7 +536,7 @@ func (c *ParserComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map[st // downstream chunker / tokenizer get stable chunk IDs. parsed, err := buildPagesFromBytes(ctx, pages, dispatched.DocType) if err != nil { - return nil, fmt.Errorf("Parser: %w", err) + return nil, fmt.Errorf("parser: %w", err) } sortPagesByNumber(parsed) lang, _ := getString(inputs, "lang") @@ -613,7 +613,7 @@ func buildPagesFromBytes(ctx context.Context, pages [][]byte, docType string) ([ // map. The accepted shapes are: // // []byte — the in-process caller's normal form -// string — UTF-8 text (json callers' normal form) +// string — UTF-8 text (JSON callers' normal form) // nil / absent — returns an empty page (not an error) // // A non-UTF-8 string is rejected with a clear error so a caller @@ -642,7 +642,7 @@ func readParserBinary(ctx context.Context, db *gorm.DB, inputs map[string]any) ( if docID, ok := getString(inputs, "doc_id"); ok && docID != "" { ref, err := ResolveDocumentStorage(ctx, db, docID) if err != nil { - return nil, fmt.Errorf("Parser: resolve doc_id %q: %w", docID, err) + return nil, fmt.Errorf("parser: resolve doc_id %q: %w", docID, err) } return FetchBinary(ctx, ref.Bucket, ref.Path) } diff --git a/internal/ingestion/component/parser_dispatch_pdf_vision_cgo_test.go b/internal/ingestion/component/parser_dispatch_pdf_vision_cgo_test.go index 62d6bae540..043a8a88c6 100644 --- a/internal/ingestion/component/parser_dispatch_pdf_vision_cgo_test.go +++ b/internal/ingestion/component/parser_dispatch_pdf_vision_cgo_test.go @@ -12,6 +12,8 @@ import ( modelModule "ragflow/internal/entity/models" "ragflow/internal/ingestion/component/schema" + + "gorm.io/gorm" ) func TestDispatch_PDFVisionJSON_RealPDFFixture(t *testing.T) { @@ -27,7 +29,7 @@ func TestDispatch_PDFVisionJSON_RealPDFFixture(t *testing.T) { pdfVisionPromptLoader = func(name string) (string, error) { return "Describe page {{ page }}.", nil } - pdfVisionModelResolver = func(tenantID string, modelID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, error) { + pdfVisionModelResolver = func(ctx context.Context, db *gorm.DB, tenantID string, modelID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, error) { if tenantID != "tenant-vision" || modelID != "CustomVLM" { t.Fatalf("resolver got tenant/model %q/%q", tenantID, modelID) } diff --git a/internal/ingestion/component/parser_dispatch_test.go b/internal/ingestion/component/parser_dispatch_test.go index 3506b37cec..b0e866c167 100644 --- a/internal/ingestion/component/parser_dispatch_test.go +++ b/internal/ingestion/component/parser_dispatch_test.go @@ -43,6 +43,8 @@ import ( modelModule "ragflow/internal/entity/models" "ragflow/internal/ingestion/component/schema" "ragflow/internal/utility" + + "gorm.io/gorm" ) type captureSetupConfigurer struct { @@ -424,7 +426,7 @@ func TestDispatch_PDFVisionJSON_UsesTenantAwareModel(t *testing.T) { {PageNumber: 2, WidthPts: 120, HeightPts: 240, ImageURL: "data:image/png;base64,bbb"}, }, nil } - pdfVisionModelResolver = func(tenantID string, modelID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, error) { + pdfVisionModelResolver = func(ctx context.Context, db *gorm.DB, tenantID string, modelID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, error) { if tenantID != "tenant-1" || modelID != "CustomVLM" { return nil, "", nil, fmt.Errorf("resolver got tenant/model %q/%q", tenantID, modelID) } @@ -504,7 +506,7 @@ func TestDispatch_PDFVisionJSON_PreservesEmptyPages(t *testing.T) { {PageNumber: 2, WidthPts: 120, HeightPts: 240, ImageURL: "data:image/png;base64,bbb"}, }, nil } - pdfVisionModelResolver = func(string, string) (modelModule.ModelDriver, string, *modelModule.APIConfig, error) { + pdfVisionModelResolver = func(ctx context.Context, db *gorm.DB, tenantID string, modelID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, error) { return nil, "resolved-vlm", nil, nil } call := 0 @@ -562,7 +564,7 @@ func TestDispatch_PDFMinerUMarkdown_UsesConfiguredBackend(t *testing.T) { defer func() { resolveTenantModelByType = origResolver }() baseURL := server.URL apiKey := "" - resolveTenantModelByType = func(tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { + resolveTenantModelByType = func(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { return &mineruTestDriver{}, "mineru-model", &modelModule.APIConfig{ApiKey: &apiKey, BaseURL: &baseURL}, 0, nil } diff --git a/internal/ingestion/component/pdf_vision_dispatch.go b/internal/ingestion/component/pdf_vision_dispatch.go index 3bd6fbf04a..b285e154f1 100644 --- a/internal/ingestion/component/pdf_vision_dispatch.go +++ b/internal/ingestion/component/pdf_vision_dispatch.go @@ -19,6 +19,8 @@ import ( modelModule "ragflow/internal/entity/models" "ragflow/internal/ingestion/component/schema" "ragflow/internal/utility" + + "gorm.io/gorm" ) type pdfVisionPage struct { @@ -44,6 +46,7 @@ var ( func maybeDispatchPDFVision( ctx context.Context, + db *gorm.DB, fileType utility.FileType, filename string, binary []byte, @@ -69,9 +72,9 @@ func maybeDispatchPDFVision( tenantID := getStringOr(inputs, "tenant_id", "") if tenantID == "" { return parserDispatchResult{}, true, - fmt.Errorf("Parser: mineru requires tenant_id") + fmt.Errorf("parser: MinerU requires tenant_id") } - res, err := dispatchMinerUPDF(filename, binary, tenantID, setup) + res, err := dispatchMinerUPDF(ctx, db, filename, binary, tenantID, setup) if err != nil { return parserDispatchResult{}, true, err } @@ -85,9 +88,9 @@ func maybeDispatchPDFVision( tenantID := getStringOr(inputs, "tenant_id", "") if tenantID == "" { return parserDispatchResult{}, true, fmt.Errorf( - `Parser: pdf parse_method %q requires tenant_id to resolve IMAGE2TEXT model`, modelID) + `parser: pdf parse_method %q requires tenant_id to resolve VLM model`, modelID) } - res, err := dispatchPDFVision(ctx, filename, binary, tenantID, modelID, setup) + res, err := dispatchPDFVision(ctx, db, filename, binary, tenantID, modelID, setup) if err != nil { return parserDispatchResult{}, true, err } @@ -96,21 +99,23 @@ func maybeDispatchPDFVision( // dispatchMinerUPDF submits a PDF to the tenant's MinerU OCR model // via the streaming /file_parse endpoint and returns parsed sections. -// Mirrors Python's mineru_parser.py:parse_pdf which POSTs with +// Mirrors Python's mineru_parser.py:parse_PDF which POSTs with // stream=True and reads the zip response body directly (no polling). func dispatchMinerUPDF( + ctx context.Context, + db *gorm.DB, _ string, binary []byte, tenantID string, setup schema.ParserSetup, ) (parserDispatchResult, error) { - driver, _, apiConfig, _, err := resolveTenantModelByType(tenantID, entity.ModelTypeOCR) + driver, _, apiConfig, _, err := resolveTenantModelByType(ctx, db, tenantID, entity.ModelTypeOCR) if err != nil { - return parserDispatchResult{}, fmt.Errorf("Parser: mineru model: %w", err) + return parserDispatchResult{}, fmt.Errorf("parser: MinerU model: %w", err) } if !isMinerUDriver(driver) { return parserDispatchResult{}, fmt.Errorf( - "Parser: mineru requires a MinerU OCR model; found %q. Please add a MinerU OCR model to your tenant.", driver.Name()) + "parser: MinerU requires a MinerU OCR model; found %q. Please add a MinerU OCR model to your tenant", driver.Name()) } baseURL := "" @@ -130,12 +135,12 @@ func dispatchMinerUPDF( zipBytes, err := mineruStreamParse(apiURL, apiConfig.ApiKey, binary, parseMethod, mineruLang, backend) if err != nil { - return parserDispatchResult{}, fmt.Errorf("Parser: mineru stream: %w", err) + return parserDispatchResult{}, fmt.Errorf("parser: MinerU stream: %w", err) } sections, err := mineruExtractSections(zipBytes) if err != nil { - return parserDispatchResult{}, fmt.Errorf("Parser: mineru extract: %w", err) + return parserDispatchResult{}, fmt.Errorf("parser: MinerU extract: %w", err) } var parts []string @@ -277,7 +282,7 @@ func mineruExtractSections(zipBytes []byte) ([]string, error) { } var items []map[string]any - if err := json.Unmarshal(contentList, &items); err != nil { + if err = json.Unmarshal(contentList, &items); err != nil { return nil, fmt.Errorf("parse content_list.json: %w", err) } @@ -293,12 +298,12 @@ func mineruExtractSections(zipBytes []byte) ([]string, error) { if tb, ok := item["table_body"].(string); ok { sections = append(sections, tb) } - for _, cap := range stringSlice(item["table_caption"]) { - sections = append(sections, cap) + for _, caption := range stringSlice(item["table_caption"]) { + sections = append(sections, caption) } case "image": - for _, cap := range stringSlice(item["image_caption"]) { - sections = append(sections, cap) + for _, caption := range stringSlice(item["image_caption"]) { + sections = append(sections, caption) } if desc, ok := item["vlm_description"].(string); ok && desc != "" { sections = append(sections, desc) @@ -386,6 +391,7 @@ func isNamedPDFParseMethod(raw string) bool { func dispatchPDFVision( ctx context.Context, + db *gorm.DB, filename string, binary []byte, tenantID string, @@ -394,15 +400,15 @@ func dispatchPDFVision( ) (parserDispatchResult, error) { renderedPages, err := pdfVisionPageRenderer(binary) if err != nil { - return parserDispatchResult{}, fmt.Errorf("Parser: pdf vision render: %w", err) + return parserDispatchResult{}, fmt.Errorf("parser: pdf vision render: %w", err) } - driver, resolvedModelName, apiConfig, err := pdfVisionModelResolver(tenantID, modelID) + driver, resolvedModelName, apiConfig, err := pdfVisionModelResolver(ctx, db, tenantID, modelID) if err != nil { - return parserDispatchResult{}, fmt.Errorf("Parser: pdf vision model %q: %w", modelID, err) + return parserDispatchResult{}, fmt.Errorf("parser: pdf vision model %q: %w", modelID, err) } promptTemplate, err := pdfVisionPromptLoader("vision_llm_describe_prompt") if err != nil { - return parserDispatchResult{}, fmt.Errorf("Parser: load vision prompt: %w", err) + return parserDispatchResult{}, fmt.Errorf("parser: load vision prompt: %w", err) } items := make([]map[string]any, 0, len(renderedPages)) @@ -411,7 +417,7 @@ func dispatchPDFVision( prompt := renderPDFVisionPrompt(promptTemplate, page.PageNumber) resp, err := pdfVisionChatInvoker(ctx, driver, resolvedModelName, buildPDFVisionMessages(prompt, page.ImageURL), apiConfig) if err != nil { - return parserDispatchResult{}, fmt.Errorf("Parser: pdf vision page %d: %w", page.PageNumber, err) + return parserDispatchResult{}, fmt.Errorf("parser: pdf vision page %d: %w", page.PageNumber, err) } text := extractPDFVisionAnswer(resp) positions := [][]any{{page.PageNumber, 0.0, page.WidthPts, 0.0, page.HeightPts}} @@ -451,7 +457,7 @@ func dispatchPDFVision( Markdown: strings.TrimSpace(strings.Join(markdownParts, "\n\n")), }, nil default: - return parserDispatchResult{}, fmt.Errorf("Parser: unsupported PDF output_format %q for vision parse_method %q", outputFormat, modelID) + return parserDispatchResult{}, fmt.Errorf("parser: unsupported PDF output_format %q for vision parse_method %q", outputFormat, modelID) } } @@ -473,11 +479,13 @@ func extractPDFVisionAnswer(resp *modelModule.ChatResponse) string { } func defaultPDFVisionModelResolver( + ctx context.Context, + db *gorm.DB, tenantID string, modelID string, ) (modelModule.ModelDriver, string, *modelModule.APIConfig, error) { if strings.TrimSpace(modelID) == "" { - driver, modelName, apiConfig, _, err := resolveTenantModelByType(tenantID, entity.ModelTypeImage2Text) + driver, modelName, apiConfig, _, err := resolveTenantModelByType(ctx, db, tenantID, entity.ModelTypeImage2Text) return driver, modelName, apiConfig, err } driver, modelName, apiConfig, _, err := resolveModelConfig(tenantID, entity.ModelTypeImage2Text, modelID) diff --git a/internal/ingestion/task/task_context.go b/internal/ingestion/task/task_context.go index 1acd1f3a14..bcb800f678 100644 --- a/internal/ingestion/task/task_context.go +++ b/internal/ingestion/task/task_context.go @@ -71,7 +71,7 @@ func LoadFromIngestionTask(ctx context.Context, ingestionTask *entity.IngestionT return nil, fmt.Errorf("error when load knowledgebase %s: %w", doc.KbID, err) } - tenant, err := dao.NewTenantDAO().GetByID(kb.TenantID) + tenant, err := dao.NewTenantDAO().GetByID(ctx, dao.DB, kb.TenantID) if err != nil || tenant == nil { return nil, fmt.Errorf("error when load tenant %s: %w", kb.TenantID, err) } diff --git a/internal/service/chat.go b/internal/service/chat.go index be5b7b51c3..6a2b03351e 100644 --- a/internal/service/chat.go +++ b/internal/service/chat.go @@ -208,7 +208,7 @@ type CreateChatRequest struct { } func (s *ChatService) Create(ctx context.Context, userID string, req map[string]interface{}) (map[string]interface{}, common.ErrorCode, error) { - tenant, err := s.tenantDAO.GetByID(userID) + tenant, err := s.tenantDAO.GetByID(ctx, dao.DB, userID) if err != nil { return nil, common.CodeDataError, errors.New("tenant not found") } @@ -848,7 +848,7 @@ func (s *ChatService) updateChatREST(ctx context.Context, userID, chatID string, if err != nil { return nil, err } - if _, err = s.tenantDAO.GetByID(userID); err != nil { + if _, err = s.tenantDAO.GetByID(ctx, dao.DB, userID); err != nil { return nil, errors.New("tenant not found") } diff --git a/internal/service/chat_session.go b/internal/service/chat_session.go index 490c14ce49..dd3ae48a24 100644 --- a/internal/service/chat_session.go +++ b/internal/service/chat_session.go @@ -1546,7 +1546,7 @@ func (s *ChatSessionService) ChatCompletions( dialog.LLMID = llmID dialog.LLMSetting = genConfig } else if dialog.LLMID == "" { - tenant, err := dao.NewTenantDAO().GetByID(dialog.TenantID) + tenant, err := dao.NewTenantDAO().GetByID(ctx, dao.DB, dialog.TenantID) if err != nil || tenant.LLMID == "" { return fail(errors.New("no default chat model for tenant")) } diff --git a/internal/service/chunk/chunk.go b/internal/service/chunk/chunk.go index 1c3342046f..8af33cbdd8 100644 --- a/internal/service/chunk/chunk.go +++ b/internal/service/chunk/chunk.go @@ -244,7 +244,7 @@ func (s *ChunkService) RetrievalTest(ctx context.Context, req *service.Retrieval // If no chatID from search_config, or chatModel not found, use tenant default if chatModelForFilter == nil { tenantSvc := service.NewTenantService() - modelName, err := tenantSvc.GetDefaultModelName(tenantIDs[0], entity.ModelTypeChat) + modelName, err := tenantSvc.GetDefaultModelName(ctx, tenantIDs[0], entity.ModelTypeChat) if err != nil || modelName == "" { common.Warn("Failed to get tenant default chat model name for meta_data_filter", zap.Error(err)) } else { @@ -290,7 +290,7 @@ func (s *ChunkService) RetrievalTest(ctx context.Context, req *service.Retrieval tenantSvc := service.NewTenantService() modelProviderSvc := service.NewModelProviderService() var err error - llmModelName, err = tenantSvc.GetDefaultModelName(tenantIDs[0], entity.ModelTypeChat) + llmModelName, err = tenantSvc.GetDefaultModelName(ctx, tenantIDs[0], entity.ModelTypeChat) if err != nil || llmModelName == "" { common.Warn("Failed to get default chat model name for LLM transformations", zap.Error(err)) } else { diff --git a/internal/service/dataset/crud.go b/internal/service/dataset/crud.go index d666014ee6..56a9d0a436 100644 --- a/internal/service/dataset/crud.go +++ b/internal/service/dataset/crud.go @@ -29,7 +29,7 @@ func (d *DatasetService) CreateDataset(ctx context.Context, req *service.CreateD return nil, common.CodeDataError, fmt.Errorf("Dataset name length is %d which is large than %d", len(name), entity.DatasetNameLimit) } - tenant, err := d.tenantDAO.GetByID(tenantID) + tenant, err := d.tenantDAO.GetByID(ctx, dao.DB, tenantID) if err != nil || tenant == nil { return nil, common.CodeDataError, errors.New("tenant not found") } @@ -365,7 +365,7 @@ func (d *DatasetService) ListDatasets(ctx context.Context, id, name string, page } } if len(tenantIDs) == 0 { - joinedTenants, err := d.tenantDAO.GetJoinedTenantsByUserID(userID) + joinedTenants, err := d.tenantDAO.GetJoinedTenantsByUserID(ctx, dao.DB, userID) if err != nil { return nil, 0, common.CodeServerError, errors.New("database operation failed") } @@ -402,8 +402,8 @@ func (d *DatasetService) ListDatasets(ctx context.Context, id, name string, page return data, total, common.CodeSuccess, nil } -func (d *DatasetService) ListDatasetFilters(userID string) (map[string]interface{}, common.ErrorCode, error) { - joinedTenants, err := d.tenantDAO.GetJoinedTenantsByUserID(userID) +func (d *DatasetService) ListDatasetFilters(ctx context.Context, userID string) (map[string]interface{}, common.ErrorCode, error) { + joinedTenants, err := d.tenantDAO.GetJoinedTenantsByUserID(ctx, dao.DB, userID) if err != nil { return nil, common.CodeServerError, errors.New("database operation failed") } diff --git a/internal/service/dataset/permission.go b/internal/service/dataset/permission.go index 31b920f0a0..ec5cce6d97 100644 --- a/internal/service/dataset/permission.go +++ b/internal/service/dataset/permission.go @@ -34,7 +34,7 @@ func (d *DatasetService) GetKnowledgebaseByID(ctx context.Context, datasetID str } // CheckKBTeamPermission checks if a user has team-level permission for the KB. -func (d *DatasetService) CheckKBTeamPermission(kb *entity.Knowledgebase, userID string) bool { +func (d *DatasetService) CheckKBTeamPermission(ctx context.Context, kb *entity.Knowledgebase, userID string) bool { if kb == nil { return false } @@ -44,7 +44,7 @@ func (d *DatasetService) CheckKBTeamPermission(kb *entity.Knowledgebase, userID if kb.Permission != string(entity.TenantPermissionTeam) { return false } - joinedTenants, err := d.tenantDAO.GetJoinedTenantsByUserID(userID) + joinedTenants, err := d.tenantDAO.GetJoinedTenantsByUserID(ctx, dao.DB, userID) if err != nil { return false } diff --git a/internal/service/document/file2document.go b/internal/service/document/file2document.go index bacb3586bf..4cdcd3ff95 100644 --- a/internal/service/document/file2document.go +++ b/internal/service/document/file2document.go @@ -142,7 +142,7 @@ func (s *File2DocumentService) LinkToDatasets(ctx context.Context, userID string // ── 5. Validate KB permissions ──────────────────────────────────────────── for _, kb := range kbMap { - if !service.HasKBTeamPermission(kb, userID, dao.NewTenantDAO()) { + if !service.HasKBTeamPermission(ctx, kb, userID, dao.NewTenantDAO()) { return ErrLinkNoAuthorization } } diff --git a/internal/service/file_permission.go b/internal/service/file_permission.go index 0cb26be104..72194d50f3 100644 --- a/internal/service/file_permission.go +++ b/internal/service/file_permission.go @@ -27,7 +27,7 @@ func CheckFileTeamPermission(ctx context.Context, fileDAO *dao.FileDAO, file *en if err != nil || kb == nil { continue } - if HasKBTeamPermission(kb, userID, dao.NewTenantDAO()) { + if HasKBTeamPermission(ctx, kb, userID, dao.NewTenantDAO()) { return true } } diff --git a/internal/service/mcp.go b/internal/service/mcp.go index 82a30cd6a2..f1c75b6be3 100644 --- a/internal/service/mcp.go +++ b/internal/service/mcp.go @@ -139,7 +139,7 @@ func (s *MCPService) CreateMCPServer(ctx context.Context, tenantID string, req C return nil, common.CodeDataError, errors.New("invalid url") } - if _, err = s.tenantDAO.GetByID(tenantID); err != nil { + if _, err = s.tenantDAO.GetByID(ctx, dao.DB, tenantID); err != nil { return nil, common.CodeDataError, errors.New("tenant not found") } diff --git a/internal/service/model_service.go b/internal/service/model_service.go index 207d59cec2..e8c334c896 100644 --- a/internal/service/model_service.go +++ b/internal/service/model_service.go @@ -1421,8 +1421,8 @@ func (m *ModelProviderService) ShowTask(ctx context.Context, providerName, insta // to ListTenantDefaultModels (which only enumerates the 6-7 default // tenant fields and returned `[]` for any tenant without defaults), // breaking the front-end's "View Models" list entirely. -func (m *ModelProviderService) ListTenantAddedModels(userID, ownerTenantID, modelTypeFilter string) ([]map[string]interface{}, common.ErrorCode, error) { - tenant, code, err := m.resolveModelListTenant(userID, ownerTenantID) +func (m *ModelProviderService) ListTenantAddedModels(ctx context.Context, userID, ownerTenantID, modelTypeFilter string) ([]map[string]interface{}, common.ErrorCode, error) { + tenant, code, err := m.resolveModelListTenant(ctx, userID, ownerTenantID) if err != nil { return nil, code, err } @@ -1629,7 +1629,7 @@ func (m *ModelProviderService) ListTenantAddedModels(userID, ownerTenantID, mode return added, common.CodeSuccess, nil } -func (m *ModelProviderService) resolveModelListTenant(userID, ownerTenantID string) (*entity.Tenant, common.ErrorCode, error) { +func (m *ModelProviderService) resolveModelListTenant(ctx context.Context, userID, ownerTenantID string) (*entity.Tenant, common.ErrorCode, error) { if ownerTenantID == "" { tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") if err != nil { @@ -1656,7 +1656,7 @@ func (m *ModelProviderService) resolveModelListTenant(userID, ownerTenantID stri } } - tenant, err := m.tenantDAO.GetByID(ownerTenantID) + tenant, err := m.tenantDAO.GetByID(ctx, dao.DB, ownerTenantID) if err != nil { if dao.IsNotFoundErr(err) { return nil, common.CodeNotFound, fmt.Errorf("tenant %s not found", ownerTenantID) @@ -3253,7 +3253,7 @@ func (m *ModelProviderService) GetTenantDefaultModelByType(ctx context.Context, return nil, "", nil, 0, fmt.Errorf("OCR model name is required") } - tenant, err := m.tenantDAO.GetByID(tenantID) + tenant, err := m.tenantDAO.GetByID(ctx, dao.DB, tenantID) if err != nil { return nil, "", nil, 0, fmt.Errorf("failed to get tenant: %s type %s: %w", tenantID, modelType, err) } diff --git a/internal/service/oauth_login.go b/internal/service/oauth_login.go index 8077a867dd..0501b22199 100644 --- a/internal/service/oauth_login.go +++ b/internal/service/oauth_login.go @@ -263,18 +263,18 @@ func (s *UserService) registerOAuthUser(ctx context.Context, channel string, inf if err := s.userDAO.Create(user); err != nil { return nil, common.CodeServerError, fmt.Errorf("Failed to register %s: %w", info.Email, err) } - if err := tenantDAO.Create(tenant); err != nil { + if err := tenantDAO.Create(ctx, dao.DB, tenant); err != nil { _ = s.userDAO.DeleteByID(userID) return nil, common.CodeServerError, fmt.Errorf("Failed to register %s: %w", info.Email, err) } if err := userTenantDAO.Create(userTenant); err != nil { _ = s.userDAO.DeleteByID(userID) - _ = tenantDAO.Delete(userID) + _ = tenantDAO.Delete(ctx, dao.DB, userID) return nil, common.CodeServerError, fmt.Errorf("Failed to register %s: %w", info.Email, err) } if err := fileDAO.Create(ctx, dao.DB, rootFile); err != nil { _ = s.userDAO.DeleteByID(userID) - _ = tenantDAO.Delete(userID) + _ = tenantDAO.Delete(ctx, dao.DB, userID) _ = userTenantDAO.Delete(userTenantID) return nil, common.CodeServerError, fmt.Errorf("Failed to register %s: %w", info.Email, err) } diff --git a/internal/service/related_question.go b/internal/service/related_question.go index a64adfef00..dd8267c6b7 100644 --- a/internal/service/related_question.go +++ b/internal/service/related_question.go @@ -33,7 +33,7 @@ func GenerateRelatedQuestions(ctx context.Context, tenantID, question, searchID return nil, fmt.Errorf("model provider service not configured") } searchConfig := relatedQuestionsSearchConfig(ctx, searchID, searchSvc) - modelID := relatedQuestionsModelID(tenantID, searchConfig, tenantSvc) + modelID := relatedQuestionsModelID(ctx, tenantID, searchConfig, tenantSvc) prompt, err := LoadPrompt("related_question") if err != nil { return nil, err @@ -72,12 +72,12 @@ func relatedQuestionsSearchConfigFromDetail(detail map[string]interface{}) map[s return map[string]interface{}{} } -func relatedQuestionsModelID(tenantID string, searchConfig map[string]interface{}, tenantSvc *TenantService) string { +func relatedQuestionsModelID(ctx context.Context, tenantID string, searchConfig map[string]interface{}, tenantSvc *TenantService) string { modelID, _ := searchConfig["chat_id"].(string) if modelID != "" || tenantSvc == nil { return modelID } - defaultModel, err := tenantSvc.GetDefaultModelName(tenantID, entity.ModelTypeChat) + defaultModel, err := tenantSvc.GetDefaultModelName(ctx, tenantID, entity.ModelTypeChat) if err == nil { modelID = defaultModel } diff --git a/internal/service/search.go b/internal/service/search.go index 4d75392c78..1ef5992ca8 100644 --- a/internal/service/search.go +++ b/internal/service/search.go @@ -406,7 +406,7 @@ func (s *SearchService) PrepareCompletion(ctx context.Context, userID, searchID tenantSvc = NewTenantService() } var defaultModelName string - defaultModelName, err = tenantSvc.GetDefaultModelName(userID, entity.ModelTypeChat) + defaultModelName, err = tenantSvc.GetDefaultModelName(ctx, userID, entity.ModelTypeChat) if err == nil { modelID = strings.TrimSpace(defaultModelName) } diff --git a/internal/service/skill_space.go b/internal/service/skill_space.go index f075cf2264..cd5bbd487e 100644 --- a/internal/service/skill_space.go +++ b/internal/service/skill_space.go @@ -252,7 +252,7 @@ func (s *SkillSpaceService) CreateSpace(ctx context.Context, req *CreateSpaceReq // Create default search config for this space defaultEmbdID := req.EmbdID if defaultEmbdID == "" { - tenant, err := s.tenantDAO.GetByID(req.TenantID) + tenant, err := s.tenantDAO.GetByID(ctx, dao.DB, req.TenantID) if err == nil && tenant != nil && tenant.EmbdID != "" { defaultEmbdID = tenant.EmbdID common.Info("Using tenant default embedding model", zap.String("tenantID", req.TenantID), zap.String("embdID", defaultEmbdID)) diff --git a/internal/service/team_permission.go b/internal/service/team_permission.go index 57b932b477..6cc770b843 100644 --- a/internal/service/team_permission.go +++ b/internal/service/team_permission.go @@ -1,6 +1,7 @@ package service import ( + "context" "ragflow/internal/dao" "ragflow/internal/entity" ) @@ -8,7 +9,7 @@ import ( // HasKBTeamPermission mirrors Python check_kb_team_permission: // direct owner access is always allowed; otherwise the KB must be team-shared // and the caller must be a joined normal member of the owner tenant. -func HasKBTeamPermission(kb *entity.Knowledgebase, userID string, tenantDAO *dao.TenantDAO) bool { +func HasKBTeamPermission(ctx context.Context, kb *entity.Knowledgebase, userID string, tenantDAO *dao.TenantDAO) bool { if kb == nil { return false } @@ -18,7 +19,7 @@ func HasKBTeamPermission(kb *entity.Knowledgebase, userID string, tenantDAO *dao if kb.Permission != string(entity.TenantPermissionTeam) { return false } - joinedTenants, err := tenantDAO.GetJoinedTenantsByUserID(userID) + joinedTenants, err := tenantDAO.GetJoinedTenantsByUserID(ctx, dao.DB, userID) if err != nil { return false } diff --git a/internal/service/tenant.go b/internal/service/tenant.go index 5b3fa12b0a..84bf6fdc5d 100644 --- a/internal/service/tenant.go +++ b/internal/service/tenant.go @@ -72,8 +72,8 @@ type TenantInfoResponse struct { } // GetTenantInfo get tenant information for the current user (owner tenant) -func (s *TenantService) GetTenantInfo(userID string) (*TenantInfoResponse, error) { - tenantInfos, err := s.tenantDAO.GetInfoByUserID(userID) +func (s *TenantService) GetTenantInfo(ctx context.Context, userID string) (*TenantInfoResponse, error) { + tenantInfos, err := s.tenantDAO.GetInfoByUserID(ctx, dao.DB, userID) if err != nil { return nil, err } @@ -463,8 +463,8 @@ func factoryModelTypeName(modelType string) string { // GetDefaultModelName returns the full default model ID for a tenant and model type // Format: modelName@instanceName@providerName or modelName@providerName // Returns empty string if no default model is set -func (s *TenantService) GetDefaultModelName(tenantID string, modelType entity.ModelType) (string, error) { - tenant, err := s.tenantDAO.GetByID(tenantID) +func (s *TenantService) GetDefaultModelName(ctx context.Context, tenantID string, modelType entity.ModelType) (string, error) { + tenant, err := s.tenantDAO.GetByID(ctx, dao.DB, tenantID) if err != nil { return "", err } @@ -592,9 +592,9 @@ func rsplitN(s, sep string, n int) []string { return result } -func (s *TenantService) ListTenantDefaultModels(userID string) ([]ModelItem, error) { +func (s *TenantService) ListTenantDefaultModels(ctx context.Context, userID string) ([]ModelItem, error) { - tenantInfos, err := s.tenantDAO.GetInfoByUserID(userID) + tenantInfos, err := s.tenantDAO.GetInfoByUserID(ctx, dao.DB, userID) if err != nil { return nil, err } @@ -750,9 +750,9 @@ func (s *TenantService) checkModelAvailable(tenantID, providerName, instanceName return nil } -func (s *TenantService) SetTenantDefaultModels(userID, modelProvider, modelInstance, modelName, modelType, modelID string) error { +func (s *TenantService) SetTenantDefaultModels(ctx context.Context, userID, modelProvider, modelInstance, modelName, modelType, modelID string) error { - tenantInfos, err := s.tenantDAO.GetInfoByUserID(userID) + tenantInfos, err := s.tenantDAO.GetInfoByUserID(ctx, dao.DB, userID) if err != nil { return err } @@ -832,12 +832,12 @@ func (s *TenantService) SetTenantDefaultModels(userID, modelProvider, modelInsta return fmt.Errorf("model provider, instance and name must be specified together") } - err = s.tenantDAO.Update(ownedTenant.TenantID, map[string]interface{}{ + err = s.tenantDAO.Update(ctx, dao.DB, ownedTenant.TenantID, map[string]interface{}{ modelTypeID: defaultModel, tenantModelTypeID: tenantModelID, }) - return nil + return err } // Tenant member role constants. diff --git a/internal/service/tenant_test.go b/internal/service/tenant_test.go index 80d22b3237..5caaa2ad65 100644 --- a/internal/service/tenant_test.go +++ b/internal/service/tenant_test.go @@ -238,7 +238,8 @@ func TestSetTenantDefaultModels_WithModelID(t *testing.T) { // 4. Run SetTenantDefaultModels s := NewTenantService() // Set chat model using modelID, explicitly passing "default" as instance name to bypass pre-existing checkModelAvailable panic - err = s.SetTenantDefaultModels(userID, "", "default", "", "chat", modelID) + ctx := t.Context() + err = s.SetTenantDefaultModels(ctx, userID, "", "default", "", "chat", modelID) if err != nil { t.Fatalf("SetTenantDefaultModels failed: %v", err) } diff --git a/internal/service/user.go b/internal/service/user.go index c4541e0891..99e6690ead 100644 --- a/internal/service/user.go +++ b/internal/service/user.go @@ -896,7 +896,7 @@ func (s *UserService) SetTenantInfo(ctx context.Context, userID string, req *Set updates = tenantLLMService.EnsureTenantModelIDForParams(tenantID, updates) if len(updates) > 0 { - if err := tenantDAO.Update(tenantID, updates); err != nil { + if err := tenantDAO.Update(ctx, dao.DB, tenantID, updates); err != nil { return common.CodeExceptionError, err } }