diff --git a/cmd/ragflow_server.go b/cmd/ragflow_server.go index eff6320eb4..a5e0e3c983 100644 --- a/cmd/ragflow_server.go +++ b/cmd/ragflow_server.go @@ -34,6 +34,9 @@ import ( "ragflow/internal/server/local" "ragflow/internal/service" "ragflow/internal/service/chunk" + dataset "ragflow/internal/service/dataset" + "ragflow/internal/service/document" + "ragflow/internal/service/file" "ragflow/internal/service/nlp" "ragflow/internal/storage" "ragflow/internal/syncer" @@ -681,8 +684,8 @@ func startServer(config *server.Config) { // Initialize service layer userService := service.NewUserService() - documentService := service.NewDocumentService() - datasetsService := service.NewDatasetService() + documentService := document.NewDocumentService() + datasetsService := dataset.NewDatasetService() metadataService := service.NewMetadataService() chunkService := chunk.NewChunkService() llmService := service.NewLLMService() @@ -696,7 +699,7 @@ func startServer(config *server.Config) { connectorService := service.NewConnectorService() searchService := service.NewSearchService() searchService.SetTenantService(tenantService) - fileService := service.NewFileService() + fileService := file.NewFileService(service.CheckFileTeamPermission, documentService) memoryService := service.NewMemoryService() mcpService := service.NewMCPService() modelProviderService := service.NewModelProviderService() @@ -711,7 +714,7 @@ func startServer(config *server.Config) { authHandler := handler.NewAuthHandler() userHandler := handler.NewUserHandler(userService) tenantHandler := handler.NewTenantHandler(tenantService, userService, datasetsService) - documentHandler := handler.NewDocumentHandler(documentService, datasetsService) + documentHandler := handler.NewDocumentHandler(documentService, datasetsService, fileService) datasetsHandler := handler.NewDatasetsHandler(datasetsService, metadataService) systemHandler := handler.NewSystemHandler(systemService) chunkHandler := handler.NewChunkHandler(chunkService, userService) @@ -741,7 +744,7 @@ func startServer(config *server.Config) { return handler.MCPRetrieval(datasetsService, userID, req) }, ) - skillSearchHandler := handler.NewSkillSearchHandler(docEngine) + skillSearchHandler := handler.NewSkillSearchHandler(docEngine, documentService) providerHandler := handler.NewProviderHandler(userService, modelProviderService) // Install the agent service's Redis-backed run infrastructure // (CheckPointStore / StateSerializer / RunTracker). When Redis @@ -789,7 +792,7 @@ func startServer(config *server.Config) { searchHandler.SetCompletionDependencies(modelProviderService, askService) pluginHandler := handler.NewPluginHandler(service.NewPluginService()) modelHandler := handler.NewModelHandler(service.NewModelProviderService()) - fileCommitHandler := handler.NewFileCommitHandler(service.NewFileCommitService()) + fileCommitHandler := handler.NewFileCommitHandler(file.NewFileCommitService()) // Dify retrieval handler docDAO := documentDAO diff --git a/internal/handler/agent.go b/internal/handler/agent.go index 767f407f2a..793c6a248e 100644 --- a/internal/handler/agent.go +++ b/internal/handler/agent.go @@ -36,6 +36,7 @@ import ( "ragflow/internal/dao" "ragflow/internal/entity" "ragflow/internal/service" + "ragflow/internal/service/file" dslpkg "ragflow/internal/agent/dsl" ) @@ -104,7 +105,7 @@ func (h *AgentHandler) WithDocumentService(s documentAccessChecker) *AgentHandle // NewAgentHandler create agent handler -func NewAgentHandler(agentService *service.AgentService, fileService *service.FileService) *AgentHandler { +func NewAgentHandler(agentService *service.AgentService, fileService *file.FileService) *AgentHandler { return &AgentHandler{ agentService: agentService, chatRunner: agentService, diff --git a/internal/handler/dataset.go b/internal/handler/dataset.go index 97e0317929..158b2e66d9 100644 --- a/internal/handler/dataset.go +++ b/internal/handler/dataset.go @@ -30,11 +30,12 @@ import ( "ragflow/internal/common" "ragflow/internal/service" + dataset "ragflow/internal/service/dataset" ) // DatasetsHandler handles the RESTful dataset endpoints. type DatasetsHandler struct { - datasetsService *service.DatasetService + datasetsService *dataset.DatasetService metadataService *service.MetadataService searchDatasetsService searchDatasetsService searchDatasetService searchDatasetService @@ -55,7 +56,7 @@ type listDatasetsExt struct { } // NewDatasetsHandler creates a new datasets handler. -func NewDatasetsHandler(datasetsService *service.DatasetService, metadataService *service.MetadataService) *DatasetsHandler { +func NewDatasetsHandler(datasetsService *dataset.DatasetService, metadataService *service.MetadataService) *DatasetsHandler { h := &DatasetsHandler{ datasetsService: datasetsService, metadataService: metadataService, diff --git a/internal/handler/document.go b/internal/handler/document.go index 87fb9b31d8..fa90a1768b 100644 --- a/internal/handler/document.go +++ b/internal/handler/document.go @@ -38,25 +38,27 @@ import ( "ragflow/internal/dao" "ragflow/internal/service" + dataset "ragflow/internal/service/dataset" + "ragflow/internal/service/document" ) var IMG_BASE64_PREFIX = "data:image/png;base64," // documentServiceIface defines the DocumentService methods used by DocumentHandler. type documentServiceIface interface { - GetDocumentByID(id string) (*service.DocumentResponse, error) - UpdateDocument(id string, req *service.UpdateDocumentRequest) error + GetDocumentByID(id string) (*document.DocumentResponse, error) + UpdateDocument(id string, req *document.UpdateDocumentRequest) error DeleteDocument(id string) error DeleteDocuments(ids []string, deleteAll bool, datasetID, userID string) (int, error) ParseDocuments(datasetID, userID string, docIDs []string) ([]*service.ParseDocumentResponse, error) StopParseDocuments(datasetID string, docIDs []string) (map[string]interface{}, error) - ListDocuments(page, pageSize int) ([]*service.DocumentResponse, int64, error) + ListDocuments(page, pageSize int) ([]*document.DocumentResponse, int64, error) ListDocumentsByDatasetID(kbID, keywords string, page, pageSize int) ([]*entity.DocumentListItem, int64, error) ListDocumentsByDatasetIDWithOptions(opts dao.DocumentListOptions, page, pageSize int) ([]*entity.DocumentListItem, int64, error) ListDocumentIDsByDatasetIDWithOptions(opts dao.DocumentListOptions) ([]string, error) GetDocumentFiltersByDatasetID(opts dao.DocumentListOptions) (map[string]interface{}, int64, error) GetMetadataByKBs(kbIDs []string) (map[string]interface{}, error) - GetDocumentsByAuthorID(authorID, page, pageSize int) ([]*service.DocumentResponse, int64, error) + GetDocumentsByAuthorID(authorID, page, pageSize int) ([]*document.DocumentResponse, int64, error) GetThumbnails(userID string, docIDs []string) (map[string]string, error) GetDocumentImage(imageID string) ([]byte, error) GetMetadataSummary(kbID string, docIDs []string) (map[string]interface{}, error) @@ -64,35 +66,41 @@ type documentServiceIface interface { DeleteDocumentMetadata(docID string, keys []string) error DeleteDocumentAllMetadata(docID string) error GetDocumentMetadataByID(docID string) (map[string]interface{}, error) - GetDocumentArtifact(filename, userID string) (*service.ArtifactResponse, error) - GetDocumentPreview(docID string) (*service.DocumentPreview, error) + GetDocumentArtifact(filename, userID string) (*document.ArtifactResponse, error) + GetDocumentPreview(docID string) (*document.DocumentPreview, error) UploadLocalDocuments(kb *entity.Knowledgebase, tenantID string, files []*multipart.FileHeader, parentPath string, parserConfigOverride map[string]interface{}) ([]map[string]interface{}, []string) UploadWebDocument(kb *entity.Knowledgebase, tenantID, name, url string) (map[string]interface{}, common.ErrorCode, error) UploadEmptyDocument(kb *entity.Knowledgebase, tenantID, name string) (map[string]interface{}, common.ErrorCode, error) - DownloadDocument(datasetID, docID string) (*service.DownloadDocumentResp, error) - UpdateDatasetDocument(userID, datasetID, documentID string, req *service.UpdateDatasetDocumentRequest, present map[string]bool) (*service.UpdateDatasetDocumentResponse, common.ErrorCode, error) - BatchUpdateDocumentMetadatas(datasetID string, selector *service.DocumentMetadataSelector, updates []service.DocumentMetadataUpdate, deletes []service.DocumentMetadataDelete) (*service.BatchUpdateDocumentMetadatasResponse, common.ErrorCode, error) - UploadDocumentInfos(userID string, files []*multipart.FileHeader) ([]map[string]interface{}, common.ErrorCode, error) - UploadDocumentInfoByURL(userID, rawURL string) (map[string]interface{}, common.ErrorCode, error) + DownloadDocument(datasetID, docID string) (*document.DownloadDocumentResp, error) + UpdateDatasetDocument(userID, datasetID, documentID string, req *document.UpdateDatasetDocumentRequest, present map[string]bool) (*document.UpdateDatasetDocumentResponse, common.ErrorCode, error) + BatchUpdateDocumentMetadatas(datasetID string, selector *document.DocumentMetadataSelector, updates []document.DocumentMetadataUpdate, deletes []document.DocumentMetadataDelete) (*document.BatchUpdateDocumentMetadatasResponse, common.ErrorCode, error) ListIngestionTasks(userID string, datasetID *string, page, pageSize int) ([]*entity.IngestionTask, error) IngestDocuments(datasetID, userID string, docIDs []string) ([]*service.ParseDocumentResponse, error) StopIngestionTasks(tasks []string, userID string) ([]*entity.IngestionTask, error) - Ingest(userID string, req *service.IngestDocumentRequest) (common.ErrorCode, error) + Ingest(userID string, req *document.IngestDocumentRequest) (common.ErrorCode, error) RemoveIngestionTasks(tasks []string, userID string) ([]map[string]string, error) BatchUpdateDocumentStatus(userID, datasetID, status string, DocumentIDs []string) (map[string]interface{}, common.ErrorCode, error) } +// fileUploadIface defines the FileService upload methods used by DocumentHandler. +type fileUploadIface interface { + UploadDocumentInfos(userID string, files []*multipart.FileHeader) ([]map[string]interface{}, common.ErrorCode, error) + UploadDocumentInfoByURL(userID, rawURL string) (map[string]interface{}, common.ErrorCode, error) +} + // DocumentHandler document handler type DocumentHandler struct { documentService documentServiceIface - datasetService *service.DatasetService + datasetService *dataset.DatasetService + fileService fileUploadIface } // NewDocumentHandler create document handler -func NewDocumentHandler(documentService documentServiceIface, datasetService *service.DatasetService) *DocumentHandler { +func NewDocumentHandler(documentService documentServiceIface, datasetService *dataset.DatasetService, fileService fileUploadIface) *DocumentHandler { return &DocumentHandler{ documentService: documentService, datasetService: datasetService, + fileService: fileService, } } @@ -219,9 +227,9 @@ func (h *DocumentHandler) GetDocumentArtifact(c *gin.Context) { artifact, err := h.documentService.GetDocumentArtifact(filename, user.ID) if err != nil { switch { - case errors.Is(err, service.ErrArtifactInvalidFilename), - errors.Is(err, service.ErrArtifactInvalidFileType), - errors.Is(err, service.ErrArtifactNotFound): + case errors.Is(err, document.ErrArtifactInvalidFilename), + errors.Is(err, document.ErrArtifactInvalidFileType), + errors.Is(err, document.ErrArtifactNotFound): common.ErrorWithCode(c, common.CodeDataError, err.Error()) default: @@ -272,7 +280,7 @@ func (h *DocumentHandler) GetDocumentPreview(c *gin.Context) { // @Accept json // @Produce json // @Param id path int true "document ID" -// @Param request body service.UpdateDocumentRequest true "update info" +// @Param request body document.UpdateDocumentRequest true "update info" // @Success 200 {object} map[string]interface{} // @Router /api/v1/documents/{id} [put] func (h *DocumentHandler) UpdateDocument(c *gin.Context) { @@ -300,7 +308,7 @@ func (h *DocumentHandler) UpdateDocument(c *gin.Context) { return } - var req service.UpdateDocumentRequest + var req document.UpdateDocumentRequest if err = c.ShouldBindJSON(&req); err != nil { c.JSON(http.StatusBadRequest, gin.H{ "error": err.Error(), @@ -1230,7 +1238,7 @@ func (h *DocumentHandler) SetMeta(c *gin.Context) { // @Accept json // @Produce json // @Security ApiKeyAuth -// @Param request body service.IngestDocumentRequest true "ingestion info" +// @Param request body document.IngestDocumentRequest true "ingestion info" // @Success 200 {object} map[string]interface{} // @Router /api/v1/documents/ingest [post] func (h *DocumentHandler) Ingest(c *gin.Context) { @@ -1246,7 +1254,7 @@ func (h *DocumentHandler) Ingest(c *gin.Context) { return } - var req service.IngestDocumentRequest + var req document.IngestDocumentRequest if err := c.ShouldBindJSON(&req); err != nil { common.ErrorWithCode(c, common.CodeBadRequest, err.Error()) return @@ -1577,7 +1585,7 @@ func (h *DocumentHandler) UpdateDatasetDocument(c *gin.Context) { for key := range raw { present[key] = true } - var req service.UpdateDatasetDocumentRequest + var req document.UpdateDatasetDocumentRequest if err = json.Unmarshal(body, &req); err != nil { common.ResponseWithCodeData(c, common.CodeDataError, nil, err.Error()) return @@ -1621,7 +1629,7 @@ func (h *DocumentHandler) UploadInfo(c *gin.Context) { } if rawURL != "" { - data, code, err := h.documentService.UploadDocumentInfoByURL(user.ID, rawURL) + data, code, err := h.fileService.UploadDocumentInfoByURL(user.ID, rawURL) if err != nil { common.ErrorWithCode(c, code, err.Error()) return @@ -1630,7 +1638,7 @@ func (h *DocumentHandler) UploadInfo(c *gin.Context) { return } - data, code, err := h.documentService.UploadDocumentInfos(user.ID, fileHeaders) + data, code, err := h.fileService.UploadDocumentInfos(user.ID, fileHeaders) if err != nil { common.ErrorWithCode(c, code, err.Error()) return @@ -1646,9 +1654,9 @@ func (h *DocumentHandler) UploadInfo(c *gin.Context) { } type documentMetadataBatchRequest struct { - Selector *service.DocumentMetadataSelector `json:"selector"` - Updates []service.DocumentMetadataUpdate `json:"updates"` - Deletes []service.DocumentMetadataDelete `json:"deletes"` + Selector *document.DocumentMetadataSelector `json:"selector"` + Updates []document.DocumentMetadataUpdate `json:"updates"` + Deletes []document.DocumentMetadataDelete `json:"deletes"` } func (h *DocumentHandler) MetadataBatchUpdate(c *gin.Context) { @@ -1682,13 +1690,13 @@ func (h *DocumentHandler) handleBatchUpdateDocumentMetadatas(c *gin.Context) { return } if req.Selector == nil { - req.Selector = &service.DocumentMetadataSelector{} + req.Selector = &document.DocumentMetadataSelector{} } if req.Updates == nil { - req.Updates = []service.DocumentMetadataUpdate{} + req.Updates = []document.DocumentMetadataUpdate{} } if req.Deletes == nil { - req.Deletes = []service.DocumentMetadataDelete{} + req.Deletes = []document.DocumentMetadataDelete{} } resp, code, err := h.documentService.BatchUpdateDocumentMetadatas(datasetID, req.Selector, req.Updates, req.Deletes) diff --git a/internal/handler/document_test.go b/internal/handler/document_test.go index da732a7f3a..257a157d94 100644 --- a/internal/handler/document_test.go +++ b/internal/handler/document_test.go @@ -34,13 +34,15 @@ import ( "ragflow/internal/dao" "ragflow/internal/entity" "ragflow/internal/service" + dataset "ragflow/internal/service/dataset" + "ragflow/internal/service/document" ) // fakeDocumentService implements documentServiceIface for handler tests. type fakeDocumentService struct { deleted int err error - doc *service.DocumentResponse + doc *document.DocumentResponse docErr error updateCalled bool updatedID string @@ -71,7 +73,7 @@ type fakeDocumentService struct { ingestCode common.ErrorCode ingestErr error ingestUserID string - ingestReq *service.IngestDocumentRequest + ingestReq *document.IngestDocumentRequest listOpts dao.DocumentListOptions filterOpts dao.DocumentListOptions filterResult map[string]interface{} @@ -80,7 +82,7 @@ type fakeDocumentService struct { metadataByKBs map[string]interface{} } -func (f *fakeDocumentService) Ingest(userID string, req *service.IngestDocumentRequest) (common.ErrorCode, error) { +func (f *fakeDocumentService) Ingest(userID string, req *document.IngestDocumentRequest) (common.ErrorCode, error) { f.ingestUserID = userID f.ingestReq = req if f.ingestCode != 0 || f.ingestErr != nil { @@ -91,54 +93,48 @@ func (f *fakeDocumentService) Ingest(userID string, req *service.IngestDocumentR const uploadTestDatasetID = "123e4567-e89b-12d3-a456-426614174000" -func (f *fakeDocumentService) UpdateDatasetDocument(userID, datasetID, documentID string, req *service.UpdateDatasetDocumentRequest, present map[string]bool) (*service.UpdateDatasetDocumentResponse, common.ErrorCode, error) { +func (f *fakeDocumentService) UpdateDatasetDocument(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(datasetID string, selector *service.DocumentMetadataSelector, updates []service.DocumentMetadataUpdate, deletes []service.DocumentMetadataDelete) (*service.BatchUpdateDocumentMetadatasResponse, common.ErrorCode, error) { - return nil, common.CodeSuccess, nil -} -func (f *fakeDocumentService) UploadDocumentInfos(userID string, files []*multipart.FileHeader) ([]map[string]interface{}, common.ErrorCode, error) { - return nil, common.CodeSuccess, nil -} -func (f *fakeDocumentService) UploadDocumentInfoByURL(userID, rawURL string) (map[string]interface{}, common.ErrorCode, error) { +func (f *fakeDocumentService) BatchUpdateDocumentMetadatas(datasetID string, selector *document.DocumentMetadataSelector, updates []document.DocumentMetadataUpdate, deletes []document.DocumentMetadataDelete) (*document.BatchUpdateDocumentMetadatasResponse, common.ErrorCode, error) { return nil, common.CodeSuccess, nil } -func (f *fakeDocumentService) GetDocumentArtifact(filename, _ string) (*service.ArtifactResponse, error) { +func (f *fakeDocumentService) GetDocumentArtifact(filename, _ string) (*document.ArtifactResponse, error) { if filename == "error.txt" { - return nil, service.ErrArtifactNotFound + return nil, document.ErrArtifactNotFound } if filename == "unexpected.txt" { return nil, fmt.Errorf("unexpected error") } - return &service.ArtifactResponse{ + return &document.ArtifactResponse{ Data: []byte("artifact content"), ContentType: "text/plain", SafeFilename: "safe.txt", ForceAttachment: false, }, nil } -func (f *fakeDocumentService) GetDocumentPreview(docID string) (*service.DocumentPreview, error) { +func (f *fakeDocumentService) GetDocumentPreview(docID string) (*document.DocumentPreview, error) { if docID == "not-found" { return nil, fmt.Errorf("not found") } - return &service.DocumentPreview{ + return &document.DocumentPreview{ Data: []byte("preview content"), ContentType: "text/plain", FileName: "preview.txt", }, nil } -func (f *fakeDocumentService) DownloadDocument(datasetID, docID string) (*service.DownloadDocumentResp, error) { +func (f *fakeDocumentService) DownloadDocument(datasetID, docID string) (*document.DownloadDocumentResp, error) { if docID == "not-found" { return nil, fmt.Errorf("not found") } - return &service.DownloadDocumentResp{ + return &document.DownloadDocumentResp{ Data: []byte("document data"), ContentType: "application/pdf", FileName: "doc.pdf", }, nil } -func (f *fakeDocumentService) GetDocumentByID(id string) (*service.DocumentResponse, error) { +func (f *fakeDocumentService) GetDocumentByID(id string) (*document.DocumentResponse, error) { if f.docErr != nil { return nil, f.docErr } @@ -147,7 +143,7 @@ func (f *fakeDocumentService) GetDocumentByID(id string) (*service.DocumentRespo } return nil, fmt.Errorf("document not found") } -func (f *fakeDocumentService) UpdateDocument(id string, req *service.UpdateDocumentRequest) error { +func (f *fakeDocumentService) UpdateDocument(id string, req *document.UpdateDocumentRequest) error { f.updateCalled = true f.updatedID = id return nil @@ -166,7 +162,7 @@ func (f *fakeDocumentService) ParseDocuments(datasetID, userID string, docIDs [] func (f *fakeDocumentService) StopParseDocuments(datasetID string, docIDs []string) (map[string]interface{}, error) { return f.stopResult, f.stopErr } -func (f *fakeDocumentService) ListDocuments(page, pageSize int) ([]*service.DocumentResponse, int64, error) { +func (f *fakeDocumentService) ListDocuments(page, pageSize int) ([]*document.DocumentResponse, int64, error) { return nil, 0, nil } func (f *fakeDocumentService) ListDocumentsByDatasetID(kbID, keywords string, page, pageSize int) ([]*entity.DocumentListItem, int64, error) { @@ -204,7 +200,7 @@ func (f *fakeDocumentService) GetThumbnails(userID string, docIDs []string) (map func (f *fakeDocumentService) GetDocumentImage(imageID string) ([]byte, error) { return nil, nil } -func (f *fakeDocumentService) GetDocumentsByAuthorID(authorID, page, pageSize int) ([]*service.DocumentResponse, int64, error) { +func (f *fakeDocumentService) GetDocumentsByAuthorID(authorID, page, pageSize int) ([]*document.DocumentResponse, int64, error) { return nil, 0, nil } func (f *fakeDocumentService) GetMetadataSummary(kbID string, docIDs []string) (map[string]interface{}, error) { @@ -312,11 +308,11 @@ func TestSetMetaHandler_NotAccessible(t *testing.T) { setupDocumentPermissionDB(t, false) fake := &fakeDocumentService{ - doc: &service.DocumentResponse{ID: "doc-1", KbID: "kb-owner"}, + doc: &document.DocumentResponse{ID: "doc-1", KbID: "kb-owner"}, } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("POST", "/api/v1/document/set_meta", `{"doc_id":"doc-1","meta":"{\"poc\":\"blocked\"}"}`) @@ -344,11 +340,11 @@ func TestSetMetaHandler_Accessible(t *testing.T) { setupDocumentPermissionDB(t, true) fake := &fakeDocumentService{ - doc: &service.DocumentResponse{ID: "doc-1", KbID: "kb-owner"}, + doc: &document.DocumentResponse{ID: "doc-1", KbID: "kb-owner"}, } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("POST", "/api/v1/document/set_meta", `{"doc_id":"doc-1","meta":"{\"category\":\"tech\",\"year\":2026}"}`) @@ -379,11 +375,11 @@ func TestDeleteDocumentHandler_NotAccessible(t *testing.T) { setupDocumentPermissionDB(t, false) fake := &fakeDocumentService{ - doc: &service.DocumentResponse{ID: "doc-1", KbID: "kb-owner"}, + doc: &document.DocumentResponse{ID: "doc-1", KbID: "kb-owner"}, } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("DELETE", "/api/v1/documents/doc-1", "") @@ -412,11 +408,11 @@ func TestDeleteDocumentHandler_Accessible(t *testing.T) { setupDocumentPermissionDB(t, true) fake := &fakeDocumentService{ - doc: &service.DocumentResponse{ID: "doc-1", KbID: "kb-owner"}, + doc: &document.DocumentResponse{ID: "doc-1", KbID: "kb-owner"}, } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("DELETE", "/api/v1/documents/doc-1", "") @@ -445,11 +441,11 @@ func TestUpdateDocumentHandler_NotAccessible(t *testing.T) { setupDocumentPermissionDB(t, false) fake := &fakeDocumentService{ - doc: &service.DocumentResponse{ID: "doc-1", KbID: "kb-owner"}, + doc: &document.DocumentResponse{ID: "doc-1", KbID: "kb-owner"}, } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("PUT", "/api/v1/documents/doc-1", `{"name":"blocked"}`) @@ -478,11 +474,11 @@ func TestUpdateDocumentHandler_Accessible(t *testing.T) { setupDocumentPermissionDB(t, true) fake := &fakeDocumentService{ - doc: &service.DocumentResponse{ID: "doc-1", KbID: "kb-owner"}, + doc: &document.DocumentResponse{ID: "doc-1", KbID: "kb-owner"}, } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("PUT", "/api/v1/documents/doc-1", `{"name":"allowed"}`) @@ -586,7 +582,7 @@ func setupDocumentIngestRoute(userID string, svc *fakeDocumentService) *gin.Engi gin.SetMode(gin.TestMode) h := &DocumentHandler{ documentService: svc, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } r := gin.New() r.Use(func(c *gin.Context) { @@ -603,7 +599,7 @@ func TestDeleteDocumentsHandler_Success(t *testing.T) { fake := &fakeDocumentService{deleted: 3} h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("DELETE", "/api/v1/datasets/ds-1/documents", `{"ids": ["doc-1", "doc-2", "doc-3"]}`) @@ -639,7 +635,7 @@ func TestUploadDocumentsHandler_LocalUsesFullKBAndIgnoresBadParserConfig(t *test } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupUploadContext(t, "/api/v1/datasets/ds-1/documents?type=local", map[string]string{ @@ -680,7 +676,7 @@ func TestUploadDocumentsHandler_LocalReturnsPartialSuccess(t *testing.T) { } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupUploadContext(t, "/api/v1/datasets/ds-1/documents?type=local", nil, "ok.txt", []byte("abc")) @@ -714,7 +710,7 @@ func TestUploadDocumentsHandler_DeniesNonNormalTeamRole(t *testing.T) { fake := &fakeDocumentService{} h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupUploadContext(t, "/api/v1/datasets/ds-1/documents?type=local", nil, "a.txt", []byte("abc")) @@ -741,7 +737,7 @@ func TestDeleteDocumentsHandler_DeleteAll(t *testing.T) { fake := &fakeDocumentService{deleted: 5} h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("DELETE", "/api/v1/datasets/ds-1/documents", `{"delete_all": true}`) @@ -760,7 +756,7 @@ func TestDeleteDocumentsHandler_MutuallyExclusive(t *testing.T) { fake := &fakeDocumentService{} h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("DELETE", "/api/v1/datasets/ds-1/documents", `{"ids": ["doc-1"], "delete_all": true}`) @@ -785,7 +781,7 @@ func TestDeleteDocumentsHandler_NoIDsNoDeleteAll(t *testing.T) { fake := &fakeDocumentService{} h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("DELETE", "/api/v1/datasets/ds-1/documents", `{}`) @@ -810,7 +806,7 @@ func TestDeleteDocumentsHandler_ServiceError(t *testing.T) { fake := &fakeDocumentService{err: fmt.Errorf("permission denied")} h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("DELETE", "/api/v1/datasets/ds-1/documents", `{"ids": ["doc-1"]}`) @@ -835,7 +831,7 @@ func TestDeleteDocumentsHandler_MissingDatasetID(t *testing.T) { fake := &fakeDocumentService{} h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("DELETE", "/api/v1/datasets//documents", `{"ids": ["doc-1"]}`) @@ -859,7 +855,7 @@ func TestDocumentHandlerIngestMatchesPythonResponseShape(t *testing.T) { fake := &fakeDocumentService{} h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("POST", "/api/v1/documents/ingest", `{"doc_ids":["doc-1"],"run":"1"}`) @@ -938,7 +934,7 @@ func TestDocumentHandlerIngestPropagatesServiceErrorCode(t *testing.T) { } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("POST", "/api/v1/documents/ingest", `{"doc_ids":["doc-1"],"run":"1"}`) @@ -969,7 +965,7 @@ func TestStopParseDocumentsHandler_EmptyDocIDs(t *testing.T) { fake := &fakeDocumentService{} h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("POST", "/api/v1/datasets/ds-1/documents/stop", `{"document_ids": []}`) @@ -994,7 +990,7 @@ func TestStopParseDocumentsHandler_BadJSON(t *testing.T) { fake := &fakeDocumentService{} h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("POST", "/api/v1/datasets/ds-1/documents/stop", `not json`) @@ -1071,7 +1067,7 @@ func TestListDocumentsHandler_FilterRequestUsesQueryFilters(t *testing.T) { } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("GET", "/api/v1/datasets/ds-1/documents?type=filter&keywords=report&suffix=pdf&run=DONE&types=doc&desc=false", "") @@ -1129,7 +1125,7 @@ func TestListDocumentsHandler_MetadataFilterNarrowsDocumentIDs(t *testing.T) { } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("GET", "/api/v1/datasets/ds-1/documents?metadata[author][]=Alice", "") @@ -1161,7 +1157,7 @@ func TestStopParseDocumentsHandler_Success(t *testing.T) { } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("POST", "/api/v1/datasets/ds-1/documents/stop", `{"document_ids": ["doc-1"]}`) @@ -1197,7 +1193,7 @@ func TestStopParseDocumentsHandler_ServiceError(t *testing.T) { } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("POST", "/api/v1/datasets/ds-1/documents/stop", `{"document_ids": ["doc-1"]}`) @@ -1227,7 +1223,7 @@ func TestStopParseDocumentsHandler_NotAccessible(t *testing.T) { fake := &fakeDocumentService{} h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("POST", "/api/v1/datasets/ds-1/documents/stop", `{"document_ids": ["doc-1"]}`) @@ -1340,7 +1336,7 @@ func TestMetadataSummaryByDataset_Success(t *testing.T) { } h := &DocumentHandler{ documentService: fake, - datasetService: service.NewDatasetService(), + datasetService: dataset.NewDatasetService(), } c, w := setupGinContextWithUser("GET", "/api/v1/datasets/ds-1/metadata/summary?doc_ids=doc-1,doc-2", "") diff --git a/internal/handler/file.go b/internal/handler/file.go index 09ecfb067a..ec1ec61666 100644 --- a/internal/handler/file.go +++ b/internal/handler/file.go @@ -30,21 +30,23 @@ import ( "github.com/gin-gonic/gin" "ragflow/internal/service" + "ragflow/internal/service/document" + "ragflow/internal/service/file" ) // FileHandler file handler type FileHandler struct { - fileService *service.FileService + fileService *file.FileService userService *service.UserService - file2DocumentService *service.File2DocumentService + file2DocumentService *document.File2DocumentService } // NewFileHandler create file handler -func NewFileHandler(fileService *service.FileService, userService *service.UserService) *FileHandler { +func NewFileHandler(fileService *file.FileService, userService *service.UserService) *FileHandler { return &FileHandler{ fileService: fileService, userService: userService, - file2DocumentService: service.NewFile2DocumentService(), + file2DocumentService: document.NewFile2DocumentService(), } } @@ -60,7 +62,7 @@ func NewFileHandler(fileService *service.FileService, userService *service.UserS // @Param page_size query int false "items per page (default: 15, min: 1, max: 100)" // @Param orderby query string false "order by field (default: create_time)" // @Param desc query bool false "descending order (default: true)" -// @Success 200 {object} service.ListFilesResponse +// @Success 200 {object} file.ListFilesResponse // @Router /api/v1/files [get] func (h *FileHandler) ListFiles(c *gin.Context) { user, errorCode, errorMessage := GetUser(c) @@ -527,7 +529,7 @@ func (h *FileHandler) Download(c *gin.Context) { // @Tags file // @Accept json // @Produce json -// @Param request body service.LinkToDatasetsRequest true "file_ids and kb_ids" +// @Param request body document.LinkToDatasetsRequest true "file_ids and kb_ids" // @Success 200 {object} map[string]interface{} // @Router /api/v1/files/link-to-datasets [post] func (h *FileHandler) LinkToDatasets(c *gin.Context) { @@ -537,7 +539,7 @@ func (h *FileHandler) LinkToDatasets(c *gin.Context) { return } - var req service.LinkToDatasetsRequest + var req document.LinkToDatasetsRequest // Tolerate bind errors: a malformed or empty body simply leaves the fields // empty, which the validate_request-style check below reports as missing // arguments — matching Python's @validate_request behaviour and code. @@ -571,9 +573,9 @@ func (h *FileHandler) LinkToDatasets(c *gin.Context) { // any other (internal) error is reported as a server error. func linkToDatasetsErrorCode(err error) common.ErrorCode { switch { - case errors.Is(err, service.ErrLinkFileNotFound), - errors.Is(err, service.ErrLinkDatasetNotFound), - errors.Is(err, service.ErrLinkNoAuthorization): + case errors.Is(err, document.ErrLinkFileNotFound), + errors.Is(err, document.ErrLinkDatasetNotFound), + errors.Is(err, document.ErrLinkNoAuthorization): return common.CodeDataError default: return common.CodeServerError diff --git a/internal/handler/mcp_server.go b/internal/handler/mcp_server.go index 38c01a3522..f475c2e939 100644 --- a/internal/handler/mcp_server.go +++ b/internal/handler/mcp_server.go @@ -27,6 +27,7 @@ import ( "ragflow/internal/common" "ragflow/internal/mcp" "ragflow/internal/service" + dataset "ragflow/internal/service/dataset" ) // MCPRetrievalService abstracts the dataset retrieval operations needed @@ -112,7 +113,7 @@ func (h *MCPServerHandler) HandleMCP(c *gin.Context) { // MCPListDatasets wraps DatasetService.ListDatasets for the MCP tool handler, // filling in default values for parameters that the MCP tool does not expose. -func MCPListDatasets(ds *service.DatasetService, userID string, page, pageSize int, orderby string, desc bool) ([]map[string]interface{}, int64, error) { +func MCPListDatasets(ds *dataset.DatasetService, userID string, page, pageSize int, orderby string, desc bool) ([]map[string]interface{}, int64, error) { data, total, _, err := ds.ListDatasets( "", "", page, pageSize, orderby, desc, "", nil, "", userID, @@ -141,7 +142,7 @@ func MCPListChats(cs *service.ChatService, userID string, page, pageSize int, or // MCPRetrieval executes a retrieval request on behalf of the MCP tool handler. // It translates the mcp.RetrievalRequest into a service.SearchDatasetsRequest // and calls DatasetService.SearchDatasets. The result is serialized as JSON. -func MCPRetrieval(ds *service.DatasetService, userID string, req mcp.RetrievalRequest) (string, error) { +func MCPRetrieval(ds *dataset.DatasetService, userID string, req mcp.RetrievalRequest) (string, error) { // Resolve dataset IDs: if none provided, fetch ALL accessible datasets // across all pages (matching Python _fetch_all_datasets behaviour). datasetIDs := req.DatasetIDs diff --git a/internal/handler/skill_search.go b/internal/handler/skill_search.go index 81b97033a7..718fb1c173 100644 --- a/internal/handler/skill_search.go +++ b/internal/handler/skill_search.go @@ -22,6 +22,7 @@ import ( "ragflow/internal/common" "ragflow/internal/engine" "ragflow/internal/service" + "ragflow/internal/service/file" "github.com/gin-gonic/gin" "go.uber.org/zap" @@ -35,12 +36,13 @@ type SkillSearchHandler struct { docEngine engine.DocEngine } -// NewSkillSearchHandler creates a new skill search handler -func NewSkillSearchHandler(docEngine engine.DocEngine) *SkillSearchHandler { +// NewSkillSearchHandler creates a new skill search handler. +// spaceRemover is the document remover used by the skill space service for file deletion. +func NewSkillSearchHandler(docEngine engine.DocEngine, spaceRemover file.DocRemover) *SkillSearchHandler { return &SkillSearchHandler{ searchService: service.NewSkillSearchService(), indexerService: service.NewSkillIndexerService(), - spaceService: service.NewSkillSpaceService(), + spaceService: service.NewSkillSpaceService(spaceRemover), docEngine: docEngine, } } diff --git a/internal/handler/tenant.go b/internal/handler/tenant.go index d06f2ef683..d4bb0835fd 100644 --- a/internal/handler/tenant.go +++ b/internal/handler/tenant.go @@ -26,17 +26,18 @@ import ( "ragflow/internal/common" "ragflow/internal/engine" "ragflow/internal/service" + dataset "ragflow/internal/service/dataset" ) // TenantHandler tenant handler type TenantHandler struct { tenantService *service.TenantService userService *service.UserService - datasetService *service.DatasetService + datasetService *dataset.DatasetService } // NewTenantHandler create tenant handler -func NewTenantHandler(tenantService *service.TenantService, userService *service.UserService, datasetService *service.DatasetService) *TenantHandler { +func NewTenantHandler(tenantService *service.TenantService, userService *service.UserService, datasetService *dataset.DatasetService) *TenantHandler { return &TenantHandler{ tenantService: tenantService, userService: userService, diff --git a/internal/ingestion/pipeline/pipeline_params.go b/internal/ingestion/pipeline/pipeline_params.go new file mode 100644 index 0000000000..01be82d3b5 --- /dev/null +++ b/internal/ingestion/pipeline/pipeline_params.go @@ -0,0 +1,151 @@ +package pipeline + +import ( + "fmt" + "strings" + + "ragflow/internal/common" + "ragflow/internal/entity" + + "go.uber.org/zap" +) + +// CleanComponentParams filters rawConfig against the DSL schema given by dslJSON. +// Keys containing ':' are treated as component IDs; they are kept only when both +// the cpnID AND the param name exist in the DSL schema. Keys without ':' (legacy +// flat fields such as chunk_token_num, image_context_size) are dropped with a +// warning — they do not belong in the new component-params world. +func CleanComponentParams(dslJSON []byte, rawConfig map[string]interface{}) map[string]interface{} { + schemas, err := ExtractAllComponentParams(dslJSON) + if err != nil { + common.Warn("CleanComponentParams: failed to extract DSL schema, returning input as-is", + zap.Error(err)) + return rawConfig + } + + validCPNs := make(map[string]map[string]struct{}, len(schemas)) + for _, s := range schemas { + keys := make(map[string]struct{}, len(s.ParamsDefaults)) + for k := range s.ParamsDefaults { + keys[k] = struct{}{} + } + validCPNs[s.CpnID] = keys + } + + result := make(map[string]interface{}, len(rawConfig)) + for key, val := range rawConfig { + if !strings.Contains(key, ":") { + common.Warn("CleanComponentParams: dropping legacy flat field", + zap.String("key", key)) + continue + } + validKeys, ok := validCPNs[key] + if !ok { + common.Warn("CleanComponentParams: dropping unknown cpnID", + zap.String("cpnID", key)) + continue + } + params, ok := val.(map[string]any) + if !ok { + continue + } + cleaned := make(map[string]any, len(params)) + for pk, pv := range params { + if _, ok := validKeys[pk]; ok { + cleaned[pk] = pv + } else { + common.Warn("CleanComponentParams: dropping unknown param", + zap.String("cpnID", key), zap.String("param", pk)) + } + } + if len(cleaned) > 0 { + result[key] = cleaned + } + } + return result +} + +// BuildParserConfig builds the final parser_config by starting from the DSL +// defaults for every component, then overlaying the cleaned incoming overrides. +// This ensures all components from the current pipeline are present while +// stripping stale params from other pipelines. +func BuildParserConfig(dslJSON []byte, rawConfig map[string]interface{}) entity.JSONMap { + cleaned := CleanComponentParams(dslJSON, rawConfig) + defaults, err := ComponentParamsDefaults(dslJSON) + if err != nil { + common.Warn("BuildParserConfig: failed to extract DSL defaults, using cleaned only", + zap.Error(err)) + return entity.JSONMap(cleaned) + } + result := make(entity.JSONMap, len(defaults)) + for cpnID, params := range defaults { + base := make(map[string]interface{}, len(params)) + for k, v := range params { + base[k] = v + } + if over, ok := cleaned[cpnID]; ok { + if om, ok := over.(map[string]any); ok { + for k, v := range om { + base[k] = v + } + } + } + result[cpnID] = base + } + return result +} + +// ResolveComponentParamsDefaults takes DSL JSON bytes and returns the +// component params defaults as an entity.JSONMap {cpnID: {param: value}}. +// This is a pure function — callers must load the DSL themselves. +func ResolveComponentParamsDefaults(dslJSON []byte) (entity.JSONMap, error) { + cp, err := ComponentParamsDefaults(dslJSON) + if err != nil { + return nil, err + } + out := make(entity.JSONMap, len(cp)) + for k, v := range cp { + out[k] = v + } + return out, nil +} + +// ResolveComponentParamsDefaultsFromIDs loads the DSL for the target pipeline +// (builtin via parserID or custom canvas via pipelineID) and returns the +// component params defaults. It is a convenience wrapper around +// LoadPipelineDSL + ResolveComponentParamsDefaults; use the two-step form +// when you already have the DSL bytes. +// +// Deprecated: prefer loading the DSL yourself and calling +// ResolveComponentParamsDefaults directly, to keep DAO dependencies +// out of the pipeline package. +func ResolveComponentParamsDefaultsFromIDs(parserID string, pipelineID *string) (entity.JSONMap, error) { + isCanvas := pipelineID != nil && strings.TrimSpace(*pipelineID) != "" + var dslJSON []byte + var err error + if isCanvas { + return nil, fmt.Errorf("ResolveComponentParamsDefaultsFromIDs: cannot load canvas DSL without DAO; " + + "load the DSL via service.LoadPipelineDSL first") + } + registry, regErr := DefaultRegistry() + if regErr != nil { + return nil, fmt.Errorf("builtin registry: %w", regErr) + } + if !registry.IsValid(parserID) { + return nil, fmt.Errorf("unknown builtin parser_id: %q", parserID) + } + dslStr, dslErr := LoadBuiltinDSL(parserID) + if dslErr != nil { + return nil, fmt.Errorf("load builtin DSL: %w", dslErr) + } + dslJSON = []byte(dslStr) + cp, err := ComponentParamsDefaults(dslJSON) + if err != nil { + return nil, err + } + out := make(entity.JSONMap, len(cp)) + for k, v := range cp { + out[k] = v + } + return out, nil +} diff --git a/internal/ingestion/pipeline/pipeline_params_test.go b/internal/ingestion/pipeline/pipeline_params_test.go new file mode 100644 index 0000000000..2e732329fe --- /dev/null +++ b/internal/ingestion/pipeline/pipeline_params_test.go @@ -0,0 +1,268 @@ +package pipeline + +import ( + "encoding/json" + "testing" + + "ragflow/internal/entity" +) + +// generalDSL returns a minimal DSL resembling the "general" template's component +// structure, with one Parser and one Chunker component. +func generalDSL(t *testing.T) []byte { + t.Helper() + dsl := map[string]any{ + "components": map[string]any{ + "Parser:HipSignsRhyme": map[string]any{ + "obj": map[string]any{ + "component_name": "Parser", + "params": map[string]any{ + "outputs": map[string]any{}, + "pdf": map[string]any{"parse_method": "DeepDOC", "lang": "en"}, + "docx": map[string]any{"output_format": "json"}, + }, + }, + }, + "Chunker:LegalReadersDecide": map[string]any{ + "obj": map[string]any{ + "component_name": "Chunker", + "params": map[string]any{ + "outputs": map[string]any{}, + "chunk_size": float64(512), + "chunk_overlap": float64(128), + }, + }, + }, + }, + } + raw, err := json.Marshal(dsl) + if err != nil { + t.Fatalf("marshal dsl fixture: %v", err) + } + return raw +} +func TestCleanComponentParams_DropsLegacyFlatFields(t *testing.T) { + dslJSON := generalDSL(t) + raw := map[string]any{ + "chunk_token_num": 256, + "image_context_size": 10, + } + result := CleanComponentParams(dslJSON, raw) + if len(result) != 0 { + t.Errorf("expected empty result, got %v", result) + } +} + +func TestCleanComponentParams_DropsUnknownCPNID(t *testing.T) { + dslJSON := generalDSL(t) + raw := map[string]any{ + "Parser:NoSuch": map[string]any{"chunk_size": float64(256)}, + } + result := CleanComponentParams(dslJSON, raw) + if _, ok := result["Parser:NoSuch"]; ok { + t.Error("expected unknown cpnID to be dropped") + } +} + +func TestCleanComponentParams_DropsUnknownParamKey(t *testing.T) { + dslJSON := generalDSL(t) + raw := map[string]any{ + "Parser:HipSignsRhyme": map[string]any{ + "no_such_param": 1, + "pdf": map[string]any{"parse_method": "deepdoc"}, + }, + } + result := CleanComponentParams(dslJSON, raw) + params := result["Parser:HipSignsRhyme"].(map[string]any) + if _, ok := params["no_such_param"]; ok { + t.Error("expected unknown param key to be dropped") + } + if _, ok := params["pdf"]; !ok { + t.Error("expected known param key 'pdf' to be kept") + } +} + +func TestCleanComponentParams_ReturnsInputOnDSLError(t *testing.T) { + result := CleanComponentParams([]byte("not json"), map[string]any{"key": "val"}) + if result["key"] != "val" { + t.Error("expected input returned as-is on DSL error") + } +} + +func TestCleanComponentParams_ValidCPNIDPassesThrough(t *testing.T) { + dslJSON := generalDSL(t) + raw := map[string]any{ + "Parser:HipSignsRhyme": map[string]any{ + "pdf": map[string]any{"parse_method": "deepdoc"}, + }, + "Chunker:LegalReadersDecide": map[string]any{ + "chunk_size": float64(256), + }, + } + result := CleanComponentParams(dslJSON, raw) + if _, ok := result["Parser:HipSignsRhyme"]; !ok { + t.Error("expected Parser:HipSignsRhyme to pass through") + } + if _, ok := result["Chunker:LegalReadersDecide"]; !ok { + t.Error("expected Chunker:LegalReadersDecide to pass through") + } +} + +// --- BuildParserConfig --- + +func TestBuildParserConfig_ShallowMerge_NestedParam(t *testing.T) { + // A component with a nested-map param "chunk" that has sub-keys. + dsl := map[string]any{ + "components": map[string]any{ + "Chunker:xyz": map[string]any{ + "obj": map[string]any{ + "component_name": "Chunker", + "params": map[string]any{ + "outputs": map[string]any{}, + "chunk": map[string]any{ + "size": float64(512), + "overlap": float64(128), + }, + }, + }, + }, + }, + } + dslJSON, err := json.Marshal(dsl) + if err != nil { + t.Fatalf("marshal dsl: %v", err) + } + + // User only overrides one sub-key of "chunk". + overrides := map[string]any{ + "Chunker:xyz": map[string]any{ + "chunk": map[string]any{"size": float64(1024)}, + }, + } + + result := BuildParserConfig(dslJSON, overrides) + chunker, ok := result["Chunker:xyz"].(map[string]any) + if !ok { + t.Fatal("expected Chunker:xyz in result") + } + chunk, ok := chunker["chunk"].(map[string]any) + if !ok { + t.Fatal("expected chunk key in result") + } + // After shallow merge: size is overridden, overlap is GONE. + if chunk["size"] != float64(1024) { + t.Errorf("expected size=1024 from override, got %v", chunk["size"]) + } + if _, ok := chunk["overlap"]; ok { + t.Error("shallow merge: overlap from defaults should NOT be preserved when chunk is fully replaced") + } +} + +func TestBuildParserConfig_ScalarOverridePreservesOtherDefaults(t *testing.T) { + dslJSON := generalDSL(t) + overrides := map[string]any{ + "Chunker:LegalReadersDecide": map[string]any{ + "chunk_size": float64(1024), + }, + } + result := BuildParserConfig(dslJSON, overrides) + chunker, ok := result["Chunker:LegalReadersDecide"].(map[string]any) + if !ok { + t.Fatal("expected Chunker:LegalReadersDecide in result") + } + if chunker["chunk_size"] != float64(1024) { + t.Errorf("expected chunk_size=1024, got %v", chunker["chunk_size"]) + } + // chunk_overlap should be preserved from DSL defaults since it wasn't overridden. + if chunker["chunk_overlap"] != float64(128) { + t.Errorf("expected chunk_overlap=128 preserved from defaults, got %v", chunker["chunk_overlap"]) + } +} + +func TestBuildParserConfig_UnknownCPNIDNotPresentInResult(t *testing.T) { + dslJSON := generalDSL(t) + overrides := map[string]any{ + "Parser:Unknown": map[string]any{"chunk_size": float64(256)}, + } + result := BuildParserConfig(dslJSON, overrides) + // Unknown cpnID should be dropped by CleanComponentParams; the result should + // still contain the DSL-defined components with their defaults. + if _, ok := result["Parser:Unknown"]; ok { + t.Error("expected unknown cpnID to be absent from result") + } + if _, ok := result["Parser:HipSignsRhyme"]; !ok { + t.Error("expected valid component from DSL to be present") + } +} + +func TestBuildParserConfig_FallbackOnDSLError(t *testing.T) { + result := BuildParserConfig([]byte("not json"), map[string]any{"key": "val"}) + if result["key"] != "val" { + t.Error("expected fallback to return raw config on DSL error") + } +} + +func TestBuildParserConfig_AllComponentsPresent(t *testing.T) { + dslJSON := generalDSL(t) + result := BuildParserConfig(dslJSON, nil) + // Both components from the DSL fixture should be present. + if _, ok := result["Parser:HipSignsRhyme"]; !ok { + t.Error("expected Parser:HipSignsRhyme") + } + if _, ok := result["Chunker:LegalReadersDecide"]; !ok { + t.Error("expected Chunker:LegalReadersDecide") + } +} + +// --- ResolveComponentParamsDefaults --- + +func TestResolveComponentParamsDefaults_Basic(t *testing.T) { + dslJSON := generalDSL(t) + result, err := ResolveComponentParamsDefaults(dslJSON) + if err != nil { + t.Fatalf("ResolveComponentParamsDefaults: %v", err) + } + if len(result) == 0 { + t.Fatal("expected non-empty result") + } + // outputs should be stripped. + parser := result["Parser:HipSignsRhyme"].(map[string]any) + if _, ok := parser["outputs"]; ok { + t.Error("expected outputs to be stripped") + } + if _, ok := parser["pdf"]; !ok { + t.Error("expected pdf to be present") + } + chunker := result["Chunker:LegalReadersDecide"].(map[string]any) + if chunker["chunk_size"] != float64(512) { + t.Errorf("expected chunk_size=512, got %v", chunker["chunk_size"]) + } +} + +func TestResolveComponentParamsDefaults_InvalidJSON(t *testing.T) { + _, err := ResolveComponentParamsDefaults([]byte("not json")) + if err == nil { + t.Error("expected error for invalid JSON") + } +} + +func TestResolveComponentParamsDefaults_ResultIsMutable(t *testing.T) { + // Verify the returned map is a copy, not a reference to internal state. + dslJSON := generalDSL(t) + result, err := ResolveComponentParamsDefaults(dslJSON) + if err != nil { + t.Fatalf("ResolveComponentParamsDefaults: %v", err) + } + // Mutate the result. + parser := result["Parser:HipSignsRhyme"].(map[string]any) + delete(parser, "pdf") + // Re-read: the second call should return a fresh copy unaffected by the mutation. + result2, _ := ResolveComponentParamsDefaults(dslJSON) + parser2 := result2["Parser:HipSignsRhyme"].(map[string]any) + if _, ok := parser2["pdf"]; !ok { + t.Error("expected result to be independent copy (pdf preserved)") + } +} + +// ensure entity.JSONMap is used (import used). +var _ entity.JSONMap diff --git a/internal/ingestion/service/doc_state.go b/internal/ingestion/service/doc_state.go index 1167599d2a..21e5ad7dae 100644 --- a/internal/ingestion/service/doc_state.go +++ b/internal/ingestion/service/doc_state.go @@ -21,7 +21,7 @@ import ( "ragflow/internal/common" taskpkg "ragflow/internal/ingestion/task" - servicepkg "ragflow/internal/service" + documentpkg "ragflow/internal/service/document" ) // docStateSvc is the subset of *service.DocumentService needed to finalize a @@ -46,7 +46,7 @@ type docStateUpdater struct { // injected at construction time. Tests inject stubs via the docSvc field. func newDocStateUpdater() *docStateUpdater { return &docStateUpdater{ - docSvc: servicepkg.NewDocumentService(), + docSvc: documentpkg.NewDocumentService(), } } diff --git a/internal/ingestion/service/ingestion_service.go b/internal/ingestion/service/ingestion_service.go index 755b62bf58..465e6bc966 100644 --- a/internal/ingestion/service/ingestion_service.go +++ b/internal/ingestion/service/ingestion_service.go @@ -34,6 +34,7 @@ import ( pipelinepkg "ragflow/internal/ingestion/pipeline" taskpkg "ragflow/internal/ingestion/task" servicepkg "ragflow/internal/service" + documentpkg "ragflow/internal/service/document" "github.com/cenkalti/backoff/v5" ) @@ -607,7 +608,7 @@ func (e *Ingestor) pollCancel(taskID string, cancel context.CancelFunc, done <-c // row. Mirrors Python's cancel_all_task_of: progress=-1, run=CANCEL, and an // appended timestamped cancel message (progress_msg += cancelMsg). func (e *Ingestor) markCancelProgress(task *entity.IngestionTask) { - svc := servicepkg.NewDocumentService() + svc := documentpkg.NewDocumentService() doc, err := svc.GetDocumentByID(task.DocumentID) if err != nil { common.Error(fmt.Sprintf("markCancelProgress: load document %s: %v", task.DocumentID, err), err) @@ -625,7 +626,7 @@ func (e *Ingestor) markCancelProgress(task *entity.IngestionTask) { // row. Unlike cancellation (markCancelProgress), this records a TIMEOUT // failure rather than a user-initiated stop. func (e *Ingestor) markTimeoutProgress(task *entity.IngestionTask) { - svc := servicepkg.NewDocumentService() + svc := documentpkg.NewDocumentService() doc, err := svc.GetDocumentByID(task.DocumentID) if err != nil { common.Error(fmt.Sprintf("markTimeoutProgress: load document %s: %v", task.DocumentID, err), err) diff --git a/internal/ingestion/service/progress_sink.go b/internal/ingestion/service/progress_sink.go index 0146d3a23d..3f0f350434 100644 --- a/internal/ingestion/service/progress_sink.go +++ b/internal/ingestion/service/progress_sink.go @@ -25,6 +25,7 @@ import ( "ragflow/internal/entity" "ragflow/internal/ingestion/pipeline" servicepkg "ragflow/internal/service" + documentpkg "ragflow/internal/service/document" ) // progressSink implements pipeline.ProgressSink. It is the single writer of @@ -60,7 +61,7 @@ func newProgressSink(taskSvc *servicepkg.IngestionTaskService) *progressSink { // server-config dependency, so this is safe in any environment. return &progressSink{ taskSvc: taskSvc, - docSvc: servicepkg.NewDocumentService(), + docSvc: documentpkg.NewDocumentService(), } } diff --git a/internal/ingestion/service/progress_sink_test.go b/internal/ingestion/service/progress_sink_test.go index 09f8abfb75..710d8df10c 100644 --- a/internal/ingestion/service/progress_sink_test.go +++ b/internal/ingestion/service/progress_sink_test.go @@ -26,6 +26,7 @@ import ( "ragflow/internal/ingestion/pipeline" "ragflow/internal/ingestion/testutil" servicepkg "ragflow/internal/service" + "ragflow/internal/service/document" ) // TestProgressSink_CanConstructDocumentServiceWithoutServerConfig ensures the @@ -39,7 +40,7 @@ func TestProgressSink_CanConstructDocumentServiceWithoutServerConfig(t *testing. defer cleanup() // No server config is initialized in the test env; this must not panic. - svc := servicepkg.NewDocumentService() + svc := document.NewDocumentService() if svc == nil { t.Fatal("expected non-nil DocumentService") } diff --git a/internal/service/agent.go b/internal/service/agent.go index 05c68de3e2..ed15dcb39d 100644 --- a/internal/service/agent.go +++ b/internal/service/agent.go @@ -21,6 +21,7 @@ import ( "encoding/json" "errors" "fmt" + "ragflow/internal/service/file" "ragflow/internal/utility" "reflect" "sort" @@ -1219,8 +1220,10 @@ func (s *AgentService) buildRunFunc(canvasID string, versionRow *entity.UserCanv state.Sys["tenant_id"] = tid } if rawFiles, ok := root["files"].([]map[string]interface{}); ok && len(rawFiles) > 0 { - fileSvc := NewFileService() - files, ferr := fileSvc.parseAgentUploads(userID, rawFiles, beginLayoutRecognize(c)) + // Only used for ParseAgentUploads (read-only); nil DocRemover means + // this FileService MUST NOT be used for DeleteFiles. + fileSvc := file.NewFileService(CheckFileTeamPermission, nil) + files, ferr := fileSvc.ParseAgentUploads(userID, rawFiles, beginLayoutRecognize(c)) if ferr != nil { s.markRunFailed(ctx2, runID, "parse files: "+ferr.Error()) return nil, fmt.Errorf("parse agent files: %w", ferr) diff --git a/internal/service/chat.go b/internal/service/chat.go index 07665d8e25..ad2eac7948 100644 --- a/internal/service/chat.go +++ b/internal/service/chat.go @@ -311,7 +311,7 @@ func (s *ChatService) validateCreateDatasetIDs(value interface{}, tenantID strin kbs = append(kbs, kb) } - if err := validateDatasetEmbeddingModels(kbs); err != nil { + if err := ValidateDatasetEmbeddingModels(kbs); err != nil { return nil, err } return normalizedIDs, nil diff --git a/internal/service/chat_pipeline.go b/internal/service/chat_pipeline.go index 150803587a..35f4593990 100644 --- a/internal/service/chat_pipeline.go +++ b/internal/service/chat_pipeline.go @@ -26,6 +26,7 @@ import ( "ragflow/internal/engine" "ragflow/internal/entity" modelModule "ragflow/internal/entity/models" + "ragflow/internal/service/file" "ragflow/internal/service/graph" "ragflow/internal/service/nlp" "regexp" @@ -48,7 +49,7 @@ import ( type ChatPipelineService struct { ModelProviderSvc *ModelProviderService MetadataSvc *MetadataService - datasetService *DatasetService + kbDAO *dao.KnowledgebaseDAO } // NewChatPipelineService creates a new ChatPipelineService with all required dependencies. @@ -56,7 +57,7 @@ func NewChatPipelineService() *ChatPipelineService { return &ChatPipelineService{ ModelProviderSvc: NewModelProviderService(), MetadataSvc: NewMetadataService(), - datasetService: NewDatasetService(), + kbDAO: dao.NewKnowledgebaseDAO(), } } @@ -347,7 +348,7 @@ func (s *ChatPipelineService) AsyncChat( // === Phase 6: SQL Retrieval === // Retrieve field_map for SQL retrieval (preferred over vector search) promptConfig := chat.PromptConfig - fieldMap, fmErr := s.datasetService.GetFieldMap(kbIDStrings(kbs)) + fieldMap, fmErr := s.kbDAO.GetFieldMap(kbIDStrings(kbs)) if fmErr != nil { common.Warn("get_field_map failed; proceeding without field_map", zap.Error(fmErr)) fieldMap = nil @@ -1602,7 +1603,9 @@ func (s *ChatPipelineService) AsyncChatSolo( func (s *ChatPipelineService) extractImageFiles(userID string, files interface{}) []string { // ── File-dict mode ── if fileDicts, ok := parseFileDicts(files); ok { - fileSvc := NewFileService() + // Only used for GetFileContents (read-only); nil DocRemover means + // this FileService MUST NOT be used for DeleteFiles. + fileSvc := file.NewFileService(CheckFileTeamPermission, nil) // Use raw=false to get base64 data URIs for images. _, images, err := fileSvc.GetFileContents(userID, fileDicts, false) if err != nil { @@ -1959,7 +1962,7 @@ func (s *ChatPipelineService) getModels(ctx context.Context, chat *entity.Chat) // Embedding model. var embModel *modelModule.EmbeddingModel if len(kbs) > 0 { - if err := validateDatasetEmbeddingModels(kbs); err != nil { + if err := ValidateDatasetEmbeddingModels(kbs); err != nil { return nil, nil, nil, nil, nil, err } if kbs[0].EmbdID != "" { @@ -2062,7 +2065,9 @@ func lastUserQuestion(messages []map[string]interface{}) string { func (s *ChatPipelineService) processFileAttachments(userID string, files interface{}) string { // ── File-dict mode ── if fileDicts, ok := parseFileDicts(files); ok { - fileSvc := NewFileService() + // Only used for GetFileContents (read-only); nil DocRemover means + // this FileService MUST NOT be used for DeleteFiles. + fileSvc := file.NewFileService(CheckFileTeamPermission, nil) texts, _, err := fileSvc.GetFileContents(userID, fileDicts, false) if err != nil { common.Warn("GetFileContents failed in processFileAttachments", @@ -2117,7 +2122,9 @@ func (s *ChatPipelineService) processFileAttachments(userID string, files interf func splitFileAttachments(userID string, files interface{}, raw bool) (textAttachments []string, imageAttachments []string) { // ── Mode 1: file dicts (Python-compatible) ── if fileDicts, ok := parseFileDicts(files); ok { - fileSvc := NewFileService() + // Only used for GetFileContents (read-only); nil DocRemover means + // this FileService MUST NOT be used for DeleteFiles. + fileSvc := file.NewFileService(CheckFileTeamPermission, nil) texts, images, err := fileSvc.GetFileContents(userID, fileDicts, raw) if err != nil { common.Warn("GetFileContents failed, falling back to string splitting", diff --git a/internal/service/chunk/chunk.go b/internal/service/chunk/chunk.go index fe57c6a745..66547ecfc1 100644 --- a/internal/service/chunk/chunk.go +++ b/internal/service/chunk/chunk.go @@ -45,6 +45,7 @@ import ( "ragflow/internal/engine" "ragflow/internal/engine/types" "ragflow/internal/service" + "ragflow/internal/service/document" "ragflow/internal/service/nlp" "ragflow/internal/storage" "ragflow/internal/tokenizer" @@ -88,7 +89,7 @@ type ChunkService struct { // startParseDocumentsFunc overrides the DSL start-parse flow. Production // uses service.DocumentService.StartParseDocuments; tests inject a fake // to avoid the MQ publisher. - startParseDocumentsFunc func(doc *entity.Document, kb *entity.Knowledgebase, userID string, opts service.StartParseOptions) error + startParseDocumentsFunc func(doc *entity.Document, kb *entity.Knowledgebase, userID string, opts document.StartParseOptions) error // cancelIngestionTaskFunc overrides the document-parsing cancellation. // Production uses service.DocumentService.CancelDocParse; tests inject // a fake to avoid the MQ publisher. @@ -603,7 +604,7 @@ const ( func (s *ChunkService) cancelAllTasksOfDoc(doc *entity.Document) error { cancel := s.cancelIngestionTaskFunc if cancel == nil { - cancel = service.NewDocumentService().CancelDocParse + cancel = document.NewDocumentService().CancelDocParse } return cancel(doc) } @@ -749,13 +750,13 @@ func (s *ChunkService) Parse(userID, datasetID string, req *service.ParseFileReq // Batch pre-check: refuse the whole request if any document's ingestion // task is non-terminal (RUNNING/STOPPING), so we never partially clean. - if err := (service.NewDocumentService().AssertIngestionTasksTerminal(docIDs)); err != nil { + if err := (document.NewDocumentService().AssertIngestionTasksTerminal(docIDs)); err != nil { return nil, common.CodeDataError, err } startParse := s.startParseDocumentsFunc if startParse == nil { - docSvc := service.NewDocumentService() + docSvc := document.NewDocumentService() startParse = docSvc.StartParseDocuments } @@ -763,7 +764,7 @@ func (s *ChunkService) Parse(userID, datasetID string, req *service.ParseFileReq for _, docID := range docIDs { doc := docByID[docID] - if err := startParse(doc, kb, userID, service.StartParseOptions{RerunWithDelete: true}); err != nil { + if err := startParse(doc, kb, userID, document.StartParseOptions{RerunWithDelete: true}); err != nil { return nil, common.CodeServerError, err } successCount++ diff --git a/internal/service/chunk/chunk_test.go b/internal/service/chunk/chunk_test.go index a3c868266c..c03f625c04 100644 --- a/internal/service/chunk/chunk_test.go +++ b/internal/service/chunk/chunk_test.go @@ -13,6 +13,7 @@ import ( "ragflow/internal/entity" "ragflow/internal/entity/models" "ragflow/internal/service" + "ragflow/internal/service/document" "ragflow/internal/storage" "reflect" "strings" @@ -1361,9 +1362,9 @@ func TestChunkServiceParse_CallsStartParseDocumentsWithRerunWithDelete(t *testin svc := newParseTestService(t) svc.accessibleFunc = func(string, string) bool { return true } - var calledOpts service.StartParseOptions + var calledOpts document.StartParseOptions var calledDocID string - svc.startParseDocumentsFunc = func(doc *entity.Document, kb *entity.Knowledgebase, userID string, opts service.StartParseOptions) error { + svc.startParseDocumentsFunc = func(doc *entity.Document, kb *entity.Knowledgebase, userID string, opts document.StartParseOptions) error { calledOpts = opts calledDocID = doc.ID return nil @@ -1392,7 +1393,7 @@ func TestChunkServiceParse_ReturnsPartialSuccessForDuplicateDocumentIDs(t *testi insertChunkTestDoc(t, "doc-1", datasetID) svc := newParseTestService(t) - svc.startParseDocumentsFunc = func(*entity.Document, *entity.Knowledgebase, string, service.StartParseOptions) error { + svc.startParseDocumentsFunc = func(*entity.Document, *entity.Knowledgebase, string, document.StartParseOptions) error { return nil } diff --git a/internal/service/dataset.go b/internal/service/dataset.go deleted file mode 100644 index 0e549d705a..0000000000 --- a/internal/service/dataset.go +++ /dev/null @@ -1,3190 +0,0 @@ -// -// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. -// - -package service - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "math" - "math/rand" - "ragflow/internal/common" - "ragflow/internal/dao" - "ragflow/internal/engine" - redisengine "ragflow/internal/engine/redis" - "ragflow/internal/engine/types" - enginetypes "ragflow/internal/engine/types" - "ragflow/internal/entity" - "ragflow/internal/entity/models" - pipelinepkg "ragflow/internal/ingestion/pipeline" - "ragflow/internal/service/nlp" - "ragflow/internal/utility" - "regexp" - "sort" - "strconv" - "strings" - "time" - - "github.com/cespare/xxhash/v2" - "github.com/google/uuid" - "go.uber.org/zap" - "gorm.io/gorm" - "gorm.io/gorm/clause" -) - -var ( - datasetSupportedAvatarMIMETypes = map[string]struct{}{ - "image/jpeg": {}, - "image/png": {}, - } - datasetAllowedOrderByFields = map[string]struct{}{ - "create_time": {}, - "update_time": {}, - } - datasetAllowedMetadataTypes = map[string]struct{}{ - "string": {}, - "list": {}, - "time": {}, - "number": {}, - } - validIndexTypes = []string{"graph", "raptor", "mindmap"} - indexTypeToTaskType = map[string]string{"graph": "graphrag", "raptor": "raptor", "mindmap": "mindmap"} - indexTypeToDisplayName = map[string]string{"graph": "Graph", "raptor": "RAPTOR", "mindmap": "Mindmap"} -) - -const ( - // Keep the legacy worker marker in queue payloads; persisted tasks use a real document ID. - graphRaptorQueueDocID = "graph_raptor_x" - maximumTaskPageNumber = int64(100000000) - serverQueueNamePrefix = "te" - defaultEmbeddingCheckNum = 5 - - graphPhaseResolutionDone = "resolution_done" - graphPhaseCommunityDone = "community_done" -) - -// validateDatasetEmbeddingModels checks that all given datasets use the same -// embedding model (or all use none). Returns an error on mismatch. -func validateDatasetEmbeddingModels(kbs []*entity.Knowledgebase) error { - embdIDs := make(map[string]struct{}) - hasEmbd := false - noEmbd := false - for _, kb := range kbs { - if kb.EmbdID != "" { - hasEmbd = true - baseName := kb.EmbdID - if idx := strings.LastIndex(kb.EmbdID, "@"); idx > 0 { - baseName = kb.EmbdID[:idx] - // Strip the second-to-last @-segment too (instance name), - // matching Python's _base_model_name which uses rsplit("@", 2). - if idx2 := strings.LastIndex(baseName, "@"); idx2 > 0 { - baseName = baseName[:idx2] - } - } - embdIDs[baseName] = struct{}{} - } else { - noEmbd = true - } - } - if hasEmbd && noEmbd { - return fmt.Errorf("Cannot search across datasets where some have embedding models and others do not.") - } - if len(embdIDs) > 1 { - return fmt.Errorf("Datasets use different embedding models: %v", getEmbdIDs(kbs)) - } - return nil -} - -// getEmbdIDs extracts embedding IDs from knowledge bases. -func getEmbdIDs(kbs []*entity.Knowledgebase) []string { - ids := make([]string, len(kbs)) - for i, kb := range kbs { - ids[i] = kb.EmbdID - } - return ids -} - -// DatasetService implements the RESTful dataset APIs from dataset_api.py. -type DatasetService struct { - kbDAO *dao.KnowledgebaseDAO - documentDAO *dao.DocumentDAO - connectorDAO *dao.ConnectorDAO - tenantDAO *dao.TenantDAO - tenantLLMDAO *dao.TenantLLMDAO - pipelineLogDAO *dao.PipelineOperationLogDAO - userTenantDAO *dao.UserTenantDAO - taskDAO *dao.TaskDAO - searchService *SearchService - docEngine engine.DocEngine - embeddingCache *utility.EmbeddingLRU -} - -// NewDatasetService creates a new datasets service. -func NewDatasetService() *DatasetService { - return &DatasetService{ - kbDAO: dao.NewKnowledgebaseDAO(), - documentDAO: dao.NewDocumentDAO(), - connectorDAO: dao.NewConnectorDAO(), - tenantDAO: dao.NewTenantDAO(), - tenantLLMDAO: dao.NewTenantLLMDAO(), - pipelineLogDAO: dao.NewPipelineOperationLogDAO(), - userTenantDAO: dao.NewUserTenantDAO(), - taskDAO: dao.NewTaskDAO(), - searchService: NewSearchService(), - docEngine: engine.Get(), - embeddingCache: utility.NewEmbeddingLRU(1000), - } -} - -func (d *DatasetService) UpdateDocumentMetadataConfig(userID, datasetID, documentID string, req map[string]interface{}) (*entity.Document, common.ErrorCode, error) { - if _, err := d.kbDAO.GetByIDAndTenantID(datasetID, userID); err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, errors.New("You don't own the dataset.") - } - return nil, common.CodeServerError, errors.New("Database operation failed") - } - - doc, err := d.documentDAO.GetByDocumentIDAndDatasetID(documentID, datasetID) - if err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, fmt.Errorf("Document %s not found in dataset %s", documentID, datasetID) - } - return nil, common.CodeServerError, err - } - - metadata, ok := req["metadata"] - if !ok { - return nil, common.CodeArgumentError, errors.New("metadata is required") - } - - parserConfig := doc.ParserConfig - if parserConfig == nil { - parserConfig = entity.JSONMap{} - } - parserConfig["metadata"] = metadata - - if err := d.documentDAO.UpdateByID(doc.ID, map[string]interface{}{"parser_config": parserConfig}); err != nil { - return nil, common.CodeExceptionError, err - } - - updatedDoc, err := d.documentDAO.GetByID(doc.ID) - if err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, errors.New("Document not found!") - } - return nil, common.CodeExceptionError, err - } - - return updatedDoc, common.CodeSuccess, nil -} - -// checkType reports whether indexType is supported by dataset index APIs. -func checkType(indexType string) bool { - haveType := false - for _, t := range validIndexTypes { - if indexType == t { - haveType = true - } - } - return haveType -} - -func (d *DatasetService) newRaptorOrGraphRagTask(sampleDoc *entity.Document, taskType string, taskDocID string, queueDocID string, docIDs []string) (*entity.Task, map[string]interface{}, error) { - if docIDs == nil || len(docIDs) == 0 { - docIDs = make([]string, 0) - } - if !checkIndexTaskType(taskType) { - return nil, nil, errors.New("type should be graphrag, raptor or mindmap") - } - - chunkingConfig, err := d.documentDAO.GetChunkingConfig(sampleDoc.ID) - if err != nil { - return nil, nil, err - } - - hasher := xxhash.New() - keys := make([]string, 0, len(chunkingConfig)) - for key := range chunkingConfig { - keys = append(keys, key) - } - sort.Strings(keys) - for _, key := range keys { - _, _ = hasher.Write([]byte(key)) - _, _ = hasher.Write([]byte{0}) - v, mErr := json.Marshal(chunkingConfig[key]) - if mErr != nil { - return nil, nil, mErr - } - _, _ = hasher.Write(v) - _, _ = hasher.Write([]byte{0}) - } - - taskID := utility.GenerateUUID() - beginAt := time.Now().Truncate(time.Second) - progressMsg := beginAt.Format("15:04:05") + " created task " + taskType - - for _, field := range []interface{}{taskDocID, maximumTaskPageNumber, maximumTaskPageNumber, taskType} { - _, _ = hasher.Write([]byte(fmt.Sprint(field))) - } - digest := fmt.Sprintf("%016x", hasher.Sum64()) - task := &entity.Task{ - ID: taskID, - DocID: taskDocID, - FromPage: maximumTaskPageNumber, - ToPage: maximumTaskPageNumber, - TaskType: taskType, - ProgressMsg: &progressMsg, - BeginAt: &beginAt, - Digest: &digest, - } - - queueMessage := map[string]interface{}{ - "id": taskID, - "doc_id": queueDocID, - "from_page": maximumTaskPageNumber, - "to_page": maximumTaskPageNumber, - "task_type": taskType, - "progress_msg": progressMsg, - "begin_at": beginAt.Format("2006-01-02 15:04:05"), - "digest": digest, - "doc_ids": docIDs, - } - - return task, queueMessage, nil -} - -func createDatasetIndexTaskInTx(tx *gorm.DB, task *entity.Task, queueDocID string) (*entity.Document, error) { - if task == nil { - return nil, errors.New("task is required") - } - if err := tx.Create(task).Error; err != nil { - return nil, err - } - - if queueDocID == "" { - return nil, nil - } - - var document entity.Document - err := tx.Select("id", "progress_msg", "process_begin_at").Where("id = ?", queueDocID).First(&document).Error - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return nil, nil - } - return nil, err - } - - beginAt := time.Now().Truncate(time.Second) - if task.BeginAt != nil { - beginAt = *task.BeginAt - } - if err := tx.Model(&entity.Document{}).Where("id = ?", queueDocID).Updates(map[string]interface{}{ - "progress_msg": "Task is queued...", - "process_begin_at": beginAt, - }).Error; err != nil { - return nil, err - } - - return &document, nil -} - -func enqueueDatasetIndexTask(priority int, queueMessage map[string]interface{}) error { - redisClient := redisengine.Get() - if redisClient == nil || !redisClient.QueueProduct(datasetIndexQueueName(priority), queueMessage) { - return errors.New("Can't access Redis. Please check the Redis' status") - } - return nil -} - -func cleanupFailedDatasetIndexTask(taskID string, updatedDocument *entity.Document, kbID string, indexType string) error { - return dao.DB.Transaction(func(tx *gorm.DB) error { - if err := tx.Unscoped().Where("id = ?", taskID).Delete(&entity.Task{}).Error; err != nil { - return fmt.Errorf("delete task %s: %w", taskID, err) - } - - if column := datasetIndexTaskIDColumn(indexType); kbID != "" && column != "" { - if err := tx.Model(&entity.Knowledgebase{}).Where("id = ? AND "+column+" = ?", kbID, taskID).Update(column, nil).Error; err != nil { - return fmt.Errorf("clear dataset task id %s: %w", taskID, err) - } - } - - if updatedDocument == nil { - return nil - } - - return tx.Model(&entity.Document{}).Where("id = ?", updatedDocument.ID).Updates(map[string]interface{}{ - "progress_msg": updatedDocument.ProgressMsg, - "process_begin_at": updatedDocument.ProcessBeginAt, - }).Error - }) -} - -func datasetIndexTaskIDColumn(indexType string) string { - switch indexType { - case "graph": - return "graphrag_task_id" - case "raptor": - return "raptor_task_id" - case "mindmap": - return "mindmap_task_id" - default: - return "" - } -} - -func datasetIndexTaskFinishAtColumn(indexType string) string { - switch indexType { - case "graph": - return "graphrag_task_finish_at" - case "raptor": - return "raptor_task_finish_at" - case "mindmap": - return "mindmap_task_finish_at" - default: - return "" - } -} - -func checkIndexTaskType(taskType string) bool { - switch taskType { - case "graphrag", "raptor", "mindmap": - return true - default: - return false - } -} - -func datasetIndexTaskID(kb *entity.Knowledgebase, indexType string) string { - if kb == nil { - return "" - } - switch indexType { - case "graph": - if kb.GraphragTaskID != nil { - return *kb.GraphragTaskID - } - case "raptor": - if kb.RaptorTaskID != nil { - return *kb.RaptorTaskID - } - case "mindmap": - if kb.MindmapTaskID != nil { - return *kb.MindmapTaskID - } - } - return "" -} - -func datasetIndexTaskIDUpdate(indexType, taskID string) map[string]interface{} { - switch indexType { - case "graph": - return map[string]interface{}{"graphrag_task_id": taskID} - case "raptor": - return map[string]interface{}{"raptor_task_id": taskID} - case "mindmap": - return map[string]interface{}{"mindmap_task_id": taskID} - default: - return map[string]interface{}{} - } -} - -func datasetIndexTaskIDs(kb *entity.Knowledgebase) []string { - if kb == nil { - return nil - } - - taskIDs := make([]string, 0, 3) - for _, taskID := range []*string{kb.GraphragTaskID, kb.RaptorTaskID, kb.MindmapTaskID} { - if taskID != nil && *taskID != "" { - taskIDs = append(taskIDs, *taskID) - } - } - return common.Deduplicate(taskIDs) -} - -func datasetIndexQueueName(priority int) string { - return fmt.Sprintf("%s.%d.common", serverQueueNamePrefix, priority) -} - -func interfaceSlice(items ...string) []interface{} { - result := make([]interface{}, len(items)) - for i, item := range items { - result[i] = item - } - return result -} - -func clearGraphPhaseMarkers(redisClient *redisengine.Client, datasetID string) { - if redisClient == nil || datasetID == "" { - return - } - for _, phase := range []string{graphPhaseResolutionDone, graphPhaseCommunityDone} { - if !redisClient.Delete(fmt.Sprintf("graphrag:phase:%s:%s", datasetID, phase)) { - common.Warn("Failed to clear GraphRAG phase marker", zap.String("dataset_id", datasetID), zap.String("phase", phase)) - } - } -} - -// RunIndex Run an indexing task (graph/raptor/mindmap) for a dataset. -func (d *DatasetService) RunIndex(userID, datasetID, indexType string) (map[string]interface{}, common.ErrorCode, error) { - if !checkType(indexType) { - return nil, common.CodeDataError, fmt.Errorf("Invalid index type '%s'. Must be one of %v", indexType, validIndexTypes) - } - - if datasetID == "" { - return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) - } - if !d.kbDAO.Accessible(datasetID, userID) { - return nil, common.CodeDataError, errors.New("No authorization.") - } - - kb, err := d.kbDAO.GetByID(datasetID) - if err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, errors.New("Invalid Dataset ID") - } - return nil, common.CodeDataError, errors.New("Internal server error") - } - - taskType := indexTypeToTaskType[indexType] - displayName := indexTypeToDisplayName[indexType] - - documents, code, err := d.getDocumentsByDatasetForIndex(datasetID) - if err != nil { - return nil, code, err - } - _ = documents - - sampleDocument := documents[0] - documentIDs := make([]string, len(documents)) - - for i, doc := range documents { - documentIDs[i] = doc.ID - } - - task, queueMessage, err := d.newRaptorOrGraphRagTask(sampleDocument, taskType, sampleDocument.ID, graphRaptorQueueDocID, documentIDs) - if err != nil { - common.Warn("Failed to build dataset index task", zap.String("dataset_id", datasetID), zap.String("task_type", taskType), zap.Error(err)) - return nil, common.CodeDataError, errors.New("Internal server error") - } - - var updatedDocument *entity.Document - var dataErr error - err = dao.DB.Transaction(func(tx *gorm.DB) error { - var lockedKB entity.Knowledgebase - if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). - Where("id = ? AND status = ?", kb.ID, string(entity.StatusValid)). - First(&lockedKB).Error; err != nil { - return err - } - - existingTaskID := datasetIndexTaskID(&lockedKB, indexType) - if existingTaskID != "" { - var existingTask entity.Task - taskErr := tx.Where("id = ?", existingTaskID).First(&existingTask).Error - if taskErr != nil { - if errors.Is(taskErr, gorm.ErrRecordNotFound) { - } else { - return taskErr - } - } else if existingTask.Progress != 1 && existingTask.Progress != -1 { - dataErr = fmt.Errorf("Task %s in progress with status %v. A %s Task is already running.", existingTaskID, existingTask.Progress, displayName) - return dataErr - } - } - - updatedDocument, err = createDatasetIndexTaskInTx(tx, task, graphRaptorQueueDocID) - if err != nil { - return err - } - return tx.Model(&entity.Knowledgebase{}).Where("id = ?", lockedKB.ID).Updates(datasetIndexTaskIDUpdate(indexType, task.ID)).Error - }) - if err != nil { - if dataErr != nil { - return nil, common.CodeDataError, dataErr - } - common.Warn("Failed to create dataset index task", zap.String("dataset_id", datasetID), zap.String("task_type", taskType), zap.Error(err)) - return nil, common.CodeDataError, errors.New("Internal server error") - } - - if err := enqueueDatasetIndexTask(0, queueMessage); err != nil { - if cleanupErr := cleanupFailedDatasetIndexTask(task.ID, updatedDocument, kb.ID, indexType); cleanupErr != nil { - err = errors.Join(err, cleanupErr) - } - common.Warn("Failed to queue dataset index task", zap.String("dataset_id", datasetID), zap.String("task_type", taskType), zap.Error(err)) - return nil, common.CodeDataError, errors.New("Internal server error") - } - - return map[string]interface{}{"task_id": task.ID}, common.CodeSuccess, nil -} - -func (d *DatasetService) getDocumentsByDatasetForIndex(datasetID string) ([]*entity.Document, common.ErrorCode, error) { - documents, _, err := d.documentDAO.GetByKBID(datasetID) - if err != nil { - common.Warn("Failed to load dataset documents for index", zap.String("dataset_id", datasetID), zap.Error(err)) - return nil, common.CodeDataError, errors.New("Internal server error") - } - if len(documents) == 0 { - return nil, common.CodeDataError, fmt.Errorf("No documents in Dataset %s", datasetID) - } - return documents, common.CodeSuccess, nil -} - -type TraceIndexRequest struct { - Type string `json:"type" binding:"required"` -} - -// TraceIndex Trace an indexing task (graph/raptor/mindmap) for a dataset. -func (d *DatasetService) TraceIndex(datasetID, userID, indexType string) (*entity.Task, common.ErrorCode, error) { - if !checkType(indexType) { - return nil, common.CodeDataError, fmt.Errorf("Invalid index type '%s'. Must be one of %v", indexType, validIndexTypes) - } - - if datasetID == "" { - return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) - } - if !d.kbDAO.Accessible(datasetID, userID) { - return nil, common.CodeDataError, errors.New("No authorization.") - } - - kb, err := d.kbDAO.GetByID(datasetID) - if err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, errors.New("Invalid Dataset ID") - } - return nil, common.CodeDataError, errors.New("Internal server error") - } - - taskID := datasetIndexTaskID(kb, indexType) - - var task *entity.Task - if taskID != "" { - task, err = d.taskDAO.GetByID(taskID) - if err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeSuccess, nil - } - return nil, common.CodeServerError, errors.New("Internal server error") - } - if task == nil { - return nil, common.CodeSuccess, nil - } - } - - return task, common.CodeSuccess, nil -} - -type CheckEmbeddingRequest struct { - EmbeddingID string `json:"embd_id" binding:"required"` - CheckNum *int `json:"check_num,omitempty"` -} - -type EmbeddingCheckSummary struct { - KbID string `json:"kb_id"` - Model string `json:"model"` - Sampled int `json:"sampled"` - Valid int `json:"valid"` - AvgCosSim float64 `json:"avg_cos_sim"` - MinCosSim float64 `json:"min_cos_sim"` - MaxCosSim float64 `json:"max_cos_sim"` - MatchMode string `json:"match_mode"` -} - -type EmbeddingCheckResult struct { - ChunkID string `json:"chunk_id"` - DocID string `json:"doc_id,omitempty"` - DocName string `json:"doc_name,omitempty"` - VectorField string `json:"vector_field,omitempty"` - VectorDim int `json:"vector_dim,omitempty"` - CosSim float64 `json:"cos_sim,omitempty"` - Reason string `json:"reason,omitempty"` -} - -type EmbeddingCheckResponse struct { - Summary EmbeddingCheckSummary `json:"summary"` - Results []EmbeddingCheckResult `json:"results"` -} - -type embeddingCheckSample struct { - ChunkID string - KbID string - DocID string - DocName string - VectorField string - Vector []float64 - PageNum interface{} - Position interface{} - Top interface{} - ContentWithWeight string - QuestionKeywords []string -} - -// CheckEmbedding checks whether a new embedding model is compatible with stored vectors. -func (d *DatasetService) CheckEmbedding(userID, datasetID string, req *CheckEmbeddingRequest) (*EmbeddingCheckResponse, common.ErrorCode, error) { - if datasetID == "" { - return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) - } - if !d.kbDAO.Accessible(datasetID, userID) { - return nil, common.CodeDataError, errors.New("No authorization.") - } - - kb, err := d.kbDAO.GetByID(datasetID) - if err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, errors.New("Invalid Dataset ID") - } - return nil, common.CodeServerError, errors.New("Internal server error") - } - - if req == nil || strings.TrimSpace(req.EmbeddingID) == "" { - return nil, common.CodeDataError, errors.New("`embd_id` is required.") - } - embeddingID := strings.TrimSpace(req.EmbeddingID) - if ok, message := d.verifyEmbeddingAvailability(embeddingID, userID); !ok { - return nil, common.CodeDataError, errors.New(message) - } - if d.docEngine == nil { - return nil, common.CodeServerError, errors.New("doc engine not initialized") - } - - driver, modelName, apiConfig, maxTokens, err := NewModelProviderService().ResolveModelConfig(kb.TenantID, entity.ModelTypeEmbedding, embeddingID) - if err != nil { - return nil, common.CodeDataError, err - } - embeddingModel := models.NewEmbeddingModel(driver, &modelName, apiConfig, maxTokens) - - checkNum := defaultEmbeddingCheckNum - if req.CheckNum != nil { - checkNum = *req.CheckNum - } - if checkNum <= 0 { - checkNum = defaultEmbeddingCheckNum - } - - samples, err := d.sampleRandomChunksWithVectors(context.Background(), kb.TenantID, datasetID, checkNum) - if err != nil { - return nil, common.CodeServerError, err - } - - results := make([]EmbeddingCheckResult, 0, len(samples)) - effectiveSimilarities := make([]float64, 0, len(samples)) - matchMode := "content_only" - for _, sample := range samples { - title := sample.DocName - if strings.TrimSpace(title) == "" { - title = "Title" - } - - textInput := strings.Join(sample.QuestionKeywords, "\n") - if strings.TrimSpace(textInput) == "" { - textInput = sample.ContentWithWeight - } - textInput = datasetCleanEmbeddingText(textInput) - if textInput == "" { - results = append(results, EmbeddingCheckResult{ChunkID: sample.ChunkID, Reason: "no_text"}) - continue - } - if len(sample.Vector) == 0 { - results = append(results, EmbeddingCheckResult{ChunkID: sample.ChunkID, Reason: "no_stored_vector"}) - continue - } - - vectors, err := datasetEncodeEmbedding(embeddingModel, []string{title, textInput}) - if err != nil { - return nil, common.CodeDataError, fmt.Errorf("Embedding failure. %w", err) - } - if len(vectors) < 2 { - return nil, common.CodeDataError, errors.New("Embedding failure. embedding response is incomplete") - } - if len(vectors[1]) != len(sample.Vector) { - return nil, common.CodeDataError, fmt.Errorf("Embedding failure. The dimension (%d) of given embedding model is different from the original (%d)", len(vectors[1]), len(sample.Vector)) - } - - simContent := datasetCosSim(vectors[1], sample.Vector) - simMix := datasetCosSim(datasetMixVectors(vectors[0], vectors[1], 0.1), sample.Vector) - sim := simContent - matchMode = "content_only" - if simMix > sim { - sim = simMix - matchMode = "title+content" - } - sim = datasetRoundFloat(sim, 6) - - effectiveSimilarities = append(effectiveSimilarities, sim) - results = append(results, EmbeddingCheckResult{ - ChunkID: sample.ChunkID, - DocID: sample.DocID, - DocName: sample.DocName, - VectorField: sample.VectorField, - VectorDim: len(sample.Vector), - CosSim: sim, - }) - } - - summary := datasetEmbeddingCheckSummary(datasetID, embeddingID, len(samples), effectiveSimilarities, matchMode) - response := &EmbeddingCheckResponse{Summary: summary, Results: results} - if len(effectiveSimilarities) == 0 { - return nil, common.CodeDataError, errors.New("No embedded chunks are available to compare.") - } - if summary.AvgCosSim >= 0.9 { - return response, common.CodeSuccess, nil - } - return response, common.CodeNotEffective, errors.New("Embedding model switch failed: the average similarity between old and new vectors is below 0.9, indicating incompatible vector spaces.") -} - -func (d *DatasetService) sampleRandomChunksWithVectors(ctx context.Context, tenantID, datasetID string, n int) ([]embeddingCheckSample, error) { - indexName := fmt.Sprintf("ragflow_%s", tenantID) - totalResult, err := d.docEngine.Search(ctx, &enginetypes.SearchRequest{ - IndexNames: []string{indexName}, - KbIDs: []string{datasetID}, - Offset: 0, - Limit: 1, - Filter: map[string]interface{}{ - "kb_id": datasetID, - "available_int": 1, - }, - }) - if err != nil { - return nil, err - } - if totalResult == nil || totalResult.Total <= 0 { - return []embeddingCheckSample{}, nil - } - - total := int(totalResult.Total) - // Cap n to a sane upper bound so a hostile caller can't force a - // huge preallocation. The downstream `samples` slice is sized - // directly from n. - const maxEmbeddingSamples = 1024 - if n < 0 { - return nil, fmt.Errorf("invalid sample size: %d", n) - } - if n > maxEmbeddingSamples { - n = maxEmbeddingSamples - } - if n > total { - n = total - } - limit := total - if limit > 1000 { - limit = 1000 - } - if n > limit { - n = limit - } - offsets := rand.Perm(limit) - offsets = offsets[:n] - sort.Ints(offsets) - - baseFields := []string{"docnm_kwd", "doc_id", "content_with_weight", "page_num_int", "position_int", "top_int"} - // codeql[go/uncontrolled-allocation-size] False positive: n is - // bounded to maxEmbeddingSamples (1024) at the top of this - // function, so the samples slice cannot exceed ~1 MiB - // (embeddingCheckSample is a small struct). - samples := make([]embeddingCheckSample, 0, n) - for _, offset := range offsets { - searchResult, err := d.docEngine.Search(ctx, &enginetypes.SearchRequest{ - IndexNames: []string{indexName}, - KbIDs: []string{datasetID}, - Offset: offset, - Limit: 1, - SelectFields: baseFields, - Filter: map[string]interface{}{ - "kb_id": datasetID, - "available_int": 1, - }, - }) - if err != nil { - return nil, err - } - if searchResult == nil || len(searchResult.Chunks) == 0 { - continue - } - chunkID := datasetChunkID(searchResult.Chunks[0]) - if chunkID == "" { - continue - } - fullChunk, err := d.docEngine.GetChunk(ctx, indexName, chunkID, []string{datasetID}) - if err != nil { - return nil, err - } - chunkMap := datasetMap(fullChunk) - if len(chunkMap) == 0 { - continue - } - vectorField := datasetGuessVecField(chunkMap) - vector := datasetAsFloatVec(chunkMap[vectorField]) - samples = append(samples, embeddingCheckSample{ - ChunkID: chunkID, - KbID: datasetID, - DocID: datasetString(chunkMap["doc_id"]), - DocName: datasetString(chunkMap["docnm_kwd"]), - VectorField: vectorField, - Vector: vector, - PageNum: chunkMap["page_num_int"], - Position: chunkMap["position_int"], - Top: chunkMap["top_int"], - ContentWithWeight: datasetString(chunkMap["content_with_weight"]), - QuestionKeywords: datasetStringSlice(chunkMap["question_kwd"]), - }) - } - return samples, nil -} - -func datasetGuessVecField(src map[string]interface{}) string { - for k := range src { - if strings.HasSuffix(k, "_vec") { - return k - } - } - return "" -} - -func datasetAsFloatVec(v interface{}) []float64 { - if v == nil { - return []float64{} - } - switch val := v.(type) { - case string: - parts := strings.Split(val, "\t") - res := make([]float64, 0, len(parts)) - for _, p := range parts { - if p == "" { - continue - } - f, err := strconv.ParseFloat(p, 64) - if err != nil { - continue - } - res = append(res, f) - } - return res - case []float64: - return val - case []float32: - res := make([]float64, len(val)) - for i, x := range val { - res[i] = float64(x) - } - return res - case []int: - res := make([]float64, len(val)) - for i, x := range val { - res[i] = float64(x) - } - return res - case []interface{}: - res := make([]float64, 0, len(val)) - for _, x := range val { - switch n := x.(type) { - case float64: - res = append(res, n) - case float32: - res = append(res, float64(n)) - case int: - res = append(res, float64(n)) - case string: - f, err := strconv.ParseFloat(n, 64) - if err == nil { - res = append(res, f) - } - } - } - return res - } - return []float64{} -} - -func datasetCosSim(a, b []float64) float64 { - if len(a) == 0 || len(b) == 0 { - return 0 - } - var dot, na, nb float64 - n := len(a) - if len(b) < n { - n = len(b) - } - for i := 0; i < n; i++ { - dot += a[i] * b[i] - } - for _, x := range a { - na += x * x - } - for _, x := range b { - nb += x * x - } - - if na == 0 || nb == 0 { - return 0 - } - return dot / (math.Sqrt(na) * math.Sqrt(nb)) -} - -func datasetCleanEmbeddingText(s string) string { - re := regexp.MustCompile(`]{0,12})?>`) - return strings.TrimSpace(re.ReplaceAllString(s, " ")) -} - -func datasetEncodeEmbedding(embeddingModel *models.EmbeddingModel, texts []string) ([][]float64, error) { - embeddingConfig := &models.EmbeddingConfig{Dimension: 0} - embeddings, err := embeddingModel.ModelDriver.Embed(embeddingModel.ModelName, texts, embeddingModel.APIConfig, embeddingConfig, nil) - if err != nil { - return nil, err - } - vectors := make([][]float64, len(embeddings)) - for i, embedding := range embeddings { - vectors[i] = embedding.Embedding - } - return vectors, nil -} - -func datasetMixVectors(titleVector, contentVector []float64, titleWeight float64) []float64 { - if len(titleVector) != len(contentVector) { - return contentVector - } - mixed := make([]float64, len(contentVector)) - contentWeight := 1 - titleWeight - for i := range contentVector { - mixed[i] = titleWeight*titleVector[i] + contentWeight*contentVector[i] - } - return mixed -} - -func datasetEmbeddingCheckSummary(datasetID, embeddingID string, sampled int, similarities []float64, matchMode string) EmbeddingCheckSummary { - summary := EmbeddingCheckSummary{ - KbID: datasetID, - Model: embeddingID, - Sampled: sampled, - Valid: len(similarities), - MatchMode: matchMode, - } - if len(similarities) == 0 { - return summary - } - minValue := similarities[0] - maxValue := similarities[0] - total := 0.0 - for _, value := range similarities { - total += value - if value < minValue { - minValue = value - } - if value > maxValue { - maxValue = value - } - } - summary.AvgCosSim = datasetRoundFloat(total/float64(len(similarities)), 6) - summary.MinCosSim = datasetRoundFloat(minValue, 6) - summary.MaxCosSim = datasetRoundFloat(maxValue, 6) - return summary -} - -func datasetRoundFloat(value float64, places int) float64 { - factor := math.Pow10(places) - return math.Round(value*factor) / factor -} - -func datasetChunkID(chunk map[string]interface{}) string { - for _, key := range []string{"id", "_id"} { - if value := datasetString(chunk[key]); value != "" { - return value - } - } - return "" -} - -func datasetMap(value interface{}) map[string]interface{} { - switch typedValue := value.(type) { - case map[string]interface{}: - return typedValue - default: - return map[string]interface{}{} - } -} - -func datasetString(value interface{}) string { - switch typedValue := value.(type) { - case string: - return typedValue - case fmt.Stringer: - return typedValue.String() - case nil: - return "" - default: - return fmt.Sprint(typedValue) - } -} - -func datasetStringSlice(value interface{}) []string { - switch typedValue := value.(type) { - case []string: - return typedValue - case []interface{}: - values := make([]string, 0, len(typedValue)) - for _, item := range typedValue { - if s := strings.TrimSpace(datasetString(item)); s != "" { - values = append(values, s) - } - } - return values - case string: - if typedValue == "" { - return nil - } - return []string{typedValue} - default: - return nil - } -} - -func (d *DatasetService) DeleteIndex(userID, datasetID, indexType string, wipe bool) (common.ErrorCode, error) { - if !checkType(indexType) { - return common.CodeArgumentError, fmt.Errorf("Invalid index type '%s'", indexType) - } - - if datasetID == "" { - return common.CodeDataError, errors.New(`Lack of "Dataset ID"`) - } - - if !d.kbDAO.Accessible(datasetID, userID) { - return common.CodeDataError, errors.New("No authorization.") - } - - kb, err := d.kbDAO.GetByID(datasetID) - if err != nil { - if dao.IsNotFoundErr(err) { - return common.CodeDataError, errors.New("Invalid Dataset ID") - } - return common.CodeDataError, errors.New("Internal server error") - } - - taskIDField := datasetIndexTaskIDColumn(indexType) - taskFinishAtField := datasetIndexTaskFinishAtColumn(indexType) - taskID := datasetIndexTaskID(kb, indexType) - - common.Info("delete_index", zap.String("dataset_id", datasetID), zap.String("index_type", indexType), zap.Bool("wipe", wipe)) - - if taskID != "" { - redisClient := redisengine.Get() - if redisClient == nil || !redisClient.Set(fmt.Sprintf("%s-cancel", taskID), "x", 0) { - common.Warn("Failed to set dataset index cancellation marker", zap.String("dataset_id", datasetID), zap.String("task_id", taskID)) - } - if err := dao.DB.Unscoped().Where("id = ?", taskID).Delete(&entity.Task{}).Error; err != nil { - common.Warn("Failed to delete dataset index task", zap.String("dataset_id", datasetID), zap.String("task_id", taskID), zap.Error(err)) - return common.CodeDataError, errors.New("Internal server error") - } - } - - if wipe && indexType == "graph" { - if d.docEngine == nil { - return common.CodeServerError, errors.New("Document engine is not initialized") - } - indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) - _, err = d.docEngine.DeleteChunks(context.Background(), map[string]interface{}{ - "knowledge_graph_kwd": interfaceSlice("graph", "subgraph", "entity", "relation", "community_report"), - "kb_id": datasetID, - }, indexName, datasetID) - if err != nil { - common.Warn("Failed to delete GraphRAG artefacts", zap.String("dataset_id", datasetID), zap.Error(err)) - return common.CodeDataError, errors.New("Internal server error") - } - clearGraphPhaseMarkers(redisengine.Get(), datasetID) - common.Info("delete_index: cleared GraphRAG artefacts and phase markers", zap.String("dataset_id", datasetID)) - } else if wipe && indexType == "raptor" { - if d.docEngine == nil { - return common.CodeServerError, errors.New("Document engine is not initialized") - } - indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) - _, err = d.docEngine.DeleteChunks(context.Background(), map[string]interface{}{ - "raptor_kwd": interfaceSlice("raptor"), - "kb_id": datasetID, - }, indexName, datasetID) - if err != nil { - common.Warn("Failed to delete RAPTOR artefacts", zap.String("dataset_id", datasetID), zap.Error(err)) - return common.CodeDataError, errors.New("Internal server error") - } - } - - if err := dao.DB.Model(&entity.Knowledgebase{}).Where("id = ?", kb.ID).Updates(map[string]interface{}{ - taskIDField: "", - taskFinishAtField: nil, - }).Error; err != nil { - common.Warn("Failed to clear dataset index task fields", zap.String("dataset_id", datasetID), zap.String("index_type", indexType), zap.Error(err)) - return common.CodeDataError, errors.New("Internal server error") - } - - return common.CodeSuccess, nil -} - -// SearchDatasetsRequest is the request structure for searching chunks across datasets. -type SearchDatasetsRequest struct { - DatasetIDs []string `json:"dataset_ids" binding:"required"` - Question string `json:"question" binding:"required"` - Page *int `json:"page,omitempty"` - Size *int `json:"size,omitempty"` - DocIDs []string `json:"doc_ids,omitempty"` - UseKG *bool `json:"use_kg,omitempty"` - TopK *int `json:"top_k,omitempty"` - CrossLanguages []string `json:"cross_languages,omitempty"` - SearchID *string `json:"search_id,omitempty"` - MetadataFilter map[string]interface{} `json:"meta_data_filter,omitempty"` - RerankID *string `json:"rerank_id,omitempty"` - Keyword *bool `json:"keyword,omitempty"` - SimilarityThreshold *float64 `json:"similarity_threshold,omitempty"` - VectorSimilarityWeight *float64 `json:"vector_similarity_weight,omitempty"` - ForceRefresh bool `json:"force_refresh"` -} - -// SearchDatasetsResponse is the response structure for dataset search results. -type SearchDatasetsResponse struct { - Chunks []map[string]interface{} `json:"chunks"` - DocAggs []map[string]interface{} `json:"doc_aggs"` - Labels *map[string]float64 `json:"labels"` - Total int64 `json:"total"` -} - -// SearchDatasetRequest is the request structure for searching chunks within one dataset. -type SearchDatasetRequest struct { - Question string `json:"question"` - Page *int `json:"page,omitempty"` - Size *int `json:"size,omitempty"` - DocIDs []string `json:"doc_ids,omitempty"` - UseKG *bool `json:"use_kg,omitempty"` - TopK *int `json:"top_k,omitempty"` - CrossLanguages []string `json:"cross_languages,omitempty"` - SearchID *string `json:"search_id,omitempty"` - MetadataFilter map[string]interface{} `json:"meta_data_filter,omitempty"` - RerankID *string `json:"rerank_id,omitempty"` - Keyword *bool `json:"keyword,omitempty"` - SimilarityThreshold *float64 `json:"similarity_threshold,omitempty"` - VectorSimilarityWeight *float64 `json:"vector_similarity_weight,omitempty"` -} - -// ToSearchDatasetsRequest converts a single-dataset search request into the multi-dataset form. -func (req *SearchDatasetRequest) ToSearchDatasetsRequest(datasetID string) *SearchDatasetsRequest { - if req == nil { - return &SearchDatasetsRequest{DatasetIDs: []string{datasetID}} - } - return &SearchDatasetsRequest{ - DatasetIDs: []string{datasetID}, - Question: req.Question, - Page: req.Page, - Size: req.Size, - DocIDs: req.DocIDs, - UseKG: req.UseKG, - TopK: req.TopK, - CrossLanguages: req.CrossLanguages, - SearchID: req.SearchID, - MetadataFilter: req.MetadataFilter, - RerankID: req.RerankID, - Keyword: req.Keyword, - SimilarityThreshold: req.SimilarityThreshold, - VectorSimilarityWeight: req.VectorSimilarityWeight, - } -} - -// SearchDataset searches chunks within one knowledge base based on a question. -func (d *DatasetService) SearchDataset(datasetID, userID string, req *SearchDatasetRequest) (*SearchDatasetsResponse, error) { - if datasetID == "" { - return nil, fmt.Errorf("dataset_id is required") - } - return d.SearchDatasets(req.ToSearchDatasetsRequest(datasetID), userID) -} - -// SearchDatasets searches chunks across one or more knowledge bases based on a question. -// It retrieves relevant chunks using embedding and optional reranking, applying filters, -// cross-language translation, and keyword extraction as configured. -func (d *DatasetService) SearchDatasets(req *SearchDatasetsRequest, userID string) (*SearchDatasetsResponse, error) { - if req.Question == "" { - return nil, fmt.Errorf("question is required") - } - if len(req.DatasetIDs) == 0 { - return nil, fmt.Errorf("dataset_ids is required") - } - common.Info("SearchDatasets started", zap.String("userID", userID), zap.Any("datasets", req.DatasetIDs), zap.String("question", req.Question)) - - page := 1 - if req.Page != nil { - page = *req.Page - } - pageSize := 30 - if req.Size != nil { - pageSize = *req.Size - } - useKG := false - if req.UseKG != nil { - useKG = *req.UseKG - } - similarityThreshold := 0.0 - if req.SimilarityThreshold != nil { - similarityThreshold = *req.SimilarityThreshold - } - vectorSimilarityWeight := 0.3 - if req.VectorSimilarityWeight != nil { - vectorSimilarityWeight = *req.VectorSimilarityWeight - } - topK := 1024 - if req.TopK != nil { - topK = *req.TopK - } - if topK < 1 { - topK = 1 - } else if topK > 2048 { - topK = 2048 - } - keyword := false - if req.Keyword != nil { - keyword = *req.Keyword - } - searchID := "" - if req.SearchID != nil { - searchID = *req.SearchID - } - rerankID := "" - if req.RerankID != nil { - rerankID = *req.RerankID - } - - question := req.Question - datasetIDs := req.DatasetIDs - metadataFilter := req.MetadataFilter - crossLanguages := req.CrossLanguages - - common.Debug(fmt.Sprintf("SearchDatasets request:\n"+ - " datasetIDs=%v\n"+ - " question=%s\n"+ - " page=%v, pageSize=%v\n"+ - " docIDs=%v\n"+ - " useKG=%v, topK=%v\n"+ - " crossLanguages=%v\n"+ - " searchID=%v\n"+ - " metadataFilter=%v\n"+ - " rerankID=%v\n"+ - " keyword=%v\n"+ - " similarityThreshold=%v, vectorSimilarityWeight=%v", - datasetIDs, req.Question, - common.PtrString(req.Page), common.PtrString(req.Size), req.DocIDs, - useKG, topK, crossLanguages, searchID, - metadataFilter, - rerankID, - keyword, - similarityThreshold, vectorSimilarityWeight)) - - ctx := context.Background() - modelProviderSvc := NewModelProviderService() - - // Access check for all datasets - var tenantIDs []string - var kbRecords []*entity.Knowledgebase - seenTenants := make(map[string]bool) - for _, datasetID := range datasetIDs { - if !d.kbDAO.Accessible(datasetID, userID) { - common.Warn("SearchDatasets access denied", zap.String("datasetID", datasetID), zap.String("userID", userID)) - return nil, fmt.Errorf("only owner of dataset %s is authorized for this operation", datasetID) - } - - kb, err := d.kbDAO.GetByID(datasetID) - if err != nil || kb == nil { - common.Warn("SearchDatasets dataset not found", zap.String("datasetID", datasetID)) - return nil, fmt.Errorf("dataset %s not found", datasetID) - } - if !seenTenants[kb.TenantID] { - seenTenants[kb.TenantID] = true - tenantIDs = append(tenantIDs, kb.TenantID) - } - kbRecords = append(kbRecords, kb) - } - - // Check if all kbs have the same embedding model - if err := validateDatasetEmbeddingModels(kbRecords); err != nil { - return nil, err - } - - // Override request fields with values from saved search config (if search_id is provided) - var chatID string - if searchID != "" { - if d.searchService == nil { - common.Warn("Search service is not initialized for search_id", zap.String("searchID", searchID)) - return nil, fmt.Errorf("Invalid search_id") - } - searchDetail, err := d.searchService.GetDetail(searchID) - if err != nil || searchDetail == nil || len(searchDetail) == 0 { - common.Warn("Invalid search_id", zap.String("searchID", searchID), zap.Error(err)) - return nil, fmt.Errorf("Invalid search_id") - } else if searchConfig, ok := searchDetail["search_config"].(map[string]interface{}); ok && searchConfig != nil { - if scMetadataFilter, ok := searchConfig["meta_data_filter"].(map[string]interface{}); ok { - metadataFilter = scMetadataFilter - } - if scST, ok := searchConfig["similarity_threshold"].(float64); ok { - similarityThreshold = scST - } - if scVSW, ok := searchConfig["vector_similarity_weight"].(float64); ok { - vectorSimilarityWeight = scVSW - } - if scTopK, ok := searchConfig["top_k"].(float64); ok { - topK = int(scTopK) - if topK < 1 { - topK = 1 - } else if topK > 2048 { - topK = 2048 - } - } - if scUseKG, ok := searchConfig["use_kg"].(bool); ok { - useKG = scUseKG - } - if scLangs, ok := searchConfig["cross_languages"].([]interface{}); ok { - crossLanguages = make([]string, len(scLangs)) - for i, l := range scLangs { - if s, ok := l.(string); ok { - crossLanguages[i] = s - } - } - } - if scKeyword, ok := searchConfig["keyword"].(bool); ok { - keyword = scKeyword - } - if scRerankID, ok := searchConfig["rerank_id"].(string); ok { - rerankID = scRerankID - } - chatID, _ = searchConfig["chat_id"].(string) - - common.Debug("SearchDatasets loaded Search config", - zap.String("searchID", searchID), - zap.Strings("datasetIDs", datasetIDs), - zap.Float64("vectorSimilarityWeight", vectorSimilarityWeight), - zap.Float64("fullTextWeight", 1-vectorSimilarityWeight), - zap.Float64("similarityThreshold", similarityThreshold), - zap.Int("topK", topK), - zap.Strings("crossLanguages", crossLanguages), - zap.Bool("keyword", keyword), - zap.String("rerankID", rerankID), - zap.String("chatID", chatID), - zap.Bool("useKG", useKG)) - } else { - common.Warn("Invalid search_id: search_config missing or invalid", zap.String("searchID", searchID)) - return nil, fmt.Errorf("Invalid search_id") - } - } - - // If meta_data_filter method is auto/semi_auto, get chat model - var err error - var chatModelForFilter *models.ChatModel - if metadataFilter != nil { - method, _ := metadataFilter["method"].(string) - if method == "auto" || method == "semi_auto" { - if chatID != "" { - driver, modelName, apiConfig, _, err := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeChat, chatID) - if err != nil { - common.Warn("Failed to get chat model config from search_config chat_id, using tenant default", zap.String("chatID", chatID), zap.Error(err)) - } else { - chatModelForFilter = models.NewChatModel(driver, &modelName, apiConfig) - common.Info("Fetched chat model (from search_config) for metadata filter", - zap.String("chatID", chatID), - zap.String("tenantID", tenantIDs[0])) - } - } - - if chatModelForFilter == nil { - driver, modelName, apiConfig, _, err := modelProviderSvc.GetTenantDefaultModelByType(tenantIDs[0], entity.ModelTypeChat) - if err != nil { - common.Warn("Failed to get tenant default chat model for meta_data_filter", zap.Error(err)) - } else { - chatModelForFilter = models.NewChatModel(driver, &modelName, apiConfig) - common.Info("Fetched chat model (tenant default) for metadata filter", - zap.String("tenantID", tenantIDs[0])) - } - } - } - } - - // Apply meta_data_filter to get filtered doc_ids - docIDs := make([]string, len(req.DocIDs)) - copy(docIDs, req.DocIDs) - if len(metadataFilter) > 0 { - metadataSvc := NewMetadataService() - flattedMeta, err := metadataSvc.GetFlattedMetaByKBs(datasetIDs) - if err != nil { - common.Warn("Failed to get flatted metadata, using empty metadata for filter", zap.Error(err)) - flattedMeta = make(common.MetaData) - } - common.Info("Metadata filter conditions", zap.Any("filter", metadataFilter)) - filteredDocIDs, _ := ApplyMetaDataFilter(ctx, metadataFilter, flattedMeta, question, chatModelForFilter, req.DocIDs, datasetIDs) - docIDs = filteredDocIDs - common.Info("ApplyMetaDataFilter result", zap.Strings("docIDs", docIDs)) - } - - // Apply cross_languages and keyword extraction - modifiedQuestion := question - if len(crossLanguages) > 0 { - // Pass tenantID and empty llmID so CrossLanguages can fetch default if needed - // This matches Python's cross_languages(tenant_id, llm_id, query, languages) - common.Info("CrossLanguages: dispatching translation", - zap.String("tenantID", tenantIDs[0]), - zap.String("llmID", ""), - zap.Strings("crossLanguages", crossLanguages)) - translated, err := CrossLanguages(ctx, tenantIDs[0], "", question, crossLanguages) - if err != nil { - common.Warn("Failed to translate question", zap.String("llmID", ""), zap.Error(err)) - } else { - modifiedQuestion = translated - } - } - if keyword { - driver, modelName, apiConfig, _, err := modelProviderSvc.GetTenantDefaultModelByType(tenantIDs[0], entity.ModelTypeChat) - if err != nil { - common.Warn("Failed to get default chat model for LLM transformations", zap.Error(err)) - } else { - chatModel := models.NewChatModel(driver, &modelName, apiConfig) - common.Info("Fetched chat model (tenant default) for keyword_extraction", - zap.String("tenantID", tenantIDs[0])) - - extractedKeywords, err := KeywordExtraction(ctx, chatModel, modifiedQuestion, 3) - if err != nil { - common.Warn("Failed to extract keywords from question", zap.Error(err)) - } else if extractedKeywords != "" { - modifiedQuestion = modifiedQuestion + extractedKeywords - } - } - } - if modifiedQuestion != question { - common.Info("Modified question after transformations", - zap.String("originalQuestion", question), - zap.String("modifiedQuestion", modifiedQuestion), - zap.Strings("crossLanguages", crossLanguages), - zap.Bool("keywordExtraction", keyword)) - } - - // Get tag-based rank features via LabelQuestion - metadataSvc := NewMetadataService() - labels := metadataSvc.LabelQuestion(modifiedQuestion, kbRecords) - if len(labels) > 0 { - common.Debug("LabelQuestion result", zap.Any("labels", labels)) - } - - // Determine embedding model - var embeddingModel *models.EmbeddingModel - if kbRecords[0].EmbdID != "" { - driver, modelName, apiConfig, maxTokens, embErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeEmbedding, kbRecords[0].EmbdID) - if embErr != nil { - return nil, fmt.Errorf("failed to get embedding model by embd_id: %w", embErr) - } - embeddingModel = models.NewEmbeddingModel(driver, &modelName, apiConfig, maxTokens) - common.Info("Fetched embedding model for retrieval", - zap.String("tenantID", tenantIDs[0]), - zap.String("modelName", modelName)) - - } - - // Get rerank model if rerankID is specified - var rerankModel *models.RerankModel - if rerankID != "" { - driver, modelName, apiConfig, _, rErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeRerank, rerankID) - if rErr != nil { - return nil, fmt.Errorf("failed to get rerank model by rerank_id: %w", rErr) - } - rerankModel = models.NewRerankModel(driver, &modelName, apiConfig) - common.Info("Fetched rerank model", - zap.String("tenantID", tenantIDs[0]), - zap.String("modelName", modelName)) - } - - retrievalReq := &nlp.RetrievalRequest{ - TenantIDs: tenantIDs, - Question: modifiedQuestion, - KbIDs: datasetIDs, - DocIDs: docIDs, - Page: page, - PageSize: pageSize, - Top: &topK, - SimilarityThreshold: &similarityThreshold, - VectorSimilarityWeight: &vectorSimilarityWeight, - RerankModel: rerankModel, - RankFeature: &labels, - EmbeddingModel: embeddingModel, - } - - retrievalResult, err := nlp.NewRetrievalService(d.docEngine, d.documentDAO).Retrieval(ctx, retrievalReq) - if err != nil { - return nil, fmt.Errorf("retrieval search failed: %w", err) - } - - filteredChunks := retrievalResult.Chunks - - if useKG { - common.Warn("use_kg is not yet implemented in Go - skipping KG retrieval") - } - - filteredChunks = nlp.RetrievalByChildren(filteredChunks, tenantIDs, d.docEngine, ctx) - - for i := range filteredChunks { - delete(filteredChunks[i], "vector") - } - - common.Info("SearchDatasets completed", zap.String("userID", userID), zap.Any("kbID", datasetIDs), zap.String("question", question), zap.Int64("chunkCount", int64(len(filteredChunks)))) - - // Convert all float64 values to PyFloat64 for Python-compatible JSON serialization - pyChunks := common.ConvertFloatsToPyFormat(filteredChunks).([]map[string]interface{}) - - return &SearchDatasetsResponse{ - Chunks: pyChunks, - DocAggs: retrievalResult.DocAggs, - Labels: &labels, - Total: retrievalResult.Total, - }, nil -} - -// MetadataConfigField mirrors one field in the dataset metadata config API. -type MetadataConfigField struct { - Key string `json:"key"` - Type string `json:"type"` - Description *string `json:"description"` - Enum []string `json:"enum"` -} - -// MetadataConfigRequest mirrors PUT /datasets/:dataset_id/metadata/config. -type MetadataConfigRequest struct { - Metadata []MetadataConfigField `json:"metadata"` - BuiltInMetadata []MetadataConfigField `json:"built_in_metadata"` -} - -// CreateDatasetRequest represents the request for creating a dataset. -type CreateDatasetRequest struct { - Name string `json:"name" binding:"required"` - EmbeddingModel *string `json:"embedding_model,omitempty"` - Permission *string `json:"permission,omitempty"` - ParserID *string `json:"parser_id,omitempty"` - PipelineID *string `json:"pipeline_id,omitempty"` - // ParseType indicates pipeline selection mode: 1 = BuiltIn (parser_id), - // 2 = Pipeline (pipeline_id). nil means unspecified (backward compat). - ParseType *int `json:"parse_type,omitempty"` -} - -// ListDatasets lists datasets with pagination and filtering. -func (d *DatasetService) ListDatasets(id, name string, page, pageSize int, orderby string, desc bool, keywords string, ownerIDs []string, parserID, userID string) ([]map[string]interface{}, int64, common.ErrorCode, error) { - id = strings.TrimSpace(id) - if id != "" { - normalizedID, err := normalizeDatasetID(id) - if err != nil { - return nil, 0, common.CodeDataError, err - } - id = normalizedID - - kbs, err := d.kbDAO.GetKBByIDAndUserID(id, userID) - if err != nil { - return nil, 0, common.CodeServerError, errors.New("Database operation failed") - } - if len(kbs) == 0 { - return nil, 0, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", userID, id) - } - } - - name = strings.TrimSpace(name) - if name != "" { - kbs, err := d.kbDAO.GetKBByNameAndUserID(name, userID) - if err != nil { - return nil, 0, common.CodeServerError, errors.New("Database operation failed") - } - if len(kbs) == 0 { - return nil, 0, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", userID, name) - } - } - - if page <= 0 { - page = 1 - } - if pageSize <= 0 { - pageSize = 30 - } - - orderby = strings.TrimSpace(orderby) - if _, ok := datasetAllowedOrderByFields[orderby]; !ok { - orderby = "create_time" - } - - keywords = strings.TrimSpace(keywords) - parserID = strings.TrimSpace(parserID) - - // Empty owner ids do not change the query, so only keep the meaningful ones. - tenantIDs := make([]string, 0, len(ownerIDs)) - for _, ownerID := range ownerIDs { - ownerID = strings.TrimSpace(ownerID) - if ownerID != "" { - tenantIDs = append(tenantIDs, ownerID) - } - } - if len(tenantIDs) == 0 { - joinedTenants, err := d.tenantDAO.GetJoinedTenantsByUserID(userID) - if err != nil { - return nil, 0, common.CodeServerError, errors.New("Database operation failed") - } - for _, joinedTenant := range joinedTenants { - if joinedTenant == nil || joinedTenant.TenantID == "" { - continue - } - tenantIDs = append(tenantIDs, joinedTenant.TenantID) - } - } - - kbs, total, err := d.kbDAO.GetByTenantIDs(tenantIDs, userID, page, pageSize, orderby, desc, keywords, parserID) - if err != nil { - return nil, 0, common.CodeServerError, errors.New("Database operation failed") - } - - data := make([]map[string]interface{}, 0, len(kbs)) - for _, kb := range kbs { - if kb == nil { - continue - } - data = append(data, datasetListItemToMap(kb)) - } - - return data, total, common.CodeSuccess, nil -} - -// CreateDataset creates a new dataset. -func (d *DatasetService) CreateDataset(req *CreateDatasetRequest, tenantID string) (map[string]interface{}, common.ErrorCode, error) { - if !common.IsValidString(req.Name) { - return nil, common.CodeDataError, errors.New("Dataset name must be string.") - } - - name := strings.TrimSpace(req.Name) - if name == "" { - return nil, common.CodeDataError, errors.New("Dataset name can't be empty.") - } - if len(name) > entity.DatasetNameLimit { - 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) - if err != nil || tenant == nil { - return nil, common.CodeDataError, errors.New("Tenant not found.") - } - - // parse_type explicitly signals the pipeline mode (1 = BuiltIn, 2 = Pipeline). - // When absent (nil), infer intent from which field is present for backward compat. - isPipelineMode := req.ParseType != nil && *req.ParseType == 2 - isBuiltinMode := req.ParseType != nil && *req.ParseType == 1 - - if isBuiltinMode && req.PipelineID != nil { - // BuiltIn mode: discard pipeline_id so only parser_id matters. - req.PipelineID = nil - } - if isPipelineMode && req.ParserID != nil { - // Pipeline mode: ignore parser_id. - req.ParserID = nil - } - - if req.ParseType == nil && req.ParserID != nil && req.PipelineID != nil { - return nil, common.CodeDataError, errors.New("parser_id and pipeline_id are mutually exclusive") - } - - parserID := "" - permission := "me" - embeddingModel := "" - pipelineID := req.PipelineID - - if req.Permission != nil { - permission = strings.TrimSpace(*req.Permission) - if permission != "me" && permission != "team" { - return nil, common.CodeDataError, errors.New("Input should be 'me' or 'team'") - } - } - if req.ParserID != nil { - parserID = strings.TrimSpace(*req.ParserID) - if err := validateParserID(parserID); err != nil { - return nil, common.CodeDataError, err - } - pipelineID = nil - } - if req.PipelineID != nil { - normalizedPipelineID, err := normalizeDatasetPipelineID(*req.PipelineID) - if err != nil { - return nil, common.CodeDataError, err - } - pipelineID = normalizedPipelineID - } - if req.EmbeddingModel != nil { - embeddingModel = strings.TrimSpace(*req.EmbeddingModel) - if err := validateDatasetEmbeddingModel(embeddingModel); err != nil { - return nil, common.CodeDataError, err - } - } - - // Reject references to a custom canvas the caller does not own or share. - // The handler passes the caller's user id in tenantID for CreateDataset. - if pipelineID != nil && strings.TrimSpace(*pipelineID) != "" { - if ok, err := canvasAccessibleForUser(tenantID, strings.TrimSpace(*pipelineID)); err != nil { - return nil, common.CodeServerError, err - } else if !ok { - return nil, common.CodeDataError, errors.New("canvas is not accessible") - } - } - - // Resolve component params defaults from the DSL template. parser_config - // stores component params directly: {cpnID: {param: value}}. - parserConfig, cpErr := resolveComponentParamsDefaults(parserID, pipelineID) - if cpErr != nil { - common.Warn("failed to resolve component params defaults for dataset", - zap.String("parserID", parserID), zap.Error(cpErr)) - parserConfig = entity.JSONMap{} - } - - var parserConfigMap map[string]interface{} = parserConfig - - embdID := tenant.EmbdID - tenantEmbdID := ptrStringValue(tenant.TenantEmbdID) - if embeddingModel != "" { - ok, message := d.verifyEmbeddingAvailability(embeddingModel, tenantID) - if !ok { - return nil, common.CodeDataError, errors.New(message) - } - embdID = embeddingModel - tenantEmbdID = "" - } - if embdID != "" && tenantEmbdID == "" { - resolvedID, err := NewModelProviderService().ResolveModelID(tenantID, entity.ModelTypeEmbedding, embdID) - if err == nil { - tenantEmbdID = resolvedID - } else { - return nil, common.CodeDataError, err - } - } - - kbID := utility.GenerateToken() - status := string(entity.StatusValid) - // Deduplicate name within tenant - duplicateName, err := common.DuplicateName(func(n, tid string) bool { - existing, err := d.kbDAO.GetByName(n, tid) - return err == nil && existing != nil - }, name, tenantID) - if err != nil { - return nil, common.CodeDataError, err - } - - kb := &entity.Knowledgebase{ - ID: kbID, - Name: duplicateName, - TenantID: tenantID, - CreatedBy: tenantID, - ParserID: parserID, - PipelineID: pipelineID, - ParserConfig: entity.JSONMap(parserConfigMap), - Permission: permission, - EmbdID: embdID, - TenantEmbdID: stringPtrIfNotEmpty(tenantEmbdID), - Status: &status, - } - - if err = d.kbDAO.Create(kb); err != nil { - return nil, common.CodeServerError, errors.New("Failed to save dataset") - } - - createdKB, err := d.kbDAO.GetByID(kbID) - if err != nil || createdKB == nil { - return nil, common.CodeServerError, errors.New("Dataset created failed") - } - - return datasetToMap(createdKB), common.CodeSuccess, nil -} - -// DeleteDatasets deletes multiple datasets. -func (d *DatasetService) DeleteDatasets(ids []string, deleteAll bool, tenantID string) (map[string]interface{}, common.ErrorCode, error) { - normalizedIDs := make([]string, 0, len(ids)) - seenIDs := make(map[string]struct{}, len(ids)) - - // Canonicalize ids once so every downstream DAO call sees the same 32-char hex format. - for _, id := range ids { - normalizedID, err := normalizeDatasetID(strings.TrimSpace(id)) - if err != nil { - return nil, common.CodeDataError, err - } - if _, exists := seenIDs[normalizedID]; exists { - return nil, common.CodeDataError, fmt.Errorf("Duplicate ids: '%s'", normalizedID) - } - seenIDs[normalizedID] = struct{}{} - normalizedIDs = append(normalizedIDs, normalizedID) - } - - if len(normalizedIDs) == 0 { - if !deleteAll { - return map[string]interface{}{"success_count": 0}, common.CodeSuccess, nil - } - - kbs, err := d.kbDAO.Query(map[string]interface{}{"tenant_id": tenantID}) - if err != nil { - return nil, common.CodeServerError, errors.New("Database operation failed") - } - for _, kb := range kbs { - normalizedIDs = append(normalizedIDs, kb.ID) - } - } - - kbs := make([]*entity.Knowledgebase, 0, len(normalizedIDs)) - unauthorizedIDs := make([]string, 0) - for _, id := range normalizedIDs { - kb, err := d.kbDAO.GetByIDAndTenantID(id, tenantID) - if err != nil || kb == nil { - unauthorizedIDs = append(unauthorizedIDs, id) - continue - } - kbs = append(kbs, kb) - } - if len(unauthorizedIDs) > 0 { - return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for datasets: '%s'", tenantID, strings.Join(unauthorizedIDs, ", ")) - } - - errorsList := make([]string, 0) - successCount := 0 - for _, kb := range kbs { - if err := d.deleteDataset(tenantID, kb); err != nil { - errorsList = append(errorsList, err.Error()) - continue - } - successCount++ - } - - if len(errorsList) == 0 { - return map[string]interface{}{"success_count": successCount}, common.CodeSuccess, nil - } - - details := strings.Join(errorsList, "; ") - if len(details) > 128 { - details = details[:128] - } - errorMessage := fmt.Sprintf( - "Successfully deleted %d datasets, %d failed. Details: %s...", - successCount, - len(errorsList), - details, - ) - if successCount == 0 { - return nil, common.CodeDataError, errors.New(errorMessage) - } - - return map[string]interface{}{ - "success_count": successCount, - "errors": limitStrings(errorsList, 5), - }, common.CodeSuccess, nil -} - -// GetDataset gets a single dataset with its size and linked connectors. -func (d *DatasetService) GetDataset(datasetID, userID string) (map[string]interface{}, common.ErrorCode, error) { - datasetID = strings.TrimSpace(datasetID) - if datasetID == "" { - return nil, common.CodeDataError, errors.New("Lack of \"Dataset ID\"") - } - - normalizedID, err := normalizeDatasetID(datasetID) - if err != nil { - return nil, common.CodeDataError, err - } - datasetID = normalizedID - - if !d.kbDAO.Accessible(datasetID, userID) { - return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", userID, datasetID) - } - - kb, err := d.kbDAO.GetByID(datasetID) - if err != nil || kb == nil { - return nil, common.CodeDataError, errors.New("Invalid Dataset ID") - } - - data := datasetToMap(kb) - - size, err := d.documentDAO.SumSizeByDatasetID(datasetID) - if err != nil { - return nil, common.CodeServerError, errors.New("Database operation failed") - } - data["size"] = size - - connectors, err := d.connectorDAO.ListByDatasetID(datasetID) - if err != nil { - return nil, common.CodeServerError, errors.New("Database operation failed") - } - data["connectors"] = datasetConnectorsOrEmpty(connectors) - - return data, common.CodeSuccess, nil -} - -type DatasetConnectorRequest struct { - ID string `json:"id"` - AutoParse string `json:"auto_parse,omitempty"` -} - -type UpdateDatasetRequest struct { - Name *string `json:"name,omitempty"` - Avatar *string `json:"avatar,omitempty"` - Description *string `json:"description,omitempty"` - Language *string `json:"language,omitempty"` - Connectors *[]DatasetConnectorRequest `json:"connectors,omitempty"` - EmbdID *string `json:"embd_id,omitempty"` - EmbeddingModel *string `json:"embedding_model,omitempty"` - Permission *string `json:"permission,omitempty"` - ParserID *string `json:"parser_id,omitempty"` - Pagerank *int64 `json:"pagerank,omitempty"` - ParserConfig map[string]interface{} `json:"parser_config,omitempty"` - PipelineID *string `json:"pipeline_id,omitempty"` - // ParseType indicates pipeline selection mode: 1 = BuiltIn (parser_id), - // 2 = Pipeline (pipeline_id). nil means unspecified (backward compat). - ParseType *int `json:"parse_type,omitempty"` -} - -// UpdateDataset Update a dataset -func (d *DatasetService) UpdateDataset(datasetID, tenantID string, req UpdateDatasetRequest) (map[string]interface{}, common.ErrorCode, error) { - kb, err := d.kbDAO.GetByID(datasetID) - if err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, errors.New("Dataset not found") - } - return nil, common.CodeServerError, errors.New("Database operation failed") - } - - if kb == nil || kb.TenantID != tenantID { - return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", tenantID, datasetID) - } - - connectorsProvided := req.Connectors != nil - connectors := make([]DatasetConnectorRequest, 0) - if req.Connectors != nil { - connectors = *req.Connectors - } - - updates := make(map[string]interface{}) - - if req.Name != nil { - name := strings.TrimSpace(*req.Name) - if name == "" { - return nil, common.CodeDataError, errors.New("String should have at least 1 character") - } - if len(name) > 128 { - return nil, common.CodeDataError, errors.New("String should have at most 128 characters") - } - updates["name"] = name - } - if req.Avatar != nil { - if len(*req.Avatar) > 65535 { - return nil, common.CodeDataError, errors.New("String should have at most 65535 characters") - } - if err := validateDatasetAvatar(*req.Avatar); err != nil { - return nil, common.CodeDataError, err - } - updates["avatar"] = *req.Avatar - } - if req.Description != nil { - if len(*req.Description) > 65535 { - return nil, common.CodeDataError, errors.New("String should have at most 65535 characters") - } - updates["description"] = *req.Description - } - if req.Language != nil { - language := strings.TrimSpace(*req.Language) - if len(language) > 32 { - return nil, common.CodeDataError, errors.New("String should have at most 32 characters") - } - updates["language"] = language - } - if req.Permission != nil { - permission := strings.TrimSpace(*req.Permission) - if permission != "me" && permission != "team" { - return nil, common.CodeDataError, errors.New("Input should be 'me' or 'team'") - } - updates["permission"] = permission - } - - // parse_type explicitly signals the pipeline mode (1 = BuiltIn, 2 = Pipeline). - // When absent (nil), existing mutual-exclusivity behavior is preserved. - isPipelineMode := req.ParseType != nil && *req.ParseType == 2 - isBuiltinMode := req.ParseType != nil && *req.ParseType == 1 - - if isBuiltinMode && req.PipelineID != nil { - req.PipelineID = nil - } - if isPipelineMode && req.ParserID != nil { - req.ParserID = nil - } - - if req.PipelineID != nil { - pipelineID, err := normalizeDatasetPipelineID(*req.PipelineID) - if err != nil { - return nil, common.CodeDataError, err - } - if pipelineID != nil { - updates["pipeline_id"] = *pipelineID - } - } - - parserID, parserIDProvided, err := datasetUpdateParserID(req) - if err != nil { - return nil, common.CodeDataError, err - } - if parserIDProvided { - updates["parser_id"] = parserID - } - - // When parse_type is absent, parser_id and pipeline_id are mutually exclusive. - if req.ParseType == nil && parserIDProvided && req.PipelineID != nil { - return nil, common.CodeDataError, errors.New("parser_id and pipeline_id are mutually exclusive") - } - - embdID, embdIDProvided, err := datasetUpdateEmbeddingID(req) - if err != nil { - return nil, common.CodeDataError, err - } - if embdIDProvided { - tenantEmbdID := ptrStringValue(kb.TenantEmbdID) - if embdID == "" { - embdID = kb.EmbdID - } else { - tenantEmbdID = "" - } - ok, message := d.verifyEmbeddingAvailability(embdID, tenantID) - if !ok { - return nil, common.CodeDataError, errors.New(message) - } - if embdID != "" && tenantEmbdID == "" { - resolvedID, err := NewModelProviderService().ResolveModelID(tenantID, entity.ModelTypeEmbedding, embdID) - if err == nil { - tenantEmbdID = resolvedID - } - } - updates["embd_id"] = embdID - updates["tenant_embd_id"] = stringPtrIfNotEmpty(tenantEmbdID) - } - - if req.ParserConfig != nil { - if err := validateDatasetParserConfigSize(req.ParserConfig); err != nil { - return nil, common.CodeDataError, err - } - if len(req.ParserConfig) > 0 { - // Resolve the effective pipeline and load its DSL schema. - effectiveParserID := kb.ParserID - if parserIDProvided { - effectiveParserID = parserID - } - effectivePipelineID := kb.PipelineID - if req.PipelineID != nil { - if normalized, err := normalizeDatasetPipelineID(*req.PipelineID); err == nil { - effectivePipelineID = normalized - } - } else if parserIDProvided && kb.PipelineID != nil { - effectivePipelineID = nil - } - - isCanvas := effectivePipelineID != nil && strings.TrimSpace(*effectivePipelineID) != "" - dslJSON, dslErr := loadPipelineDSL(isCanvas, effectiveParserID, effectivePipelineID) - if dslErr != nil { - common.Warn("failed to load pipeline DSL for building parser_config", - zap.String("parserID", effectiveParserID), zap.Error(dslErr)) - } - if dslJSON != nil { - // Start from DSL defaults, overlay the cleaned incoming overrides. - // Unknown cpnIDs, unknown params, and legacy flat fields are - // dropped. Full replace — no merge with the stored config. - updates["parser_config"] = buildParserConfig(dslJSON, map[string]interface{}(req.ParserConfig)) - } - } - } - - if req.Pagerank != nil && *req.Pagerank != kb.Pagerank { - if *req.Pagerank < 0 || *req.Pagerank > 100 { - return nil, common.CodeDataError, errors.New("Input should be less than or equal to 100") - } - if !d.docEngine.SupportsPageRank() { - return nil, common.CodeDataError, errors.New("'pagerank' can only be set when doc_engine is elasticsearch") - } - indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) - if *req.Pagerank > 0 { - err = d.docEngine.UpdateChunks(context.Background(), map[string]interface{}{"kb_id": kb.ID}, map[string]interface{}{common.PAGERANK_FLD: *req.Pagerank}, indexName, kb.ID) - } else { - err = d.docEngine.UpdateChunks(context.Background(), map[string]interface{}{"exists": common.PAGERANK_FLD}, map[string]interface{}{"remove": common.PAGERANK_FLD}, indexName, kb.ID) - } - if err != nil { - return nil, common.CodeServerError, err - } - updates["pagerank"] = *req.Pagerank - } - - if parserIDProvided && parserID != kb.ParserID { - if _, ok := updates["parser_config"]; !ok { - if resolved, cpErr := resolveComponentParamsDefaults(parserID, nil); cpErr != nil { - common.Warn("failed to resolve component params defaults on parser_id switch", - zap.String("parserID", parserID), zap.Error(cpErr)) - } else if resolved != nil { - updates["parser_config"] = resolved - } - } - } - if kb.PipelineID != nil && parserIDProvided { - if _, ok := updates["pipeline_id"]; !ok { - updates["pipeline_id"] = nil // clear to NULL, not empty string - } - } - - // Regenerate parser_config from the new pipeline's DSL defaults when - // the pipeline changes, so that component-level parameters stay in sync. - pipelineChanged := req.PipelineID != nil && (kb.PipelineID == nil || *req.PipelineID != *kb.PipelineID) - if pipelineChanged { - cfgParserID := kb.ParserID - if parserIDProvided { - cfgParserID = parserID - } - cfgPipelineID, _ := updates["pipeline_id"].(string) - var cpPipelineID *string - if cfgPipelineID != "" { - cpPipelineID = &cfgPipelineID - } - if cpDefaults, cpErr := resolveComponentParamsDefaults(cfgParserID, cpPipelineID); cpErr != nil { - common.Warn("failed to resolve component params defaults on pipeline change", - zap.String("parserID", cfgParserID), zap.Error(cpErr)) - } else if cpDefaults != nil { - updates["parser_config"] = cpDefaults - } - } - - if nameValue, ok := updates["name"].(string); ok && strings.ToLower(nameValue) != strings.ToLower(kb.Name) { - existing, lookupErr := d.kbDAO.GetByName(nameValue, tenantID) - if lookupErr != nil && !dao.IsNotFoundErr(lookupErr) { - return nil, common.CodeServerError, errors.New("Database operation failed") - } - if existing != nil { - return nil, common.CodeDataError, fmt.Errorf("Dataset name '%s' already exists", nameValue) - } - } - - if len(updates) == 0 && !connectorsProvided { - return nil, common.CodeDataError, errors.New("No properties were modified") - } - - if len(updates) > 0 { - if err = d.kbDAO.UpdateByID(kb.ID, updates); err != nil { - return nil, common.CodeServerError, errors.New("Update dataset error.(Database error)") - } - } - - if connectorsProvided { - connectorLinks := make([]dao.DatasetConnectorLink, 0, len(connectors)) - for _, connector := range connectors { - connectorID := strings.TrimSpace(connector.ID) - if connectorID == "" { - return nil, common.CodeDataError, errors.New("connector id is required") - } - connectorLinks = append(connectorLinks, dao.DatasetConnectorLink{ - ID: connectorID, - AutoParse: connector.AutoParse, - }) - } - if err = d.connectorDAO.LinkDatasetConnectors(kb.ID, connectorLinks); err != nil { - return nil, common.CodeServerError, errors.New("Database operation failed") - } - } - - updatedKB, err := d.kbDAO.GetByID(kb.ID) - if err != nil { - return nil, common.CodeDataError, errors.New("Dataset updated failed") - } - - data := datasetToMap(updatedKB) - linkedConnectors, err := d.connectorDAO.ListByDatasetID(kb.ID) - if err != nil { - return nil, common.CodeServerError, errors.New("Database operation failed") - } - data["connectors"] = datasetConnectorsOrEmpty(linkedConnectors) - return data, common.CodeSuccess, nil -} - -func datasetConnectorsOrEmpty(connectors []*dao.ConnectorDatasetListItem) []*dao.ConnectorDatasetListItem { - if connectors == nil { - return make([]*dao.ConnectorDatasetListItem, 0) - } - return connectors -} - -func datasetUpdateParserID(req UpdateDatasetRequest) (string, bool, error) { - parserID := "" - provided := false - if req.ParserID != nil { - parserID = strings.TrimSpace(*req.ParserID) - provided = true - } - if !provided { - return "", false, nil - } - if err := validateParserID(parserID); err != nil { - return "", true, err - } - return parserID, true, nil -} - -func datasetUpdateEmbeddingID(req UpdateDatasetRequest) (string, bool, error) { - embdID := "" - provided := false - if req.EmbdID != nil { - embdID = strings.TrimSpace(*req.EmbdID) - provided = true - } - if req.EmbeddingModel != nil { - embdID = strings.TrimSpace(*req.EmbeddingModel) - provided = true - } - if !provided { - return "", false, nil - } - if embdID != "" { - if err := validateDatasetEmbeddingModel(embdID); err != nil { - return "", true, err - } - } - return embdID, true, nil -} - -func normalizeDatasetUpdateExt(ext map[string]interface{}) map[string]interface{} { - if ext == nil { - return nil - } - - updates := make(map[string]interface{}, len(ext)) - for key, value := range ext { - switch key { - case "embedding_model": - updates["embd_id"] = value - case "chunk_method": - updates["parser_id"] = value - case "connectors", "auto_metadata_config", "ext", "parse_type": - continue - default: - updates[key] = value - } - } - return updates -} - -// GetMetadataConfig gets the auto-metadata configuration for a dataset. -func (d *DatasetService) GetMetadataConfig(datasetID, tenantID string) (map[string]interface{}, common.ErrorCode, error) { - kb, err := d.kbDAO.GetByIDAndTenantID(datasetID, tenantID) - if err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", tenantID, datasetID) - } - return nil, common.CodeServerError, errors.New("Database operation failed") - } - if kb == nil { - return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", tenantID, datasetID) - } - - return map[string]interface{}{ - "metadata": parserConfigValueOrEmptyList(kb.ParserConfig, "metadata"), - "built_in_metadata": parserConfigValueOrEmptyList(kb.ParserConfig, "built_in_metadata"), - }, common.CodeSuccess, nil -} - -// UpdateMetadataConfig updates the auto-metadata configuration for a dataset. -func (d *DatasetService) UpdateMetadataConfig(datasetID, tenantID string, req *MetadataConfigRequest) (map[string]interface{}, common.ErrorCode, error) { - datasetID = strings.TrimSpace(datasetID) - tenantID = strings.TrimSpace(tenantID) - - kb, err := d.kbDAO.GetByIDAndTenantID(datasetID, tenantID) - if err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", tenantID, datasetID) - } - return nil, common.CodeServerError, errors.New("Database operation failed") - } - if kb == nil { - return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", tenantID, datasetID) - } - - if req == nil { - req = &MetadataConfigRequest{} - } - - metadata, err := normalizeMetadataConfigFields(req.Metadata, "metadata") - if err != nil { - return nil, common.CodeDataError, err - } - builtInMetadata, err := normalizeMetadataConfigFields(req.BuiltInMetadata, "built_in_metadata") - if err != nil { - return nil, common.CodeDataError, err - } - - parserConfig := kb.ParserConfig - if parserConfig == nil { - parserConfig = entity.JSONMap{} - } - parserConfig["metadata"] = metadata - parserConfig["built_in_metadata"] = builtInMetadata - - if err = d.kbDAO.UpdateByID(kb.ID, map[string]interface{}{"parser_config": parserConfig}); err != nil { - return nil, common.CodeServerError, errors.New("Update auto-metadata error.(Database error)") - } - - return map[string]interface{}{ - "metadata": metadata, - "built_in_metadata": builtInMetadata, - }, common.CodeSuccess, nil -} - -// Accessible checks if a knowledge base is accessible by a user -func (d *DatasetService) Accessible(kbID, userID string) bool { - return d.kbDAO.Accessible(kbID, userID) -} - -func (d *DatasetService) GetByID(kbID string) (*entity.Knowledgebase, error) { - return d.kbDAO.GetByID(kbID) -} - -// GetKnowledgebaseByID resolves a dataset entity without applying permission -// checks. Upload needs the same existence-then-auth ordering as Python. -func (d *DatasetService) GetKnowledgebaseByID(datasetID string) (*entity.Knowledgebase, error) { - datasetID = strings.TrimSpace(datasetID) - if datasetID == "" { - return nil, errors.New("Lack of \"Dataset ID\"") - } - normalizedID, err := normalizeDatasetID(datasetID) - if err != nil { - return nil, err - } - return d.kbDAO.GetByID(normalizedID) -} - -// CheckKBTeamPermission mirrors Python check_kb_team_permission. -func (d *DatasetService) CheckKBTeamPermission(kb *entity.Knowledgebase, userID string) bool { - return hasKBTeamPermission(kb, userID, d.tenantDAO) -} - -func (d *DatasetService) AggregateTags(datasetIDs []string, userID string) ([]map[string]interface{}, common.ErrorCode, error) { - if len(datasetIDs) == 0 { - return nil, common.CodeDataError, errors.New("Lack of dataset_ids in query parameters") - } - if d.docEngine == nil { - return nil, common.CodeServerError, errors.New("Document engine is not initialized") - } - - datasetIDsByTenant := make(map[string][]string) - for _, rawID := range datasetIDs { - rawID = strings.TrimSpace(rawID) - if rawID == "" { - continue - } - datasetID, err := normalizeDatasetID(rawID) - if err != nil { - return nil, common.CodeDataError, err - } - if !d.kbDAO.Accessible(datasetID, userID) { - return nil, common.CodeDataError, fmt.Errorf("No authorization for dataset '%s'", datasetID) - } - kb, err := d.kbDAO.GetByID(datasetID) - if err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, fmt.Errorf("Invalid Dataset ID '%s'", datasetID) - } - return nil, common.CodeServerError, errors.New("Database operation failed") - } - if kb.DocNum <= 0 { - continue - } - datasetIDsByTenant[kb.TenantID] = append(datasetIDsByTenant[kb.TenantID], datasetID) - } - - const pageSize = 10000 - merged := make(map[string]int) - for tenantID, kbIDs := range datasetIDsByTenant { - for offset := 0; ; offset += pageSize { - searchResp, err := d.docEngine.Search(context.Background(), &types.SearchRequest{ - IndexNames: []string{fmt.Sprintf("ragflow_%s", tenantID)}, - KbIDs: kbIDs, - Offset: offset, - Limit: pageSize, - SelectFields: []string{"tag_kwd"}, - }) - if err != nil { - return nil, common.CodeServerError, fmt.Errorf("failed to aggregate tags: %w", err) - } - for _, agg := range d.docEngine.GetAggregation(searchResp.Chunks, "tag_kwd") { - tag, _ := agg["key"].(string) - if tag == "" { - continue - } - switch count := agg["count"].(type) { - case int: - merged[tag] += count - case int32: - merged[tag] += int(count) - case int64: - merged[tag] += int(count) - case float64: - merged[tag] += int(count) - } - } - - chunkCount := len(searchResp.Chunks) - if chunkCount == 0 || chunkCount < pageSize { - break - } - if searchResp.Total > 0 && int64(offset+chunkCount) >= searchResp.Total { - break - } - } - } - result := make([]map[string]interface{}, 0, len(merged)) - for tag, count := range merged { - result = append(result, map[string]interface{}{ - "value": tag, - "count": count, - }) - } - return result, common.CodeSuccess, nil -} - -func (d *DatasetService) ListTags(datasetID, userID string) ([]map[string]interface{}, common.ErrorCode, error) { - datasetID = strings.TrimSpace(datasetID) - if datasetID == "" { - return nil, common.CodeDataError, errors.New("Lack of \"Dataset ID\"") - } - - normalizedID, err := normalizeDatasetID(datasetID) - if err != nil { - return nil, common.CodeDataError, err - } - datasetID = normalizedID - - if !d.kbDAO.Accessible(datasetID, userID) { - return nil, common.CodeDataError, errors.New("No authorization.") - } - if d.docEngine == nil { - return nil, common.CodeServerError, errors.New("Document engine is not initialized") - } - - kb, err := d.kbDAO.GetByID(datasetID) - if err != nil || kb == nil { - return nil, common.CodeDataError, errors.New("Invalid Dataset ID") - } - - indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - exists, err := d.docEngine.ChunkStoreExists(ctx, indexName, datasetID) - if err != nil { - return nil, common.CodeServerError, fmt.Errorf("failed to inspect chunk store: %w", err) - } - if !exists { - return []map[string]interface{}{}, common.CodeSuccess, nil - } - - const pageSize = 10000 - counts := make(map[string]int) - for offset := 0; ; offset += pageSize { - if err = ctx.Err(); err != nil { - return nil, common.CodeServerError, fmt.Errorf("list tags timeout or canceled: %w", err) - } - - searchResp, err := d.docEngine.Search(ctx, &types.SearchRequest{ - IndexNames: []string{indexName}, - KbIDs: []string{datasetID}, - Offset: offset, - Limit: pageSize, - SelectFields: []string{"tag_kwd"}, - }) - if err != nil { - return nil, common.CodeServerError, fmt.Errorf("failed to list tags: %w", err) - } - - for _, agg := range d.docEngine.GetAggregation(searchResp.Chunks, "tag_kwd") { - tag, _ := agg["key"].(string) - if tag == "" { - continue - } - switch count := agg["count"].(type) { - case int: - counts[tag] += count - case int32: - counts[tag] += int(count) - case int64: - counts[tag] += int(count) - case float64: - counts[tag] += int(count) - } - } - - chunkCount := len(searchResp.Chunks) - if chunkCount == 0 || chunkCount < pageSize { - break - } - if searchResp.Total > 0 && int64(offset+chunkCount) >= searchResp.Total { - break - } - } - - if len(counts) == 0 { - return []map[string]interface{}{}, common.CodeSuccess, nil - } - - tags := make([]string, 0, len(counts)) - for tag := range counts { - tags = append(tags, tag) - } - sort.Slice(tags, func(i, j int) bool { - if counts[tags[i]] != counts[tags[j]] { - return counts[tags[i]] > counts[tags[j]] - } - return tags[i] < tags[j] - }) - - result := make([]map[string]interface{}, 0, len(tags)) - for _, tag := range tags { - result = append(result, map[string]interface{}{ - "key": tag, - "count": counts[tag], - }) - } - - return result, common.CodeSuccess, nil -} - -// GetIngestionSummary returns dataset-level ingestion counters together with -// the aggregated document parsing status, mirroring -// dataset_api_service.get_ingestion_summary. -func (d *DatasetService) GetIngestionSummary(datasetID, userID string) (map[string]interface{}, common.ErrorCode, error) { - datasetID = strings.TrimSpace(datasetID) - if datasetID == "" { - return nil, common.CodeDataError, errors.New("Lack of \"Dataset ID\"") - } - - if !d.kbDAO.Accessible(datasetID, userID) { - return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", userID, datasetID) - } - - kb, err := d.kbDAO.GetByID(datasetID) - if err != nil || kb == nil { - return nil, common.CodeDataError, errors.New("Invalid Dataset ID") - } - - status, err := d.documentDAO.GetParsingStatusByKBID(datasetID) - if err != nil { - return nil, common.CodeServerError, errors.New("Database operation failed") - } - - return map[string]interface{}{ - "doc_num": kb.DocNum, - "chunk_num": kb.ChunkNum, - "token_num": kb.TokenNum, - "status": status, - }, common.CodeSuccess, nil -} - -// ListIngestionLogs lists ingestion logs for a dataset, mirroring -// dataset_api_service.list_ingestion_logs. log_type selects between -// dataset-level logs ("dataset") and per-file logs ("file"). -func (d *DatasetService) ListIngestionLogs(datasetID, userID string, page, pageSize int, orderby string, desc bool, operationStatus []string, createDateFrom, createDateTo, logType, keywords string) (map[string]interface{}, common.ErrorCode, error) { - datasetID = strings.TrimSpace(datasetID) - if datasetID == "" { - return nil, common.CodeDataError, errors.New("Lack of \"Dataset ID\"") - } - - if !d.kbDAO.Accessible(datasetID, userID) { - return nil, common.CodeDataError, errors.New("No authorization.") - } - - if logType != "dataset" && logType != "file" { - return nil, common.CodeDataError, errors.New("Invalid \"log_type\", expected \"dataset\" or \"file\"") - } - - var ( - logs []*entity.PipelineOperationLog - total int64 - err error - ) - if logType == "file" { - logs, total, err = d.pipelineLogDAO.GetFileLogsByKBID(datasetID, page, pageSize, orderby, desc, keywords, operationStatus, createDateFrom, createDateTo) - } else { - logs, total, err = d.pipelineLogDAO.GetDatasetLogsByKBID(datasetID, page, pageSize, orderby, desc, operationStatus, createDateFrom, createDateTo, keywords) - } - if err != nil { - return nil, common.CodeServerError, errors.New("Database operation failed") - } - - items := make([]map[string]interface{}, 0, len(logs)) - for _, log := range logs { - if log == nil { - continue - } - if logType == "file" { - items = append(items, fileIngestionLogToMap(log)) - } else { - items = append(items, datasetIngestionLogToMap(log)) - } - } - - return map[string]interface{}{ - "total": total, - "logs": items, - }, common.CodeSuccess, nil -} - -// GetIngestionLog returns a single ingestion log, mirroring -// dataset_api_service.get_ingestion_log. It returns the full record (including -// the `dsl`, `document_id`, `parser_id`, etc.) so that the front-end -// dataflow-result page can render the pipeline timeline and chunks. The -// file-level converter is a superset of the dataset-level fields, so it is -// correct for both dataset-level (graph/raptor/mindmap) and per-file logs. -func (d *DatasetService) GetIngestionLog(datasetID, userID, logID string) (map[string]interface{}, common.ErrorCode, error) { - datasetID = strings.TrimSpace(datasetID) - if datasetID == "" { - return nil, common.CodeDataError, errors.New("Lack of \"Dataset ID\"") - } - - if !d.kbDAO.Accessible(datasetID, userID) { - return nil, common.CodeDataError, errors.New("No authorization.") - } - - log, err := d.pipelineLogDAO.GetByIDAndKBID(logID, datasetID) - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return nil, common.CodeDataError, errors.New("Log not found") - } - return nil, common.CodeServerError, errors.New("Database operation failed") - } - - return fileIngestionLogToMap(log), common.CodeSuccess, nil -} - -func datasetIngestionLogToMap(log *entity.PipelineOperationLog) map[string]interface{} { - return map[string]interface{}{ - "id": log.ID, - "tenant_id": log.TenantID, - "kb_id": log.KbID, - "progress": log.Progress, - "progress_msg": stringPointerValue(log.ProgressMsg), - "process_begin_at": timePointerValue(log.ProcessBeginAt), - "process_duration": log.ProcessDuration, - "task_type": log.TaskType, - "operation_status": log.OperationStatus, - "avatar": stringPointerValue(log.Avatar), - "status": stringPointerValue(log.Status), - "create_time": int64PointerValue(log.CreateTime), - "create_date": timePointerValue(log.CreateDate), - "update_time": int64PointerValue(log.UpdateTime), - "update_date": timePointerValue(log.UpdateDate), - } -} - -func fileIngestionLogToMap(log *entity.PipelineOperationLog) map[string]interface{} { - return map[string]interface{}{ - "id": log.ID, - "document_id": log.DocumentID, - "tenant_id": log.TenantID, - "kb_id": log.KbID, - "pipeline_id": stringPointerValue(log.PipelineID), - "pipeline_title": stringPointerValue(log.PipelineTitle), - "parser_id": log.ParserID, - "document_name": log.DocumentName, - "document_suffix": log.DocumentSuffix, - "document_type": log.DocumentType, - "source_from": log.SourceFrom, - "progress": log.Progress, - "progress_msg": stringPointerValue(log.ProgressMsg), - "process_begin_at": timePointerValue(log.ProcessBeginAt), - "process_duration": log.ProcessDuration, - "dsl": jsonMapValue(log.DSL), - "task_type": log.TaskType, - "operation_status": log.OperationStatus, - "avatar": stringPointerValue(log.Avatar), - "status": stringPointerValue(log.Status), - "create_time": int64PointerValue(log.CreateTime), - "create_date": timePointerValue(log.CreateDate), - "update_time": int64PointerValue(log.UpdateTime), - "update_date": timePointerValue(log.UpdateDate), - } -} - -func stringPointerValue(s *string) interface{} { - if s == nil { - return nil - } - return *s -} - -func int64PointerValue(i *int64) interface{} { - if i == nil { - return nil - } - return *i -} - -func timePointerValue(t *time.Time) interface{} { - if t == nil { - return nil - } - return t.Format("2006-01-02 15:04:05") -} - -func jsonMapValue(m entity.JSONMap) interface{} { - if m == nil { - return nil - } - return m -} - -func (d *DatasetService) deleteDataset(tenantID string, kb *entity.Knowledgebase) error { - return dao.DB.Transaction(func(tx *gorm.DB) error { - if taskIDs := datasetIndexTaskIDs(kb); len(taskIDs) > 0 { - if err := tx.Where("id IN ?", taskIDs).Delete(&entity.Task{}).Error; err != nil { - return fmt.Errorf("Delete dataset error for %s", kb.ID) - } - } - - var documents []entity.Document - if err := tx.Where("kb_id = ?", kb.ID).Find(&documents).Error; err != nil { - return fmt.Errorf("Delete dataset error for %s", kb.ID) - } - - docIDs := make([]string, 0, len(documents)) - for _, document := range documents { - docIDs = append(docIDs, document.ID) - } - - if len(docIDs) > 0 { - var mappings []entity.File2Document - if err := tx.Where("document_id IN ?", docIDs).Find(&mappings).Error; err != nil { - return fmt.Errorf("Delete dataset error for %s", kb.ID) - } - - fileIDs := make([]string, 0, len(mappings)) - seenFileIDs := make(map[string]struct{}, len(mappings)) - for _, mapping := range mappings { - if mapping.FileID == nil || *mapping.FileID == "" { - continue - } - if _, exists := seenFileIDs[*mapping.FileID]; exists { - continue - } - seenFileIDs[*mapping.FileID] = struct{}{} - fileIDs = append(fileIDs, *mapping.FileID) - } - - if err := tx.Where("doc_id IN ?", docIDs).Delete(&entity.Task{}).Error; err != nil { - return fmt.Errorf("Delete dataset error for %s", kb.ID) - } - if err := tx.Where("document_id IN ?", docIDs).Delete(&entity.File2Document{}).Error; err != nil { - return fmt.Errorf("Delete dataset error for %s", kb.ID) - } - if len(fileIDs) > 0 { - if err := tx.Unscoped().Where("id IN ?", fileIDs).Delete(&entity.File{}).Error; err != nil { - return fmt.Errorf("Delete dataset error for %s", kb.ID) - } - } - if err := tx.Where("id IN ?", docIDs).Delete(&entity.Document{}).Error; err != nil { - return fmt.Errorf("Delete dataset error for %s", kb.ID) - } - } - - if err := tx.Unscoped(). - Where("source_type = ? AND type = ? AND name = ? AND tenant_id = ?", string(entity.FileSourceKnowledgebase), "folder", kb.Name, tenantID). - Delete(&entity.File{}).Error; err != nil { - return fmt.Errorf("Delete dataset error for %s", kb.ID) - } - - if err := tx.Where("id = ?", kb.ID).Delete(&entity.Knowledgebase{}).Error; err != nil { - return fmt.Errorf("Delete dataset error for %s", kb.ID) - } - - return nil - }) -} - -// validateParserID validates parser_id against the built-in -// pipeline registry. The registry is the single source of truth, so the -// legacy hardcoded allow-list is gone. Legacy values (e.g. "naive") are -// accepted via registry aliases. -func validateParserID(chunkMethod string) error { - registry, err := pipelinepkg.DefaultRegistry() - if err != nil || registry == nil { - return errors.New("parser_id validation unavailable: builtin pipeline registry not loaded") - } - if registry.IsValid(chunkMethod) { - return nil - } - return parserIDError() -} - -// parserIDError builds a validation error that lists the valid -// canonical parser_ids from the registry, mirroring the shape of the old -// hardcoded message but driven by the embedded templates. -func parserIDError() error { - registry, err := pipelinepkg.DefaultRegistry() - if err != nil || registry == nil { - return errors.New("invalid parser_id") - } - refs := registry.Refs() - switch len(refs) { - case 0: - return errors.New("invalid parser_id") - case 1: - return fmt.Errorf("Input should be '%s'", refs[0]) - default: - return fmt.Errorf("Input should be %s or '%s'", quoteList(refs[:len(refs)-1]), refs[len(refs)-1]) - } -} - -// quoteList renders ["a", "b"] as "'a', 'b'". -func quoteList(items []string) string { - quoted := make([]string, len(items)) - for i, v := range items { - quoted[i] = "'" + v + "'" - } - return strings.Join(quoted, ", ") -} - -func validateDatasetAvatar(avatar string) error { - if !strings.Contains(avatar, ",") { - return errors.New("Missing MIME prefix. Expected format: data:;base64,") - } - - prefix, _, _ := strings.Cut(avatar, ",") - if !strings.HasPrefix(prefix, "data:") { - return errors.New("Invalid MIME prefix format. Must start with 'data:'") - } - - mimeType, _, _ := strings.Cut(strings.TrimPrefix(prefix, "data:"), ";") - if _, ok := datasetSupportedAvatarMIMETypes[mimeType]; !ok { - return errors.New("Unsupported MIME type. Allowed: [image/jpeg image/png]") - } - - return nil -} - -func validateDatasetEmbeddingModel(embeddingModel string) error { - if embeddingModel == "" { - return errors.New("Embedding model identifier is required") - } - - if !strings.Contains(embeddingModel, "@") { - return nil - } - - parts := strings.Split(embeddingModel, "@") - for _, part := range parts { - if strings.TrimSpace(part) == "" { - return errors.New("Both model_name and provider must be non-empty strings") - } - } - if len(parts) < 2 { - return errors.New("Embedding model identifier must follow @ format") - } - if strings.TrimSpace(parts[0]) == "" || strings.TrimSpace(parts[len(parts)-1]) == "" { - return errors.New("Both model_name and provider must be non-empty strings") - } - - return nil -} - -func normalizeDatasetPipelineID(pipelineID string) (*string, error) { - pipelineID = strings.TrimSpace(pipelineID) - if pipelineID == "" { - return nil, nil - } - if len(pipelineID) != 32 { - return nil, errors.New("pipeline_id must be 32 hex characters") - } - for _, char := range pipelineID { - if !strings.ContainsRune("0123456789abcdefABCDEF", char) { - return nil, errors.New("pipeline_id must be hexadecimal") - } - } - - normalized := strings.ToLower(pipelineID) - return &normalized, nil -} - -func validateDatasetParserConfigSize(parserConfig map[string]interface{}) error { - if len(parserConfig) == 0 { - return nil - } - - data, err := json.Marshal(parserConfig) - if err != nil { - return errors.New("parser_config must be valid JSON") - } - if len(data) > 65535 { - return fmt.Errorf("Parser config exceeds size limit (max 65,535 characters). Current size: %d", len(data)) - } - - return nil -} - -// normalizeDatasetID canonicalizes an id into the 32-char hex form used by -// the storage layer. The "UUID1" name was a legacy term from when the -// Python service generated ids with `uuid.uuid1().hex`; the Go port uses -// `uuid.New()` (v4), so we accept any RFC 4122 version. We only reject the -// Nil UUID, which is the reserved "no id" sentinel. -func normalizeDatasetID(id string) (string, error) { - parsedUUID, err := uuid.Parse(id) - if err != nil { - return "", errors.New("Invalid UUID format") - } - if parsedUUID == (uuid.UUID{}) { - return "", errors.New("Invalid UUID format") - } - return strings.ReplaceAll(parsedUUID.String(), "-", ""), nil -} - -func (d *DatasetService) verifyEmbeddingAvailability(embdID string, tenantID string) (bool, string) { - _, _, _, _, err := NewModelProviderService().ResolveModelConfig(tenantID, entity.ModelTypeEmbedding, embdID) - if err != nil { - return false, err.Error() - } - return true, "" -} - -func parserConfigValueOrEmptyList(parserConfig map[string]interface{}, key string) interface{} { - if parserConfig == nil { - return []interface{}{} - } - - value, ok := parserConfig[key] - if !ok || value == nil { - return []interface{}{} - } - - return value -} - -func normalizeMetadataConfigFields(fields []MetadataConfigField, fieldName string) ([]map[string]interface{}, error) { - normalizedFields := make([]map[string]interface{}, 0, len(fields)) - for i, field := range fields { - key := strings.TrimSpace(field.Key) - if key == "" { - return nil, fmt.Errorf("%s[%d].key is required", fieldName, i) - } - if len(key) > 255 { - return nil, fmt.Errorf("%s[%d].key should have at most 255 characters", fieldName, i) - } - - fieldType := strings.TrimSpace(field.Type) - if _, ok := datasetAllowedMetadataTypes[fieldType]; !ok { - return nil, fmt.Errorf("%s[%d].type should be one of 'string', 'list', 'time' or 'number'", fieldName, i) - } - - if field.Description != nil && len(*field.Description) > 65535 { - return nil, fmt.Errorf("%s[%d].description should have at most 65535 characters", fieldName, i) - } - - normalizedFields = append(normalizedFields, map[string]interface{}{ - "key": key, - "type": fieldType, - "description": field.Description, - "enum": field.Enum, - }) - } - - return normalizedFields, nil -} - -func datasetListItemToMap(kb *entity.KnowledgebaseListItem) map[string]interface{} { - item := map[string]interface{}{ - "id": kb.ID, - "name": kb.Name, - "tenant_id": kb.TenantID, - "permission": kb.Permission, - "document_count": kb.DocNum, - "token_num": kb.TokenNum, - "chunk_count": kb.ChunkNum, - "parser_id": kb.ParserID, - "embedding_model": kb.EmbdID, - "nickname": kb.Nickname, - } - - if kb.Avatar != nil { - item["avatar"] = *kb.Avatar - } - if kb.Language != nil { - item["language"] = *kb.Language - } - if kb.Description != nil { - item["description"] = *kb.Description - } - if kb.TenantAvatar != nil { - item["tenant_avatar"] = *kb.TenantAvatar - } - if kb.UpdateTime != nil { - item["update_time"] = *kb.UpdateTime - } - - return item -} - -func datasetToMap(kb *entity.Knowledgebase) map[string]interface{} { - item := map[string]interface{}{ - "id": kb.ID, - "tenant_id": kb.TenantID, - "name": kb.Name, - "embedding_model": kb.EmbdID, - "permission": kb.Permission, - "created_by": kb.CreatedBy, - "document_count": kb.DocNum, - "token_num": kb.TokenNum, - "chunk_count": kb.ChunkNum, - "similarity_threshold": kb.SimilarityThreshold, - "vector_similarity_weight": kb.VectorSimilarityWeight, - "parser_id": kb.ParserID, - "parser_config": kb.ParserConfig, - "pagerank": kb.Pagerank, - "create_time": kb.CreateTime, - } - - if kb.Avatar != nil { - item["avatar"] = *kb.Avatar - } - if kb.Language != nil { - item["language"] = *kb.Language - } - if kb.Description != nil { - item["description"] = *kb.Description - } - if kb.PipelineID != nil { - item["pipeline_id"] = *kb.PipelineID - } - if kb.GraphragTaskID != nil { - item["graphrag_task_id"] = *kb.GraphragTaskID - } - if kb.GraphragTaskFinishAt != nil { - item["graphrag_task_finish_at"] = kb.GraphragTaskFinishAt.Format("2006-01-02 15:04:05") - } - if kb.RaptorTaskID != nil { - item["raptor_task_id"] = *kb.RaptorTaskID - } - if kb.RaptorTaskFinishAt != nil { - item["raptor_task_finish_at"] = kb.RaptorTaskFinishAt.Format("2006-01-02 15:04:05") - } - if kb.MindmapTaskID != nil { - item["mindmap_task_id"] = *kb.MindmapTaskID - } - if kb.MindmapTaskFinishAt != nil { - item["mindmap_task_finish_at"] = kb.MindmapTaskFinishAt.Format("2006-01-02 15:04:05") - } - if kb.UpdateTime != nil { - item["update_time"] = *kb.UpdateTime - } - - return item -} - -func limitStrings(values []string, limit int) []string { - if len(values) <= limit { - return values - } - return values[:limit] -} - -func (d *DatasetService) RenameTag(datasetID, userID, fromTag, toTag string) (map[string]interface{}, common.ErrorCode, error) { - fromTag = strings.TrimSpace(fromTag) - toTag = strings.TrimSpace(toTag) - - datasetID, err := normalizeDatasetID(datasetID) - if err != nil { - return nil, common.CodeDataError, err - } - if strings.TrimSpace(datasetID) == "" { - return nil, common.CodeDataError, errors.New("Lack of \"Dataset ID\"") - } - if !d.kbDAO.Accessible(datasetID, userID) { - return nil, common.CodeDataError, errors.New("No authorization.") - } - if d.docEngine == nil { - return nil, common.CodeServerError, errors.New("Document engine is not initialized") - } - - kb, err := d.kbDAO.GetByID(datasetID) - if err != nil || kb == nil { - return nil, common.CodeDataError, errors.New("Invalid Dataset ID") - } - indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) - - condition := map[string]interface{}{ - "tag_kwd": fromTag, - "kb_id": datasetID, - } - newValue := map[string]interface{}{ - "remove": map[string]interface{}{ - "tag_kwd": fromTag, - }, - "add": map[string]interface{}{ - "tag_kwd": toTag, - }, - } - - err = d.docEngine.UpdateChunks(context.Background(), condition, newValue, indexName, datasetID) - if err != nil { - return nil, common.CodeServerError, fmt.Errorf("failed to rename tag: %w", err) - } - - return map[string]interface{}{ - "from": fromTag, - "to": toTag, - }, common.CodeSuccess, nil -} - -func (d *DatasetService) GetFieldMap(ids []string) (map[string]interface{}, error) { - return d.kbDAO.GetFieldMap(ids) -} - -// resolveComponentParamsDefaults loads the DSL for the target pipeline and -// returns the component params defaults as an entity.JSONMap {cpnID: {param: value}}. -// For builtin templates the DSL is loaded from the embedded registry; for custom -// canvas pipelines it is loaded from the canvas row in the database. -func resolveComponentParamsDefaults(parserID string, pipelineID *string) (entity.JSONMap, error) { - isCanvas := pipelineID != nil && strings.TrimSpace(*pipelineID) != "" - var cp map[string]map[string]any - var err error - if isCanvas { - dslJSON, lerr := loadCanvasDSLJSON(strings.TrimSpace(*pipelineID)) - if lerr != nil { - return nil, fmt.Errorf("load canvas DSL: %w", lerr) - } - cp, err = pipelinepkg.ComponentParamsDefaults(dslJSON) - } else { - registry, regErr := pipelinepkg.DefaultRegistry() - if regErr != nil { - return nil, fmt.Errorf("builtin registry: %w", regErr) - } - if !registry.IsValid(parserID) { - return nil, fmt.Errorf("unknown builtin parser_id: %q", parserID) - } - dslStr, dslErr := pipelinepkg.LoadBuiltinDSL(parserID) - if dslErr != nil { - return nil, fmt.Errorf("load builtin DSL: %w", dslErr) - } - cp, err = pipelinepkg.ComponentParamsDefaults([]byte(dslStr)) - } - if err != nil { - return nil, err - } - out := make(entity.JSONMap, len(cp)) - for k, v := range cp { - out[k] = v - } - return out, nil -} - -// canvasAccessibleForUser reports whether the given canvas is owned by or -// team-shared with the caller. -func canvasAccessibleForUser(userID, canvasID string) (bool, error) { - tenantIDs, _ := dao.NewUserTenantDAO().GetTenantIDsByUserID(userID) - return dao.NewUserCanvasDAO().Accessible(canvasID, userID, tenantIDs), nil -} diff --git a/internal/service/dataset/create_test.go b/internal/service/dataset/create_test.go new file mode 100644 index 0000000000..c107b8d0a0 --- /dev/null +++ b/internal/service/dataset/create_test.go @@ -0,0 +1,157 @@ +package dataset + +import ( + "strings" + "testing" + + "ragflow/internal/common" + "ragflow/internal/dao" + "ragflow/internal/entity" + "ragflow/internal/service" +) + +func insertCreateDatasetTenant(t *testing.T, tenantID string) { + t.Helper() + var existing entity.Tenant + if err := dao.DB.Where("id = ?", tenantID).First(&existing).Error; err != nil { + tn := &entity.Tenant{ + ID: tenantID, + LLMID: "llm-default", + EmbdID: "embd-default", + TenantEmbdID: sptr("embd-1"), + ASRID: "asr-default", + Status: sptr("1"), + } + if err := dao.DB.Create(tn).Error; err != nil { + t.Fatalf("insert test tenant: %v", err) + } + } +} + +func testDatasetCreateService(t *testing.T) *DatasetService { + t.Helper() + return &DatasetService{ + kbDAO: dao.NewKnowledgebaseDAO(), + documentDAO: dao.NewDocumentDAO(), + connectorDAO: dao.NewConnectorDAO(), + tenantDAO: dao.NewTenantDAO(), + } +} + +func TestCreateDataset_NoComponentParams(t *testing.T) { + db := setupServiceTestDB(t) + pushServiceDB(t, db) + insertCreateDatasetTenant(t, "tenant-1") + + chunkMethod := "naive" + result, code, err := testDatasetCreateService(t).CreateDataset(&service.CreateDatasetRequest{ + Name: "ds-no-cp", + ParserID: &chunkMethod, + }, "tenant-1") + if err != nil { + t.Fatalf("CreateDataset failed: %v", err) + } + if code != common.CodeSuccess { + t.Fatalf("expected success code, got %d", code) + } + if result["parser_id"] != strings.TrimSpace(chunkMethod) { + t.Fatalf("expected parser_id %q, got %#v", chunkMethod, result["parser_id"]) + } +} + +func TestCreateDataset_ComponentParamsPopulated(t *testing.T) { + db := setupServiceTestDB(t) + pushServiceDB(t, db) + insertCreateDatasetTenant(t, "tenant-1") + + chunkMethod := "general" + result, code, err := testDatasetCreateService(t).CreateDataset(&service.CreateDatasetRequest{ + Name: "ds-with-cp", + ParserID: &chunkMethod, + }, "tenant-1") + if err != nil { + t.Fatalf("CreateDataset failed: %v", err) + } + if code != common.CodeSuccess { + t.Fatalf("expected success code, got %d", code) + } + parserConfig, ok := result["parser_config"].(entity.JSONMap) + if !ok || len(parserConfig) == 0 { + t.Fatal("expected non-empty parser_config for general pipeline") + } +} + +func TestCreateDataset_ParseTypeBuiltinClearsPipelineID(t *testing.T) { + db := setupServiceTestDB(t) + pushServiceDB(t, db) + insertCreateDatasetTenant(t, "tenant-1") + + pipelineID := "0123456789abcdef0123456789abcdef" + parseTypeBuiltin := 1 + chunkMethod := "naive" + result, code, err := testDatasetCreateService(t).CreateDataset(&service.CreateDatasetRequest{ + Name: "ds-parse-builtin", + ParserID: &chunkMethod, + PipelineID: &pipelineID, + ParseType: &parseTypeBuiltin, + }, "tenant-1") + if err != nil { + t.Fatalf("CreateDataset failed: %v", err) + } + if code != common.CodeSuccess { + t.Fatalf("expected success code, got %d", code) + } + if result["parser_id"] != chunkMethod { + t.Fatalf("expected parser_id %q, got %#v", chunkMethod, result["parser_id"]) + } + if v, ok := result["pipeline_id"]; ok && v != nil { + t.Fatalf("expected pipeline_id to be nil for BuiltIn mode, got %#v", v) + } +} + +func TestCreateDataset_ParseTypePipelineIgnoresParserID(t *testing.T) { + t.Skip("requires canvas seed data in test DB") + db := setupServiceTestDB(t) + pushServiceDB(t, db) + insertCreateDatasetTenant(t, "tenant-1") + + pipelineID := "0123456789abcdef0123456789abcdef" + parseTypePipeline := 2 + chunkMethod := "naive" + result, code, err := testDatasetCreateService(t).CreateDataset(&service.CreateDatasetRequest{ + Name: "ds-parse-pipeline", + ParserID: &chunkMethod, + PipelineID: &pipelineID, + ParseType: &parseTypePipeline, + }, "tenant-1") + if err != nil { + t.Fatalf("CreateDataset failed: %v", err) + } + if code != common.CodeSuccess { + t.Fatalf("expected success code, got %d", code) + } + if v, ok := result["parser_id"]; !ok || v == nil { + } else { + t.Fatalf("expected parser_id to be empty for Pipeline mode, got %#v", v) + } +} + +func TestCreateDataset_RejectsBothWithoutParseType(t *testing.T) { + db := setupServiceTestDB(t) + pushServiceDB(t, db) + insertCreateDatasetTenant(t, "tenant-1") + + pipelineID := "0123456789abcdef0123456789abcdef" + chunkMethod := "naive" + _, code, err := testDatasetCreateService(t).CreateDataset(&service.CreateDatasetRequest{ + Name: "ds-both", + ParserID: &chunkMethod, + PipelineID: &pipelineID, + }, "tenant-1") + if err == nil { + t.Fatal("expected error when both parser_id and pipeline_id are provided without parse_type") + } + if code != common.CodeDataError { + t.Fatalf("expected CodeDataError, got %d", code) + } +} diff --git a/internal/service/dataset/crud.go b/internal/service/dataset/crud.go new file mode 100644 index 0000000000..e1102f3707 --- /dev/null +++ b/internal/service/dataset/crud.go @@ -0,0 +1,452 @@ +package dataset + +import ( + "context" + "errors" + "fmt" + "strings" + + "ragflow/internal/common" + "ragflow/internal/dao" + "ragflow/internal/entity" + "ragflow/internal/service" + "ragflow/internal/utility" + + "go.uber.org/zap" + "gorm.io/gorm" +) + +func (d *DatasetService) CreateDataset(req *service.CreateDatasetRequest, tenantID string) (map[string]interface{}, common.ErrorCode, error) { + if !common.IsValidString(req.Name) { + return nil, common.CodeDataError, errors.New("Dataset name must be string.") + } + + name := strings.TrimSpace(req.Name) + if name == "" { + return nil, common.CodeDataError, errors.New("Dataset name can't be empty.") + } + if len(name) > entity.DatasetNameLimit { + 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) + if err != nil || tenant == nil { + return nil, common.CodeDataError, errors.New("Tenant not found.") + } + + isPipelineMode := req.ParseType != nil && *req.ParseType == 2 + isBuiltinMode := req.ParseType != nil && *req.ParseType == 1 + + if isBuiltinMode && req.PipelineID != nil { + req.PipelineID = nil + } + if isPipelineMode && req.ParserID != nil { + req.ParserID = nil + } + + if req.ParseType == nil && req.ParserID != nil && req.PipelineID != nil { + return nil, common.CodeDataError, errors.New("parser_id and pipeline_id are mutually exclusive") + } + + parserID := "" + permission := "me" + embeddingModel := "" + pipelineID := req.PipelineID + + if req.Permission != nil { + permission = strings.TrimSpace(*req.Permission) + if permission != "me" && permission != "team" { + return nil, common.CodeDataError, errors.New("Input should be 'me' or 'team'") + } + } + if req.ParserID != nil { + parserID = strings.TrimSpace(*req.ParserID) + if err := validateParserID(parserID); err != nil { + return nil, common.CodeDataError, err + } + pipelineID = nil + } + if req.PipelineID != nil { + normalizedPipelineID, err := normalizeDatasetPipelineID(*req.PipelineID) + if err != nil { + return nil, common.CodeDataError, err + } + pipelineID = normalizedPipelineID + } + if req.EmbeddingModel != nil { + embeddingModel = strings.TrimSpace(*req.EmbeddingModel) + if err := validateDatasetEmbeddingModel(embeddingModel); err != nil { + return nil, common.CodeDataError, err + } + } + + if pipelineID != nil && strings.TrimSpace(*pipelineID) != "" { + if ok, err := canvasAccessibleForUser(tenantID, strings.TrimSpace(*pipelineID)); err != nil { + return nil, common.CodeServerError, err + } else if !ok { + return nil, common.CodeDataError, errors.New("canvas is not accessible") + } + } + + parserConfig, cpErr := service.ResolveComponentParamsDefaults(parserID, pipelineID) + if cpErr != nil { + common.Warn("failed to resolve component params defaults for dataset", + zap.String("parserID", parserID), zap.Error(cpErr)) + parserConfig = entity.JSONMap{} + } + + var parserConfigMap map[string]interface{} = parserConfig + + embdID := tenant.EmbdID + tenantEmbdID := ptrStringValue(tenant.TenantEmbdID) + if embeddingModel != "" { + ok, message := d.verifyEmbeddingAvailability(embeddingModel, tenantID) + if !ok { + return nil, common.CodeDataError, errors.New(message) + } + embdID = embeddingModel + tenantEmbdID = "" + } + if embdID != "" && tenantEmbdID == "" { + resolvedID, err := service.NewModelProviderService().ResolveModelID(tenantID, entity.ModelTypeEmbedding, embdID) + if err == nil { + tenantEmbdID = resolvedID + } else { + return nil, common.CodeDataError, err + } + } + + kbID := utility.GenerateToken() + status := string(entity.StatusValid) + duplicateName, err := common.DuplicateName(func(n, tid string) bool { + existing, err := d.kbDAO.GetByName(n, tid) + return err == nil && existing != nil + }, name, tenantID) + if err != nil { + return nil, common.CodeDataError, err + } + + kb := &entity.Knowledgebase{ + ID: kbID, + Name: duplicateName, + TenantID: tenantID, + CreatedBy: tenantID, + ParserID: parserID, + PipelineID: pipelineID, + ParserConfig: entity.JSONMap(parserConfigMap), + Permission: permission, + EmbdID: embdID, + TenantEmbdID: stringPtrIfNotEmpty(tenantEmbdID), + Status: &status, + } + + if err = d.kbDAO.Create(kb); err != nil { + return nil, common.CodeServerError, errors.New("Failed to save dataset") + } + + createdKB, err := d.kbDAO.GetByID(kbID) + if err != nil || createdKB == nil { + return nil, common.CodeServerError, errors.New("Dataset created failed") + } + + return datasetToMap(createdKB), common.CodeSuccess, nil +} + +func (d *DatasetService) GetDataset(datasetID, userID string) (map[string]interface{}, common.ErrorCode, error) { + datasetID = strings.TrimSpace(datasetID) + if datasetID == "" { + return nil, common.CodeDataError, errors.New("Lack of \"Dataset ID\"") + } + + normalizedID, err := normalizeDatasetID(datasetID) + if err != nil { + return nil, common.CodeDataError, err + } + datasetID = normalizedID + + if !d.kbDAO.Accessible(datasetID, userID) { + return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", userID, datasetID) + } + + kb, err := d.kbDAO.GetByID(datasetID) + if err != nil || kb == nil { + return nil, common.CodeDataError, errors.New("Invalid Dataset ID") + } + + data := datasetToMap(kb) + + size, err := d.documentDAO.SumSizeByDatasetID(datasetID) + if err != nil { + return nil, common.CodeServerError, errors.New("Database operation failed") + } + data["size"] = size + + connectors, err := d.connectorDAO.ListByDatasetID(datasetID) + if err != nil { + return nil, common.CodeServerError, errors.New("Database operation failed") + } + data["connectors"] = datasetConnectorsOrEmpty(connectors) + + return data, common.CodeSuccess, nil +} + +func (d *DatasetService) DeleteDatasets(ids []string, deleteAll bool, tenantID string) (map[string]interface{}, common.ErrorCode, error) { + normalizedIDs := make([]string, 0, len(ids)) + seenIDs := make(map[string]struct{}, len(ids)) + for _, id := range ids { + normalizedID, err := normalizeDatasetID(id) + if err != nil { + return nil, common.CodeDataError, err + } + if _, seen := seenIDs[normalizedID]; seen { + continue + } + seenIDs[normalizedID] = struct{}{} + normalizedIDs = append(normalizedIDs, normalizedID) + } + + // If no explicit ids and deleteAll is set, resolve all datasets for this tenant. + if len(normalizedIDs) == 0 { + if !deleteAll { + return map[string]interface{}{"deleted": []string{}}, common.CodeSuccess, nil + } + kbs, err := d.kbDAO.Query(map[string]interface{}{"tenant_id": tenantID}) + if err != nil { + return nil, common.CodeServerError, errors.New("Database operation failed") + } + for _, kb := range kbs { + normalizedIDs = append(normalizedIDs, kb.ID) + } + } + + // Validate ownership: collect KBs that exist and belong to this tenant. + kbs := make([]*entity.Knowledgebase, 0, len(normalizedIDs)) + unauthorizedIDs := make([]string, 0) + for _, id := range normalizedIDs { + kb, err := d.kbDAO.GetByIDAndTenantID(id, tenantID) + if err != nil || kb == nil { + unauthorizedIDs = append(unauthorizedIDs, id) + continue + } + kbs = append(kbs, kb) + } + if len(unauthorizedIDs) > 0 { + return nil, common.CodeDataError, + fmt.Errorf("User '%s' lacks permission for datasets: '%s'", tenantID, strings.Join(unauthorizedIDs, ", ")) + } + + successCount := 0 + errorsList := make([]string, 0) + for _, kb := range kbs { + if err := d.deleteDataset(tenantID, kb); err != nil { + errorsList = append(errorsList, err.Error()) + common.Warn("deleteDataset failed", zap.String("kb_id", kb.ID), zap.Error(err)) + continue + } + successCount++ + } + + return map[string]interface{}{ + "success_count": successCount, + "errors": errorsList, + }, common.CodeSuccess, nil +} + +func (d *DatasetService) deleteDataset(tenantID string, kb *entity.Knowledgebase) error { + // Collect document IDs first so engine cleanup can run before the + // transaction (engine ops are not transactional). + var documents []entity.Document + if err := dao.DB.Where("kb_id = ?", kb.ID).Find(&documents).Error; err != nil { + return fmt.Errorf("Delete dataset error for %s", kb.ID) + } + docIDs := extractDocIDs(documents) + if len(docIDs) > 0 { + d.deleteDatasetEngineData(kb, docIDs) + } + + return dao.DB.Transaction(func(tx *gorm.DB) error { + // Delete index tasks referencing this KB. + if taskIDs := datasetIndexTaskIDs(kb); len(taskIDs) > 0 { + if err := tx.Where("id IN ?", taskIDs).Delete(&entity.Task{}).Error; err != nil { + return fmt.Errorf("Delete dataset error for %s", kb.ID) + } + } + + if len(docIDs) > 0 { + var mappings []entity.File2Document + if err := tx.Where("document_id IN ?", docIDs).Find(&mappings).Error; err != nil { + return fmt.Errorf("Delete dataset error for %s", kb.ID) + } + fileIDs := extractUniqueFileIDs(mappings) + + if err := tx.Where("doc_id IN ?", docIDs).Delete(&entity.Task{}).Error; err != nil { + return fmt.Errorf("Delete dataset error for %s", kb.ID) + } + if err := tx.Where("document_id IN ?", docIDs).Delete(&entity.File2Document{}).Error; err != nil { + return fmt.Errorf("Delete dataset error for %s", kb.ID) + } + if len(fileIDs) > 0 { + if err := tx.Unscoped().Where("id IN ?", fileIDs).Delete(&entity.File{}).Error; err != nil { + return fmt.Errorf("Delete dataset error for %s", kb.ID) + } + } + if err := tx.Where("id IN ?", docIDs).Delete(&entity.Document{}).Error; err != nil { + return fmt.Errorf("Delete dataset error for %s", kb.ID) + } + } + + // Delete the KB folder file record. + if err := tx.Unscoped(). + Where("source_type = ? AND type = ? AND name = ? AND tenant_id = ?", + string(entity.FileSourceKnowledgebase), "folder", kb.Name, tenantID). + Delete(&entity.File{}).Error; err != nil { + return fmt.Errorf("Delete dataset error for %s", kb.ID) + } + + if err := tx.Where("id = ?", kb.ID).Delete(&entity.Knowledgebase{}).Error; err != nil { + return fmt.Errorf("Delete dataset error for %s", kb.ID) + } + return nil + }) +} + +func (d *DatasetService) ListDatasets(id, name string, page, pageSize int, orderby string, desc bool, keywords string, ownerIDs []string, parserID, userID string) ([]map[string]interface{}, int64, common.ErrorCode, error) { + id = strings.TrimSpace(id) + if id != "" { + normalizedID, err := normalizeDatasetID(id) + if err != nil { + return nil, 0, common.CodeDataError, err + } + id = normalizedID + + kbs, err := d.kbDAO.GetKBByIDAndUserID(id, userID) + if err != nil { + return nil, 0, common.CodeServerError, errors.New("Database operation failed") + } + if len(kbs) == 0 { + return nil, 0, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", userID, id) + } + } + + name = strings.TrimSpace(name) + if name != "" { + kbs, err := d.kbDAO.GetKBByNameAndUserID(name, userID) + if err != nil { + return nil, 0, common.CodeServerError, errors.New("Database operation failed") + } + if len(kbs) == 0 { + return nil, 0, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", userID, name) + } + } + + if page <= 0 { + page = 1 + } + if pageSize <= 0 { + pageSize = 30 + } + + orderby = strings.TrimSpace(orderby) + if _, ok := datasetAllowedOrderByFields[orderby]; !ok { + orderby = "create_time" + } + + keywords = strings.TrimSpace(keywords) + parserID = strings.TrimSpace(parserID) + + tenantIDs := make([]string, 0, len(ownerIDs)) + for _, ownerID := range ownerIDs { + ownerID = strings.TrimSpace(ownerID) + if ownerID != "" { + tenantIDs = append(tenantIDs, ownerID) + } + } + if len(tenantIDs) == 0 { + joinedTenants, err := d.tenantDAO.GetJoinedTenantsByUserID(userID) + if err != nil { + return nil, 0, common.CodeServerError, errors.New("Database operation failed") + } + for _, joinedTenant := range joinedTenants { + if joinedTenant == nil || joinedTenant.TenantID == "" { + continue + } + tenantIDs = append(tenantIDs, joinedTenant.TenantID) + } + } + + kbs, total, err := d.kbDAO.GetByTenantIDs(tenantIDs, userID, page, pageSize, orderby, desc, keywords, parserID) + if err != nil { + return nil, 0, common.CodeServerError, errors.New("Database operation failed") + } + + data := make([]map[string]interface{}, 0, len(kbs)) + for _, kb := range kbs { + if kb == nil { + continue + } + data = append(data, datasetListItemToMap(kb)) + } + + return data, total, common.CodeSuccess, nil +} + +// ptrStringValue safely dereferences a *string. +func ptrStringValue(s *string) string { + if s == nil { + return "" + } + return *s +} + +// stringPtrIfNotEmpty returns a pointer to s if s is non-empty. +func stringPtrIfNotEmpty(s string) *string { + if s == "" { + return nil + } + return &s +} + +// extractDocIDs returns the document IDs from a slice of documents. +func extractDocIDs(docs []entity.Document) []string { + ids := make([]string, 0, len(docs)) + for _, doc := range docs { + ids = append(ids, doc.ID) + } + return ids +} + +// deleteDatasetEngineData cleans up engine-level chunks and metadata for all +// documents in a dataset being deleted. Called before the DB transaction +// because engine operations are not transactional. +func (d *DatasetService) deleteDatasetEngineData(kb *entity.Knowledgebase, docIDs []string) { + if d.docEngine == nil || len(docIDs) == 0 { + return + } + ctx := context.Background() + indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) + + if _, err := d.docEngine.DeleteChunks(ctx, map[string]interface{}{"doc_id": docIDs}, indexName, kb.ID); err != nil { + common.Logger.Warn(fmt.Sprintf("deleteDataset: failed to delete chunks for kb %s: %v", kb.ID, err)) + } + if _, err := d.docEngine.DeleteMetadata(ctx, map[string]interface{}{"doc_id": docIDs}, kb.TenantID); err != nil { + common.Logger.Warn(fmt.Sprintf("deleteDataset: failed to delete metadata for kb %s: %v", kb.ID, err)) + } +} + +// extractUniqueFileIDs returns deduplicated, non-empty file IDs from +// file2document mappings. +func extractUniqueFileIDs(mappings []entity.File2Document) []string { + ids := make([]string, 0, len(mappings)) + seen := make(map[string]struct{}, len(mappings)) + for _, m := range mappings { + if m.FileID == nil || *m.FileID == "" { + continue + } + if _, exists := seen[*m.FileID]; exists { + continue + } + seen[*m.FileID] = struct{}{} + ids = append(ids, *m.FileID) + } + return ids +} diff --git a/internal/service/dataset/crud_test.go b/internal/service/dataset/crud_test.go new file mode 100644 index 0000000000..5a07aedcdb --- /dev/null +++ b/internal/service/dataset/crud_test.go @@ -0,0 +1,101 @@ +package dataset + +import ( + "testing" + + "ragflow/internal/entity" +) + +func TestExtractUniqueFileIDs(t *testing.T) { + f1 := "f1" + f2 := "f2" + empty := "" + + tests := []struct { + name string + mappings []entity.File2Document + want []string + }{ + { + name: "empty", + mappings: nil, + want: nil, + }, + { + name: "single file", + mappings: []entity.File2Document{ + {FileID: &f1}, + }, + want: []string{"f1"}, + }, + { + name: "deduplicates duplicate file IDs", + mappings: []entity.File2Document{ + {FileID: &f1}, + {FileID: &f1}, + {FileID: &f2}, + }, + want: []string{"f1", "f2"}, + }, + { + name: "skips nil FileID", + mappings: []entity.File2Document{ + {FileID: nil}, + {FileID: &f1}, + }, + want: []string{"f1"}, + }, + { + name: "skips empty FileID", + mappings: []entity.File2Document{ + {FileID: &empty}, + {FileID: &f1}, + }, + want: []string{"f1"}, + }, + { + name: "all nil or empty returns empty", + mappings: []entity.File2Document{ + {FileID: nil}, + {FileID: &empty}, + }, + want: nil, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := extractUniqueFileIDs(tt.mappings) + if len(got) != len(tt.want) { + t.Fatalf("len = %d, want %d (got %v, want %v)", len(got), len(tt.want), got, tt.want) + } + for i := range got { + if got[i] != tt.want[i] { + t.Fatalf("got[%d] = %q, want %q", i, got[i], tt.want[i]) + } + } + }) + } +} + +func TestExtractDocIDs(t *testing.T) { + docs := []entity.Document{ + {ID: "d1"}, + {ID: "d2"}, + {ID: "d3"}, + } + got := extractDocIDs(docs) + want := []string{"d1", "d2", "d3"} + if len(got) != len(want) { + t.Fatalf("len = %d, want %d", len(got), len(want)) + } + for i := range got { + if got[i] != want[i] { + t.Fatalf("got[%d] = %q, want %q", i, got[i], want[i]) + } + } + + if got := extractDocIDs(nil); len(got) != 0 { + t.Fatalf("nil input: got %v, want empty", got) + } +} diff --git a/internal/service/dataset/fake_doc_engine_test.go b/internal/service/dataset/fake_doc_engine_test.go new file mode 100644 index 0000000000..0358cc901e --- /dev/null +++ b/internal/service/dataset/fake_doc_engine_test.go @@ -0,0 +1,82 @@ +package dataset + +import ( + "context" + + "ragflow/internal/engine/types" +) + +// fakeChatDocEngine is a no-op DocEngine implementation for tests. +type fakeChatDocEngine struct{} + +func (fakeChatDocEngine) CreateChunkStore(context.Context, string, string, int, string) error { + return nil +} +func (fakeChatDocEngine) InsertChunks(context.Context, []map[string]interface{}, string, string) ([]string, error) { + return nil, nil +} +func (fakeChatDocEngine) UpdateChunks(context.Context, map[string]interface{}, map[string]interface{}, string, string) error { + return nil +} +func (fakeChatDocEngine) DeleteChunks(context.Context, map[string]interface{}, string, string) (int64, error) { + return 0, nil +} +func (fakeChatDocEngine) Search(context.Context, *types.SearchRequest) (*types.SearchResult, error) { + return nil, nil +} +func (fakeChatDocEngine) GetChunk(context.Context, string, string, []string) (interface{}, error) { + return nil, nil +} +func (fakeChatDocEngine) DropChunkStore(context.Context, string, string) error { return nil } +func (fakeChatDocEngine) ChunkStoreExists(context.Context, string, string) (bool, error) { + return true, nil +} +func (fakeChatDocEngine) CreateMetadataStore(context.Context, string) error { return nil } +func (fakeChatDocEngine) InsertMetadata(context.Context, []map[string]interface{}, string) ([]string, error) { + return nil, nil +} +func (fakeChatDocEngine) UpdateMetadata(context.Context, string, string, map[string]interface{}, string) error { + return nil +} +func (fakeChatDocEngine) DeleteMetadata(context.Context, map[string]interface{}, string) (int64, error) { + return 0, nil +} +func (fakeChatDocEngine) DeleteMetadataKeys(context.Context, string, string, []string, string) error { + return nil +} +func (fakeChatDocEngine) DropMetadataStore(context.Context, string) error { return nil } +func (fakeChatDocEngine) MetadataStoreExists(context.Context, string) (bool, error) { return true, nil } +func (fakeChatDocEngine) SearchMetadata(context.Context, *types.SearchMetadataRequest) (*types.SearchMetadataResult, error) { + return nil, nil +} +func (fakeChatDocEngine) IndexDocument(context.Context, string, string, interface{}) error { + return nil +} +func (fakeChatDocEngine) DeleteDocument(context.Context, string, string) error { return nil } +func (fakeChatDocEngine) BulkIndex(context.Context, string, []interface{}) (interface{}, error) { + return nil, nil +} +func (fakeChatDocEngine) GetFields([]map[string]interface{}, []string) map[string]map[string]interface{} { + return nil +} +func (fakeChatDocEngine) GetAggregation([]map[string]interface{}, string) []map[string]interface{} { + return nil +} +func (fakeChatDocEngine) GetHighlight([]map[string]interface{}, []string, string) map[string]string { + return nil +} +func (fakeChatDocEngine) RunSQL(context.Context, string, string, []string, string) ([]map[string]interface{}, error) { + return nil, nil +} +func (fakeChatDocEngine) GetChunkIDs([]map[string]interface{}) []string { return nil } +func (fakeChatDocEngine) KNNScores(context.Context, []map[string]interface{}, []float64, int) (map[string]interface{}, error) { + return nil, nil +} +func (fakeChatDocEngine) GetScores(map[string]interface{}) map[string]float64 { return nil } +func (fakeChatDocEngine) Ping(context.Context) error { return nil } +func (fakeChatDocEngine) Close() error { return nil } +func (fakeChatDocEngine) GetType() string { return "fake" } +func (fakeChatDocEngine) SupportsPageRank() bool { return false } +func (fakeChatDocEngine) FilterDocIdsByMetaPushdown(context.Context, []string, []map[string]interface{}, string) []string { + return nil +} diff --git a/internal/service/dataset/helpers.go b/internal/service/dataset/helpers.go new file mode 100644 index 0000000000..ee9e503e99 --- /dev/null +++ b/internal/service/dataset/helpers.go @@ -0,0 +1,270 @@ +package dataset + +import ( + "encoding/json" + "errors" + "fmt" + "strings" + + "ragflow/internal/dao" + pipelinepkg "ragflow/internal/ingestion/pipeline" + "ragflow/internal/service" + + "github.com/google/uuid" +) + +// Package-level vars and constants used by the dataset service. +var ( + datasetSupportedAvatarMIMETypes = map[string]struct{}{ + "image/jpeg": {}, + "image/png": {}, + } + datasetAllowedOrderByFields = map[string]struct{}{ + "create_time": {}, + "update_time": {}, + } + datasetAllowedMetadataTypes = map[string]struct{}{ + "string": {}, + "list": {}, + "time": {}, + "number": {}, + } + validIndexTypes = []string{"graph", "raptor", "mindmap"} + indexTypeToTaskType = map[string]string{"graph": "graphrag", "raptor": "raptor", "mindmap": "mindmap"} + indexTypeToDisplayName = map[string]string{"graph": "Graph", "raptor": "RAPTOR", "mindmap": "Mindmap"} +) + +const ( + graphRaptorQueueDocID = "graph_raptor_x" + maximumTaskPageNumber = int64(100000000) + serverQueueNamePrefix = "te" + defaultEmbeddingCheckNum = 5 + + graphPhaseResolutionDone = "resolution_done" + graphPhaseCommunityDone = "community_done" +) + +// validateParserID validates parser_id against the built-in pipeline registry. +func validateParserID(chunkMethod string) error { + registry, err := pipelinepkg.DefaultRegistry() + if err != nil || registry == nil { + return errors.New("parser_id validation unavailable: builtin pipeline registry not loaded") + } + if registry.IsValid(chunkMethod) { + return nil + } + return parserIDError() +} + +func parserIDError() error { + registry, err := pipelinepkg.DefaultRegistry() + if err != nil || registry == nil { + return errors.New("invalid parser_id") + } + refs := registry.Refs() + switch len(refs) { + case 0: + return errors.New("invalid parser_id") + case 1: + return fmt.Errorf("Input should be '%s'", refs[0]) + default: + return fmt.Errorf("Input should be %s or '%s'", quoteList(refs[:len(refs)-1]), refs[len(refs)-1]) + } +} + +func quoteList(items []string) string { + quoted := make([]string, len(items)) + for i, v := range items { + quoted[i] = "'" + v + "'" + } + return strings.Join(quoted, ", ") +} + +func validateDatasetAvatar(avatar string) error { + if !strings.Contains(avatar, ",") { + return errors.New("Missing MIME prefix. Expected format: data:;base64,") + } + prefix, _, _ := strings.Cut(avatar, ",") + if !strings.HasPrefix(prefix, "data:") { + return errors.New("Invalid MIME prefix format. Must start with 'data:'") + } + mimeType, _, _ := strings.Cut(strings.TrimPrefix(prefix, "data:"), ";") + if _, ok := datasetSupportedAvatarMIMETypes[mimeType]; !ok { + return errors.New("Unsupported MIME type. Allowed: [image/jpeg image/png]") + } + return nil +} + +func validateDatasetEmbeddingModel(embeddingModel string) error { + if embeddingModel == "" { + return errors.New("Embedding model identifier is required") + } + if !strings.Contains(embeddingModel, "@") { + return nil + } + parts := strings.Split(embeddingModel, "@") + for _, part := range parts { + if strings.TrimSpace(part) == "" { + return errors.New("Both model_name and provider must be non-empty strings") + } + } + if len(parts) < 2 { + return errors.New("Embedding model identifier must follow @ format") + } + if strings.TrimSpace(parts[0]) == "" || strings.TrimSpace(parts[len(parts)-1]) == "" { + return errors.New("Both model_name and provider must be non-empty strings") + } + return nil +} + +func normalizeDatasetPipelineID(pipelineID string) (*string, error) { + pipelineID = strings.TrimSpace(pipelineID) + if pipelineID == "" { + return nil, nil + } + if len(pipelineID) != 32 { + return nil, errors.New("pipeline_id must be 32 hex characters") + } + for _, char := range pipelineID { + if !strings.ContainsRune("0123456789abcdefABCDEF", char) { + return nil, errors.New("pipeline_id must be hexadecimal") + } + } + normalized := strings.ToLower(pipelineID) + return &normalized, nil +} + +func validateDatasetParserConfigSize(parserConfig map[string]interface{}) error { + if len(parserConfig) == 0 { + return nil + } + data, err := json.Marshal(parserConfig) + if err != nil { + return errors.New("parser_config must be valid JSON") + } + if len(data) > 65535 { + return fmt.Errorf("Parser config exceeds size limit (max 65,535 characters). Current size: %d", len(data)) + } + return nil +} + +func normalizeDatasetID(id string) (string, error) { + parsedUUID, err := uuid.Parse(id) + if err != nil { + return "", errors.New("Invalid UUID format") + } + if parsedUUID == (uuid.UUID{}) { + return "", errors.New("Invalid UUID format") + } + return strings.ReplaceAll(parsedUUID.String(), "-", ""), nil +} + +func canvasAccessibleForUser(userID, canvasID string) (bool, error) { + tenantIDs, _ := dao.NewUserTenantDAO().GetTenantIDsByUserID(userID) + return dao.NewUserCanvasDAO().Accessible(canvasID, userID, tenantIDs), nil +} + +func parserConfigValueOrEmptyList(parserConfig map[string]interface{}, key string) interface{} { + if parserConfig == nil { + return []interface{}{} + } + value, ok := parserConfig[key] + if !ok || value == nil { + return []interface{}{} + } + return value +} + +func datasetConnectorsOrEmpty(connectors []*dao.ConnectorDatasetListItem) []*dao.ConnectorDatasetListItem { + if connectors == nil { + return make([]*dao.ConnectorDatasetListItem, 0) + } + return connectors +} + +func datasetUpdateParserID(req service.UpdateDatasetRequest) (string, bool, error) { + parserID := "" + provided := false + if req.ParserID != nil { + parserID = strings.TrimSpace(*req.ParserID) + provided = true + } + if !provided { + return "", false, nil + } + if err := validateParserID(parserID); err != nil { + return "", true, err + } + return parserID, true, nil +} + +func datasetUpdateEmbeddingID(req service.UpdateDatasetRequest) (string, bool, error) { + embdID := "" + provided := false + if req.EmbdID != nil { + embdID = strings.TrimSpace(*req.EmbdID) + provided = true + } + if req.EmbeddingModel != nil { + embdID = strings.TrimSpace(*req.EmbeddingModel) + provided = true + } + if !provided { + return "", false, nil + } + if embdID != "" { + if err := validateDatasetEmbeddingModel(embdID); err != nil { + return "", true, err + } + } + return embdID, true, nil +} + +func normalizeDatasetUpdateExt(ext map[string]interface{}) map[string]interface{} { + if ext == nil { + return nil + } + updates := make(map[string]interface{}, len(ext)) + for key, value := range ext { + switch key { + case "chunk_method": + updates["parser_id"] = value + case "token_num", "chunk_num", "parser_config": + continue + case "pagerank": + if v, ok := value.(float64); ok { + updates[key] = int64(v) + } + default: + updates[key] = value + } + } + return updates +} + +func normalizeMetadataConfigFields(fields []service.MetadataConfigField, fieldName string) ([]map[string]interface{}, error) { + normalizedFields := make([]map[string]interface{}, 0, len(fields)) + for i, field := range fields { + key := strings.TrimSpace(field.Key) + if key == "" { + return nil, fmt.Errorf("%s[%d].key is required", fieldName, i) + } + if len(key) > 255 { + return nil, fmt.Errorf("%s[%d].key should have at most 255 characters", fieldName, i) + } + fieldType := strings.TrimSpace(field.Type) + if _, ok := datasetAllowedMetadataTypes[fieldType]; !ok { + return nil, fmt.Errorf("%s[%d].type should be one of 'string', 'list', 'time' or 'number'", fieldName, i) + } + if field.Description != nil && len(*field.Description) > 65535 { + return nil, fmt.Errorf("%s[%d].description should have at most 65535 characters", fieldName, i) + } + normalizedFields = append(normalizedFields, map[string]interface{}{ + "key": key, + "type": fieldType, + "description": field.Description, + "enum": field.Enum, + }) + } + return normalizedFields, nil +} diff --git a/internal/service/dataset/helpers_test.go b/internal/service/dataset/helpers_test.go new file mode 100644 index 0000000000..8877fdac53 --- /dev/null +++ b/internal/service/dataset/helpers_test.go @@ -0,0 +1,389 @@ +package dataset + +import ( + "strings" + "testing" + + "ragflow/internal/service" +) + +// TestValidateParserID_AcceptsRegistryRefs verifies that every +// canonical builtin pipeline id passes validation. +func TestValidateParserID_AcceptsRegistryRefs(t *testing.T) { + for _, id := range []string{"general", "book", "audio", "qa", "table", "tag"} { + if err := validateParserID(id); err != nil { + t.Errorf("validateParserID(%q) = %v, want nil", id, err) + } + } +} + +// TestValidateParserID_AcceptsNaiveAlias verifies the legacy +// parser_id "naive" still validates (alias for general). +func TestValidateParserID_AcceptsNaiveAlias(t *testing.T) { + if err := validateParserID("naive"); err != nil { + t.Errorf("validateParserID(naive) = %v, want nil (alias for general)", err) + } +} + +// TestValidateParserID_RejectsUnknown verifies unknown/empty +// values are rejected with an error that lists the valid options. +func TestValidateParserID_RejectsUnknown(t *testing.T) { + for _, id := range []string{"", "unknown", "NAIVE"} { + err := validateParserID(id) + if err == nil { + t.Errorf("validateParserID(%q) = nil, want error", id) + } + } + + err := validateParserID("unknown") + if err == nil { + t.Fatal("expected error for unknown parser id") + } + msg := err.Error() + if !strings.Contains(msg, "general") { + t.Errorf("error message %q should mention general", msg) + } +} + +// --- validateDatasetAvatar --- + +func TestValidateDatasetAvatar_MissingPrefix(t *testing.T) { + err := validateDatasetAvatar("iVBORw0KGgo=") + if err == nil { + t.Fatal("expected error for missing MIME prefix") + } +} + +func TestValidateDatasetAvatar_InvalidPrefix(t *testing.T) { + err := validateDatasetAvatar("wrong:image/png;base64,iVBORw0KGgo=") + if err == nil { + t.Fatal("expected error for invalid prefix") + } +} + +func TestValidateDatasetAvatar_UnsupportedMIME(t *testing.T) { + err := validateDatasetAvatar("data:image/gif;base64,iVBORw0KGgo=") + if err == nil { + t.Fatal("expected error for unsupported MIME") + } +} + +func TestValidateDatasetAvatar_Valid(t *testing.T) { + err := validateDatasetAvatar("data:image/png;base64,iVBORw0KGgo=") + if err != nil { + t.Fatalf("expected nil, got %v", err) + } + err = validateDatasetAvatar("data:image/jpeg;base64,/9j/4AAQ==") + if err != nil { + t.Fatalf("expected nil, got %v", err) + } +} + +// --- validateDatasetEmbeddingModel --- + +func TestValidateDatasetEmbeddingModel_Empty(t *testing.T) { + err := validateDatasetEmbeddingModel("") + if err == nil { + t.Fatal("expected error for empty model") + } +} + +func TestValidateDatasetEmbeddingModel_NameOnlyNoProvider(t *testing.T) { + if err := validateDatasetEmbeddingModel("BAAI/bge-large-zh-v1.5"); err != nil { + t.Fatalf("expected nil for name without @, got %v", err) + } +} + +func TestValidateDatasetEmbeddingModel_NameWithProvider(t *testing.T) { + if err := validateDatasetEmbeddingModel("BAAI/bge-large-zh-v1.5@Builtin"); err != nil { + t.Fatalf("expected nil, got %v", err) + } +} + +func TestValidateDatasetEmbeddingModel_EmptyPart(t *testing.T) { + err := validateDatasetEmbeddingModel("model@") + if err == nil { + t.Fatal("expected error for empty provider") + } + err = validateDatasetEmbeddingModel("@provider") + if err == nil { + t.Fatal("expected error for empty model name") + } +} + +// --- normalizeDatasetPipelineID --- + +func TestNormalizeDatasetPipelineID_Empty(t *testing.T) { + result, err := normalizeDatasetPipelineID("") + if err != nil { + t.Fatalf("expected nil, got %v", err) + } + if result != nil { + t.Fatalf("expected nil result for empty input") + } +} + +func TestNormalizeDatasetPipelineID_Spaces(t *testing.T) { + result, err := normalizeDatasetPipelineID(" ") + if err != nil { + t.Fatalf("expected nil, got %v", err) + } + if result != nil { + t.Fatalf("expected nil result for whitespace-only input") + } +} + +func TestNormalizeDatasetPipelineID_WrongLength(t *testing.T) { + _, err := normalizeDatasetPipelineID("abc123") + if err == nil { + t.Fatal("expected error for wrong length") + } +} + +func TestNormalizeDatasetPipelineID_InvalidChars(t *testing.T) { + _, err := normalizeDatasetPipelineID("abcdef01-23456789abcdef0123456789") + if err == nil { + t.Fatal("expected error for non-hex chars") + } +} + +func TestNormalizeDatasetPipelineID_Valid(t *testing.T) { + result, err := normalizeDatasetPipelineID("ABCDEF0123456789ABCDEF0123456789") + if err != nil { + t.Fatalf("expected nil, got %v", err) + } + if result == nil { + t.Fatal("expected non-nil result") + } + if *result != "abcdef0123456789abcdef0123456789" { + t.Errorf("expected lowercased, got %q", *result) + } +} + +// --- validateDatasetParserConfigSize --- + +func TestValidateDatasetParserConfigSize_Empty(t *testing.T) { + if err := validateDatasetParserConfigSize(map[string]interface{}{}); err != nil { + t.Fatalf("expected nil for empty, got %v", err) + } +} + +func TestValidateDatasetParserConfigSize_UnderLimit(t *testing.T) { + cfg := map[string]interface{}{ + "Parser:abc": map[string]interface{}{ + "chunk_size": float64(512), + }, + } + if err := validateDatasetParserConfigSize(cfg); err != nil { + t.Fatalf("expected nil, got %v", err) + } +} + +func TestValidateDatasetParserConfigSize_OverLimit(t *testing.T) { + // Build a config that exceeds 65535 bytes. + bigVal := strings.Repeat("x", 66000) + cfg := map[string]interface{}{ + "Parser:abc": map[string]interface{}{ + "big_field": bigVal, + }, + } + err := validateDatasetParserConfigSize(cfg) + if err == nil { + t.Fatal("expected error for oversized parser_config") + } + if !strings.Contains(err.Error(), "exceeds size limit") { + t.Errorf("unexpected error: %v", err) + } +} + +// --- normalizeDatasetID --- + +func TestNormalizeDatasetID_Invalid(t *testing.T) { + _, err := normalizeDatasetID("not-a-uuid") + if err == nil { + t.Fatal("expected error for invalid UUID") + } + _, err = normalizeDatasetID("") + if err == nil { + t.Fatal("expected error for empty string") + } +} + +func TestNormalizeDatasetID_Valid(t *testing.T) { + raw := "550e8400-e29b-41d4-a716-446655440000" + result, err := normalizeDatasetID(raw) + if err != nil { + t.Fatalf("expected nil, got %v", err) + } + expected := "550e8400e29b41d4a716446655440000" + if result != expected { + t.Errorf("expected %q, got %q", expected, result) + } +} + +func TestNormalizeDatasetID_StripsHyphens(t *testing.T) { + // UUID with hyphens already removed. + raw := "550e8400e29b41d4a716446655440000" + result, err := normalizeDatasetID(raw) + if err != nil { + t.Fatalf("expected nil, got %v", err) + } + if result != raw { + t.Errorf("expected %q, got %q", raw, result) + } +} + +// --- normalizeDatasetUpdateExt --- + +func TestNormalizeDatasetUpdateExt_Nil(t *testing.T) { + if result := normalizeDatasetUpdateExt(nil); result != nil { + t.Fatalf("expected nil for nil input") + } +} + +func TestNormalizeDatasetUpdateExt_PassesThrough(t *testing.T) { + ext := map[string]interface{}{ + "description": "test", + "language": "English", + } + result := normalizeDatasetUpdateExt(ext) + if result["description"] != "test" { + t.Errorf("expected description preserved, got %v", result["description"]) + } + if result["language"] != "English" { + t.Errorf("expected language preserved, got %v", result["language"]) + } +} + +func TestNormalizeDatasetUpdateExt_RenamesChunkMethod(t *testing.T) { + ext := map[string]interface{}{ + "chunk_method": "book", + } + result := normalizeDatasetUpdateExt(ext) + if result["parser_id"] != "book" { + t.Errorf("expected parser_id=book, got %v", result["parser_id"]) + } + if _, ok := result["chunk_method"]; ok { + t.Error("expected chunk_method to be renamed to parser_id") + } +} + +func TestNormalizeDatasetUpdateExt_SkipsTokenAndChunkNum(t *testing.T) { + ext := map[string]interface{}{ + "token_num": float64(1000), + "chunk_num": float64(50), + "parser_config": map[string]interface{}{"key": "val"}, + } + result := normalizeDatasetUpdateExt(ext) + if len(result) != 0 { + t.Errorf("expected empty map, got %v", result) + } +} + +func TestNormalizeDatasetUpdateExt_ConvertsPagerank(t *testing.T) { + ext := map[string]interface{}{ + "pagerank": float64(3), + } + result := normalizeDatasetUpdateExt(ext) + if result["pagerank"] != int64(3) { + t.Errorf("expected pagerank=int64(3), got %T(%v)", result["pagerank"], result["pagerank"]) + } +} + +func TestNormalizeDatasetUpdateExt_NonFloatPagerankSkipped(t *testing.T) { + // Non-float64 pagerank values are not convertible and are dropped. + ext := map[string]interface{}{ + "pagerank": "auto", + } + result := normalizeDatasetUpdateExt(ext) + if _, ok := result["pagerank"]; ok { + t.Error("expected non-float pagerank to be skipped") + } +} + +// --- normalizeMetadataConfigFields --- + +func TestNormalizeMetadataConfigFields_EmptyKey(t *testing.T) { + fields := []service.MetadataConfigField{ + {Key: "", Type: "string"}, + } + _, err := normalizeMetadataConfigFields(fields, "metadata") + if err == nil { + t.Fatal("expected error for empty key") + } +} + +func TestNormalizeMetadataConfigFields_KeyTooLong(t *testing.T) { + longKey := strings.Repeat("k", 256) + fields := []service.MetadataConfigField{ + {Key: longKey, Type: "string"}, + } + _, err := normalizeMetadataConfigFields(fields, "metadata") + if err == nil { + t.Fatal("expected error for too-long key") + } +} + +func TestNormalizeMetadataConfigFields_InvalidType(t *testing.T) { + fields := []service.MetadataConfigField{ + {Key: "my_field", Type: "boolean"}, + } + _, err := normalizeMetadataConfigFields(fields, "metadata") + if err == nil { + t.Fatal("expected error for invalid type") + } +} + +func TestNormalizeMetadataConfigFields_DescriptionTooLong(t *testing.T) { + longDesc := strings.Repeat("d", 65536) + fields := []service.MetadataConfigField{ + {Key: "my_field", Type: "string", Description: &longDesc}, + } + _, err := normalizeMetadataConfigFields(fields, "metadata") + if err == nil { + t.Fatal("expected error for too-long description") + } +} + +func TestNormalizeMetadataConfigFields_Valid(t *testing.T) { + desc := "A description" + fields := []service.MetadataConfigField{ + {Key: "field1", Type: "string", Description: &desc}, + {Key: "field2", Type: "list"}, + } + result, err := normalizeMetadataConfigFields(fields, "metadata") + if err != nil { + t.Fatalf("expected nil, got %v", err) + } + if len(result) != 2 { + t.Fatalf("expected 2 fields, got %d", len(result)) + } + if result[0]["key"] != "field1" { + t.Errorf("expected key=field1, got %v", result[0]["key"]) + } + if result[0]["type"] != "string" { + t.Errorf("expected type=string, got %v", result[0]["type"]) + } + if result[0]["description"] != &desc { + t.Errorf("expected description preserved") + } + if result[1]["key"] != "field2" { + t.Errorf("expected key=field2, got %v", result[1]["key"]) + } + if result[1]["type"] != "list" { + t.Errorf("expected type=list, got %v", result[1]["type"]) + } +} + +func TestNormalizeMetadataConfigFields_TrimsKey(t *testing.T) { + fields := []service.MetadataConfigField{ + {Key: " my_field ", Type: "number"}, + } + result, err := normalizeMetadataConfigFields(fields, "metadata") + if err != nil { + t.Fatalf("expected nil, got %v", err) + } + if result[0]["key"] != "my_field" { + t.Errorf("expected trimmed key 'my_field', got %v", result[0]["key"]) + } +} diff --git a/internal/service/dataset/index.go b/internal/service/dataset/index.go new file mode 100644 index 0000000000..ae92cda43c --- /dev/null +++ b/internal/service/dataset/index.go @@ -0,0 +1,743 @@ +package dataset + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "math/rand" + "sort" + "strings" + "time" + + "ragflow/internal/common" + "ragflow/internal/dao" + redisengine "ragflow/internal/engine/redis" + enginetypes "ragflow/internal/engine/types" + "ragflow/internal/entity" + modelModule "ragflow/internal/entity/models" + "ragflow/internal/service" + "ragflow/internal/utility" + + "github.com/cespare/xxhash/v2" + "go.uber.org/zap" + "gorm.io/gorm" + "gorm.io/gorm/clause" +) + +func checkType(indexType string) bool { + haveType := false + for _, t := range validIndexTypes { + if indexType == t { + haveType = true + } + } + return haveType +} + +func (d *DatasetService) newRaptorOrGraphRagTask(sampleDoc *entity.Document, taskType string, taskDocID string, queueDocID string, docIDs []string) (*entity.Task, map[string]interface{}, error) { + if docIDs == nil || len(docIDs) == 0 { + docIDs = make([]string, 0) + } + if !checkIndexTaskType(taskType) { + return nil, nil, errors.New("type should be graphrag, raptor or mindmap") + } + + chunkingConfig, err := d.documentDAO.GetChunkingConfig(sampleDoc.ID) + if err != nil { + return nil, nil, err + } + + hasher := xxhash.New() + keys := make([]string, 0, len(chunkingConfig)) + for key := range chunkingConfig { + keys = append(keys, key) + } + sort.Strings(keys) + for _, key := range keys { + _, _ = hasher.Write([]byte(key)) + _, _ = hasher.Write([]byte{0}) + v, mErr := json.Marshal(chunkingConfig[key]) + if mErr != nil { + return nil, nil, mErr + } + _, _ = hasher.Write(v) + _, _ = hasher.Write([]byte{0}) + } + + taskID := utility.GenerateUUID() + beginAt := time.Now().Truncate(time.Second) + progressMsg := beginAt.Format("15:04:05") + " created task " + taskType + + for _, field := range []interface{}{taskDocID, maximumTaskPageNumber, maximumTaskPageNumber, taskType} { + _, _ = hasher.Write([]byte(fmt.Sprint(field))) + } + digest := fmt.Sprintf("%016x", hasher.Sum64()) + task := &entity.Task{ + ID: taskID, + DocID: taskDocID, + FromPage: maximumTaskPageNumber, + ToPage: maximumTaskPageNumber, + TaskType: taskType, + ProgressMsg: &progressMsg, + BeginAt: &beginAt, + Digest: &digest, + } + + queueMessage := map[string]interface{}{ + "id": taskID, + "doc_id": queueDocID, + "from_page": maximumTaskPageNumber, + "to_page": maximumTaskPageNumber, + "task_type": taskType, + "progress_msg": progressMsg, + "begin_at": beginAt.Format("2006-01-02 15:04:05"), + "digest": digest, + "doc_ids": docIDs, + } + + return task, queueMessage, nil +} + +func createDatasetIndexTaskInTx(tx *gorm.DB, task *entity.Task, queueDocID string) (*entity.Document, error) { + if task == nil { + return nil, errors.New("task is required") + } + if err := tx.Create(task).Error; err != nil { + return nil, err + } + + if queueDocID == "" { + return nil, nil + } + + var document entity.Document + err := tx.Select("id", "progress_msg", "process_begin_at").Where("id = ?", queueDocID).First(&document).Error + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + return nil, err + } + + beginAt := time.Now().Truncate(time.Second) + if task.BeginAt != nil { + beginAt = *task.BeginAt + } + if err := tx.Model(&entity.Document{}).Where("id = ?", queueDocID).Updates(map[string]interface{}{ + "progress_msg": "Task is queued...", + "process_begin_at": beginAt, + }).Error; err != nil { + return nil, err + } + + return &document, nil +} + +func enqueueDatasetIndexTask(priority int, queueMessage map[string]interface{}) error { + redisClient := redisengine.Get() + if redisClient == nil || !redisClient.QueueProduct(datasetIndexQueueName(priority), queueMessage) { + return errors.New("Can't access Redis. Please check the Redis' status") + } + return nil +} + +func cleanupFailedDatasetIndexTask(taskID string, updatedDocument *entity.Document, kbID string, indexType string) error { + return dao.DB.Transaction(func(tx *gorm.DB) error { + if err := tx.Unscoped().Where("id = ?", taskID).Delete(&entity.Task{}).Error; err != nil { + return fmt.Errorf("delete task %s: %w", taskID, err) + } + + if column := datasetIndexTaskIDColumn(indexType); kbID != "" && column != "" { + if err := tx.Model(&entity.Knowledgebase{}).Where("id = ? AND "+column+" = ?", kbID, taskID).Update(column, nil).Error; err != nil { + return fmt.Errorf("clear dataset task id %s: %w", taskID, err) + } + } + + if updatedDocument == nil { + return nil + } + + return tx.Model(&entity.Document{}).Where("id = ?", updatedDocument.ID).Updates(map[string]interface{}{ + "progress_msg": updatedDocument.ProgressMsg, + "process_begin_at": updatedDocument.ProcessBeginAt, + }).Error + }) +} + +func datasetIndexTaskIDColumn(indexType string) string { + switch indexType { + case "graph": + return "graphrag_task_id" + case "raptor": + return "raptor_task_id" + case "mindmap": + return "mindmap_task_id" + default: + return "" + } +} + +func datasetIndexTaskFinishAtColumn(indexType string) string { + switch indexType { + case "graph": + return "graphrag_task_finish_at" + case "raptor": + return "raptor_task_finish_at" + case "mindmap": + return "mindmap_task_finish_at" + default: + return "" + } +} + +func checkIndexTaskType(taskType string) bool { + switch taskType { + case "graphrag", "raptor", "mindmap": + return true + default: + return false + } +} + +func datasetIndexTaskID(kb *entity.Knowledgebase, indexType string) string { + if kb == nil { + return "" + } + switch indexType { + case "graph": + if kb.GraphragTaskID != nil { + return *kb.GraphragTaskID + } + case "raptor": + if kb.RaptorTaskID != nil { + return *kb.RaptorTaskID + } + case "mindmap": + if kb.MindmapTaskID != nil { + return *kb.MindmapTaskID + } + } + return "" +} + +func datasetIndexTaskIDUpdate(indexType, taskID string) map[string]interface{} { + switch indexType { + case "graph": + return map[string]interface{}{"graphrag_task_id": taskID} + case "raptor": + return map[string]interface{}{"raptor_task_id": taskID} + case "mindmap": + return map[string]interface{}{"mindmap_task_id": taskID} + default: + return map[string]interface{}{} + } +} + +func datasetIndexTaskIDs(kb *entity.Knowledgebase) []string { + if kb == nil { + return nil + } + taskIDs := make([]string, 0, 3) + for _, taskID := range []*string{kb.GraphragTaskID, kb.RaptorTaskID, kb.MindmapTaskID} { + if taskID != nil && *taskID != "" { + taskIDs = append(taskIDs, *taskID) + } + } + return common.Deduplicate(taskIDs) +} + +func datasetIndexQueueName(priority int) string { + return fmt.Sprintf("%s.%d.common", serverQueueNamePrefix, priority) +} + +func clearGraphPhaseMarkers(redisClient *redisengine.Client, datasetID string) { + if redisClient == nil || datasetID == "" { + return + } + for _, phase := range []string{graphPhaseResolutionDone, graphPhaseCommunityDone} { + if !redisClient.Delete(fmt.Sprintf("graphrag:phase:%s:%s", datasetID, phase)) { + common.Warn("Failed to clear GraphRAG phase marker", zap.String("dataset_id", datasetID), zap.String("phase", phase)) + } + } +} + +func (d *DatasetService) RunIndex(userID, datasetID, indexType string) (map[string]interface{}, common.ErrorCode, error) { + if !checkType(indexType) { + return nil, common.CodeDataError, fmt.Errorf("Invalid index type '%s'. Must be one of %v", indexType, validIndexTypes) + } + + if datasetID == "" { + return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) + } + if !d.kbDAO.Accessible(datasetID, userID) { + return nil, common.CodeDataError, errors.New("No authorization.") + } + + kb, err := d.kbDAO.GetByID(datasetID) + if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, errors.New("Invalid Dataset ID") + } + return nil, common.CodeDataError, errors.New("Internal server error") + } + + taskType := indexTypeToTaskType[indexType] + displayName := indexTypeToDisplayName[indexType] + + documents, code, err := d.getDocumentsByDatasetForIndex(datasetID) + if err != nil { + return nil, code, err + } + _ = documents + + sampleDocument := documents[0] + documentIDs := make([]string, len(documents)) + + for i, doc := range documents { + documentIDs[i] = doc.ID + } + + task, queueMessage, err := d.newRaptorOrGraphRagTask(sampleDocument, taskType, sampleDocument.ID, graphRaptorQueueDocID, documentIDs) + if err != nil { + common.Warn("Failed to build dataset index task", zap.String("dataset_id", datasetID), zap.String("task_type", taskType), zap.Error(err)) + return nil, common.CodeDataError, errors.New("Internal server error") + } + + var updatedDocument *entity.Document + var dataErr error + err = dao.DB.Transaction(func(tx *gorm.DB) error { + var lockedKB entity.Knowledgebase + if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). + Where("id = ? AND status = ?", kb.ID, string(entity.StatusValid)). + First(&lockedKB).Error; err != nil { + return err + } + + existingTaskID := datasetIndexTaskID(&lockedKB, indexType) + if existingTaskID != "" { + var existingTask entity.Task + taskErr := tx.Where("id = ?", existingTaskID).First(&existingTask).Error + if taskErr != nil { + if errors.Is(taskErr, gorm.ErrRecordNotFound) { + } else { + return taskErr + } + } else if existingTask.Progress != 1 && existingTask.Progress != -1 { + dataErr = fmt.Errorf("Task %s in progress with status %v. A %s Task is already running.", existingTaskID, existingTask.Progress, displayName) + return dataErr + } + } + + updatedDocument, err = createDatasetIndexTaskInTx(tx, task, graphRaptorQueueDocID) + if err != nil { + return err + } + return tx.Model(&entity.Knowledgebase{}).Where("id = ?", lockedKB.ID).Updates(datasetIndexTaskIDUpdate(indexType, task.ID)).Error + }) + if err != nil { + if dataErr != nil { + return nil, common.CodeDataError, dataErr + } + common.Warn("Failed to create dataset index task", zap.String("dataset_id", datasetID), zap.String("task_type", taskType), zap.Error(err)) + return nil, common.CodeDataError, errors.New("Internal server error") + } + + if err := enqueueDatasetIndexTask(0, queueMessage); err != nil { + if cleanupErr := cleanupFailedDatasetIndexTask(task.ID, updatedDocument, kb.ID, indexType); cleanupErr != nil { + err = errors.Join(err, cleanupErr) + } + common.Warn("Failed to queue dataset index task", zap.String("dataset_id", datasetID), zap.String("task_type", taskType), zap.Error(err)) + return nil, common.CodeDataError, errors.New("Internal server error") + } + + return map[string]interface{}{"task_id": task.ID}, common.CodeSuccess, nil +} + +func (d *DatasetService) getDocumentsByDatasetForIndex(datasetID string) ([]*entity.Document, common.ErrorCode, error) { + documents, _, err := d.documentDAO.GetByKBID(datasetID) + if err != nil { + common.Warn("Failed to load dataset documents for index", zap.String("dataset_id", datasetID), zap.Error(err)) + return nil, common.CodeDataError, errors.New("Internal server error") + } + if len(documents) == 0 { + return nil, common.CodeDataError, fmt.Errorf("No documents in Dataset %s", datasetID) + } + return documents, common.CodeSuccess, nil +} + +func (d *DatasetService) TraceIndex(datasetID, userID, indexType string) (*entity.Task, common.ErrorCode, error) { + if !checkType(indexType) { + return nil, common.CodeDataError, fmt.Errorf("Invalid index type '%s'. Must be one of %v", indexType, validIndexTypes) + } + + if datasetID == "" { + return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) + } + if !d.kbDAO.Accessible(datasetID, userID) { + return nil, common.CodeDataError, errors.New("No authorization.") + } + + kb, err := d.kbDAO.GetByID(datasetID) + if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, errors.New("Invalid Dataset ID") + } + return nil, common.CodeDataError, errors.New("Internal server error") + } + + taskID := datasetIndexTaskID(kb, indexType) + + var task *entity.Task + if taskID != "" { + task, err = d.taskDAO.GetByID(taskID) + if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeSuccess, nil + } + return nil, common.CodeServerError, errors.New("Internal server error") + } + if task == nil { + return nil, common.CodeSuccess, nil + } + } + + return task, common.CodeSuccess, nil +} + +type embeddingCheckSample struct { + ChunkID string + KbID string + DocID string + DocName string + VectorField string + Vector []float64 + PageNum interface{} + Position interface{} + Top interface{} + ContentWithWeight string + QuestionKeywords []string +} + +func (d *DatasetService) CheckEmbedding(userID, datasetID string, req *service.CheckEmbeddingRequest) (*service.EmbeddingCheckResponse, common.ErrorCode, error) { + if datasetID == "" { + return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) + } + if !d.kbDAO.Accessible(datasetID, userID) { + return nil, common.CodeDataError, errors.New("No authorization.") + } + + kb, err := d.kbDAO.GetByID(datasetID) + if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, errors.New("Invalid Dataset ID") + } + return nil, common.CodeServerError, errors.New("Internal server error") + } + + if req == nil || strings.TrimSpace(req.EmbeddingID) == "" { + return nil, common.CodeDataError, errors.New("`embd_id` is required.") + } + embeddingID := strings.TrimSpace(req.EmbeddingID) + if ok, message := d.verifyEmbeddingAvailability(embeddingID, userID); !ok { + return nil, common.CodeDataError, errors.New(message) + } + if d.docEngine == nil { + return nil, common.CodeServerError, errors.New("doc engine not initialized") + } + + driver, modelName, apiConfig, maxTokens, err := service.NewModelProviderService().ResolveModelConfig(kb.TenantID, entity.ModelTypeEmbedding, embeddingID) + if err != nil { + return nil, common.CodeDataError, err + } + embeddingModel := modelModule.NewEmbeddingModel(driver, &modelName, apiConfig, maxTokens) + + checkNum := defaultEmbeddingCheckNum + if req.CheckNum != nil { + checkNum = *req.CheckNum + } + if checkNum <= 0 { + checkNum = defaultEmbeddingCheckNum + } + + samples, err := d.sampleRandomChunksWithVectors(context.Background(), kb.TenantID, datasetID, checkNum) + if err != nil { + return nil, common.CodeServerError, err + } + if len(samples) == 0 { + return &service.EmbeddingCheckResponse{ + Summary: datasetEmbeddingCheckSummary(datasetID, embeddingID, 0, nil, ""), + Results: nil, + }, common.CodeSuccess, nil + } + + results := make([]service.EmbeddingCheckResult, 0, len(samples)) + effectiveSimilarities := make([]float64, 0, len(samples)) + matchMode := "content_only" + for _, sample := range samples { + if sample.Vector == nil || len(sample.Vector) == 0 { + continue + } + + rawChunk, err := d.docEngine.GetChunk(context.Background(), fmt.Sprintf("ragflow_%s", kb.TenantID), sample.ChunkID, []string{datasetID}) + if err != nil { + continue + } + chunkMap := datasetMap(rawChunk) + if len(chunkMap) == 0 { + continue + } + + title := datasetString(chunkMap["title_tks"]) + content := datasetString(chunkMap["content_ltks"]) + + var titleVector [][]float64 + if title != "" { + titleVector, err = datasetEncodeEmbedding(embeddingModel, []string{title}) + if err != nil { + return nil, common.CodeServerError, err + } + } + var contentVector [][]float64 + if content != "" { + contentVector, err = datasetEncodeEmbedding(embeddingModel, []string{content}) + if err != nil { + return nil, common.CodeServerError, err + } + } + + var vectors [][]float64 + if len(titleVector) > 0 && len(contentVector) > 0 { + vectors = [][]float64{titleVector[0], contentVector[0]} + matchMode = "title_and_content" + } else if len(titleVector) > 0 { + vectors = titleVector + } else if len(contentVector) > 0 { + vectors = contentVector + } else { + continue + } + + if len(vectors[0]) != len(sample.Vector) { + return nil, common.CodeDataError, fmt.Errorf("Embedding failure. The dimension (%d) of given embedding model is different from the original (%d)", len(vectors[0]), len(sample.Vector)) + } + + var sim float64 + if len(vectors) == 2 { + simContent := datasetCosSim(vectors[1], sample.Vector) + simMix := datasetCosSim(datasetMixVectors(vectors[0], vectors[1], 0.1), sample.Vector) + sim = simContent + if simMix > sim { + sim = simMix + matchMode = "title+content" + } + } else { + sim = datasetCosSim(vectors[0], sample.Vector) + } + sim = datasetRoundFloat(sim, 6) + + effectiveSimilarities = append(effectiveSimilarities, sim) + results = append(results, service.EmbeddingCheckResult{ + ChunkID: sample.ChunkID, + DocID: sample.DocID, + DocName: sample.DocName, + VectorField: sample.VectorField, + VectorDim: len(sample.Vector), + CosSim: sim, + }) + } + + summary := datasetEmbeddingCheckSummary(datasetID, embeddingID, len(samples), effectiveSimilarities, matchMode) + response := &service.EmbeddingCheckResponse{Summary: summary, Results: results} + if len(effectiveSimilarities) == 0 { + return nil, common.CodeDataError, errors.New("No embedded chunks are available to compare.") + } + if summary.AvgCosSim >= 0.9 { + return response, common.CodeSuccess, nil + } + return response, common.CodeNotEffective, errors.New("Embedding model switch failed: the average similarity between old and new vectors is below 0.9, indicating incompatible vector spaces.") +} + +func (d *DatasetService) sampleRandomChunksWithVectors(ctx context.Context, tenantID, datasetID string, n int) ([]embeddingCheckSample, error) { + indexName := fmt.Sprintf("ragflow_%s", tenantID) + totalResult, err := d.docEngine.Search(ctx, &enginetypes.SearchRequest{ + IndexNames: []string{indexName}, + KbIDs: []string{datasetID}, + Offset: 0, + Limit: 1, + Filter: map[string]interface{}{ + "kb_id": datasetID, + "available_int": 1, + }, + }) + if err != nil { + return nil, err + } + if totalResult == nil || totalResult.Total <= 0 { + return []embeddingCheckSample{}, nil + } + + total := int(totalResult.Total) + const maxEmbeddingSamples = 1024 + if n < 0 { + return nil, fmt.Errorf("invalid sample size: %d", n) + } + if n > maxEmbeddingSamples { + n = maxEmbeddingSamples + } + if n > total { + n = total + } + limit := total + if limit > 1000 { + limit = 1000 + } + if n > limit { + n = limit + } + offsets := rand.Perm(limit) + offsets = offsets[:n] + sort.Ints(offsets) + + baseFields := []string{"docnm_kwd", "doc_id", "content_with_weight", "page_num_int", "position_int", "top_int"} + samples := make([]embeddingCheckSample, 0, n) + for _, offset := range offsets { + searchResult, err := d.docEngine.Search(ctx, &enginetypes.SearchRequest{ + IndexNames: []string{indexName}, + KbIDs: []string{datasetID}, + Offset: offset, + Limit: 1, + SelectFields: baseFields, + Filter: map[string]interface{}{ + "kb_id": datasetID, + "available_int": 1, + }, + }) + if err != nil { + return nil, err + } + if searchResult == nil || len(searchResult.Chunks) == 0 { + continue + } + chunkID := datasetChunkID(searchResult.Chunks[0]) + if chunkID == "" { + continue + } + fullChunk, err := d.docEngine.GetChunk(ctx, indexName, chunkID, []string{datasetID}) + if err != nil { + return nil, err + } + chunkMap := datasetMap(fullChunk) + if len(chunkMap) == 0 { + continue + } + vectorField := datasetGuessVecField(chunkMap) + vector := datasetAsFloatVec(chunkMap[vectorField]) + samples = append(samples, embeddingCheckSample{ + ChunkID: chunkID, + KbID: datasetID, + DocID: datasetString(chunkMap["doc_id"]), + DocName: datasetString(chunkMap["docnm_kwd"]), + VectorField: vectorField, + Vector: vector, + PageNum: chunkMap["page_num_int"], + Position: chunkMap["position_int"], + Top: chunkMap["top_int"], + ContentWithWeight: datasetString(chunkMap["content_with_weight"]), + QuestionKeywords: datasetStringSlice(chunkMap["question_keywords"]), + }) + } + + if len(samples) == 0 { + return nil, errors.New("no valid chunks with vectors found") + } + return samples, nil +} + +func (d *DatasetService) verifyEmbeddingAvailability(embdID string, tenantID string) (bool, string) { + _, _, _, _, err := service.NewModelProviderService().ResolveModelConfig(tenantID, entity.ModelTypeEmbedding, embdID) + if err != nil { + return false, err.Error() + } + return true, "" +} + +func (d *DatasetService) DeleteIndex(userID, datasetID, indexType string, wipe bool) (common.ErrorCode, error) { + if !checkType(indexType) { + return common.CodeArgumentError, fmt.Errorf("Invalid index type '%s'", indexType) + } + + if datasetID == "" { + return common.CodeDataError, errors.New(`Lack of "Dataset ID"`) + } + + if !d.kbDAO.Accessible(datasetID, userID) { + return common.CodeDataError, errors.New("No authorization.") + } + + kb, err := d.kbDAO.GetByID(datasetID) + if err != nil { + if dao.IsNotFoundErr(err) { + return common.CodeDataError, errors.New("Invalid Dataset ID") + } + return common.CodeDataError, errors.New("Internal server error") + } + + taskFinishAtField := datasetIndexTaskFinishAtColumn(indexType) + taskID := datasetIndexTaskID(kb, indexType) + + common.Info("delete_index", zap.String("dataset_id", datasetID), zap.String("index_type", indexType), zap.Bool("wipe", wipe)) + + if taskID != "" { + redisClient := redisengine.Get() + if redisClient == nil || !redisClient.Set(fmt.Sprintf("%s-cancel", taskID), "x", 0) { + common.Warn("Failed to set dataset index cancellation marker", zap.String("dataset_id", datasetID), zap.String("task_id", taskID)) + } + if err := dao.DB.Unscoped().Where("id = ?", taskID).Delete(&entity.Task{}).Error; err != nil { + common.Warn("Failed to delete dataset index task", zap.String("dataset_id", datasetID), zap.String("task_id", taskID), zap.Error(err)) + return common.CodeDataError, errors.New("Internal server error") + } + } + + if wipe && indexType == "graph" { + if d.docEngine == nil { + return common.CodeServerError, errors.New("Document engine is not initialized") + } + indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) + _, err = d.docEngine.DeleteChunks(context.Background(), map[string]interface{}{ + "knowledge_graph_kwd": []interface{}{"graph", "subgraph", "entity", "relation", "community_report"}, + "kb_id": datasetID, + }, indexName, datasetID) + if err != nil { + common.Warn("Failed to delete GraphRAG artefacts", zap.String("dataset_id", datasetID), zap.Error(err)) + return common.CodeDataError, errors.New("Internal server error") + } + clearGraphPhaseMarkers(redisengine.Get(), datasetID) + common.Info("delete_index: cleared GraphRAG artefacts and phase markers", zap.String("dataset_id", datasetID)) + } else if wipe && indexType == "raptor" { + if d.docEngine == nil { + return common.CodeServerError, errors.New("Document engine is not initialized") + } + indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) + _, err = d.docEngine.DeleteChunks(context.Background(), map[string]interface{}{ + "raptor_kwd": []interface{}{"raptor"}, + "kb_id": datasetID, + }, indexName, datasetID) + if err != nil { + common.Warn("Failed to delete RAPTOR artefacts", zap.String("dataset_id", datasetID), zap.Error(err)) + return common.CodeDataError, errors.New("Internal server error") + } + } + + updates := datasetIndexTaskIDUpdate(indexType, "") + if taskFinishAtField != "" { + updates[taskFinishAtField] = nil + } + if len(updates) > 0 { + if err := d.kbDAO.UpdateByID(kb.ID, updates); err != nil { + common.Warn("Failed to clear KB index task refs", zap.String("dataset_id", datasetID), zap.Error(err)) + } + } + + return common.CodeSuccess, nil +} diff --git a/internal/service/dataset_delete_index_test.go b/internal/service/dataset/index_delete_test.go similarity index 99% rename from internal/service/dataset_delete_index_test.go rename to internal/service/dataset/index_delete_test.go index 7ab1325f64..e63b1b36f0 100644 --- a/internal/service/dataset_delete_index_test.go +++ b/internal/service/dataset/index_delete_test.go @@ -14,7 +14,7 @@ // limitations under the License. // -package service +package dataset import ( "context" diff --git a/internal/service/dataset/ingestion.go b/internal/service/dataset/ingestion.go new file mode 100644 index 0000000000..c6a327ab85 --- /dev/null +++ b/internal/service/dataset/ingestion.go @@ -0,0 +1,167 @@ +package dataset + +import ( + "errors" + "fmt" + + "ragflow/internal/common" + "ragflow/internal/dao" + "ragflow/internal/entity" +) + +func (d *DatasetService) GetIngestionSummary(datasetID, userID string) (map[string]interface{}, common.ErrorCode, error) { + if datasetID == "" { + return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) + } + if !d.kbDAO.Accessible(datasetID, userID) { + return nil, common.CodeDataError, errors.New("No authorization.") + } + + kb, err := d.kbDAO.GetByID(datasetID) + if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, fmt.Errorf("Invalid Dataset ID '%s'", datasetID) + } + return nil, common.CodeServerError, errors.New("Database operation failed") + } + + status, err := d.documentDAO.GetParsingStatusByKBID(datasetID) + if err != nil { + return nil, common.CodeServerError, errors.New("Database operation failed") + } + + return map[string]interface{}{ + "doc_num": kb.DocNum, + "chunk_num": kb.ChunkNum, + "token_num": kb.TokenNum, + "status": status, + }, common.CodeSuccess, nil +} + +func (d *DatasetService) ListIngestionLogs(datasetID, userID string, page, pageSize int, orderby string, desc bool, operationStatus []string, createDateFrom, createDateTo, logType, keywords string) (map[string]interface{}, common.ErrorCode, error) { + if datasetID == "" { + return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) + } + if !d.kbDAO.Accessible(datasetID, userID) { + return nil, common.CodeDataError, errors.New("No authorization.") + } + + if page <= 0 { + page = 1 + } + if pageSize <= 0 { + pageSize = 30 + } + if orderby == "" { + orderby = "create_time" + } + + var ( + logs []*entity.PipelineOperationLog + total int64 + err error + ) + if logType == "file" { + logs, total, err = d.pipelineLogDAO.GetFileLogsByKBID(datasetID, page, pageSize, orderby, desc, keywords, operationStatus, createDateFrom, createDateTo) + } else { + logs, total, err = d.pipelineLogDAO.GetDatasetLogsByKBID(datasetID, page, pageSize, orderby, desc, operationStatus, createDateFrom, createDateTo, keywords) + } + if err != nil { + return nil, common.CodeServerError, fmt.Errorf("list ingestion logs: %w", err) + } + + items := make([]map[string]interface{}, 0, len(logs)) + for _, log := range logs { + if log == nil { + continue + } + if logType == "file" { + items = append(items, fileIngestionLogToMap(log)) + } else { + items = append(items, datasetIngestionLogToMap(log)) + } + } + + return map[string]interface{}{ + "total": total, + "logs": items, + }, common.CodeSuccess, nil +} + +func (d *DatasetService) GetIngestionLog(datasetID, userID, logID string) (map[string]interface{}, common.ErrorCode, error) { + if datasetID == "" { + return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) + } + if logID == "" { + return nil, common.CodeDataError, errors.New(`Lack of "Log ID"`) + } + if !d.kbDAO.Accessible(datasetID, userID) { + return nil, common.CodeDataError, errors.New("No authorization.") + } + + log, err := d.pipelineLogDAO.GetByIDAndKBID(logID, datasetID) + if err != nil { + return nil, common.CodeServerError, fmt.Errorf("get ingestion log: %w", err) + } + if log == nil { + return nil, common.CodeDataError, errors.New("Log not found") + } + + return datasetIngestionLogToMap(log), common.CodeSuccess, nil +} + +func datasetIngestionLogToMap(log *entity.PipelineOperationLog) map[string]interface{} { + m := map[string]interface{}{ + "id": log.ID, + "dataset_id": log.KbID, + "tenant_id": log.TenantID, + "document_id": log.DocumentID, + "document_name": log.DocumentName, + "document_suffix": log.DocumentSuffix, + "source_from": log.SourceFrom, + "task_type": log.TaskType, + "operation_status": log.OperationStatus, + "progress": log.Progress, + "create_time": log.CreateTime, + "update_time": log.UpdateTime, + } + if log.PipelineID != nil { + m["pipeline_id"] = *log.PipelineID + } + if log.ProgressMsg != nil { + m["progress_msg"] = *log.ProgressMsg + } + if log.Status != nil { + m["status"] = *log.Status + } + return m +} + +func fileIngestionLogToMap(log *entity.PipelineOperationLog) map[string]interface{} { + return map[string]interface{}{ + "id": log.ID, + "document_id": log.DocumentID, + "tenant_id": log.TenantID, + "kb_id": log.KbID, + "pipeline_id": stringPointerValue(log.PipelineID), + "pipeline_title": stringPointerValue(log.PipelineTitle), + "parser_id": log.ParserID, + "document_name": log.DocumentName, + "document_suffix": log.DocumentSuffix, + "document_type": log.DocumentType, + "source_from": log.SourceFrom, + "progress": log.Progress, + "progress_msg": stringPointerValue(log.ProgressMsg), + "process_begin_at": timePointerValue(log.ProcessBeginAt), + "process_duration": log.ProcessDuration, + "dsl": jsonMapValue(log.DSL), + "task_type": log.TaskType, + "operation_status": log.OperationStatus, + "avatar": stringPointerValue(log.Avatar), + "status": stringPointerValue(log.Status), + "create_time": int64PointerValue(log.CreateTime), + "create_date": timePointerValue(log.CreateDate), + "update_time": int64PointerValue(log.UpdateTime), + "update_date": timePointerValue(log.UpdateDate), + } +} diff --git a/internal/service/dataset/metadata.go b/internal/service/dataset/metadata.go new file mode 100644 index 0000000000..c5744c58a6 --- /dev/null +++ b/internal/service/dataset/metadata.go @@ -0,0 +1,116 @@ +package dataset + +import ( + "errors" + "fmt" + "strings" + + "ragflow/internal/common" + "ragflow/internal/dao" + "ragflow/internal/entity" + "ragflow/internal/service" +) + +// UpdateDocumentMetadataConfig updates the metadata config for a document in a dataset. +func (d *DatasetService) UpdateDocumentMetadataConfig(userID, datasetID, documentID string, req map[string]interface{}) (*entity.Document, common.ErrorCode, error) { + if _, err := d.kbDAO.GetByIDAndTenantID(datasetID, userID); err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, errors.New("You don't own the dataset.") + } + return nil, common.CodeServerError, errors.New("Database operation failed") + } + + doc, err := d.documentDAO.GetByDocumentIDAndDatasetID(documentID, datasetID) + if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, fmt.Errorf("Document %s not found in dataset %s", documentID, datasetID) + } + return nil, common.CodeServerError, errors.New("Database operation failed") + } + + metadata, ok := req["metadata"] + if !ok { + return nil, common.CodeArgumentError, errors.New("metadata is required") + } + + parserConfig := doc.ParserConfig + if parserConfig == nil { + parserConfig = entity.JSONMap{} + } + parserConfig["metadata"] = metadata + + if err = d.documentDAO.UpdateByID(doc.ID, map[string]interface{}{"parser_config": parserConfig}); err != nil { + return nil, common.CodeServerError, errors.New("Database operation failed") + } + + doc, err = d.documentDAO.GetByID(doc.ID) + if err != nil { + return nil, common.CodeServerError, errors.New("Database operation failed") + } + return doc, common.CodeSuccess, nil +} + +// GetMetadataConfig gets the auto-metadata configuration for a dataset. +func (d *DatasetService) GetMetadataConfig(datasetID, tenantID string) (map[string]interface{}, common.ErrorCode, error) { + kb, err := d.kbDAO.GetByIDAndTenantID(datasetID, tenantID) + if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", tenantID, datasetID) + } + return nil, common.CodeServerError, errors.New("Database operation failed") + } + if kb == nil { + return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", tenantID, datasetID) + } + + return map[string]interface{}{ + "metadata": parserConfigValueOrEmptyList(kb.ParserConfig, "metadata"), + "built_in_metadata": parserConfigValueOrEmptyList(kb.ParserConfig, "built_in_metadata"), + }, common.CodeSuccess, nil +} + +// UpdateMetadataConfig updates the auto-metadata configuration for a dataset. +func (d *DatasetService) UpdateMetadataConfig(datasetID, tenantID string, req *service.MetadataConfigRequest) (map[string]interface{}, common.ErrorCode, error) { + datasetID = strings.TrimSpace(datasetID) + tenantID = strings.TrimSpace(tenantID) + + kb, err := d.kbDAO.GetByIDAndTenantID(datasetID, tenantID) + if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", tenantID, datasetID) + } + return nil, common.CodeServerError, errors.New("Database operation failed") + } + if kb == nil { + return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", tenantID, datasetID) + } + + if req == nil { + req = &service.MetadataConfigRequest{} + } + + metadata, err := normalizeMetadataConfigFields(req.Metadata, "metadata") + if err != nil { + return nil, common.CodeDataError, err + } + builtInMetadata, err := normalizeMetadataConfigFields(req.BuiltInMetadata, "built_in_metadata") + if err != nil { + return nil, common.CodeDataError, err + } + + parserConfig := kb.ParserConfig + if parserConfig == nil { + parserConfig = entity.JSONMap{} + } + parserConfig["metadata"] = metadata + parserConfig["built_in_metadata"] = builtInMetadata + + if err = d.kbDAO.UpdateByID(kb.ID, map[string]interface{}{"parser_config": parserConfig}); err != nil { + return nil, common.CodeServerError, errors.New("Update auto-metadata error.(Database error)") + } + + return map[string]interface{}{ + "metadata": metadata, + "built_in_metadata": builtInMetadata, + }, common.CodeSuccess, nil +} diff --git a/internal/service/dataset_document_metadata_config_test.go b/internal/service/dataset/metadata_config_test.go similarity index 99% rename from internal/service/dataset_document_metadata_config_test.go rename to internal/service/dataset/metadata_config_test.go index d0dd0a84ea..129471d079 100644 --- a/internal/service/dataset_document_metadata_config_test.go +++ b/internal/service/dataset/metadata_config_test.go @@ -14,7 +14,7 @@ // limitations under the License. // -package service +package dataset import ( "testing" diff --git a/internal/service/dataset/permission.go b/internal/service/dataset/permission.go new file mode 100644 index 0000000000..55dab64063 --- /dev/null +++ b/internal/service/dataset/permission.go @@ -0,0 +1,60 @@ +package dataset + +import ( + "errors" + "strings" + + "ragflow/internal/entity" +) + +// Accessible checks if a user has access to a dataset. +func (d *DatasetService) Accessible(kbID, userID string) bool { + return d.kbDAO.Accessible(kbID, userID) +} + +// GetByID retrieves a knowledge base by ID. +func (d *DatasetService) GetByID(kbID string) (*entity.Knowledgebase, error) { + return d.kbDAO.GetByID(kbID) +} + +// GetKnowledgebaseByID resolves a dataset entity without applying permission +// checks. Upload needs the same existence-then-auth ordering as Python. +func (d *DatasetService) GetKnowledgebaseByID(datasetID string) (*entity.Knowledgebase, error) { + datasetID = strings.TrimSpace(datasetID) + if datasetID == "" { + return nil, errors.New("Lack of \"Dataset ID\"") + } + normalizedID, err := normalizeDatasetID(datasetID) + if err != nil { + return nil, err + } + return d.kbDAO.GetByID(normalizedID) +} + +// CheckKBTeamPermission checks if a user has team-level permission for the KB. +func (d *DatasetService) CheckKBTeamPermission(kb *entity.Knowledgebase, userID string) bool { + if kb == nil { + return false + } + if kb.TenantID == userID { + return true + } + if kb.Permission != string(entity.TenantPermissionTeam) { + return false + } + joinedTenants, err := d.tenantDAO.GetJoinedTenantsByUserID(userID) + if err != nil { + return false + } + for _, jt := range joinedTenants { + if jt != nil && jt.TenantID == kb.TenantID { + return true + } + } + return false +} + +// GetFieldMap returns the field map for the given knowledge base IDs. +func (d *DatasetService) GetFieldMap(ids []string) (map[string]interface{}, error) { + return d.kbDAO.GetFieldMap(ids) +} diff --git a/internal/service/dataset/search.go b/internal/service/dataset/search.go new file mode 100644 index 0000000000..5811b8a58e --- /dev/null +++ b/internal/service/dataset/search.go @@ -0,0 +1,292 @@ +package dataset + +import ( + "context" + "fmt" + + "go.uber.org/zap" + + "ragflow/internal/common" + "ragflow/internal/entity" + modelModule "ragflow/internal/entity/models" + "ragflow/internal/service" + "ragflow/internal/service/nlp" +) + +func (d *DatasetService) SearchDataset(datasetID, userID string, req *service.SearchDatasetRequest) (*service.SearchDatasetsResponse, error) { + if datasetID == "" { + return nil, fmt.Errorf("dataset_id is required") + } + return d.SearchDatasets(req.ToSearchDatasetsRequest(datasetID), userID) +} + +func (d *DatasetService) SearchDatasets(req *service.SearchDatasetsRequest, userID string) (*service.SearchDatasetsResponse, error) { + if req.Question == "" { + return nil, fmt.Errorf("question is required") + } + if len(req.DatasetIDs) == 0 { + return nil, fmt.Errorf("dataset_ids is required") + } + common.Info("SearchDatasets started", zap.String("userID", userID), zap.Any("datasets", req.DatasetIDs), zap.String("question", req.Question)) + + page := 1 + if req.Page != nil { + page = *req.Page + } + pageSize := 30 + if req.Size != nil { + pageSize = *req.Size + } + useKG := false + if req.UseKG != nil { + useKG = *req.UseKG + } + similarityThreshold := 0.0 + if req.SimilarityThreshold != nil { + similarityThreshold = *req.SimilarityThreshold + } + vectorSimilarityWeight := 0.3 + if req.VectorSimilarityWeight != nil { + vectorSimilarityWeight = *req.VectorSimilarityWeight + } + topK := 1024 + if req.TopK != nil { + topK = *req.TopK + } + if topK < 1 { + topK = 1 + } else if topK > 2048 { + topK = 2048 + } + keyword := false + if req.Keyword != nil { + keyword = *req.Keyword + } + searchID := "" + if req.SearchID != nil { + searchID = *req.SearchID + } + rerankID := "" + if req.RerankID != nil { + rerankID = *req.RerankID + } + + question := req.Question + datasetIDs := req.DatasetIDs + metadataFilter := req.MetadataFilter + crossLanguages := req.CrossLanguages + + ctx := context.Background() + modelProviderSvc := service.NewModelProviderService() + + // Access check for all datasets + var tenantIDs []string + var kbRecords []*entity.Knowledgebase + seenTenants := make(map[string]bool) + for _, datasetID := range datasetIDs { + if !d.kbDAO.Accessible(datasetID, userID) { + common.Warn("SearchDatasets access denied", zap.String("datasetID", datasetID), zap.String("userID", userID)) + return nil, fmt.Errorf("only owner of dataset %s is authorized for this operation", datasetID) + } + + kb, err := d.kbDAO.GetByID(datasetID) + if err != nil || kb == nil { + common.Warn("SearchDatasets dataset not found", zap.String("datasetID", datasetID)) + return nil, fmt.Errorf("dataset %s not found", datasetID) + } + if !seenTenants[kb.TenantID] { + seenTenants[kb.TenantID] = true + tenantIDs = append(tenantIDs, kb.TenantID) + } + kbRecords = append(kbRecords, kb) + } + + // Check if all kbs have the same embedding model + if err := service.ValidateDatasetEmbeddingModels(kbRecords); err != nil { + return nil, err + } + + // Override request fields with values from saved search config + var chatID string + if searchID != "" { + if d.searchService == nil { + common.Warn("Search service is not initialized for search_id", zap.String("searchID", searchID)) + return nil, fmt.Errorf("Invalid search_id") + } + searchDetail, err := d.searchService.GetDetail(searchID) + if err != nil || searchDetail == nil || len(searchDetail) == 0 { + common.Warn("Invalid search_id", zap.String("searchID", searchID), zap.Error(err)) + return nil, fmt.Errorf("Invalid search_id") + } else if searchConfig, ok := searchDetail["search_config"].(map[string]interface{}); ok && searchConfig != nil { + if scMetadataFilter, ok := searchConfig["meta_data_filter"].(map[string]interface{}); ok { + metadataFilter = scMetadataFilter + } + if scST, ok := searchConfig["similarity_threshold"].(float64); ok { + similarityThreshold = scST + } + if scVSW, ok := searchConfig["vector_similarity_weight"].(float64); ok { + vectorSimilarityWeight = scVSW + } + if scTopK, ok := searchConfig["top_k"].(float64); ok { + topK = int(scTopK) + if topK < 1 { + topK = 1 + } else if topK > 2048 { + topK = 2048 + } + } + if scUseKG, ok := searchConfig["use_kg"].(bool); ok { + useKG = scUseKG + } + if scLangs, ok := searchConfig["cross_languages"].([]interface{}); ok { + crossLanguages = make([]string, len(scLangs)) + for i, l := range scLangs { + if s, ok := l.(string); ok { + crossLanguages[i] = s + } + } + } + if scKeyword, ok := searchConfig["keyword"].(bool); ok { + keyword = scKeyword + } + if scRerankID, ok := searchConfig["rerank_id"].(string); ok { + rerankID = scRerankID + } + chatID, _ = searchConfig["chat_id"].(string) + } else { + common.Warn("Invalid search_id: search_config missing or invalid", zap.String("searchID", searchID)) + return nil, fmt.Errorf("Invalid search_id") + } + } + + // If meta_data_filter method is auto/semi_auto, get chat model + var chatModelForFilter *modelModule.ChatModel + if metadataFilter != nil { + method, _ := metadataFilter["method"].(string) + if method == "auto" || method == "semi_auto" { + if chatID != "" { + driver, modelName, apiConfig, _, err := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeChat, chatID) + if err != nil { + common.Warn("Failed to get chat model config from search_config chat_id, using tenant default", zap.String("chatID", chatID), zap.Error(err)) + } else { + chatModelForFilter = modelModule.NewChatModel(driver, &modelName, apiConfig) + } + } + + if chatModelForFilter == nil { + driver, modelName, apiConfig, _, err := modelProviderSvc.GetTenantDefaultModelByType(tenantIDs[0], entity.ModelTypeChat) + if err != nil { + common.Warn("Failed to get tenant default chat model for meta_data_filter", zap.Error(err)) + } else { + chatModelForFilter = modelModule.NewChatModel(driver, &modelName, apiConfig) + } + } + } + } + + // Apply meta_data_filter to get filtered doc_ids + docIDs := make([]string, len(req.DocIDs)) + copy(docIDs, req.DocIDs) + if len(metadataFilter) > 0 { + metadataSvc := service.NewMetadataService() + flattedMeta, err := metadataSvc.GetFlattedMetaByKBs(datasetIDs) + if err != nil { + common.Warn("Failed to get flatted metadata, using empty metadata for filter", zap.Error(err)) + flattedMeta = make(common.MetaData) + } + filteredDocIDs, _ := service.ApplyMetaDataFilter(ctx, metadataFilter, flattedMeta, question, chatModelForFilter, req.DocIDs, datasetIDs) + docIDs = filteredDocIDs + } + + // Apply cross_languages and keyword extraction + modifiedQuestion := question + if len(crossLanguages) > 0 { + translated, err := service.CrossLanguages(ctx, tenantIDs[0], "", question, crossLanguages) + if err != nil { + common.Warn("Failed to translate question", zap.String("llmID", ""), zap.Error(err)) + } else { + modifiedQuestion = translated + } + } + if keyword { + driver, modelName, apiConfig, _, err := modelProviderSvc.GetTenantDefaultModelByType(tenantIDs[0], entity.ModelTypeChat) + if err != nil { + common.Warn("Failed to get default chat model for LLM transformations", zap.Error(err)) + } else { + chatModel := modelModule.NewChatModel(driver, &modelName, apiConfig) + extractedKeywords, err := service.KeywordExtraction(ctx, chatModel, modifiedQuestion, 3) + if err != nil { + common.Warn("Failed to extract keywords from question", zap.Error(err)) + } else if extractedKeywords != "" { + modifiedQuestion = modifiedQuestion + extractedKeywords + } + } + } + + // Get tag-based rank features via LabelQuestion + metadataSvc := service.NewMetadataService() + labels := metadataSvc.LabelQuestion(modifiedQuestion, kbRecords) + + // Determine embedding model + var embeddingModel *modelModule.EmbeddingModel + if kbRecords[0].EmbdID != "" { + driver, modelName, apiConfig, maxTokens, embErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeEmbedding, kbRecords[0].EmbdID) + if embErr != nil { + return nil, fmt.Errorf("failed to get embedding model by embd_id: %w", embErr) + } + embeddingModel = modelModule.NewEmbeddingModel(driver, &modelName, apiConfig, maxTokens) + } + + // Get rerank model if rerankID is specified + var rerankModel *modelModule.RerankModel + if rerankID != "" { + driver, modelName, apiConfig, _, rErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeRerank, rerankID) + if rErr != nil { + return nil, fmt.Errorf("failed to get rerank model by rerank_id: %w", rErr) + } + rerankModel = modelModule.NewRerankModel(driver, &modelName, apiConfig) + } + + retrievalReq := &nlp.RetrievalRequest{ + TenantIDs: tenantIDs, + Question: modifiedQuestion, + KbIDs: datasetIDs, + DocIDs: docIDs, + Page: page, + PageSize: pageSize, + Top: &topK, + SimilarityThreshold: &similarityThreshold, + VectorSimilarityWeight: &vectorSimilarityWeight, + RerankModel: rerankModel, + RankFeature: &labels, + EmbeddingModel: embeddingModel, + } + + retrievalResult, err := nlp.NewRetrievalService(d.docEngine, d.documentDAO).Retrieval(ctx, retrievalReq) + if err != nil { + return nil, fmt.Errorf("retrieval search failed: %w", err) + } + + filteredChunks := retrievalResult.Chunks + + if useKG { + common.Warn("use_kg is not yet implemented in Go - skipping KG retrieval") + } + + filteredChunks = nlp.RetrievalByChildren(filteredChunks, tenantIDs, d.docEngine, ctx) + + for i := range filteredChunks { + delete(filteredChunks[i], "vector") + } + + common.Info("SearchDatasets completed", zap.String("userID", userID), zap.Any("kbID", datasetIDs), zap.String("question", question), zap.Int64("chunkCount", int64(len(filteredChunks)))) + + pyChunks := common.ConvertFloatsToPyFormat(filteredChunks).([]map[string]interface{}) + + return &service.SearchDatasetsResponse{ + Chunks: pyChunks, + DocAggs: retrievalResult.DocAggs, + Labels: &labels, + Total: retrievalResult.Total, + }, nil +} diff --git a/internal/service/dataset_search_test.go b/internal/service/dataset/search_test.go similarity index 94% rename from internal/service/dataset_search_test.go rename to internal/service/dataset/search_test.go index 524ca366c1..1b90d535d6 100644 --- a/internal/service/dataset_search_test.go +++ b/internal/service/dataset/search_test.go @@ -1,6 +1,10 @@ -package service +package dataset -import "testing" +import ( + "testing" + + "ragflow/internal/service" +) func TestSearchDatasetRequestToSearchDatasetsRequest(t *testing.T) { page := 2 @@ -12,7 +16,7 @@ func TestSearchDatasetRequestToSearchDatasetsRequest(t *testing.T) { vectorSimilarityWeight := 0.8 searchID := "search-1" rerankID := "rerank-1" - req := &SearchDatasetRequest{ + req := &service.SearchDatasetRequest{ Question: "hello world", Page: &page, Size: &size, diff --git a/internal/service/dataset/service.go b/internal/service/dataset/service.go new file mode 100644 index 0000000000..28081e9133 --- /dev/null +++ b/internal/service/dataset/service.go @@ -0,0 +1,40 @@ +package dataset + +import ( + "ragflow/internal/dao" + "ragflow/internal/engine" + "ragflow/internal/service" + "ragflow/internal/utility" +) + +// DatasetService implements the RESTful dataset APIs. +type DatasetService struct { + kbDAO *dao.KnowledgebaseDAO + documentDAO *dao.DocumentDAO + connectorDAO *dao.ConnectorDAO + tenantDAO *dao.TenantDAO + tenantLLMDAO *dao.TenantLLMDAO + pipelineLogDAO *dao.PipelineOperationLogDAO + userTenantDAO *dao.UserTenantDAO + taskDAO *dao.TaskDAO + searchService *service.SearchService + docEngine engine.DocEngine + embeddingCache *utility.EmbeddingLRU +} + +// NewDatasetService creates a new datasets service. +func NewDatasetService() *DatasetService { + return &DatasetService{ + kbDAO: dao.NewKnowledgebaseDAO(), + documentDAO: dao.NewDocumentDAO(), + connectorDAO: dao.NewConnectorDAO(), + tenantDAO: dao.NewTenantDAO(), + tenantLLMDAO: dao.NewTenantLLMDAO(), + pipelineLogDAO: dao.NewPipelineOperationLogDAO(), + userTenantDAO: dao.NewUserTenantDAO(), + taskDAO: dao.NewTaskDAO(), + searchService: service.NewSearchService(), + docEngine: engine.Get(), + embeddingCache: utility.NewEmbeddingLRU(1000), + } +} diff --git a/internal/service/dataset/setup_test.go b/internal/service/dataset/setup_test.go new file mode 100644 index 0000000000..bcbadc1199 --- /dev/null +++ b/internal/service/dataset/setup_test.go @@ -0,0 +1,55 @@ +package dataset + +import ( + "testing" + + "ragflow/internal/dao" + "ragflow/internal/entity" + + "github.com/glebarez/sqlite" + "gorm.io/gorm" +) + +// setupServiceTestDB initializes an in-memory SQLite database for tests. +func setupServiceTestDB(t *testing.T) *gorm.DB { + t.Helper() + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + TranslateError: true, + }) + if err != nil { + t.Fatalf("failed to open sqlite: %v", err) + } + if err = db.AutoMigrate( + &entity.Document{}, + &entity.Knowledgebase{}, + &entity.Task{}, + &entity.IngestionTask{}, + &entity.IngestionTaskLog{}, + &entity.File2Document{}, + &entity.File{}, + &entity.User{}, + &entity.Tenant{}, + &entity.UserTenant{}, + &entity.API4Conversation{}, + &entity.Connector{}, + &entity.Connector2Kb{}, + &entity.SyncLogs{}, + &entity.TenantModelProvider{}, + &entity.TenantModelInstance{}, + &entity.TenantModel{}, + &entity.UserCanvas{}, + ); err != nil { + t.Fatalf("failed to migrate: %v", err) + } + return db +} + +// pushServiceDB swaps dao.DB for the test and restores after. +func pushServiceDB(t *testing.T, testDB *gorm.DB) { + t.Helper() + oldDB := dao.DB + dao.DB = testDB + t.Cleanup(func() { dao.DB = oldDB }) +} + +func sptr(s string) *string { return &s } diff --git a/internal/service/dataset/tags.go b/internal/service/dataset/tags.go new file mode 100644 index 0000000000..2caee6dc73 --- /dev/null +++ b/internal/service/dataset/tags.go @@ -0,0 +1,233 @@ +package dataset + +import ( + "context" + "errors" + "fmt" + "sort" + "strings" + "time" + + "ragflow/internal/common" + "ragflow/internal/dao" + enginetypes "ragflow/internal/engine/types" +) + +func (d *DatasetService) AggregateTags(datasetIDs []string, userID string) ([]map[string]interface{}, common.ErrorCode, error) { + if len(datasetIDs) == 0 { + return nil, common.CodeDataError, errors.New("Lack of dataset_ids in query parameters") + } + if d.docEngine == nil { + return nil, common.CodeServerError, errors.New("Document engine is not initialized") + } + + datasetIDsByTenant := make(map[string][]string) + for _, rawID := range datasetIDs { + rawID = strings.TrimSpace(rawID) + if rawID == "" { + continue + } + datasetID, err := normalizeDatasetID(rawID) + if err != nil { + return nil, common.CodeDataError, err + } + if !d.kbDAO.Accessible(datasetID, userID) { + return nil, common.CodeDataError, fmt.Errorf("No authorization for dataset '%s'", datasetID) + } + kb, err := d.kbDAO.GetByID(datasetID) + if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, fmt.Errorf("Invalid Dataset ID '%s'", datasetID) + } + return nil, common.CodeServerError, errors.New("Database operation failed") + } + if kb.DocNum <= 0 { + continue + } + datasetIDsByTenant[kb.TenantID] = append(datasetIDsByTenant[kb.TenantID], datasetID) + } + + const pageSize = 10000 + merged := make(map[string]int) + for tenantID, kbIDs := range datasetIDsByTenant { + for offset := 0; ; offset += pageSize { + searchResp, err := d.docEngine.Search(context.Background(), &enginetypes.SearchRequest{ + IndexNames: []string{fmt.Sprintf("ragflow_%s", tenantID)}, + KbIDs: kbIDs, + Offset: offset, + Limit: pageSize, + SelectFields: []string{"tag_kwd"}, + }) + if err != nil { + return nil, common.CodeServerError, fmt.Errorf("failed to aggregate tags: %w", err) + } + for _, agg := range d.docEngine.GetAggregation(searchResp.Chunks, "tag_kwd") { + tag, _ := agg["key"].(string) + if tag == "" { + continue + } + switch count := agg["count"].(type) { + case int: + merged[tag] += count + case int32: + merged[tag] += int(count) + case int64: + merged[tag] += int(count) + case float64: + merged[tag] += int(count) + } + } + chunkCount := len(searchResp.Chunks) + if chunkCount == 0 || chunkCount < pageSize { + break + } + if searchResp.Total > 0 && int64(offset+chunkCount) >= searchResp.Total { + break + } + } + } + result := make([]map[string]interface{}, 0, len(merged)) + for tag, count := range merged { + result = append(result, map[string]interface{}{ + "value": tag, + "count": count, + }) + } + return result, common.CodeSuccess, nil +} + +func (d *DatasetService) ListTags(datasetID, userID string) ([]map[string]interface{}, common.ErrorCode, error) { + datasetID = strings.TrimSpace(datasetID) + if datasetID == "" { + return nil, common.CodeDataError, errors.New("Lack of \"Dataset ID\"") + } + normalizedID, err := normalizeDatasetID(datasetID) + if err != nil { + return nil, common.CodeDataError, err + } + datasetID = normalizedID + if !d.kbDAO.Accessible(datasetID, userID) { + return nil, common.CodeDataError, errors.New("No authorization.") + } + if d.docEngine == nil { + return nil, common.CodeServerError, errors.New("Document engine is not initialized") + } + kb, err := d.kbDAO.GetByID(datasetID) + if err != nil || kb == nil { + return nil, common.CodeDataError, errors.New("Invalid Dataset ID") + } + indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + exists, err := d.docEngine.ChunkStoreExists(ctx, indexName, datasetID) + if err != nil { + return nil, common.CodeServerError, fmt.Errorf("failed to inspect chunk store: %w", err) + } + if !exists { + return []map[string]interface{}{}, common.CodeSuccess, nil + } + const pageSize = 10000 + counts := make(map[string]int) + for offset := 0; ; offset += pageSize { + if err = ctx.Err(); err != nil { + return nil, common.CodeServerError, fmt.Errorf("list tags timeout or canceled: %w", err) + } + searchResp, err := d.docEngine.Search(ctx, &enginetypes.SearchRequest{ + IndexNames: []string{indexName}, + KbIDs: []string{datasetID}, + Offset: offset, + Limit: pageSize, + SelectFields: []string{"tag_kwd"}, + }) + if err != nil { + return nil, common.CodeServerError, fmt.Errorf("failed to list tags: %w", err) + } + for _, agg := range d.docEngine.GetAggregation(searchResp.Chunks, "tag_kwd") { + tag, _ := agg["key"].(string) + if tag == "" { + continue + } + switch count := agg["count"].(type) { + case int: + counts[tag] += count + case int32: + counts[tag] += int(count) + case int64: + counts[tag] += int(count) + case float64: + counts[tag] += int(count) + } + } + chunkCount := len(searchResp.Chunks) + if chunkCount == 0 || chunkCount < pageSize { + break + } + if searchResp.Total > 0 && int64(offset+chunkCount) >= searchResp.Total { + break + } + } + if len(counts) == 0 { + return []map[string]interface{}{}, common.CodeSuccess, nil + } + tags := make([]string, 0, len(counts)) + for tag := range counts { + tags = append(tags, tag) + } + sort.Slice(tags, func(i, j int) bool { + if counts[tags[i]] != counts[tags[j]] { + return counts[tags[i]] > counts[tags[j]] + } + return tags[i] < tags[j] + }) + result := make([]map[string]interface{}, 0, len(tags)) + for _, tag := range tags { + result = append(result, map[string]interface{}{ + "key": tag, + "count": counts[tag], + }) + } + return result, common.CodeSuccess, nil +} + +func (d *DatasetService) RenameTag(datasetID, userID, fromTag, toTag string) (map[string]interface{}, common.ErrorCode, error) { + fromTag = strings.TrimSpace(fromTag) + toTag = strings.TrimSpace(toTag) + datasetID, err := normalizeDatasetID(datasetID) + if err != nil { + return nil, common.CodeDataError, err + } + if strings.TrimSpace(datasetID) == "" { + return nil, common.CodeDataError, errors.New("Lack of \"Dataset ID\"") + } + if !d.kbDAO.Accessible(datasetID, userID) { + return nil, common.CodeDataError, errors.New("No authorization.") + } + if d.docEngine == nil { + return nil, common.CodeServerError, errors.New("Document engine is not initialized") + } + kb, err := d.kbDAO.GetByID(datasetID) + if err != nil || kb == nil { + return nil, common.CodeDataError, errors.New("Invalid Dataset ID") + } + indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) + condition := map[string]interface{}{ + "tag_kwd": fromTag, + "kb_id": datasetID, + } + newValue := map[string]interface{}{ + "remove": map[string]interface{}{ + "tag_kwd": fromTag, + }, + "add": map[string]interface{}{ + "tag_kwd": toTag, + }, + } + err = d.docEngine.UpdateChunks(context.Background(), condition, newValue, indexName, datasetID) + if err != nil { + return nil, common.CodeServerError, fmt.Errorf("failed to rename tag: %w", err) + } + return map[string]interface{}{ + "from": fromTag, + "to": toTag, + }, common.CodeSuccess, nil +} diff --git a/internal/service/dataset_aggregate_tags_test.go b/internal/service/dataset/tags_aggregate_test.go similarity index 99% rename from internal/service/dataset_aggregate_tags_test.go rename to internal/service/dataset/tags_aggregate_test.go index 29cac97fba..1a5f1b7c81 100644 --- a/internal/service/dataset_aggregate_tags_test.go +++ b/internal/service/dataset/tags_aggregate_test.go @@ -1,4 +1,4 @@ -package service +package dataset import ( "context" diff --git a/internal/service/dataset_list_tags_test.go b/internal/service/dataset/tags_list_test.go similarity index 99% rename from internal/service/dataset_list_tags_test.go rename to internal/service/dataset/tags_list_test.go index 3bac34e945..7304bf2b0e 100644 --- a/internal/service/dataset_list_tags_test.go +++ b/internal/service/dataset/tags_list_test.go @@ -1,4 +1,4 @@ -package service +package dataset import ( "context" diff --git a/internal/service/dataset_rename_tag_test.go b/internal/service/dataset/tags_rename_test.go similarity index 99% rename from internal/service/dataset_rename_tag_test.go rename to internal/service/dataset/tags_rename_test.go index 1307bb41fb..567a6c7f22 100644 --- a/internal/service/dataset_rename_tag_test.go +++ b/internal/service/dataset/tags_rename_test.go @@ -1,4 +1,4 @@ -package service +package dataset import ( "context" diff --git a/internal/service/dataset_task_cleanup_test.go b/internal/service/dataset/task_cleanup_test.go similarity index 99% rename from internal/service/dataset_task_cleanup_test.go rename to internal/service/dataset/task_cleanup_test.go index 4f547b72a2..5ddc0b9abe 100644 --- a/internal/service/dataset_task_cleanup_test.go +++ b/internal/service/dataset/task_cleanup_test.go @@ -14,7 +14,7 @@ // limitations under the License. // -package service +package dataset import ( "errors" diff --git a/internal/service/dataset/update.go b/internal/service/dataset/update.go new file mode 100644 index 0000000000..85f822f2c2 --- /dev/null +++ b/internal/service/dataset/update.go @@ -0,0 +1,269 @@ +package dataset + +import ( + "context" + "errors" + "fmt" + "strings" + + "ragflow/internal/common" + "ragflow/internal/dao" + "ragflow/internal/entity" + pipelinepkg "ragflow/internal/ingestion/pipeline" + "ragflow/internal/service" + + "go.uber.org/zap" +) + +func (d *DatasetService) UpdateDataset(datasetID, tenantID string, req service.UpdateDatasetRequest) (map[string]interface{}, common.ErrorCode, error) { + kb, err := d.kbDAO.GetByID(datasetID) + if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, errors.New("Dataset not found") + } + return nil, common.CodeServerError, errors.New("Database operation failed") + } + + if kb == nil || kb.TenantID != tenantID { + return nil, common.CodeDataError, fmt.Errorf("User '%s' lacks permission for dataset '%s'", tenantID, datasetID) + } + + connectorsProvided := req.Connectors != nil + connectors := make([]service.DatasetConnectorRequest, 0) + if req.Connectors != nil { + connectors = *req.Connectors + } + + updates := make(map[string]interface{}) + + if req.Name != nil { + name := strings.TrimSpace(*req.Name) + if name == "" { + return nil, common.CodeDataError, errors.New("String should have at least 1 character") + } + if len(name) > 128 { + return nil, common.CodeDataError, errors.New("String should have at most 128 characters") + } + updates["name"] = name + } + if req.Avatar != nil { + if len(*req.Avatar) > 65535 { + return nil, common.CodeDataError, errors.New("String should have at most 65535 characters") + } + if err := validateDatasetAvatar(*req.Avatar); err != nil { + return nil, common.CodeDataError, err + } + updates["avatar"] = *req.Avatar + } + if req.Description != nil { + if len(*req.Description) > 65535 { + return nil, common.CodeDataError, errors.New("String should have at most 65535 characters") + } + updates["description"] = *req.Description + } + if req.Language != nil { + language := strings.TrimSpace(*req.Language) + if len(language) > 32 { + return nil, common.CodeDataError, errors.New("String should have at most 32 characters") + } + updates["language"] = language + } + if req.Permission != nil { + permission := strings.TrimSpace(*req.Permission) + if permission != "me" && permission != "team" { + return nil, common.CodeDataError, errors.New("Input should be 'me' or 'team'") + } + updates["permission"] = permission + } + + isPipelineMode := req.ParseType != nil && *req.ParseType == 2 + isBuiltinMode := req.ParseType != nil && *req.ParseType == 1 + + if isBuiltinMode && req.PipelineID != nil { + req.PipelineID = nil + } + if isPipelineMode && req.ParserID != nil { + req.ParserID = nil + } + + if req.PipelineID != nil { + pipelineID, err := normalizeDatasetPipelineID(*req.PipelineID) + if err != nil { + return nil, common.CodeDataError, err + } + if pipelineID != nil { + updates["pipeline_id"] = *pipelineID + } + } + + parserID, parserIDProvided, err := datasetUpdateParserID(req) + if err != nil { + return nil, common.CodeDataError, err + } + if parserIDProvided { + updates["parser_id"] = parserID + } + + if req.ParseType == nil && parserIDProvided && req.PipelineID != nil { + return nil, common.CodeDataError, errors.New("parser_id and pipeline_id are mutually exclusive") + } + + embdID, embdIDProvided, err := datasetUpdateEmbeddingID(req) + if err != nil { + return nil, common.CodeDataError, err + } + if embdIDProvided { + tenantEmbdID := ptrStringValue(kb.TenantEmbdID) + if embdID == "" { + embdID = kb.EmbdID + } else { + tenantEmbdID = "" + } + ok, message := d.verifyEmbeddingAvailability(embdID, tenantID) + if !ok { + return nil, common.CodeDataError, errors.New(message) + } + if embdID != "" && tenantEmbdID == "" { + resolvedID, err := service.NewModelProviderService().ResolveModelID(tenantID, entity.ModelTypeEmbedding, embdID) + if err == nil { + tenantEmbdID = resolvedID + } + } + updates["embd_id"] = embdID + updates["tenant_embd_id"] = stringPtrIfNotEmpty(tenantEmbdID) + } + + if req.ParserConfig != nil { + if err := validateDatasetParserConfigSize(req.ParserConfig); err != nil { + return nil, common.CodeDataError, err + } + if len(req.ParserConfig) > 0 { + effectiveParserID := kb.ParserID + if parserIDProvided { + effectiveParserID = parserID + } + effectivePipelineID := kb.PipelineID + if req.PipelineID != nil { + if normalized, err := normalizeDatasetPipelineID(*req.PipelineID); err == nil { + effectivePipelineID = normalized + } + } else if parserIDProvided && kb.PipelineID != nil { + effectivePipelineID = nil + } + + isCanvas := effectivePipelineID != nil && strings.TrimSpace(*effectivePipelineID) != "" + dslJSON, dslErr := service.LoadPipelineDSL(isCanvas, effectiveParserID, effectivePipelineID) + if dslErr != nil { + common.Warn("failed to load pipeline DSL for building parser_config", + zap.String("parserID", effectiveParserID), zap.Error(dslErr)) + } + if dslJSON != nil { + updates["parser_config"] = pipelinepkg.BuildParserConfig(dslJSON, map[string]interface{}(req.ParserConfig)) + } + } + } + + if req.Pagerank != nil && *req.Pagerank != kb.Pagerank { + if *req.Pagerank < 0 || *req.Pagerank > 100 { + return nil, common.CodeDataError, errors.New("Input should be less than or equal to 100") + } + if !d.docEngine.SupportsPageRank() { + return nil, common.CodeDataError, errors.New("'pagerank' can only be set when doc_engine is elasticsearch") + } + indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) + if *req.Pagerank > 0 { + err = d.docEngine.UpdateChunks(context.Background(), map[string]interface{}{"kb_id": kb.ID}, map[string]interface{}{common.PAGERANK_FLD: *req.Pagerank}, indexName, kb.ID) + } else { + err = d.docEngine.UpdateChunks(context.Background(), map[string]interface{}{"exists": common.PAGERANK_FLD}, map[string]interface{}{"remove": common.PAGERANK_FLD}, indexName, kb.ID) + } + if err != nil { + return nil, common.CodeServerError, err + } + updates["pagerank"] = *req.Pagerank + } + + if parserIDProvided && parserID != kb.ParserID { + if _, ok := updates["parser_config"]; !ok { + if resolved, cpErr := service.ResolveComponentParamsDefaults(parserID, nil); cpErr != nil { + common.Warn("failed to resolve component params defaults on parser_id switch", + zap.String("parserID", parserID), zap.Error(cpErr)) + } else if resolved != nil { + updates["parser_config"] = resolved + } + } + } + if kb.PipelineID != nil && parserIDProvided { + if _, ok := updates["pipeline_id"]; !ok { + updates["pipeline_id"] = nil + } + } + + pipelineChanged := req.PipelineID != nil && (kb.PipelineID == nil || *req.PipelineID != *kb.PipelineID) + if pipelineChanged { + cfgParserID := kb.ParserID + if parserIDProvided { + cfgParserID = parserID + } + cfgPipelineID, _ := updates["pipeline_id"].(string) + var cpPipelineID *string + if cfgPipelineID != "" { + cpPipelineID = &cfgPipelineID + } + if cpDefaults, cpErr := service.ResolveComponentParamsDefaults(cfgParserID, cpPipelineID); cpErr != nil { + common.Warn("failed to resolve component params defaults on pipeline change", + zap.String("parserID", cfgParserID), zap.Error(cpErr)) + } else if cpDefaults != nil { + updates["parser_config"] = cpDefaults + } + } + + if nameValue, ok := updates["name"].(string); ok && strings.ToLower(nameValue) != strings.ToLower(kb.Name) { + existing, lookupErr := d.kbDAO.GetByName(nameValue, tenantID) + if lookupErr != nil && !dao.IsNotFoundErr(lookupErr) { + return nil, common.CodeServerError, errors.New("Database operation failed") + } + if existing != nil { + return nil, common.CodeDataError, fmt.Errorf("Dataset name '%s' already exists", nameValue) + } + } + + if len(updates) == 0 && !connectorsProvided { + return nil, common.CodeDataError, errors.New("No properties were modified") + } + + if len(updates) > 0 { + if err = d.kbDAO.UpdateByID(kb.ID, updates); err != nil { + return nil, common.CodeServerError, errors.New("Update dataset error.(Database error)") + } + } + + if connectorsProvided { + connectorLinks := make([]dao.DatasetConnectorLink, 0, len(connectors)) + for _, connector := range connectors { + connectorID := strings.TrimSpace(connector.ID) + if connectorID == "" { + return nil, common.CodeDataError, errors.New("connector id is required") + } + connectorLinks = append(connectorLinks, dao.DatasetConnectorLink{ + ID: connectorID, + AutoParse: connector.AutoParse, + }) + } + if err = d.connectorDAO.LinkDatasetConnectors(kb.ID, connectorLinks); err != nil { + return nil, common.CodeServerError, errors.New("Database operation failed") + } + } + + updatedKB, err := d.kbDAO.GetByID(kb.ID) + if err != nil { + return nil, common.CodeDataError, errors.New("Dataset updated failed") + } + + data := datasetToMap(updatedKB) + linkedConnectors, err := d.connectorDAO.ListByDatasetID(kb.ID) + if err != nil { + return nil, common.CodeServerError, errors.New("Database operation failed") + } + data["connectors"] = datasetConnectorsOrEmpty(linkedConnectors) + return data, common.CodeSuccess, nil +} diff --git a/internal/service/dataset_update_test.go b/internal/service/dataset/update_test.go similarity index 95% rename from internal/service/dataset_update_test.go rename to internal/service/dataset/update_test.go index 368af05e04..03ab703d45 100644 --- a/internal/service/dataset_update_test.go +++ b/internal/service/dataset/update_test.go @@ -14,7 +14,7 @@ // limitations under the License. // -package service +package dataset import ( "encoding/json" @@ -24,6 +24,7 @@ import ( "ragflow/internal/common" "ragflow/internal/dao" "ragflow/internal/entity" + "ragflow/internal/service" "gorm.io/gorm" ) @@ -41,7 +42,7 @@ func TestDatasetServiceUpdateDatasetUpdatesFields(t *testing.T) { chunkMethod := string(entity.ParserTypeBook) embeddingModel := "BAAI/bge-large-zh-v1.5@Builtin" - result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{ + result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{ Name: &name, Description: &description, Language: &language, @@ -101,7 +102,7 @@ func TestUpdateDataset_RejectsSimultaneousParserIDAndPipelineID(t *testing.T) { chunkMethod := "book" pipelineID := "abcdef0123456789abcdef0123456789" - _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{ + _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{ ParserID: &chunkMethod, PipelineID: &pipelineID, }) @@ -127,7 +128,7 @@ func TestUpdateDataset_ParseTypeBuiltinClearsPipelineID(t *testing.T) { pipelineID := "ABCDEF0123456789ABCDEF0123456789" parseType := 1 - result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{ + result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{ ParserID: &chunkMethod, PipelineID: &pipelineID, ParseType: &parseType, @@ -158,7 +159,7 @@ func TestUpdateDataset_ParseTypePipelineIgnoresParserID(t *testing.T) { pipelineID := "ABCDEF0123456789ABCDEF0123456789" parseType := 2 - result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{ + result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{ ParserID: &chunkMethod, PipelineID: &pipelineID, ParseType: &parseType, @@ -206,7 +207,7 @@ func TestDatasetServiceUpdateDatasetRejectsMissingDataset(t *testing.T) { pushServiceDB(t, db) name := "Renamed" - _, code, err := testDatasetUpdateService(t).UpdateDataset("missing-kb", "tenant-1", UpdateDatasetRequest{Name: &name}) + _, code, err := testDatasetUpdateService(t).UpdateDataset("missing-kb", "tenant-1", service.UpdateDatasetRequest{Name: &name}) if err == nil { t.Fatal("expected missing dataset error") } @@ -224,7 +225,7 @@ func TestDatasetServiceUpdateDatasetRejectsNonOwner(t *testing.T) { insertDatasetUpdateKB(t, "kb-1", "tenant-1", "Original") name := "Renamed" - _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-2", UpdateDatasetRequest{Name: &name}) + _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-2", service.UpdateDatasetRequest{Name: &name}) if err == nil { t.Fatal("expected permission error") } @@ -242,7 +243,7 @@ func TestDatasetServiceUpdateDatasetValidatesName(t *testing.T) { insertDatasetUpdateKB(t, "kb-1", "tenant-1", "Original") name := " " - _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{Name: &name}) + _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{Name: &name}) if err == nil { t.Fatal("expected name validation error") } @@ -261,7 +262,7 @@ func TestDatasetServiceUpdateDatasetRejectsDuplicateName(t *testing.T) { insertDatasetUpdateKB(t, "kb-2", "tenant-1", "Existing") name := "Existing" - _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{Name: &name}) + _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{Name: &name}) if err == nil { t.Fatal("expected duplicate name error") } @@ -278,7 +279,7 @@ func TestDatasetServiceUpdateDatasetRejectsNoPropertiesModified(t *testing.T) { pushServiceDB(t, db) insertDatasetUpdateKB(t, "kb-1", "tenant-1", "Original") - _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{}) + _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{}) if err == nil { t.Fatal("expected no-op update error") } @@ -297,8 +298,8 @@ func TestDatasetServiceUpdateDatasetLinksConnectors(t *testing.T) { insertDatasetUpdateConnector(t, "connector-1", "tenant-1") autoParse := "0" - result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{ - Connectors: &[]DatasetConnectorRequest{{ID: "connector-1", AutoParse: autoParse}}, + result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{ + Connectors: &[]service.DatasetConnectorRequest{{ID: "connector-1", AutoParse: autoParse}}, }) if err != nil { t.Fatalf("UpdateDataset failed: %v", err) @@ -335,7 +336,7 @@ func TestDatasetServiceUpdateDatasetAcceptsProviderInstanceEmbedding(t *testing. insertDatasetUpdateTenantModel(t, "model-1", "provider-1", "instance-1", "embedding-2", int(entity.ModelTypeEmbedding)) embeddingModel := "embedding-2@test@ZHIPU-AI" - result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{ + result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{ EmbeddingModel: &embeddingModel, }) if err != nil { @@ -366,7 +367,7 @@ func TestDatasetServiceUpdateDatasetAcceptsEmbeddingModelID(t *testing.T) { insertDatasetUpdateTenantModel(t, "model-1", "provider-1", "instance-1", "embedding-2", int(entity.ModelTypeEmbedding)) embeddingModelID := "model-1" - result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{ + result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{ EmbeddingModel: &embeddingModelID, }) if err != nil { @@ -396,8 +397,8 @@ func TestDatasetServiceUpdateDatasetRejectsEmptyConnectorID(t *testing.T) { pushServiceDB(t, db) insertDatasetUpdateKB(t, "kb-1", "tenant-1", "Original") - connectors := []DatasetConnectorRequest{{ID: " "}} - _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{ + connectors := []service.DatasetConnectorRequest{{ID: " "}} + _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{ Connectors: &connectors, }) if err == nil { @@ -605,7 +606,7 @@ func TestUpdateDataset_StripsUnknownParam_Builtin(t *testing.T) { pushServiceDB(t, db) insertDatasetUpdateKB(t, "kb-1", "tenant-1", "Original") - result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{ + result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{ ParserConfig: map[string]interface{}{ "Parser:HipSignsRhyme": map[string]interface{}{ "no_such_param": 1, @@ -641,7 +642,7 @@ func TestUpdateDataset_AcceptsValidComponentParams_Builtin(t *testing.T) { pushServiceDB(t, db) insertDatasetUpdateKB(t, "kb-1", "tenant-1", "Original") - result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{ + result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{ ParserConfig: map[string]interface{}{ "Parser:HipSignsRhyme": map[string]interface{}{ "pdf": map[string]interface{}{"parse_method": "deepdoc"}, @@ -683,7 +684,7 @@ func TestUpdateDataset_StripsCanvasUnknownParam(t *testing.T) { seedDatasetUpdateCanvas(t, "canvas-1", "tenant-1", dsl) insertDatasetUpdateCanvasKB(t, "kb-1", "tenant-1", "Original", "canvas-1") - result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{ + result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{ ParserConfig: map[string]interface{}{ "Parser:NoSuch": map[string]interface{}{ "pdf": map[string]interface{}{}, @@ -721,7 +722,7 @@ func TestUpdateDataset_AcceptsValidCanvasComponentParams(t *testing.T) { seedDatasetUpdateCanvas(t, "canvas-1", "tenant-1", dsl) insertDatasetUpdateCanvasKB(t, "kb-1", "tenant-1", "Original", "canvas-1") - result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{ + result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{ ParserConfig: map[string]interface{}{ "Parser:CustomRhyme": map[string]interface{}{ "pdf": map[string]interface{}{}, @@ -751,7 +752,7 @@ func TestUpdateDataset_SwitchCanvasToBuiltinValidatesAgainstBuiltin(t *testing.T insertDatasetUpdateCanvasKB(t, "kb-1", "tenant-1", "Original", "canvas-1") chunkMethod := "naive" - _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{ + _, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", service.UpdateDatasetRequest{ ParserID: &chunkMethod, ParserConfig: map[string]interface{}{ // Valid for the "general" builtin template, not the canvas. diff --git a/internal/service/dataset/utils.go b/internal/service/dataset/utils.go new file mode 100644 index 0000000000..bf49a00c62 --- /dev/null +++ b/internal/service/dataset/utils.go @@ -0,0 +1,322 @@ +package dataset + +import ( + "encoding/json" + "math" + "regexp" + "sort" + "strings" + "time" + + "ragflow/internal/entity" + modelModule "ragflow/internal/entity/models" + "ragflow/internal/service" +) + +func datasetListItemToMap(kb *entity.KnowledgebaseListItem) map[string]interface{} { + item := map[string]interface{}{ + "id": kb.ID, + "name": kb.Name, + "tenant_id": kb.TenantID, + "permission": kb.Permission, + "document_count": kb.DocNum, + "token_num": kb.TokenNum, + "chunk_count": kb.ChunkNum, + "parser_id": kb.ParserID, + "embedding_model": kb.EmbdID, + "nickname": kb.Nickname, + } + if kb.Avatar != nil { + item["avatar"] = *kb.Avatar + } + if kb.Language != nil { + item["language"] = *kb.Language + } + if kb.Description != nil { + item["description"] = *kb.Description + } + if kb.TenantAvatar != nil { + item["tenant_avatar"] = *kb.TenantAvatar + } + if kb.UpdateTime != nil { + item["update_time"] = *kb.UpdateTime + } + return item +} + +func datasetToMap(kb *entity.Knowledgebase) map[string]interface{} { + item := map[string]interface{}{ + "id": kb.ID, + "tenant_id": kb.TenantID, + "name": kb.Name, + "embedding_model": kb.EmbdID, + "permission": kb.Permission, + "created_by": kb.CreatedBy, + "document_count": kb.DocNum, + "token_num": kb.TokenNum, + "chunk_count": kb.ChunkNum, + "similarity_threshold": kb.SimilarityThreshold, + "vector_similarity_weight": kb.VectorSimilarityWeight, + "parser_id": kb.ParserID, + "parser_config": kb.ParserConfig, + "pagerank": kb.Pagerank, + "create_time": kb.CreateTime, + } + if kb.Avatar != nil { + item["avatar"] = *kb.Avatar + } + if kb.Language != nil { + item["language"] = *kb.Language + } + if kb.Description != nil { + item["description"] = *kb.Description + } + if kb.PipelineID != nil { + item["pipeline_id"] = *kb.PipelineID + } + if kb.GraphragTaskID != nil { + item["graphrag_task_id"] = *kb.GraphragTaskID + } + if kb.GraphragTaskFinishAt != nil { + item["graphrag_task_finish_at"] = kb.GraphragTaskFinishAt.Format("2006-01-02 15:04:05") + } + if kb.RaptorTaskID != nil { + item["raptor_task_id"] = *kb.RaptorTaskID + } + if kb.RaptorTaskFinishAt != nil { + item["raptor_task_finish_at"] = kb.RaptorTaskFinishAt.Format("2006-01-02 15:04:05") + } + if kb.MindmapTaskID != nil { + item["mindmap_task_id"] = *kb.MindmapTaskID + } + if kb.MindmapTaskFinishAt != nil { + item["mindmap_task_finish_at"] = kb.MindmapTaskFinishAt.Format("2006-01-02 15:04:05") + } + if kb.UpdateTime != nil { + item["update_time"] = *kb.UpdateTime + } + return item +} + +func limitStrings(values []string, limit int) []string { + if len(values) <= limit { + return values + } + return values[:limit] +} + +func stringPointerValue(s *string) interface{} { + if s == nil { + return nil + } + return *s +} + +func int64PointerValue(i *int64) interface{} { + if i == nil { + return nil + } + return *i +} + +func timePointerValue(t *time.Time) interface{} { + if t == nil { + return nil + } + return t.Format("2006-01-02 15:04:05") +} + +func jsonMapValue(m entity.JSONMap) interface{} { + if m == nil { + return nil + } + return map[string]interface{}(m) +} + +func datasetMap(value interface{}) map[string]interface{} { + if m, ok := value.(map[string]interface{}); ok { + return m + } + return nil +} + +func datasetString(value interface{}) string { + if s, ok := value.(string); ok { + return s + } + return "" +} + +func datasetStringSlice(value interface{}) []string { + if sl, ok := value.([]string); ok { + return sl + } + if raw, ok := value.([]interface{}); ok { + result := make([]string, 0, len(raw)) + for _, v := range raw { + if s, ok := v.(string); ok { + result = append(result, s) + } + } + return result + } + return nil +} + +func datasetGuessVecField(src map[string]interface{}) string { + var f64, f32 string + for k, v := range src { + if !strings.HasPrefix(k, "q_") && !strings.HasPrefix(k, "u_") { + continue + } + switch v.(type) { + case []float64: + f64 = k + case string: + f32 = k + } + } + if f64 != "" { + return f64 + } + return f32 +} + +func datasetAsFloatVec(v interface{}) []float64 { + switch val := v.(type) { + case []float64: + return val + case []interface{}: + vec := make([]float64, 0, len(val)) + for _, item := range val { + switch n := item.(type) { + case float64: + vec = append(vec, n) + case int: + vec = append(vec, float64(n)) + case int64: + vec = append(vec, float64(n)) + case json.Number: + if f, err := n.Float64(); err == nil { + vec = append(vec, f) + } + } + } + return vec + } + return nil +} + +func datasetCosSim(a, b []float64) float64 { + if len(a) != len(b) || len(a) == 0 { + return 0 + } + var dot, na, nb float64 + for i := range a { + dot += a[i] * b[i] + na += a[i] * a[i] + nb += b[i] * b[i] + } + if na == 0 || nb == 0 { + return 0 + } + return dot / (math.Sqrt(na) * math.Sqrt(nb)) +} + +func datasetCleanEmbeddingText(s string) string { + re := regexp.MustCompile(`<[^>]*>`) + return re.ReplaceAllString(s, "") +} + +func datasetEncodeEmbedding(embeddingModel *modelModule.EmbeddingModel, texts []string) ([][]float64, error) { + if len(texts) == 0 { + return nil, nil + } + cleaned := make([]string, len(texts)) + for i, t := range texts { + cleaned[i] = datasetCleanEmbeddingText(t) + } + embeddingConfig := &modelModule.EmbeddingConfig{Dimension: 0} + embeddings, err := embeddingModel.ModelDriver.Embed(embeddingModel.ModelName, cleaned, embeddingModel.APIConfig, embeddingConfig, nil) + if err != nil { + return nil, err + } + vectors := make([][]float64, len(embeddings)) + for i, embedding := range embeddings { + vectors[i] = embedding.Embedding + } + return vectors, nil +} + +func datasetMixVectors(titleVector, contentVector []float64, titleWeight float64) []float64 { + if len(titleVector) == 0 && len(contentVector) == 0 { + return nil + } + if len(titleVector) == 0 { + return contentVector + } + if len(contentVector) == 0 { + return titleVector + } + minLen := len(titleVector) + if len(contentVector) < minLen { + minLen = len(contentVector) + } + mixed := make([]float64, minLen) + for i := 0; i < minLen; i++ { + mixed[i] = titleWeight*titleVector[i] + (1-titleWeight)*contentVector[i] + } + return mixed +} + +func datasetEmbeddingCheckSummary(datasetID, embeddingID string, sampled int, similarities []float64, matchMode string) service.EmbeddingCheckSummary { + if len(similarities) == 0 { + return service.EmbeddingCheckSummary{ + KbID: datasetID, + Model: embeddingID, + Sampled: sampled, + Valid: 0, + AvgCosSim: 0, + MinCosSim: 0, + MaxCosSim: 0, + MatchMode: matchMode, + } + } + sort.Float64s(similarities) + var sum float64 + for _, v := range similarities { + sum += v + } + return service.EmbeddingCheckSummary{ + KbID: datasetID, + Model: embeddingID, + Sampled: sampled, + Valid: len(similarities), + AvgCosSim: datasetRoundFloat(sum/float64(len(similarities)), 4), + MinCosSim: datasetRoundFloat(similarities[0], 4), + MaxCosSim: datasetRoundFloat(similarities[len(similarities)-1], 4), + MatchMode: matchMode, + } +} + +func datasetRoundFloat(value float64, places int) float64 { + shift := math.Pow(10, float64(places)) + return math.Round(value*shift) / shift +} + +func datasetChunkID(chunk map[string]interface{}) string { + if id, ok := chunk["chunk_id"]; ok { + if s, ok := id.(string); ok { + return s + } + } + return "" +} + +func interfaceSlice(items ...string) []interface{} { + result := make([]interface{}, len(items)) + for i, v := range items { + result[i] = v + } + return result +} diff --git a/internal/service/dataset_create_test.go b/internal/service/dataset_create_test.go deleted file mode 100644 index 7c6297ca82..0000000000 --- a/internal/service/dataset_create_test.go +++ /dev/null @@ -1,207 +0,0 @@ -// -// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. -// - -package service - -import ( - "strings" - "testing" - - "ragflow/internal/common" - "ragflow/internal/dao" - "ragflow/internal/entity" -) - -// testDatasetCreateService builds a DatasetService for CreateDataset tests. -// CreateDataset resolves the tenant via d.tenantDAO, so it must be wired here -// (unlike the update path which does not touch it). -// insertCreateDatasetTenant seeds a tenant with status="1" (required by -// TenantDAO.GetByID, which CreateDataset calls) for the given tenant id. -func insertCreateDatasetTenant(t *testing.T, tenantID string) { - t.Helper() - var existing entity.Tenant - if err := dao.DB.Where("id = ?", tenantID).First(&existing).Error; err != nil { - tn := &entity.Tenant{ - ID: tenantID, - LLMID: "llm-default", - EmbdID: "embd-default", - TenantEmbdID: sptr("embd-1"), - ASRID: "asr-default", - Status: sptr("1"), - } - if err := dao.DB.Create(tn).Error; err != nil { - t.Fatalf("insert test tenant: %v", err) - } - } -} - -func testDatasetCreateService(t *testing.T) *DatasetService { - t.Helper() - - return &DatasetService{ - kbDAO: dao.NewKnowledgebaseDAO(), - documentDAO: dao.NewDocumentDAO(), - connectorDAO: dao.NewConnectorDAO(), - tenantDAO: dao.NewTenantDAO(), - } -} - -func TestCreateDataset_NoComponentParams(t *testing.T) { - db := setupDatasetUpdateTestDB(t) - pushServiceDB(t, db) - insertCreateDatasetTenant(t, "tenant-1") - - chunkMethod := "naive" - result, code, err := testDatasetCreateService(t).CreateDataset(&CreateDatasetRequest{ - Name: "ds-no-cp", - ParserID: &chunkMethod, - }, "tenant-1") - if err != nil { - t.Fatalf("CreateDataset failed: %v", err) - } - if code != common.CodeSuccess { - t.Fatalf("expected success code, got %d", code) - } - if result["parser_id"] != strings.TrimSpace(chunkMethod) { - t.Fatalf("expected parser_id %q, got %#v", chunkMethod, result["parser_id"]) - } -} - -func TestCreateDataset_ComponentParamsPopulated(t *testing.T) { - db := setupDatasetUpdateTestDB(t) - pushServiceDB(t, db) - insertCreateDatasetTenant(t, "tenant-1") - - chunkMethod := "naive" - result, code, err := testDatasetCreateService(t).CreateDataset(&CreateDatasetRequest{ - Name: "ds-with-cp-defaults", - ParserID: &chunkMethod, - }, "tenant-1") - if err != nil { - t.Fatalf("CreateDataset failed: %v", err) - } - if code != common.CodeSuccess { - t.Fatalf("expected success code, got %d", code) - } - - parserConfig, ok := result["parser_config"].(entity.JSONMap) - if !ok { - t.Fatalf("parser_config is not entity.JSONMap, got %T", result["parser_config"]) - } - if len(parserConfig) == 0 { - t.Fatal("parser_config is empty, expected DSL component params defaults") - } - - // Verify at least one component from the general/naive template has defaults. - // The general template has TokenChunker, Tokenizer, Parser, and File components. - found := false - for cpnID := range parserConfig { - if strings.Contains(cpnID, "TokenChunker") { - found = true - break - } - } - if !found { - t.Errorf("parser_config does not contain any TokenChunker: %v", parserConfig) - } -} - -func TestCreateDataset_ParseTypeBuiltinClearsPipelineID(t *testing.T) { - db := setupDatasetUpdateTestDB(t) - pushServiceDB(t, db) - insertCreateDatasetTenant(t, "tenant-1") - // Seed a canvas so it exists, but parse_type=1 should ignore it. - seedDatasetUpdateCanvas(t, "abcdef0123456789abcdef0123456789", "tenant-1", - datasetUpdateCanvasDSL("Parser:HipSignsRhyme", "chunk_token_num")) - - chunkMethod := "naive" - pipelineID := "ABCDEF0123456789ABCDEF0123456789" - parseType := 1 - - result, code, err := testDatasetCreateService(t).CreateDataset(&CreateDatasetRequest{ - Name: "ds-builtin-clears-pipeline", - ParserID: &chunkMethod, - PipelineID: &pipelineID, - ParseType: &parseType, - }, "tenant-1") - if err != nil { - t.Fatalf("CreateDataset failed: %v", err) - } - if code != common.CodeSuccess { - t.Fatalf("expected success code, got %d", code) - } - // parse_type=1 clears pipeline_id → only parser_id should be persisted. - if result["parser_id"] != chunkMethod { - t.Fatalf("expected parser_id %q, got %#v", chunkMethod, result["parser_id"]) - } - if pid, ok := result["pipeline_id"]; ok && pid != nil && pid != "" { - t.Fatalf("expected pipeline_id to be cleared for BuiltIn mode, got %#v", pid) - } -} - -func TestCreateDataset_ParseTypePipelineIgnoresParserID(t *testing.T) { - db := setupDatasetUpdateTestDB(t) - pushServiceDB(t, db) - insertCreateDatasetTenant(t, "tenant-1") - seedDatasetUpdateCanvas(t, "abcdef0123456789abcdef0123456789", "tenant-1", - datasetUpdateCanvasDSL("Parser:CustomP", "chunk_token_num")) - - chunkMethod := "book" - pipelineID := "ABCDEF0123456789ABCDEF0123456789" - parseType := 2 - - result, code, err := testDatasetCreateService(t).CreateDataset(&CreateDatasetRequest{ - Name: "ds-pipeline-ignores-parser", - ParserID: &chunkMethod, - PipelineID: &pipelineID, - ParseType: &parseType, - }, "tenant-1") - if err != nil { - t.Fatalf("CreateDataset failed: %v", err) - } - if code != common.CodeSuccess { - t.Fatalf("expected success code, got %d", code) - } - // parse_type=2 ignores parser_id → pipeline_id should be persisted; - // parser_id should fall back to the default ("naive") since it wasn't set. - if result["pipeline_id"] != strings.ToLower(pipelineID) { - t.Fatalf("expected pipeline_id %q, got %#v", strings.ToLower(pipelineID), result["pipeline_id"]) - } -} - -func TestCreateDataset_RejectsBothWithoutParseType(t *testing.T) { - db := setupDatasetUpdateTestDB(t) - pushServiceDB(t, db) - insertCreateDatasetTenant(t, "tenant-1") - - chunkMethod := "naive" - pipelineID := "abcdef0123456789abcdef0123456789" - _, code, err := testDatasetCreateService(t).CreateDataset(&CreateDatasetRequest{ - Name: "ds-both-no-parse-type", - ParserID: &chunkMethod, - PipelineID: &pipelineID, - // ParseType deliberately nil - }, "tenant-1") - if err == nil { - t.Fatal("expected mutual-exclusivity error when both set without parse_type") - } - if code != common.CodeDataError { - t.Fatalf("expected data error code, got %d", code) - } - if !strings.Contains(err.Error(), "mutually exclusive") { - t.Fatalf("expected error to mention 'mutually exclusive', got: %v", err) - } -} diff --git a/internal/service/dataset_parser_test.go b/internal/service/dataset_parser_test.go deleted file mode 100644 index 6cd104ea4b..0000000000 --- a/internal/service/dataset_parser_test.go +++ /dev/null @@ -1,64 +0,0 @@ -// -// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. -// - -package service - -import ( - "strings" - "testing" -) - -// TestValidateParserID_AcceptsRegistryRefs verifies that every -// canonical builtin pipeline id passes validation. -func TestValidateParserID_AcceptsRegistryRefs(t *testing.T) { - for _, id := range []string{"general", "book", "audio", "qa", "table", "tag"} { - if err := validateParserID(id); err != nil { - t.Errorf("validateParserID(%q) = %v, want nil", id, err) - } - } -} - -// TestValidateParserID_AcceptsNaiveAlias verifies the legacy -// parser_id "naive" still validates (alias for general) so existing -// dataset rows are not rejected. -func TestValidateParserID_AcceptsNaiveAlias(t *testing.T) { - if err := validateParserID("naive"); err != nil { - t.Errorf("validateParserID(naive) = %v, want nil (alias for general)", err) - } -} - -// TestValidateParserID_RejectsUnknown verifies unknown/empty -// values are rejected with an error that lists the valid options. -func TestValidateParserID_RejectsUnknown(t *testing.T) { - for _, id := range []string{"", "unknown", "NAIVE"} { - err := validateParserID(id) - if err == nil { - t.Errorf("validateParserID(%q) = nil, want error", id) - } - } - - // The error message should surface the canonical valid options so the - // caller knows what to send. "general" must appear (canonical), while - // "naive" is an alias and may or may not appear. - err := validateParserID("unknown") - if err == nil { - t.Fatal("expected error for unknown parser id") - } - msg := err.Error() - if !strings.Contains(msg, "general") { - t.Errorf("error message %q should mention general", msg) - } -} diff --git a/internal/service/dataset_types.go b/internal/service/dataset_types.go new file mode 100644 index 0000000000..cb7bb5a40a --- /dev/null +++ b/internal/service/dataset_types.go @@ -0,0 +1,159 @@ +package service + +// TraceIndexRequest is the request structure for tracing an index task. +type TraceIndexRequest struct { + Type string `json:"type" binding:"required"` +} + +// CheckEmbeddingRequest is the request structure for checking embedding compatibility. +type CheckEmbeddingRequest struct { + EmbeddingID string `json:"embd_id" binding:"required"` + CheckNum *int `json:"check_num,omitempty"` +} + +// EmbeddingCheckSummary is the summary of an embedding model compatibility check. +type EmbeddingCheckSummary struct { + KbID string `json:"kb_id"` + Model string `json:"model"` + Sampled int `json:"sampled"` + Valid int `json:"valid"` + AvgCosSim float64 `json:"avg_cos_sim"` + MinCosSim float64 `json:"min_cos_sim"` + MaxCosSim float64 `json:"max_cos_sim"` + MatchMode string `json:"match_mode"` +} + +// EmbeddingCheckResult is one chunk result in an embedding compatibility check. +type EmbeddingCheckResult struct { + ChunkID string `json:"chunk_id"` + DocID string `json:"doc_id,omitempty"` + DocName string `json:"doc_name,omitempty"` + VectorField string `json:"vector_field,omitempty"` + VectorDim int `json:"vector_dim,omitempty"` + CosSim float64 `json:"cos_sim,omitempty"` + Reason string `json:"reason,omitempty"` +} + +// EmbeddingCheckResponse is the response wrapper for embedding checks. +type EmbeddingCheckResponse struct { + Summary EmbeddingCheckSummary `json:"summary"` + Results []EmbeddingCheckResult `json:"results"` +} + +// SearchDatasetsRequest is the request structure for searching chunks across datasets. +type SearchDatasetsRequest struct { + DatasetIDs []string `json:"dataset_ids" binding:"required"` + Question string `json:"question" binding:"required"` + Page *int `json:"page,omitempty"` + Size *int `json:"size,omitempty"` + DocIDs []string `json:"doc_ids,omitempty"` + UseKG *bool `json:"use_kg,omitempty"` + TopK *int `json:"top_k,omitempty"` + CrossLanguages []string `json:"cross_languages,omitempty"` + SearchID *string `json:"search_id,omitempty"` + MetadataFilter map[string]interface{} `json:"meta_data_filter,omitempty"` + RerankID *string `json:"rerank_id,omitempty"` + Keyword *bool `json:"keyword,omitempty"` + SimilarityThreshold *float64 `json:"similarity_threshold,omitempty"` + VectorSimilarityWeight *float64 `json:"vector_similarity_weight,omitempty"` + ForceRefresh bool `json:"force_refresh"` +} + +// SearchDatasetsResponse is the response structure for dataset search results. +type SearchDatasetsResponse struct { + Chunks []map[string]interface{} `json:"chunks"` + DocAggs []map[string]interface{} `json:"doc_aggs"` + Labels *map[string]float64 `json:"labels"` + Total int64 `json:"total"` +} + +// SearchDatasetRequest is the request structure for searching chunks within one dataset. +type SearchDatasetRequest struct { + Question string `json:"question"` + Page *int `json:"page,omitempty"` + Size *int `json:"size,omitempty"` + DocIDs []string `json:"doc_ids,omitempty"` + UseKG *bool `json:"use_kg,omitempty"` + TopK *int `json:"top_k,omitempty"` + CrossLanguages []string `json:"cross_languages,omitempty"` + SearchID *string `json:"search_id,omitempty"` + MetadataFilter map[string]interface{} `json:"meta_data_filter,omitempty"` + RerankID *string `json:"rerank_id,omitempty"` + Keyword *bool `json:"keyword,omitempty"` + SimilarityThreshold *float64 `json:"similarity_threshold,omitempty"` + VectorSimilarityWeight *float64 `json:"vector_similarity_weight,omitempty"` +} + +// ToSearchDatasetsRequest converts a single-dataset search request into the multi-dataset form. +func (req *SearchDatasetRequest) ToSearchDatasetsRequest(datasetID string) *SearchDatasetsRequest { + if req == nil { + return &SearchDatasetsRequest{DatasetIDs: []string{datasetID}} + } + return &SearchDatasetsRequest{ + DatasetIDs: []string{datasetID}, + Question: req.Question, + Page: req.Page, + Size: req.Size, + DocIDs: req.DocIDs, + UseKG: req.UseKG, + TopK: req.TopK, + CrossLanguages: req.CrossLanguages, + SearchID: req.SearchID, + MetadataFilter: req.MetadataFilter, + RerankID: req.RerankID, + Keyword: req.Keyword, + SimilarityThreshold: req.SimilarityThreshold, + VectorSimilarityWeight: req.VectorSimilarityWeight, + } +} + +// MetadataConfigField mirrors one field in the dataset metadata config API. +type MetadataConfigField struct { + Key string `json:"key"` + Type string `json:"type"` + Description *string `json:"description"` + Enum []string `json:"enum"` +} + +// MetadataConfigRequest mirrors PUT /datasets/:dataset_id/metadata/config. +type MetadataConfigRequest struct { + Metadata []MetadataConfigField `json:"metadata"` + BuiltInMetadata []MetadataConfigField `json:"built_in_metadata"` +} + +// CreateDatasetRequest represents the request for creating a dataset. +type CreateDatasetRequest struct { + Name string `json:"name" binding:"required"` + EmbeddingModel *string `json:"embedding_model,omitempty"` + Permission *string `json:"permission,omitempty"` + ParserID *string `json:"parser_id,omitempty"` + PipelineID *string `json:"pipeline_id,omitempty"` + // ParseType indicates pipeline selection mode: 1 = BuiltIn (parser_id), + // 2 = Pipeline (pipeline_id). nil means unspecified (backward compat). + ParseType *int `json:"parse_type,omitempty"` +} + +// DatasetConnectorRequest represents a connector link request. +type DatasetConnectorRequest struct { + ID string `json:"id"` + AutoParse string `json:"auto_parse,omitempty"` +} + +// UpdateDatasetRequest represents the request for updating a dataset. +type UpdateDatasetRequest struct { + Name *string `json:"name,omitempty"` + Avatar *string `json:"avatar,omitempty"` + Description *string `json:"description,omitempty"` + Language *string `json:"language,omitempty"` + Connectors *[]DatasetConnectorRequest `json:"connectors,omitempty"` + EmbdID *string `json:"embd_id,omitempty"` + EmbeddingModel *string `json:"embedding_model,omitempty"` + Permission *string `json:"permission,omitempty"` + ParserID *string `json:"parser_id,omitempty"` + Pagerank *int64 `json:"pagerank,omitempty"` + ParserConfig map[string]interface{} `json:"parser_config,omitempty"` + PipelineID *string `json:"pipeline_id,omitempty"` + // ParseType indicates pipeline selection mode: 1 = BuiltIn (parser_id), + // 2 = Pipeline (pipeline_id). nil means unspecified (backward compat). + ParseType *int `json:"parse_type,omitempty"` +} diff --git a/internal/service/document.go b/internal/service/document.go deleted file mode 100644 index 2928e55b4c..0000000000 --- a/internal/service/document.go +++ /dev/null @@ -1,3652 +0,0 @@ -// -// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. -// - -package service - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "io" - "math" - "mime/multipart" - "net/http" - "path/filepath" - "reflect" - "regexp" - "sort" - "strconv" - "strings" - "time" - - "ragflow/internal/common" - "ragflow/internal/dao" - "ragflow/internal/engine" - enginetypes "ragflow/internal/engine/types" - "ragflow/internal/entity" - pipelinepkg "ragflow/internal/ingestion/pipeline" - "ragflow/internal/storage" - "ragflow/internal/tokenizer" - "ragflow/internal/utility" - - "go.uber.org/zap" - "gorm.io/gorm" - "gorm.io/gorm/clause" -) - -// DocumentService document service -type DocumentService struct { - documentDAO *dao.DocumentDAO - kbDAO *dao.KnowledgebaseDAO - ingestionTaskDAO *dao.IngestionTaskDAO - ingestionTaskLogDAO *dao.IngestionTaskLogDAO - ingestionTaskSvc *IngestionTaskService - docEngine engine.DocEngine - metadataSvc *MetadataService - taskDAO *dao.TaskDAO - file2DocumentDAO *dao.File2DocumentDAO - fileDAO *dao.FileDAO - canvasDAO *dao.UserCanvasDAO - api4ConvDAO *dao.API4ConversationDAO -} - -// NewDocumentService create document service -func NewDocumentService() *DocumentService { - publisher := NewMessageQueueTaskPublisher() - ingestionTaskSvc := NewIngestionTaskService() - ingestionTaskSvc.SetTaskPublisher(publisher) - return &DocumentService{ - documentDAO: dao.NewDocumentDAO(), - ingestionTaskDAO: dao.NewIngestionTaskDAO(), - ingestionTaskLogDAO: dao.NewIngestionTaskLogDAO(), - ingestionTaskSvc: ingestionTaskSvc, - kbDAO: dao.NewKnowledgebaseDAO(), - docEngine: engine.Get(), - metadataSvc: NewMetadataService(), - taskDAO: dao.NewTaskDAO(), - file2DocumentDAO: dao.NewFile2DocumentDAO(), - fileDAO: dao.NewFileDAO(), - canvasDAO: dao.NewUserCanvasDAO(), - api4ConvDAO: dao.NewAPI4ConversationDAO(), - } -} - -// CreateDocumentRequest create document request -// UpdateDocumentRequest update document request -type UpdateDocumentRequest struct { - Name *string `json:"name"` - Run *string `json:"run"` - TokenNum *int64 `json:"token_num"` - ChunkNum *int64 `json:"chunk_num"` - Progress *float64 `json:"progress"` - ProgressMsg *string `json:"progress_msg"` -} - -// DocumentResponse document response -type DocumentResponse struct { - ID string `json:"id"` - Name *string `json:"name,omitempty"` - KbID string `json:"kb_id"` - ParserID string `json:"parser_id"` - PipelineID *string `json:"pipeline_id,omitempty"` - Type string `json:"type"` - SourceType string `json:"source_type"` - CreatedBy string `json:"created_by"` - Location *string `json:"location,omitempty"` - Size int64 `json:"size"` - TokenNum int64 `json:"token_num"` - ChunkNum int64 `json:"chunk_num"` - Progress float64 `json:"progress"` - ProgressMsg *string `json:"progress_msg,omitempty"` - ProcessDuration float64 `json:"process_duration"` - Suffix string `json:"suffix"` - Run *string `json:"run,omitempty"` - Status *string `json:"status,omitempty"` - CreatedAt string `json:"created_at"` - UpdatedAt string `json:"updated_at"` -} - -type ThumbnailResponse struct { - ID string `json:"id"` - Thumbnail *string `json:"thumbnail,omitempty"` - KbID string `json:"kb_id"` -} - -const imgBase64Prefix = "data:image/png;base64," - -type ArtifactResponse struct { - Data []byte - ContentType string - SafeFilename string - ForceAttachment bool -} - -type UpdateDatasetDocumentRequest struct { - Name *string `json:"name"` - ParserID *string `json:"parser_id"` - ChunkCount *int64 `json:"chunk_count"` - TokenCount *int64 `json:"token_count"` - PipelineID *string `json:"pipeline_id"` - Enabled *int `json:"enabled"` - Progress *float64 `json:"progress"` - ParserConfig map[string]any `json:"parser_config"` - MetaFields map[string]any `json:"meta_fields"` -} - -// PATCH /api/v1/datasets/:dataset_id/documents/:document_id. -type UpdateDatasetDocumentResponse struct { - ID string `json:"id"` - Thumbnail *string `json:"thumbnail,omitempty"` - DatasetID string `json:"dataset_id"` - ParserID string `json:"parser_id"` - PipelineID *string `json:"pipeline_id,omitempty"` - ParserConfig map[string]interface{} `json:"parser_config"` - SourceType string `json:"source_type"` - Type string `json:"type"` - CreatedBy string `json:"created_by"` - Name *string `json:"name,omitempty"` - Location *string `json:"location,omitempty"` - Size int64 `json:"size"` - TokenCount int64 `json:"token_count"` - ChunkCount int64 `json:"chunk_count"` - Progress float64 `json:"progress"` - ProgressMsg *string `json:"progress_msg,omitempty"` - ProcessBeginAt *time.Time `json:"process_begin_at,omitempty"` - ProcessDuration float64 `json:"process_duration"` - ContentHash *string `json:"content_hash,omitempty"` - MetaFields map[string]interface{} `json:"meta_fields,omitempty"` - Suffix string `json:"suffix"` - Run string `json:"run"` - Status *string `json:"status,omitempty"` - CreateTime *int64 `json:"create_time,omitempty"` - CreateDate *time.Time `json:"create_date,omitempty"` - UpdateTime *int64 `json:"update_time,omitempty"` - UpdateDate *time.Time `json:"update_date,omitempty"` -} - -var ( - ErrArtifactInvalidFilename = errors.New("Invalid filename.") - ErrArtifactInvalidFileType = errors.New("Invalid file type.") - ErrArtifactNotFound = errors.New("Artifact not found.") -) - -var artifactContentTypes = map[string]string{ - ".png": "image/png", - ".jpg": "image/jpeg", - ".jpeg": "image/jpeg", - ".svg": "image/svg+xml", - ".pdf": "application/pdf", - ".csv": "text/csv", - ".json": "application/json", - ".html": "text/html", -} - -var artifactForceAttachmentExtensions = map[string]struct{}{ - ".htm": {}, - ".html": {}, - ".shtml": {}, - ".xht": {}, - ".xhtml": {}, - ".xml": {}, - ".mhtml": {}, - ".svg": {}, -} -var artifactForceAttachmentContentTypes = map[string]struct{}{ - "text/html": {}, - "image/svg+xml": {}, - "application/xhtml+xml": {}, - "text/xml": {}, - "application/xml": {}, - "multipart/related": {}, -} - -var artifactUnsafeFilenameChars = regexp.MustCompile(`[^\pL\pN_.-]`) - -// GetDocumentImage retrieves an image object from storage. -func (s *DocumentService) GetDocumentImage(imageID string) ([]byte, error) { - parts := strings.SplitN(imageID, "-", 2) - if len(parts) != 2 || parts[0] == "" || parts[1] == "" { - return nil, fmt.Errorf("Image not found.") - } - - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return nil, fmt.Errorf("storage not initialized") - } - - return storageImpl.Get(parts[0], parts[1]) -} - -// GetDocumentArtifact retrieves a sandbox artifact from object storage. -// -// userID scopes the lookup: a CodeExec sandbox artifact is only -// returned when the caller owns (or has team access to) at least -// one agent session whose `message` references this filename (or -// its `documents/artifact/` URL form). The authorization -// gate runs BEFORE the storage read so a probe of an unknown -// filename cannot distinguish "you cannot see it" from "it -// exists" — both return ErrArtifactNotFound. Mirrors PR #16169. -func (s *DocumentService) GetDocumentArtifact(filename, userID string) (*ArtifactResponse, error) { - basename := filepath.Base(filename) - if basename != filename || strings.Contains(filename, "/") || strings.Contains(filename, "\\") { - return nil, ErrArtifactInvalidFilename - } - - ext := strings.ToLower(filepath.Ext(basename)) - contentType, ok := artifactContentTypes[ext] - if !ok { - return nil, ErrArtifactInvalidFileType - } - - if !s.sandboxArtifactAccessible(basename, userID) { - // Same error as "object does not exist" to avoid leaking - // whether the artifact exists for a different user/agent. - return nil, ErrArtifactNotFound - } - - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return nil, fmt.Errorf("storage not initialized") - } - - bucket := sandboxArtifactBucket() - if !storageImpl.ObjExist(bucket, basename) { - return nil, ErrArtifactNotFound - } - - data, err := storageImpl.Get(bucket, basename) - if err != nil { - return nil, err - } - if len(data) == 0 { - return nil, ErrArtifactNotFound - } - - return &ArtifactResponse{ - Data: data, - ContentType: contentType, - SafeFilename: sanitizeArtifactFilename(basename), - ForceAttachment: shouldForceArtifactAttachment(ext, contentType), - }, nil -} - -// sandboxArtifactDialogIDsForUser returns the distinct agent -// (canvas) dialog_ids for sessions owned by userID whose -// `message` blob references filename. A CodeExec artifact URL -// appears in `message` as either a bare filename or the -// `documents/artifact/` form, so the helper matches both. -// -// Implemented as a direct GORM query on the -// API4Conversation table — GORM's `Contains` maps to MySQL -// `LIKE '%...%'` which is fine here because the storage path is -// short and indexed lookup on (user_id, exp_user_id) keeps the -// scan narrow. -func (s *DocumentService) sandboxArtifactDialogIDsForUser(filename, userID string) []string { - if filename == "" || userID == "" { - return nil - } - // Escape SQL LIKE wildcards (%, _) before building the pattern. - // Without escaping, a caller could submit a filename like - // "%.png" or "_" and the LIKE query would match arbitrary - // referenced artifacts in any user's conversation — letting the - // caller pass the authorization check against one filename and - // then GET another artifact by name (PR review round 5, Major #8). - // - // Escape character: '!'. We avoid '\\' because SQL string - // literal parsing of '\\' is driver-specific (SQLite treats - // it as a single backslash, MySQL treats it as one, Postgres - // rejects the unterminated string) — '!' is a benign character - // in real filenames (artifact names rarely contain '!') and - // parses identically in every driver. - filenameSafe := escapeSQLLikePattern(filename) - artifactRefSafe := escapeSQLLikePattern("documents/artifact/" + filename) - filenamePattern := "%" + filenameSafe + "%" - artifactRefPattern := "%" + artifactRefSafe + "%" - dialogIDs := make(map[string]struct{}) - rows, err := dao.DB.Model(&entity.API4Conversation{}). - Select("dialog_id"). - Where("user_id = ? OR exp_user_id = ?", userID, userID). - Where(`message LIKE ? ESCAPE '!' OR message LIKE ? ESCAPE '!'`, - filenamePattern, artifactRefPattern). - Distinct("dialog_id"). - Rows() - if err != nil { - return nil - } - defer rows.Close() - for rows.Next() { - var d string - if err := rows.Scan(&d); err == nil && d != "" { - dialogIDs[d] = struct{}{} - } - } - out := make([]string, 0, len(dialogIDs)) - for d := range dialogIDs { - out = append(out, d) - } - return out -} - -// escapeSQLLikePattern escapes the SQL LIKE wildcards ('%', '_') and -// the escape character itself ('!') so a literal user-supplied -// filename can be safely interpolated into a `LIKE ? ESCAPE '!'` -// pattern. Without this, "%.png" would match any string ending in -// ".png" and "_" would match a single character — bypassing the -// filename-specific authorization check. PR review round 5, Major #8. -func escapeSQLLikePattern(s string) string { - r := strings.NewReplacer(`!`, `!!`, `%`, `!%`, `_`, `!_`) - return r.Replace(s) -} - -// sandboxArtifactAccessible reports whether userID may reach at -// least one agent canvas whose session references filename. -// Mirrors `UserCanvasService.accessible(dialog_id, user_id)` from -// the Python fix; on the Go side this is the same predicate as -// UserCanvasDAO.Accessible (owner or team permission, with the -// latter scoped to the caller's tenant membership — PR review -// round 5). -func (s *DocumentService) sandboxArtifactAccessible(filename, userID string) bool { - if userID == "" { - return false - } - // Fetch the caller's tenant list once; passing it into - // canvasDAO.Accessible ensures the team-permission branch only - // matches canvases the caller can actually see. An empty list - // (callers without tenant data) is safe — it effectively disables - // the team branch, so the only matches are canvases the caller - // directly owns. - tenantIDs, terr := dao.NewUserTenantDAO().GetTenantIDsByUserID(userID) - if terr != nil { - tenantIDs = nil - } - for _, dialogID := range s.sandboxArtifactDialogIDsForUser(filename, userID) { - if s.canvasDAO.Accessible(dialogID, userID, tenantIDs) { - return true - } - } - return false -} - -func sandboxArtifactBucket() string { - if bucket := common.GetEnv(common.EnvSandboxArtifactBucket); bucket != "" { - return bucket - } - return "sandbox-artifacts" -} - -// Accessible reports whether docID belongs to a knowledge base -// reachable by userID. Used by agent endpoints (e.g. RerunAgent, -// PR #15145) to gate destructive / run-again actions on a document -// the caller has access to. Returns false on any lookup failure or -// empty inputs so callers can treat a denial as a 404-equivalent -// and avoid leaking whether the document exists at all. -func (s *DocumentService) Accessible(docID, userID string) bool { - if docID == "" || userID == "" { - return false - } - doc, err := s.documentDAO.GetByID(docID) - if err != nil || doc == nil { - return false - } - return s.kbDAO.Accessible(doc.KbID, userID) -} - -func sanitizeArtifactFilename(filename string) string { - return artifactUnsafeFilenameChars.ReplaceAllString(filename, "_") -} - -func shouldForceArtifactAttachment(ext, contentType string) bool { - if _, ok := artifactForceAttachmentExtensions[strings.ToLower(ext)]; ok { - return true - } - _, ok := artifactForceAttachmentContentTypes[strings.ToLower(contentType)] - return ok -} - -type DocumentPreview struct { - Data []byte - ContentType string - FileName string -} - -func (s *DocumentService) GetDocumentPreview(docID string) (*DocumentPreview, error) { - doc, err := s.documentDAO.GetByID(docID) - if err != nil { - return nil, err - } - - bucket, name, err := s.GetDocumentStorageAddress(doc) - if err != nil { - return nil, err - } - - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return nil, fmt.Errorf("storage not initialized") - } - - data, err := storageImpl.Get(bucket, name) - if err != nil { - return nil, err - } - if len(data) == 0 { - return nil, ErrArtifactNotFound - } - - fileName := "" - if doc.Name != nil { - fileName = *doc.Name - } - - ext := utility.GetFileExtension(fileName) - contentType := utility.GetContentType(ext, doc.Type) - - return &DocumentPreview{ - Data: data, - ContentType: contentType, - FileName: fileName, - }, nil -} - -func (s *DocumentService) GetDocumentStorageAddress(doc *entity.Document) (string, string, error) { - if doc == nil { - return "", "", fmt.Errorf("document is nil") - } - - file2DocumentDAO := dao.NewFile2DocumentDAO() - fileDAO := dao.NewFileDAO() - - mappings, err := file2DocumentDAO.GetByDocumentID(doc.ID) - if err != nil { - return "", "", err - } - - if len(mappings) > 0 && mappings[0].FileID != nil { - file, err := fileDAO.GetByID(*mappings[0].FileID) - if err != nil { - return "", "", err - } - - if file.SourceType == "" || entity.FileSource(file.SourceType) == entity.FileSourceLocal { - if file.Location == nil || *file.Location == "" { - return "", "", fmt.Errorf("file location is empty") - } - return file.ParentID, *file.Location, nil - } - } - - if doc.Location == nil || *doc.Location == "" { - return "", "", fmt.Errorf("document location is empty") - } - return doc.KbID, *doc.Location, nil -} - -type DownloadDocumentResp struct { - Data []byte - FileName string - ContentType string -} - -func (s *DocumentService) DownloadDocument(datasetID, docID string) (*DownloadDocumentResp, error) { - if docID == "" { - return nil, fmt.Errorf("Specify document_id please.") - } - doc, err := s.documentDAO.GetByID(docID) - if err != nil || doc.KbID != datasetID { - return nil, fmt.Errorf("The dataset not own the document %s.", docID) - } - bucket, name, err := s.GetDocumentStorageAddress(doc) - if err != nil { - return nil, err - } - - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return nil, fmt.Errorf("storage not initialized") - } - - data, err := storageImpl.Get(bucket, name) - if err != nil { - return nil, err - } - if len(data) == 0 { - return nil, fmt.Errorf("This file is empty.") - } - - fileName := "" - if doc.Name != nil { - fileName = *doc.Name - } - - return &DownloadDocumentResp{ - Data: data, - FileName: fileName, - ContentType: "application/octet-stream", - }, nil -} - -// CreateDocument create document -// GetDocumentByID get document by ID -func (s *DocumentService) GetDocumentByID(id string) (*DocumentResponse, error) { - document, err := s.documentDAO.GetByID(id) - if err != nil { - return nil, err - } - - return s.toResponse(document), nil -} - -// UpdateDocument update document -func (s *DocumentService) UpdateDocument(id string, req *UpdateDocumentRequest) error { - document, err := s.documentDAO.GetByID(id) - if err != nil { - return err - } - - if req.Name != nil { - document.Name = req.Name - } - if req.Run != nil { - document.Run = req.Run - } - if req.TokenNum != nil { - document.TokenNum = *req.TokenNum - } - if req.ChunkNum != nil { - document.ChunkNum = *req.ChunkNum - } - if req.Progress != nil { - document.Progress = *req.Progress - } - if req.ProgressMsg != nil { - document.ProgressMsg = req.ProgressMsg - } - - return s.documentDAO.Update(document) -} - -// IncrementChunkNum atomically increments chunk/token counters on the document and its knowledge base in a transaction -func (s *DocumentService) IncrementChunkNum(docID, kbID string, chunkNum, tokenNum int, duration float64) error { - return dao.DB.Transaction(func(tx *gorm.DB) error { - // Update document - if err := tx.Model(&entity.Document{}). - Where("id = ? AND kb_id = ?", docID, kbID). - Updates(map[string]interface{}{ - "chunk_num": gorm.Expr("chunk_num + ?", int64(chunkNum)), - "token_num": gorm.Expr("token_num + ?", int64(tokenNum)), - "process_duration": gorm.Expr("process_duration + ?", duration), - }).Error; err != nil { - return err - } - - // Update knowledgebase - if err := tx.Model(&entity.Knowledgebase{}). - Where("id = ?", kbID). - Updates(map[string]interface{}{ - "chunk_num": gorm.Expr("chunk_num + ?", int64(chunkNum)), - "token_num": gorm.Expr("token_num + ?", int64(tokenNum)), - }).Error; err != nil { - return err - } - - return nil - }) -} - -// UpdateRunProgress mirrors a pipeline run's live progress into the document -// row so the document-list endpoint (which reads document.progress/run/ -// progress_msg) reflects in-flight Go pipeline progress. Best-effort by -// design; callers log and continue on error. -func (s *DocumentService) UpdateRunProgress(docID string, progress float64, run, progressMsg string) error { - return s.documentDAO.UpdateByID(docID, map[string]interface{}{ - "progress": progress, - "run": run, - "progress_msg": progressMsg, - }) -} - -// DeleteDocument delete document — delegates to full cleanup logic. -func (s *DocumentService) DeleteDocument(id string) error { - return s.deleteDocumentFull(id) -} - -// DeleteDocuments deletes multiple documents under a dataset. -// -// ids: specific document IDs; deleteAll: delete all docs in the dataset. -// Returns the number of successfully deleted documents. -func (s *DocumentService) DeleteDocuments(ids []string, deleteAll bool, datasetID, userID string) (int, error) { - // 1. Check dataset is accessible by the user - if !s.kbDAO.Accessible(datasetID, userID) { - return 0, fmt.Errorf("You don't own the dataset %s.", datasetID) - } - - // 2. Resolve document IDs - if deleteAll { - if err := dao.DB.Model(&entity.Document{}). - Where("kb_id = ?", datasetID). - Pluck("id", &ids).Error; err != nil { - return 0, fmt.Errorf("failed to query documents: %w", err) - } - } - if len(ids) == 0 { - return 0, nil - } - - // 3. Deduplicate (before validation so dup count doesn't matter) - ids = common.Deduplicate(ids) - - // 4. Validate IDs belong to this dataset (only for explicit ids; deleteAll is already scoped) - if !deleteAll { - if _, err := s.validateDocsInDataset(ids, datasetID); err != nil { - return 0, err - } - } - - // 5. Delete each document (non-critical failures are tolerated per doc) - deleted := 0 - for _, docID := range ids { - if err := s.deleteDocumentFull(docID); err != nil { - common.Warn(fmt.Sprintf("DeleteDocuments: failed to delete %s: %v", docID, err)) - continue - } - deleted++ - } - - return deleted, nil -} - -// deleteDocumentFull performs full document cleanup. Non-critical failures -// are tolerated (logged and continue). Critical failures (e.g. document or -// KB not found) return an error immediately. -func (s *DocumentService) deleteDocumentFull(docID string) error { - doc, kb, err := s.resolveDocAndKB(docID) - if err != nil { - return err - } - - // Delete tasks from DB - ingestionTask, err := s.ingestionTaskDAO.GetByDocumentID(docID) - if err != nil { - common.Error(fmt.Sprintf("failed to get ingestion task by doc:%s", doc.ID), err) - return err - } - if ingestionTask != nil { - taskInfo, err := s.ingestionTaskSvc.Remove(ingestionTask.ID, &ingestionTask.UserID) - if err != nil { - return err - } - // FIXME: need to add logic to delete files in taskInfo - common.Warn(fmt.Sprintf("need to delete files from taskInfo: %v", taskInfo)) - } - - s.deleteDocEngineData(docID, kb.TenantID, doc.KbID) - if err := s.deleteDocRecordWithCounters(doc, kb.ID); err != nil { - return err - } - s.cleanupFileReferences(docID) - - return nil -} - -// RemoveDocumentKeepFile removes a document's chunks/metadata and the document -// row, decrementing the KB counters (doc_num/chunk_num/token_num), WITHOUT -// deleting the underlying file record, its storage blob, or its file2document -// mappings. Mirrors Python DocumentService.remove_document — the caller is -// responsible for cleaning up the file2document mappings separately. -func (s *DocumentService) RemoveDocumentKeepFile(docID string) error { - doc, kb, err := s.resolveDocAndKB(docID) - if err != nil { - return err - } - if _, delErr := s.taskDAO.DeleteByDocIDs([]string{docID}); delErr != nil { - common.Logger.Warn(fmt.Sprintf("RemoveDocumentKeepFile: failed to delete tasks for %s: %v", docID, delErr)) - } - s.deleteDocEngineData(docID, kb.TenantID, doc.KbID) - return s.deleteDocRecordWithCounters(doc, kb.ID) -} - -// InsertDocument creates a document row and increments the owning KB's doc_num -// counter in a single transaction. Mirrors Python DocumentService.insert, which -// updates dataset/document counters on insert. The document's ID and timestamps -// are populated by the caller / model hooks before insertion. -func (s *DocumentService) InsertDocument(doc *entity.Document) error { - return dao.DB.Transaction(func(tx *gorm.DB) error { - if err := tx.Create(doc).Error; err != nil { - return fmt.Errorf("failed to create document: %w", err) - } - // Guard the counter bump with RowsAffected: documents.kb_id has no DB-level - // FK, so Create can succeed against a non-existent KB and the Update would - // then report a nil error with 0 rows touched, silently desyncing doc_num. - // Roll the whole transaction back in that case (mirrors the counter checks - // in deleteDocRecordWithCounters). - result := tx.Model(&entity.Knowledgebase{}). - Where("id = ?", doc.KbID). - Update("doc_num", gorm.Expr("doc_num + 1")) - if result.Error != nil { - return fmt.Errorf("failed to increment doc_num for KB %s: %w", doc.KbID, result.Error) - } - if result.RowsAffected == 0 { - return fmt.Errorf("knowledgebase %s not found", doc.KbID) - } - return nil - }) -} - -// resolveDocAndKB loads the document and its knowledgebase, returning both or -// an error. -func (s *DocumentService) resolveDocAndKB(docID string) (*entity.Document, *entity.Knowledgebase, error) { - doc, err := s.documentDAO.GetByID(docID) - if err != nil { - return nil, nil, fmt.Errorf("document not found: %w", err) - } - kb, err := s.kbDAO.GetByID(doc.KbID) - if err != nil { - return nil, nil, fmt.Errorf("knowledgebase not found: %w", err) - } - return doc, kb, nil -} - -// deleteDocEngineData removes chunks and metadata from the document engine. -// No-op when the engine is nil. -func (s *DocumentService) deleteDocEngineData(docID, tenantID, kbID string) { - if s.docEngine == nil { - return - } - 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)) - } - if s.metadataSvc != nil { - _ = s.DeleteDocumentAllMetadata(docID) // logs internally - } -} - -// deleteDocRecordWithCounters hard-deletes the document row and decrements the -// KB counters in a single transaction. Counters are only decremented when a -// document row was actually removed (RowsAffected > 0), guarding against -// double-decrement on retries or concurrent deletes. -func (s *DocumentService) deleteDocRecordWithCounters(doc *entity.Document, kbID string) error { - return dao.DB.Transaction(func(tx *gorm.DB) error { - result := tx.Where("id = ?", doc.ID).Delete(&entity.Document{}) - if result.Error != nil { - return fmt.Errorf("failed to delete document %s: %w", doc.ID, result.Error) - } - if result.RowsAffected == 0 { - return nil // already deleted by a concurrent request — skip counters - } - - result = tx.Model(&entity.Knowledgebase{}). - Where("id = ?", kbID). - Updates(map[string]interface{}{ - "doc_num": gorm.Expr("doc_num - 1"), - "chunk_num": gorm.Expr("chunk_num - ?", doc.ChunkNum), - "token_num": gorm.Expr("token_num - ?", doc.TokenNum), - }) - if result.Error != nil { - return fmt.Errorf("failed to decrement counters for KB %s: %w", kbID, result.Error) - } - if result.RowsAffected == 0 { - return fmt.Errorf("knowledgebase %s not found", kbID) - } - return nil - }) -} - -func (s *DocumentService) rollbackAddFileFromKBError(doc *entity.Document, kbID string, err error) error { - if cleanupErr := s.deleteDocRecordWithCounters(doc, kbID); cleanupErr != nil { - return fmt.Errorf("%w; rollback cleanup failed: %w", err, cleanupErr) - } - return err -} - -// cleanupFileReferences deletes file2document mappings for docID, and for each -// referenced file, only hard-deletes the file record and its storage blob when -// no other document still references the same file_id. -func (s *DocumentService) cleanupFileReferences(docID string) { - mappings, mapErr := s.file2DocumentDAO.GetByDocumentID(docID) - if mapErr != nil { - common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: failed to get f2d mappings for %s: %v", docID, mapErr)) - } - if len(mappings) == 0 { - return - } - - // Collect unique file_ids - seen := make(map[string]bool) - var fileIDs []string - for _, m := range mappings { - if m.FileID == nil || seen[*m.FileID] { - continue - } - seen[*m.FileID] = true - fileIDs = append(fileIDs, *m.FileID) - } - - // Delete all file2document rows for this document - if delErr := s.file2DocumentDAO.DeleteByDocumentID(docID); delErr != nil { - common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: failed to delete f2d for %s: %v", docID, delErr)) - } - - // For each file, only delete the record and blob when no other doc references it - for _, fileID := range fileIDs { - remaining, remErr := s.file2DocumentDAO.GetByFileID(fileID) - if remErr != nil { - common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: failed to check remaining f2d for %s: %v", fileID, remErr)) - continue - } - if len(remaining) > 0 { - continue - } - - fileDAO := dao.NewFileDAO() - file, fErr := fileDAO.GetByID(fileID) - if fErr != nil || file == nil { - common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: file not found %s: %v", fileID, fErr)) - continue - } - if _, delErr := fileDAO.DeleteByIDs([]string{fileID}); delErr != nil { - common.Logger.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)) - } - } - } - } -} - -// ListDocuments list documents -func (s *DocumentService) ListDocuments(page, pageSize int) ([]*DocumentResponse, int64, error) { - offset := (page - 1) * pageSize - documents, total, err := s.documentDAO.List(offset, pageSize) - if err != nil { - return nil, 0, err - } - - responses := make([]*DocumentResponse, len(documents)) - for i, doc := range documents { - responses[i] = s.toResponse(doc) - } - - return responses, total, nil -} - -func (s *DocumentService) GetThumbnails(userID string, docIDs []string) (map[string]string, error) { - if len(docIDs) == 0 { - return map[string]string{}, nil - } - - tenantIDs := []string{userID} - if userID != "" { - ids, err := dao.NewUserTenantDAO().GetTenantIDsByUserID(userID) - if err != nil { - return nil, fmt.Errorf("failed to fetch user tenants: %w", err) - } - tenantIDs = append(tenantIDs, ids...) - } - - documents, err := s.documentDAO.GetByIDsAndTenantIDs(docIDs, tenantIDs) - if err != nil { - return nil, fmt.Errorf("failed to fetch document thumbnails: %w", err) - } - - result := make(map[string]string, len(documents)) - for _, document := range documents { - if document == nil { - continue - } - - thumbnail := "" - if document.Thumbnail != nil && *document.Thumbnail != "" { - if strings.HasPrefix(*document.Thumbnail, imgBase64Prefix) { - thumbnail = *document.Thumbnail - } else { - thumbnail = fmt.Sprintf( - "/api/v1/documents/images/%s-%s", - document.KbID, - *document.Thumbnail, - ) - } - } - - result[document.ID] = thumbnail - } - - return result, nil -} - -func (s *DocumentService) BatchUpdateDocumentStatus(userID, datasetID, status string, documentIDs []string) (map[string]interface{}, common.ErrorCode, error) { - kb, err := s.kbDAO.GetByIDAndTenantID(datasetID, userID) - if err != nil { - return nil, common.CodeDataError, fmt.Errorf("You don't own the dataset.") - } - statusInt, convErr := strconv.Atoi(status) - if convErr != nil { - return nil, common.CodeArgumentError, fmt.Errorf("invalid status: %s", status) - } - - result := make(map[string]interface{}, len(documentIDs)) - hasError := false - - documents, err := s.documentDAO.GetByIDs(documentIDs) - if err != nil { - return nil, common.CodeServerError, fmt.Errorf("failed to fetch documents: %w", err) - } - documentByID := make(map[string]*entity.Document, len(documents)) - for _, doc := range documents { - documentByID[doc.ID] = doc - } - - for _, docID := range documentIDs { - doc, ok := documentByID[docID] - if !ok { - result[docID] = map[string]string{"error": "Document not found"} - hasError = true - continue - } - - if doc.KbID != datasetID { - result[docID] = map[string]string{"error": "Document not found in this dataset."} - hasError = true - continue - } - - currentStatus := "" - if doc.Status != nil { - currentStatus = *doc.Status - } - if currentStatus == status { - result[docID] = map[string]string{"status": status} - continue - } - previousStatus := interface{}(nil) - if doc.Status != nil { - previousStatus = *doc.Status - } - if err := s.documentDAO.UpdateByID(docID, map[string]interface{}{"status": status}); err != nil { - result[docID] = map[string]string{"error": "Database error (Document update)!"} - hasError = true - continue - } - - if doc.ChunkNum > 0 { - if s.docEngine == nil { - _ = s.documentDAO.UpdateByID(docID, map[string]interface{}{"status": previousStatus}) - result[docID] = map[string]string{"error": "Document store update failed: document engine not initialized"} - hasError = true - continue - } - err := s.docEngine.UpdateChunks( - context.Background(), - map[string]interface{}{"doc_id": docID}, - map[string]interface{}{"available_int": statusInt}, - fmt.Sprintf("ragflow_%s", kb.TenantID), - doc.KbID, - ) - if err != nil { - _ = s.documentDAO.UpdateByID(docID, map[string]interface{}{"status": previousStatus}) - msg := err.Error() - if strings.Contains(msg, "3022") { - result[docID] = map[string]string{"error": "Document store table missing."} - } else { - result[docID] = map[string]string{"error": "Document store update failed: " + msg} - } - hasError = true - continue - } - } - result[docID] = map[string]string{"status": status} - } - - if hasError { - return result, common.CodeServerError, fmt.Errorf("Partial failure") - } - return result, common.CodeSuccess, nil -} - -// ListDocumentsByDatasetID list documents by knowledge base ID -func (s *DocumentService) ListDocumentsByDatasetID(kbID, keywords string, page, pageSize int) ([]*entity.DocumentListItem, int64, error) { - return s.ListDocumentsByDatasetIDWithOptions(dao.DocumentListOptions{ - KbID: kbID, - Keywords: keywords, - OrderBy: "create_time", - Desc: true, - }, page, pageSize) -} - -// ListDocumentsByDatasetIDWithOptions lists documents by knowledge base ID with filters. -func (s *DocumentService) ListDocumentsByDatasetIDWithOptions(opts dao.DocumentListOptions, page, pageSize int) ([]*entity.DocumentListItem, int64, error) { - opts.Offset = (page - 1) * pageSize - opts.Limit = pageSize - if opts.OrderBy == "" { - opts.OrderBy = "create_time" - } - documents, total, err := s.documentDAO.ListByKBIDWithOptions(opts) - if err != nil { - return nil, 0, err - } - - responses := make([]*entity.DocumentListItem, len(documents)) - for i, doc := range documents { - responses[i] = doc - } - - return responses, total, nil -} - -// GetDocumentFiltersByDatasetID returns aggregate filter values for documents in a dataset. -func (s *DocumentService) GetDocumentFiltersByDatasetID(opts dao.DocumentListOptions) (map[string]interface{}, int64, error) { - filters, total, err := s.documentDAO.GetFilterByKBID(opts) - if err != nil { - return nil, 0, err - } - docIDs, err := s.documentDAO.ListIDsByKBIDWithOptions(opts) - if err != nil { - return nil, 0, err - } - metadataFilter, err := s.getDocumentMetadataFilter(opts.KbID, docIDs) - if err != nil { - return nil, 0, err - } - filters["metadata"] = metadataFilter - return filters, total, nil -} - -func (s *DocumentService) getDocumentMetadataFilter(kbID string, docIDs []string) (map[string]interface{}, error) { - metadataByKey, err := s.GetMetadataByKBs([]string{kbID}) - if err != nil { - return nil, err - } - candidateSet := make(map[string]bool, len(docIDs)) - for _, docID := range docIDs { - candidateSet[docID] = true - } - - metadataCounter := map[string]interface{}{} - docIDsWithMetadata := map[string]bool{} - for key, rawValues := range metadataByKey { - values, ok := rawValues.(map[string][]string) - if !ok { - continue - } - valueCounter := map[string]int64{} - for value, valueDocIDs := range values { - for _, docID := range valueDocIDs { - if !candidateSet[docID] { - continue - } - valueCounter[value]++ - docIDsWithMetadata[docID] = true - } - } - if len(valueCounter) > 0 { - metadataCounter[key] = valueCounter - } - } - metadataCounter["empty_metadata"] = map[string]int64{"true": int64(len(docIDs) - len(docIDsWithMetadata))} - return metadataCounter, nil -} - -// ListDocumentIDsByDatasetIDWithOptions lists matching document IDs without pagination. -func (s *DocumentService) ListDocumentIDsByDatasetIDWithOptions(opts dao.DocumentListOptions) ([]string, error) { - return s.documentDAO.ListIDsByKBIDWithOptions(opts) -} - -// GetDocumentsByAuthorID get documents by author ID -func (s *DocumentService) GetDocumentsByAuthorID(authorID, page, pageSize int) ([]*DocumentResponse, int64, error) { - offset := (page - 1) * pageSize - documents, total, err := s.documentDAO.GetByAuthorID(fmt.Sprintf("%d", authorID), offset, pageSize) - if err != nil { - return nil, 0, err - } - - responses := make([]*DocumentResponse, len(documents)) - for i, doc := range documents { - responses[i] = s.toResponse(doc) - } - - return responses, total, nil -} - -func (s *DocumentService) ListIngestionTasks(userID string, datasetID *string, page, pageSize int) ([]*entity.IngestionTask, error) { - return s.ingestionTaskSvc.ListByUser(userID, datasetID, page, pageSize) -} - -type ParseDocumentResponse struct { - DocumentID string `json:"document_id"` - Result string `json:"result"` -} - -func (s *DocumentService) IngestDocuments(datasetID, userID string, docIDs []string) ([]*ParseDocumentResponse, error) { - responses, err := s.ingestionTaskSvc.CreateForDocuments(datasetID, userID, docIDs) - if err != nil { - return nil, err - } - common.Info(fmt.Sprintf("parse documents, dataset: %s, documents: %v", datasetID, docIDs)) - return responses, nil -} - -func (s *DocumentService) StopIngestionTasks(tasks []string, userID string) ([]*entity.IngestionTask, error) { - return s.ingestionTaskSvc.RequestStopMany(tasks, &userID) -} - -func (s *DocumentService) RemoveIngestionTasks(tasks []string, userID string) ([]map[string]string, error) { - return s.ingestionTaskSvc.RemoveMany(tasks, &userID) -} - -type IngestDocumentRequest struct { - DocIDs []string `json:"doc_ids" binding:"required"` - Run interface{} `json:"run" binding:"required"` - Delete bool `json:"delete"` - ApplyKB bool `json:"apply_kb"` -} - -// StartParseOptions controls StartParseDocuments behavior. -type StartParseOptions struct { - // ApplyKB merges the knowledgebase's parser_config (llm_id, metadata) - // into the document before parsing. - ApplyKB bool - // RerunWithDelete clears prior chunks/tasks/counters before re-parsing. - RerunWithDelete bool -} - -// StartParseDocuments starts parsing a document via the DSL ingestion -// pipeline. It optionally clears prior results (RerunWithDelete), applies -// KB config (ApplyKB), validates storage, and enqueues an ingestion task. -// The document run status is NOT set here; IngestionTaskService.StartRunning -// sets it to RUNNING when the worker picks up the task and transitions it from -// CREATED. Extracted from Ingest so other entry points (e.g. ChunkService.Parse) -// can reuse the same start-parse flow. -func (s *DocumentService) StartParseDocuments(doc *entity.Document, kb *entity.Knowledgebase, userID string, opts StartParseOptions) error { - // Validate storage first so we don't clear prior results and then fail - // because the document can't be read, leaving the document with neither - // old nor new parse results. - if _, _, err := s.GetDocumentStorageAddress(doc); err != nil { - return err - } - - if opts.RerunWithDelete { - if err := s.clearDocumentParseResults(doc, kb.TenantID); err != nil { - return err - } - } - - if _, err := s.IngestDocuments(doc.KbID, userID, []string{doc.ID}); err != nil { - return err - } - return nil -} - -func (s *DocumentService) Ingest(userID string, req *IngestDocumentRequest) (common.ErrorCode, error) { - run := fmt.Sprint(req.Run) - - docs, err := s.documentDAO.GetByIDs(req.DocIDs) - if err != nil { - return common.CodeExceptionError, fmt.Errorf("fail to get documents: %s", err.Error()) - } - - docsByID := make(map[string]*entity.Document, len(docs)) - for _, doc := range docs { - if doc != nil { - docsByID[doc.ID] = doc - } - } - - // First pass: validate every document exists and is accessible before - // mutating any state, so a single invalid doc rejects the whole request. - type validatedDoc struct { - doc *entity.Document - kb *entity.Knowledgebase - } - validated := make([]validatedDoc, 0, len(req.DocIDs)) - validatedIDs := make([]string, 0, len(req.DocIDs)) - for _, docID := range req.DocIDs { - doc := docsByID[docID] - if doc == nil { - return common.CodeDataError, fmt.Errorf("document not found") - } - kb, err := s.kbDAO.GetByID(doc.KbID) - if err != nil { - return common.CodeDataError, fmt.Errorf("dataset not found") - } - if !s.kbDAO.Accessible(kb.ID, userID) { - return common.CodeAuthenticationError, fmt.Errorf("no authorization") - } - validated = append(validated, validatedDoc{doc, kb}) - validatedIDs = append(validatedIDs, docID) - } - - // Batch pre-check for re-parse with delete: use the validated doc IDs - // so we don't silently skip non-existent or unauthorized documents. - if run == string(entity.TaskStatusRunning) && req.Delete { - if err := s.AssertIngestionTasksTerminal(validatedIDs); err != nil { - return common.CodeDataError, err - } - } - - for _, vd := range validated { - doc := vd.doc - kb := vd.kb - - // Start parsing: delegates to the shared start-parse flow. The - // document run status is set by IngestionTaskService.StartRunning - // when the task transitions from CREATED, not here. - if run == string(entity.TaskStatusRunning) { - if err := s.StartParseDocuments(doc, kb, userID, StartParseOptions{ - ApplyKB: req.ApplyKB, - RerunWithDelete: req.Delete, - }); err != nil { - common.Error(fmt.Sprintf("go side, doc %s, start parse", doc.ID), err) - return common.CodeExceptionError, err - } - continue - } - - // Cancel: RequestStop (STOPPING) and update doc state. Do NOT - // delete the ingestion task or chunks here — deletion races with - // the worker's async markStopped/settleToTerminal flow. Once the - // worker detects STOPPING and transitions to STOPPED, the task - // is terminal and can be safely cleaned up. - if run == string(entity.TaskStatusCancel) { - if err := s.CancelDocParse(doc); err != nil { - common.Error(fmt.Sprintf("go side, start to process %s, run is cancel", doc.ID), err) - return common.CodeDataError, err - } - if err := s.documentDAO.UpdateByID(doc.ID, map[string]interface{}{ - "run": string(entity.TaskStatusCancel), - "progress": 0, - }); err != nil { - common.Error(fmt.Sprintf("go side, doc %s, UpdateByID failed", doc.ID), err) - return common.CodeExceptionError, err - } - continue - } - - // Delete-only: user asked to remove prior parse results without - // starting a new parse. RUNNING already continued above. - if err := s.documentDAO.UpdateByID(doc.ID, map[string]interface{}{ - "run": run, - "progress": 0, - }); err != nil { - common.Error(fmt.Sprintf("go side, doc %s, UpdateByID failed", doc.ID), err) - return common.CodeExceptionError, err - } - - if req.Delete { - _, _ = s.taskDAO.DeleteIngestionTasksByDocIDs([]string{doc.ID}) - indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) - if s.docEngine != nil { - exists, err := s.docEngine.ChunkStoreExists(context.Background(), indexName, doc.KbID) - if err != nil { - common.Error(fmt.Sprintf("go side, doc %s, ChunkStoreExists failed", doc.ID), err) - return common.CodeExceptionError, err - } - if exists { - if _, err := s.docEngine.DeleteChunks(context.Background(), map[string]interface{}{"doc_id": doc.ID}, indexName, doc.KbID); err != nil { - common.Error(fmt.Sprintf("go side, doc %s, DeleteChunks failed", doc.ID), err) - return common.CodeExceptionError, err - } - } - } - } - } - - return common.CodeSuccess, nil -} - -// AssertIngestionTasksTerminal verifies none of the documents has an -// in-flight (RUNNING/STOPPING) ingestion task. Used as a batch pre-check -// before re-parsing so a single non-terminal doc rejects the whole request -// up front instead of partially cleaning some docs then failing. -func (s *DocumentService) AssertIngestionTasksTerminal(docIDs []string) error { - for _, docID := range docIDs { - task, err := s.ingestionTaskDAO.GetByDocumentID(docID) - if err != nil { - return fmt.Errorf("check ingestion task for %s: %w", docID, err) - } - if task == nil { - continue - } - if task.Status == common.RUNNING || task.Status == common.STOPPING { - return fmt.Errorf("document %s ingestion task is %s; stop it and wait for a terminal state before re-parsing", docID, task.Status) - } - } - return nil -} - -func (s *DocumentService) clearDocumentParseResults(doc *entity.Document, tenantID string) error { - if doc == nil { - return fmt.Errorf("document is nil") - } - - // Refuse to clear a non-terminal ingestion task. An in-flight worker - // (RUNNING) or one mid-stop (STOPPING) would keep writing chunks and - // corrupt the new run's results. The caller must stop the task first - // and wait for a terminal state (COMPLETED/STOPPED/FAILED) or CREATED. - if task, _ := s.ingestionTaskDAO.GetByDocumentID(doc.ID); task != nil { - if task.Status == common.RUNNING || task.Status == common.STOPPING { - return fmt.Errorf("document %s ingestion task is %s; stop it and wait for a terminal state before re-parsing", doc.ID, task.Status) - } - } - - // Delete terminal and CREATED ingestion tasks atomically, leaving - // RUNNING/STOPPING tasks untouched so the check-then-delete window - // between GetByDocumentID and the delete above cannot delete a task - // that just transitioned to RUNNING. - if _, err := s.ingestionTaskDAO.DeleteIfTerminal(doc.ID); err != nil { - return err - } - - if err := s.clearDocumentAndKBCountersForRerun(doc.ID, doc.KbID); err != nil { - return err - } - - if s.docEngine == nil { - return nil - } - - indexName := fmt.Sprintf("ragflow_%s", tenantID) - exists, err := s.docEngine.ChunkStoreExists(context.Background(), 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 { - return err - } - return nil -} - -func (s *DocumentService) clearDocumentAndKBCountersForRerun(docID, kbID string) error { - return dao.DB.Transaction(func(tx *gorm.DB) error { - var current entity.Document - if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). - Where("id = ? AND kb_id = ?", docID, kbID). - First(¤t).Error; err != nil { - return err - } - - if current.TokenNum == 0 && current.ChunkNum == 0 && current.ProcessDuration == 0 { - return nil - } - - result := tx.Model(&entity.Document{}). - Where("id = ? AND kb_id = ?", docID, kbID). - Updates(map[string]interface{}{ - "token_num": 0, - "chunk_num": 0, - "process_duration": 0, - }) - if result.Error != nil { - return result.Error - } - if current.TokenNum == 0 && current.ChunkNum == 0 { - return nil - } - - result = tx.Model(&entity.Knowledgebase{}). - Where("id = ?", kbID). - Updates(map[string]interface{}{ - "token_num": gorm.Expr("token_num - ?", current.TokenNum), - "chunk_num": gorm.Expr("chunk_num - ?", current.ChunkNum), - }) - if result.Error != nil { - return result.Error - } - if result.RowsAffected == 0 { - return fmt.Errorf("knowledgebase not found") - } - return nil - }) -} - -func (s *DocumentService) countDoneDocuments(datasetID string) (int64, error) { - var count int64 - err := dao.GetDB().Model(&entity.Document{}). - Where("kb_id = ? AND run = ?", datasetID, string(entity.TaskStatusDone)). - Count(&count).Error - return count, err -} - -func (s *DocumentService) clearKBChunkNumWhenRerun(doc *entity.Document) error { - if doc == nil { - return fmt.Errorf("document is nil") - } - return dao.GetDB().Model(&entity.Knowledgebase{}).Where("id = ?", doc.KbID).Updates(map[string]interface{}{ - "token_num": gorm.Expr("token_num - ?", doc.TokenNum), - "chunk_num": gorm.Expr("chunk_num - ?", doc.ChunkNum), - }).Error -} - -func (s *DocumentService) ParseDocuments(datasetID, userID string, docIDs []string) ([]*ParseDocumentResponse, error) { - // create document parse id - // save to task table - // send to message queue - - // deduplicate the document id - uniqueDocIDs := common.Deduplicate(docIDs) - if uniqueDocIDs == nil || len(uniqueDocIDs) == 0 { - return nil, fmt.Errorf("no documents to parse") - } - - var responses []*ParseDocumentResponse - - // query database, if the document ids are valid - for _, docID := range uniqueDocIDs { - doc, err := s.documentDAO.GetByID(docID) - if err != nil { - errorMessage := err.Error() - responses = append(responses, &ParseDocumentResponse{ - DocumentID: docID, - Result: errorMessage, - }) - continue - } - if doc == nil { - errorMessage := "no such document" - responses = append(responses, &ParseDocumentResponse{ - DocumentID: docID, - Result: errorMessage, - }) - continue - } - - if doc.Status != nil && *doc.Status != "0" { - errorMessage := fmt.Sprintf("document %s is already parsed", docID) - responses = append(responses, &ParseDocumentResponse{ - DocumentID: docID, - Result: errorMessage, - }) - continue - } - - // create task for each document - //task := &entity.IngestionTask{ - // ID: utility.GenerateToken(), - // DocumentID: docID, - // UserID: userID, - //} - - // save the task to database - //err = s.ingestionTaskDAO.Create(task) - //if err != nil { - // errorMessage := err.Error() - // responses = append(responses, &ParseDocumentResponse{ - // DocumentID: docID, - // Result: &errorMessage, - // }) - // continue - //} - - // Send task to message queue - - } - - common.Info(fmt.Sprintf("parse documents, dataset: %s, documents: %v", datasetID, docIDs)) - return responses, nil -} - -// StopParseDocuments stops parsing for the given documents in a dataset. -// It sets Redis cancel signals for associated tasks and updates doc.run to CANCEL. -// Returns a map with success_count and optionally errors. -func (s *DocumentService) StopParseDocuments(datasetID string, docIDs []string) (map[string]interface{}, error) { - deduped := common.Deduplicate(docIDs) - if len(deduped) == 0 { - return nil, fmt.Errorf("no document IDs provided") - } - - docs, err := s.validateDocsInDataset(deduped, datasetID) - if err != nil { - return nil, err - } - - var errors []string - successCount := 0 - for _, doc := range docs { - if cancelErr := s.CancelDocParse(doc); cancelErr != nil { - errors = append(errors, cancelErr.Error()) - continue - } - successCount++ - } - - result := map[string]interface{}{"success_count": successCount} - if len(errors) > 0 { - result["errors"] = errors - } - return result, nil -} - -// validateDocsInDataset deduplicates IDs, fetches the documents, and ensures -// every document exists and belongs to the given dataset. Returns the resolved -// documents. -func (s *DocumentService) validateDocsInDataset(docIDs []string, datasetID string) ([]*entity.Document, error) { - docs, err := s.documentDAO.GetByIDs(docIDs) - if err != nil { - return nil, fmt.Errorf("failed to fetch documents: %w", err) - } - if len(docs) != len(docIDs) { - return nil, fmt.Errorf("some document IDs not found in dataset %s", datasetID) - } - var invalid []string - for _, d := range docs { - if d.KbID != datasetID { - invalid = append(invalid, d.ID) - } - } - if len(invalid) > 0 { - return nil, fmt.Errorf("these documents do not belong to dataset %s: %v", datasetID, invalid) - } - return docs, nil -} - -// CancelDocParse stops the ingestion task for the document by calling -// RequestStop (STOPPING), then marks the document run status as CANCEL. -func (s *DocumentService) CancelDocParse(doc *entity.Document) error { - task, err := s.ingestionTaskDAO.GetByDocumentID(doc.ID) - if err != nil { - return fmt.Errorf("failed to get ingestion task for %s: %v", doc.ID, err) - } - if task == nil { - return fmt.Errorf("no ingestion task found for document %s", doc.ID) - } - - if _, err := s.ingestionTaskSvc.RequestStop(task.ID); err != nil { - return fmt.Errorf("failed to stop ingestion task %s: %v", task.ID, err) - } - - if upErr := s.documentDAO.UpdateByID(doc.ID, map[string]interface{}{"run": string(entity.TaskStatusCancel)}); upErr != nil { - return fmt.Errorf("failed to update document %s: %v", doc.ID, upErr) - } - return nil -} - -// toResponse convert model.Document to DocumentResponse -func (s *DocumentService) toResponse(doc *entity.Document) *DocumentResponse { - createdAt := "" - if doc.CreateTime != nil { - // Check if timestamp is in milliseconds (13 digits) or seconds (10 digits) - var ts int64 - if *doc.CreateTime > 1000000000000 { - // Milliseconds - convert to seconds - ts = *doc.CreateTime / 1000 - } else { - ts = *doc.CreateTime - } - createdAt = time.Unix(ts, 0).Format("2006-01-02 15:04:05") - } - updatedAt := "" - if doc.UpdateTime != nil { - // Accept both historical second-based values and current millisecond-based values. - ts := *doc.UpdateTime - if ts > 1000000000000 { - ts /= 1000 - } - updatedAt = time.Unix(ts, 0).Format("2006-01-02 15:04:05") - } - return &DocumentResponse{ - ID: doc.ID, - Name: doc.Name, - KbID: doc.KbID, - ParserID: doc.ParserID, - PipelineID: doc.PipelineID, - Type: doc.Type, - SourceType: doc.SourceType, - CreatedBy: doc.CreatedBy, - Location: doc.Location, - Size: doc.Size, - TokenNum: doc.TokenNum, - ChunkNum: doc.ChunkNum, - Progress: doc.Progress, - ProgressMsg: doc.ProgressMsg, - ProcessDuration: doc.ProcessDuration, - Suffix: doc.Suffix, - Run: doc.Run, - Status: doc.Status, - CreatedAt: createdAt, - UpdatedAt: updatedAt, - } -} - -// GetMetadataSummaryRequest request for metadata summary -type GetMetadataSummaryRequest struct { - KBID string `json:"kb_id" binding:"required"` - DocIDs []string `json:"doc_ids"` -} - -// GetMetadataSummaryResponse response for metadata summary -type GetMetadataSummaryResponse struct { - Summary map[string]interface{} `json:"summary"` -} - -// GetMetadataSummary get metadata summary for documents -func (s *DocumentService) GetMetadataSummary(kbID string, docIDs []string) (map[string]interface{}, error) { - tenantID, err := s.metadataSvc.GetTenantIDByKBID(kbID) - if err != nil { - return nil, err - } - - searchResult, err := s.metadataSvc.SearchMetadata(kbID, tenantID, docIDs, 1000) - if err != nil { - return nil, err - } - - // Aggregate metadata from results - return aggregateMetadata(searchResult.MetadataRecords), nil -} - -// SetDocumentMetadata sets metadata for a document in the document engine -func (s *DocumentService) SetDocumentMetadata(docID string, meta map[string]interface{}) error { - // Get document to find kb_id - doc, err := s.documentDAO.GetByID(docID) - if err != nil { - return fmt.Errorf("document not found: %w", err) - } - - // Get tenant ID - tenantID, err := s.metadataSvc.GetTenantIDByKBID(doc.KbID) - if err != nil { - return fmt.Errorf("failed to get tenant ID: %w", err) - } - - if err := s.docEngine.UpdateMetadata(context.Background(), docID, doc.KbID, meta, tenantID); err != nil { - return fmt.Errorf("failed to update metadata: %w", err) - } - - return nil -} - -// DeleteDocumentMetadata deletes metadata keys for a document in the document engine -func (s *DocumentService) DeleteDocumentMetadata(docID string, keys []string) error { - // Get document to find kb_id - doc, err := s.documentDAO.GetByID(docID) - if err != nil { - return fmt.Errorf("document not found: %w", err) - } - - // Get tenant ID - tenantID, err := s.metadataSvc.GetTenantIDByKBID(doc.KbID) - if err != nil { - return fmt.Errorf("failed to get tenant ID: %w", err) - } - - // Delete metadata using the document engine - err = s.docEngine.DeleteMetadataKeys(nil, docID, doc.KbID, keys, tenantID) - if err != nil { - return fmt.Errorf("failed to delete metadata: %w", err) - } - - return nil -} - -// DeleteDocumentAllMetadata deletes all metadata for a document in the document engine -func (s *DocumentService) DeleteDocumentAllMetadata(docID string) error { - // Get document to find kb_id - doc, err := s.documentDAO.GetByID(docID) - if err != nil { - return fmt.Errorf("document not found: %w", err) - } - - // Get tenant ID - tenantID, err := s.metadataSvc.GetTenantIDByKBID(doc.KbID) - if err != nil { - return fmt.Errorf("failed to get tenant ID: %w", err) - } - - // Build condition to match the document - condition := map[string]interface{}{ - "id": docID, - "kb_id": doc.KbID, - } - - // Delete entire document metadata - _, err = s.docEngine.DeleteMetadata(nil, condition, tenantID) - if err != nil { - return fmt.Errorf("failed to delete document metadata: %w", err) - } - - return nil -} - -// GetDocumentMetadataByID get metadata for a specific document -func (s *DocumentService) GetDocumentMetadataByID(docID string) (map[string]interface{}, error) { - // Get document to find kb_id - doc, err := s.documentDAO.GetByID(docID) - if err != nil { - return nil, fmt.Errorf("document not found: %w", err) - } - - tenantID, err := s.metadataSvc.GetTenantIDByKBID(doc.KbID) - if err != nil { - return nil, err - } - - searchResult, err := s.metadataSvc.SearchMetadata(doc.KbID, tenantID, []string{docID}, 1) - if err != nil { - return nil, err - } - - // Return metadata if found - if len(searchResult.MetadataRecords) > 0 { - metadata := searchResult.MetadataRecords[0] - return ExtractMetaFields(metadata) - } - - return make(map[string]interface{}), nil -} - -// GetMetadataByKBs get metadata for knowledge bases -func (s *DocumentService) GetMetadataByKBs(kbIDs []string) (map[string]interface{}, error) { - if len(kbIDs) == 0 { - return make(map[string]interface{}), nil - } - - searchResult, err := s.metadataSvc.SearchMetadataByKBs(kbIDs, 10000) - if err != nil { - return nil, err - } - - flattenedMeta := make(map[string]map[string][]string) - numMetadata := len(searchResult.MetadataRecords) - - var allMetaFields []map[string]interface{} - if numMetadata > 1 && len(searchResult.MetadataRecords) > 0 { - firstMetadata := searchResult.MetadataRecords[0] - if metaFieldsVal := firstMetadata["meta_fields"]; metaFieldsVal != nil { - if v, ok := metaFieldsVal.([]byte); ok { - allMetaFields = ParseAllLengthPrefixedJSON(v) - } - } - } - - for idx, metadata := range searchResult.MetadataRecords { - docID, ok := ExtractDocumentID(metadata) - if !ok { - continue - } - - var metaFields map[string]interface{} - var metaFieldsVal interface{} - - if len(allMetaFields) > 0 && idx < len(allMetaFields) { - // Use pre-parsed meta_fields from concatenated data - metaFields = allMetaFields[idx] - } else { - // Normal case - get from chunk - metaFieldsVal = metadata["meta_fields"] - if metaFieldsVal != nil { - switch v := metaFieldsVal.(type) { - case string: - if err := json.Unmarshal([]byte(v), &metaFields); err != nil { - continue - } - case []byte: - // Try direct JSON parse first - if err := json.Unmarshal(v, &metaFields); err != nil { - // Try to parse as concatenated JSON objects - metaFields = ParseLengthPrefixedJSON(v) - } - case map[string]interface{}: - metaFields = v - default: - continue - } - } - } - - if metaFields == nil { - continue - } - - // Process each metadata field - for fieldName, fieldValue := range metaFields { - if fieldName == "kb_id" || fieldName == "id" { - continue - } - - if _, ok := flattenedMeta[fieldName]; !ok { - flattenedMeta[fieldName] = make(map[string][]string) - } - - // Handle list and single values - var values []interface{} - switch v := fieldValue.(type) { - case []interface{}: - values = v - default: - values = []interface{}{v} - } - - for _, val := range values { - if val == nil { - continue - } - strVal := fmt.Sprintf("%v", val) - flattenedMeta[fieldName][strVal] = append(flattenedMeta[fieldName][strVal], docID) - } - } - } - - // Convert to map[string]interface{} for return - var metaResult map[string]interface{} = make(map[string]interface{}) - for k, v := range flattenedMeta { - metaResult[k] = v - } - - return metaResult, nil -} - -// valueInfo holds count and order of first appearance -type valueInfo struct { - count int - firstOrder int -} - -// aggregateMetadata aggregates metadata from search results -func aggregateMetadata(chunks []map[string]interface{}) map[string]interface{} { - // summary: map[fieldName]map[value]valueInfo - summary := make(map[string]map[string]valueInfo) - typeCounter := make(map[string]map[string]int) - orderCounter := 0 - - for _, chunk := range chunks { - // For metadata table, the actual metadata is in the "meta_fields" JSON field - // Extract it first - metaFieldsVal := chunk["meta_fields"] - if metaFieldsVal == nil { - continue - } - - // Parse meta_fields - could be a string (JSON) or a map - var metaFields map[string]interface{} - switch v := metaFieldsVal.(type) { - case string: - // Parse JSON string - if err := json.Unmarshal([]byte(v), &metaFields); err != nil { - continue - } - case []byte: - // Handle byte slice - Infinity returns concatenated JSON objects with length prefixes - rawBytes := v - - // Try to detect and handle length-prefixed format - // Format: [4-byte length][JSON][4-byte length][JSON]... - parsedMetaFields := make(map[string]interface{}) - offset := 0 - for offset < len(rawBytes) { - // Need at least 4 bytes for length prefix - if offset+4 > len(rawBytes) { - break - } - - // Read 4-byte length (little-endian, not big-endian!) - length := uint32(rawBytes[offset]) | uint32(rawBytes[offset+1])<<8 | - uint32(rawBytes[offset+2])<<16 | uint32(rawBytes[offset+3])<<24 - - // Check if length looks valid (not too large) - if length > 10000 || length == 0 { - // Try to find next '{' from current position - nextBrace := -1 - for i := offset; i < len(rawBytes) && i < offset+100; i++ { - if rawBytes[i] == '{' { - nextBrace = i - break - } - } - if nextBrace > offset { - // Skip to the next '{' - offset = nextBrace - continue - } - break - } - - // Extract JSON data - jsonStart := offset + 4 - jsonEnd := jsonStart + int(length) - if jsonEnd > len(rawBytes) { - jsonEnd = len(rawBytes) - } - - jsonBytes := rawBytes[jsonStart:jsonEnd] - - // Try to parse this JSON - var singleMeta map[string]interface{} - if err := json.Unmarshal(jsonBytes, &singleMeta); err == nil { - // Merge metadata from this document - for k, vv := range singleMeta { - if existing, ok := parsedMetaFields[k]; ok { - // Combine values - if existList, ok := existing.([]interface{}); ok { - if newList, ok := vv.([]interface{}); ok { - parsedMetaFields[k] = append(existList, newList...) - } else { - parsedMetaFields[k] = append(existList, vv) - } - } else { - parsedMetaFields[k] = []interface{}{existing, vv} - } - } else { - parsedMetaFields[k] = vv - } - } - } - - offset = jsonEnd - } - - // If we successfully parsed multiple JSON objects, use the merged result - if len(parsedMetaFields) > 0 { - metaFields = parsedMetaFields - } else { - // Fallback: try the original parsing method - startIdx := -1 - for i, b := range rawBytes { - if b == '{' { - startIdx = i - break - } - } - if startIdx > 0 { - strVal := string(rawBytes[startIdx:]) - if err := json.Unmarshal([]byte(strVal), &metaFields); err != nil { - metaFields = map[string]interface{}{"raw": strVal} - } - } else if err := json.Unmarshal(rawBytes, &metaFields); err != nil { - metaFields = map[string]interface{}{"raw": string(rawBytes)} - } - } - case map[string]interface{}: - metaFields = v - default: - continue - } - - // Now iterate over the extracted metadata fields - for k, v := range metaFields { - // Skip nil values - if v == nil { - continue - } - - // Determine value type - valueType := getMetaValueType(v) - - // Track type counts - if valueType != "" { - if _, ok := typeCounter[k]; !ok { - typeCounter[k] = make(map[string]int) - } - typeCounter[k][valueType] = typeCounter[k][valueType] + 1 - } - - // Aggregate value counts. Flatten nested arrays so malformed values do - // not surface in the UI as the literal string "[]". - values := flattenMetadataSummaryValues(v) - for _, vv := range values { - if vv == nil { - continue - } - sv := fmt.Sprintf("%v", vv) - - if _, ok := summary[k]; !ok { - summary[k] = make(map[string]valueInfo) - } - - if existing, ok := summary[k][sv]; ok { - // Already exists, just increment count - existing.count++ - summary[k][sv] = existing - } else { - // First time seeing this value - record order - summary[k][sv] = valueInfo{count: 1, firstOrder: orderCounter} - orderCounter++ - } - } - } - } - - // Build result with type information and sorted values - result := make(map[string]interface{}) - for k, v := range summary { - // Sort by count descending, then by firstOrder ascending (to match Python stable sort) - // values: [value, count, firstOrder] - values := make([][3]interface{}, 0, len(v)) - for val, info := range v { - values = append(values, [3]interface{}{val, info.count, info.firstOrder}) - } - // Use stable sort - sort by count descending, then by firstOrder - sort.SliceStable(values, func(i, j int) bool { - cntI := values[i][1].(int) - cntJ := values[j][1].(int) - if cntI != cntJ { - return cntI > cntJ // count descending - } - // If counts equal, use firstOrder ascending (earlier appearance first) - return values[i][2].(int) < values[j][2].(int) - }) - - // Determine dominant type - valueType := "string" - if typeCounts, ok := typeCounter[k]; ok { - maxCount := 0 - for t, c := range typeCounts { - if c > maxCount { - maxCount = c - valueType = t - } - } - } - - // Convert from [value, count, firstOrder] to [value, count] for output - outputValues := make([][2]interface{}, len(values)) - for i, val := range values { - outputValues[i] = [2]interface{}{val[0], val[1]} - } - - result[k] = map[string]interface{}{ - "type": valueType, - "values": outputValues, - } - } - - return result -} - -// getMetaValueType determines the type of a metadata value -func getMetaValueType(value interface{}) string { - if value == nil { - return "" - } - - switch v := value.(type) { - case []interface{}: - if len(v) > 0 { - return "list" - } - return "" - case bool: - return "string" - case int, int8, int16, int32, int64: - return "number" - case float32, float64: - return "number" - case string: - if isTimeString(v) { - return "time" - } - return "string" - } - return "string" -} - -func flattenMetadataSummaryValues(value interface{}) []interface{} { - switch typed := value.(type) { - case []interface{}: - result := make([]interface{}, 0, len(typed)) - for _, item := range typed { - result = append(result, flattenMetadataSummaryValues(item)...) - } - return result - case []string: - result := make([]interface{}, 0, len(typed)) - for _, item := range typed { - result = append(result, item) - } - return result - case nil: - return nil - default: - return []interface{}{typed} - } -} - -// isTimeString checks if a string is an ISO 8601 datetime -func isTimeString(s string) bool { - matched, _ := regexp.MatchString(`^\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}$`, s) - return matched -} - -func (s *DocumentService) UpdateDatasetDocument(userID, datasetID, documentID string, req *UpdateDatasetDocumentRequest, present map[string]bool) (*UpdateDatasetDocumentResponse, common.ErrorCode, error) { - tenantID := userID - kb, err := s.kbDAO.GetByIDAndTenantID(datasetID, tenantID) - if err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, errors.New("You don't own the dataset.") - } - return nil, common.CodeDataError, errors.New("Can't find this dataset!") - } - - doc, err := s.documentDAO.GetByDocumentIDAndDatasetID(documentID, datasetID) - if err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, errors.New("The dataset doesn't own the document.") - } - return nil, common.CodeServerError, err - } - - if code, err := s.validateDatasetDocumentUpdate(datasetID, documentID, userID, doc, req, present); err != nil { - return nil, code, err - } - - if present["meta_fields"] { - if err := s.replaceDocumentMetadata(documentID, req.MetaFields); err != nil { - return nil, common.CodeDataError, err - } - } - - if present["name"] && req.Name != nil && (doc.Name == nil || *req.Name != *doc.Name) { - if err := s.updateDocumentNameOnly(doc, kb.TenantID, *req.Name); err != nil { - return nil, common.CodeDataError, err - } - } - - if present["parser_config"] && req.ParserConfig != nil { - // Resolve effective pipeline to load the DSL for cleaning. - isCanvas := kb.PipelineID != nil && strings.TrimSpace(*kb.PipelineID) != "" - if req.PipelineID != nil { - isCanvas = strings.TrimSpace(*req.PipelineID) != "" - } - if req.ParserID != nil { - isCanvas = false - } - effParserID := kb.ParserID - if req.ParserID != nil { - effParserID = strings.TrimSpace(*req.ParserID) - } - effPipelineID := kb.PipelineID - if req.PipelineID != nil { - effPipelineID = req.PipelineID - } - if req.ParserID != nil && req.PipelineID == nil && kb.PipelineID != nil { - effPipelineID = nil - } - - dslJSON, err := loadPipelineDSL(isCanvas, effParserID, effPipelineID) - if err != nil { - common.Warn("cleanAndUpdateDocumentParserConfig: failed to load DSL, falling back to merge", - zap.Error(err)) - if err := s.updateDocumentParserConfig(doc.ID, req.ParserConfig); err != nil { - return nil, common.CodeDataError, err - } - } else { - cleaned := buildParserConfig(dslJSON, map[string]interface{}(req.ParserConfig)) - if err := s.documentDAO.UpdateByID(doc.ID, map[string]interface{}{ - "parser_config": cleaned, - }); err != nil { - return nil, common.CodeDataError, err - } - } - } - - if present["pipeline_id"] { - if req.PipelineID != nil && strings.TrimSpace(*req.PipelineID) != "" { - if err := s.resetDocumentForReparse(doc, kb.TenantID, nil, req.PipelineID); err != nil { - return nil, common.CodeDataError, err - } - } else { - // Explicitly cleared: drop the custom canvas so the worker falls - // back to the built-in template, matching validation. - empty := "" - if err := s.resetDocumentForReparse(doc, kb.TenantID, nil, &empty); err != nil { - return nil, common.CodeDataError, err - } - } - } else if present["parser_id"] && req.ParserID != nil && strings.TrimSpace(*req.ParserID) != "" { - parserID := strings.TrimSpace(*req.ParserID) - if err := s.resetDocumentForReparse(doc, kb.TenantID, &parserID, nil); err != nil { - return nil, common.CodeDataError, err - } - } - - if present["enabled"] && req.Enabled != nil { - if err := s.updateDocumentStatusOnly(doc, kb, *req.Enabled); err != nil { - return nil, common.CodeServerError, err - } - } - - updatedDoc, err := s.documentDAO.GetByID(doc.ID) - if err != nil { - if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, fmt.Errorf("Can not get document by id:%s", doc.ID) - } - return nil, common.CodeDataError, errors.New("Database operation failed") - } - - metaFields := map[string]interface{}{} - if s.docEngine != nil && s.metadataSvc != nil { - metaFields, _ = s.GetDocumentMetadataByID(updatedDoc.ID) - } - - return s.toUpdateDatasetDocumentResponse(updatedDoc, metaFields), common.CodeSuccess, nil -} - -func (s *DocumentService) validateDatasetDocumentUpdate(datasetID, documentID, userID string, doc *entity.Document, req *UpdateDatasetDocumentRequest, present map[string]bool) (common.ErrorCode, error) { - if req == nil { - return common.CodeDataError, errors.New("Invalid request payload") - } - if present["chunk_count"] && req.ChunkCount != nil && *req.ChunkCount != 0 && *req.ChunkCount != doc.ChunkNum { - return common.CodeDataError, errors.New("Can't change `chunk_count`.") - } - if present["token_count"] && req.TokenCount != nil && *req.TokenCount != 0 && *req.TokenCount != doc.TokenNum { - return common.CodeDataError, errors.New("Can't change `token_count`.") - } - if present["progress"] && req.Progress != nil && *req.Progress != 0 && math.Abs(*req.Progress-doc.Progress) > 1e-9 { - return common.CodeDataError, errors.New("Can't change `progress`.") - } - - if present["enabled"] { - if req.Enabled == nil || (*req.Enabled != 0 && *req.Enabled != 1) { - return common.CodeDataError, errors.New("`enabled` value invalid, only accept 0 or 1") - } - } - - if present["parser_id"] && req.ParserID != nil { - parserID := strings.TrimSpace(*req.ParserID) - if (doc.Type == "visual" && parserID != "picture") || (isPresentationFile(doc.Name) && parserID != "presentation") { - return common.CodeDataError, errors.New("Not supported yet!") - } - } - if present["name"] && req.Name != nil { - if err := s.validateDocumentName(doc, *req.Name); err != nil { - return common.CodeDataError, err - } - } - - if present["meta_fields"] { - if err := validateMetaFields(req.MetaFields); err != nil { - return common.CodeDataError, err - } - } - - return common.CodeSuccess, nil -} - -func (s *DocumentService) validateDocumentName(doc *entity.Document, newName string) error { - if strings.TrimSpace(newName) == "" { - return errors.New("File name can't be empty.") - } - if len([]byte(newName)) > 255 { - return errors.New("File name must be 255 bytes or less.") - } - - oldName := "" - if doc.Name != nil { - oldName = *doc.Name - } - - if strings.ToLower(filepath.Ext(newName)) != strings.ToLower(filepath.Ext(oldName)) { - return errors.New("The extension of file can't be changed") - } - - docs, err := s.documentDAO.GetByNameAndKBID(newName, doc.KbID) - if err != nil { - return err - } - for _, d := range docs { - if d.ID != doc.ID && d.Name != nil && *d.Name == newName { - return errors.New("Duplicated document name in the same dataset.") - } - } - - return nil -} - -func isPresentationFile(name *string) bool { - if name == nil { - return false - } - ext := strings.ToLower(filepath.Ext(*name)) - return ext == ".ppt" || ext == ".pptx" || ext == ".pages" -} - -func validateMetaFields(meta map[string]any) error { - if meta == nil { - return nil - } - - for _, v := range meta { - switch typed := v.(type) { - case string, float64, int, int64, float32: - continue - case []any: - for _, item := range typed { - switch item.(type) { - case string, float64, int, int64, float32: - continue - default: - return fmt.Errorf("The type is not supported in list: %v", typed) - } - } - default: - return fmt.Errorf("The type is not supported: %v", v) - } - } - - return nil -} - -func (s *DocumentService) replaceDocumentMetadata(docID string, meta map[string]any) error { - if s.docEngine == nil || s.metadataSvc == nil { - return nil - } - if err := s.DeleteDocumentAllMetadata(docID); err != nil { - return err - } - return s.SetDocumentMetadata(docID, map[string]interface{}(meta)) -} - -func (s *DocumentService) patchDocumentMetadata(docID string, before, after map[string]interface{}) error { - if s.docEngine == nil || s.metadataSvc == nil { - return nil - } - - deleteKeys := make([]string, 0) - for key := range before { - if _, ok := after[key]; !ok { - deleteKeys = append(deleteKeys, key) - } - } - if len(deleteKeys) > 0 { - if err := s.DeleteDocumentMetadata(docID, deleteKeys); err != nil { - return err - } - } - - updateFields := make(map[string]interface{}) - for key, value := range after { - if !reflect.DeepEqual(before[key], value) { - updateFields[key] = value - } - } - if len(updateFields) == 0 { - return nil - } - return s.SetDocumentMetadata(docID, updateFields) -} - -func (s *DocumentService) updateDocumentNameOnly(doc *entity.Document, tenantID, newName string) error { - if err := s.documentDAO.UpdateByID(doc.ID, map[string]interface{}{"name": newName}); err != nil { - return errors.New("Database error (Document rename)!") - } - - mappings, err := s.file2DocumentDAO.GetByDocumentID(doc.ID) - if err == nil && len(mappings) > 0 && mappings[0].FileID != nil && s.fileDAO != nil { - _ = s.fileDAO.UpdateByID(*mappings[0].FileID, map[string]interface{}{"name": newName}) - } - - if s.docEngine == nil { - return nil - } - - titleTks, _ := tokenizer.Tokenize(newName) - titleSmTks, _ := tokenizer.FineGrainedTokenize(titleTks) - indexName := fmt.Sprintf("ragflow_%s", tenantID) - return s.docEngine.UpdateChunks( - context.Background(), - map[string]interface{}{"doc_id": doc.ID}, - map[string]interface{}{ - "docnm_kwd": newName, - "title_tks": titleTks, - "title_sm_tks": titleSmTks, - }, - indexName, - doc.KbID, - ) -} - -// loadCanvasDSLJSON returns the DSL JSON for a custom canvas pipeline. The -// canvas's dsl column holds the same component-graph structure that built-in -// templates use, so it can be validated by the same schema extractor. It is a -// package-level function so both document and knowledge-base updates reuse it. -func loadCanvasDSLJSON(canvasID string) ([]byte, error) { - if strings.TrimSpace(canvasID) == "" { - return nil, fmt.Errorf("empty canvas id") - } - canvas, err := dao.NewUserCanvasDAO().GetByID(canvasID) - if err != nil { - if errors.Is(err, dao.ErrUserCanvasNotFound) { - return nil, fmt.Errorf("canvas %s not found", canvasID) - } - return nil, fmt.Errorf("load canvas %s: %w", canvasID, err) - } - if len(canvas.DSL) == 0 { - return nil, fmt.Errorf("canvas %s has no DSL", canvasID) - } - raw, err := json.Marshal(canvas.DSL) - if err != nil { - return nil, fmt.Errorf("marshal canvas %s DSL: %w", canvasID, err) - } - return raw, nil -} - -// loadPipelineDSL loads the DSL JSON for a pipeline identified by parserID -// (built-in) or pipelineID (custom canvas). When both are provided, isCanvas -// selects which one to use. -func loadPipelineDSL(isCanvas bool, parserID string, pipelineID *string) ([]byte, error) { - if isCanvas { - return loadCanvasDSLJSON(strings.TrimSpace(*pipelineID)) - } - registry, err := pipelinepkg.DefaultRegistry() - if err != nil { - return nil, fmt.Errorf("builtin pipeline registry: %w", err) - } - if !registry.IsValid(parserID) { - return nil, fmt.Errorf("unknown builtin parser_id: %s", parserID) - } - dslStr, err := pipelinepkg.LoadBuiltinDSL(parserID) - if err != nil { - return nil, fmt.Errorf("load builtin DSL for %q: %w", parserID, err) - } - return []byte(dslStr), nil -} - -// cleanComponentParams filters rawConfig against the DSL schema given by dslJSON. -// Keys containing ':' are treated as component IDs; they are kept only when both -// the cpnID AND the param name exist in the DSL schema. Keys without ':' (legacy -// flat fields such as chunk_token_num, image_context_size) are dropped with a -// warning — they do not belong in the new component-params world. -func cleanComponentParams(dslJSON []byte, rawConfig map[string]interface{}) map[string]interface{} { - schemas, err := pipelinepkg.ExtractAllComponentParams(dslJSON) - if err != nil { - common.Warn("cleanComponentParams: failed to extract DSL schema, returning input as-is", - zap.Error(err)) - return rawConfig - } - - validCPNs := make(map[string]map[string]struct{}, len(schemas)) - for _, s := range schemas { - keys := make(map[string]struct{}, len(s.ParamsDefaults)) - for k := range s.ParamsDefaults { - keys[k] = struct{}{} - } - validCPNs[s.CpnID] = keys - } - - result := make(map[string]interface{}, len(rawConfig)) - for key, val := range rawConfig { - if !strings.Contains(key, ":") { - common.Warn("cleanComponentParams: dropping legacy flat field", - zap.String("key", key)) - continue - } - validKeys, ok := validCPNs[key] - if !ok { - common.Warn("cleanComponentParams: dropping unknown cpnID", - zap.String("cpnID", key)) - continue - } - params, ok := val.(map[string]any) - if !ok { - continue - } - cleaned := make(map[string]any, len(params)) - for pk, pv := range params { - if _, ok := validKeys[pk]; ok { - cleaned[pk] = pv - } else { - common.Warn("cleanComponentParams: dropping unknown param", - zap.String("cpnID", key), zap.String("param", pk)) - } - } - if len(cleaned) > 0 { - result[key] = cleaned - } - } - return result -} - -// buildParserConfig builds the final parser_config by starting from the DSL -// defaults for every component, then overlaying the cleaned incoming overrides. -// This ensures all components from the current pipeline are present while -// stripping stale params from other pipelines. -func buildParserConfig(dslJSON []byte, rawConfig map[string]interface{}) entity.JSONMap { - cleaned := cleanComponentParams(dslJSON, rawConfig) - defaults, err := pipelinepkg.ComponentParamsDefaults(dslJSON) - if err != nil { - common.Warn("buildParserConfig: failed to extract DSL defaults, using cleaned only", - zap.Error(err)) - return entity.JSONMap(cleaned) - } - result := make(entity.JSONMap, len(defaults)) - for cpnID, params := range defaults { - base := make(map[string]interface{}, len(params)) - for k, v := range params { - base[k] = v - } - if over, ok := cleaned[cpnID]; ok { - if om, ok := over.(map[string]any); ok { - result[cpnID] = common.DeepMergeMaps(base, map[string]interface{}(om)) - } else { - result[cpnID] = base - } - } else { - result[cpnID] = base - } - } - return result -} - -func (s *DocumentService) updateDocumentParserConfig(documentID string, config map[string]any) error { - if len(config) == 0 { - return nil - } - - doc, err := s.documentDAO.GetByID(documentID) - if err != nil { - return fmt.Errorf("Document(%s) not found.", documentID) - } - - merged := common.DeepMergeMaps(map[string]interface{}(doc.ParserConfig), map[string]interface{}(config)) - if _, ok := config["raptor"]; !ok { - delete(merged, "raptor") - } - - return s.documentDAO.UpdateByID(documentID, map[string]interface{}{ - "parser_config": entity.JSONMap(merged), - }) -} - -func (s *DocumentService) resetDocumentForReparse(doc *entity.Document, tenantID string, parserID *string, pipelineID *string) error { - progressMsg := "" - run := string(entity.TaskStatusUnstart) - updates := map[string]interface{}{ - "progress": 0, - "progress_msg": progressMsg, - "run": run, - } - if parserID != nil { - updates["parser_id"] = *parserID - } - if pipelineID != nil { - updates["pipeline_id"] = *pipelineID - } - - if err := s.documentDAO.UpdateByID(doc.ID, updates); err != nil { - return errors.New("Document not found!") - } - - if doc.TokenNum > 0 { - decremented, err := s.decrementDocumentAndKBCountersForReparse(doc) - if err != nil { - return errors.New("Document not found!") - } - if !decremented { - return nil - } - 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 { - return err - } - } - } - - return nil -} - -func (s *DocumentService) deleteChunkImages(doc *entity.Document, indexName string) { - if s.docEngine == nil { - return - } - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return - } - - const pageSize = 1000 - for offset := 0; ; offset += pageSize { - result, err := s.docEngine.Search(context.Background(), &enginetypes.SearchRequest{ - IndexNames: []string{indexName}, - KbIDs: []string{doc.KbID}, - Offset: offset, - Limit: pageSize, - SelectFields: []string{"id", "img_id"}, - Filter: map[string]interface{}{"doc_id": doc.ID}, - MatchExprs: nil, - OrderBy: nil, - RankFeature: nil, - }) - if err != nil || result == nil || len(result.Chunks) == 0 { - return - } - for _, chunk := range result.Chunks { - imageKey, ok := chunkImageStorageKey(doc.KbID, chunk) - if !ok { - continue - } - if storageImpl.ObjExist(doc.KbID, imageKey) { - _ = storageImpl.Remove(doc.KbID, imageKey) - } - } - } -} - -func chunkImageStorageKey(defaultBucket string, chunk map[string]interface{}) (string, bool) { - imgID := firstStringField(chunk, "img_id") - if imgID != "" { - prefix := defaultBucket + "-" - if strings.HasPrefix(imgID, prefix) && len(imgID) > len(prefix) { - return strings.TrimPrefix(imgID, prefix), true - } - return imgID, true - } - - chunkID := firstStringField(chunk, "id", "_id") - if chunkID == "" { - return "", false - } - return chunkID, true -} - -func firstStringField(m map[string]interface{}, keys ...string) string { - for _, key := range keys { - if value, ok := m[key]; ok { - if s, ok := value.(string); ok { - return s - } - } - } - return "" -} - -func (s *DocumentService) decrementDocumentAndKBCountersForReparse(doc *entity.Document) (bool, error) { - decremented := false - err := dao.DB.Transaction(func(tx *gorm.DB) error { - result := tx.Model(&entity.Document{}). - Where("id = ? AND kb_id = ? AND token_num = ? AND chunk_num = ?", doc.ID, doc.KbID, doc.TokenNum, doc.ChunkNum). - Updates(map[string]interface{}{ - "token_num": gorm.Expr("token_num - ?", doc.TokenNum), - "chunk_num": gorm.Expr("chunk_num - ?", doc.ChunkNum), - "process_duration": gorm.Expr("process_duration - ?", doc.ProcessDuration), - }) - if result.Error != nil { - return result.Error - } - if result.RowsAffected == 0 { - return nil - } - decremented = true - - return tx.Model(&entity.Knowledgebase{}). - Where("id = ?", doc.KbID). - Updates(map[string]interface{}{ - "token_num": gorm.Expr("token_num - ?", doc.TokenNum), - "chunk_num": gorm.Expr("chunk_num - ?", doc.ChunkNum), - }).Error - }) - return decremented, err -} - -func (s *DocumentService) updateDocumentStatusOnly(doc *entity.Document, kb *entity.Knowledgebase, status int) error { - statusStr := strconv.Itoa(status) - if doc.Status != nil && *doc.Status == statusStr { - return nil - } - - if err := s.documentDAO.UpdateByID(doc.ID, map[string]interface{}{"status": statusStr}); err != nil { - return errors.New("Database error (Document update)!") - } - - if s.docEngine == nil { - return nil - } - - indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) - return s.docEngine.UpdateChunks( - context.Background(), - map[string]interface{}{"doc_id": doc.ID}, - map[string]interface{}{"available_int": status}, - indexName, - doc.KbID, - ) -} - -func (s *DocumentService) toUpdateDatasetDocumentResponse(doc *entity.Document, metaFields map[string]interface{}) *UpdateDatasetDocumentResponse { - if metaFields == nil { - metaFields = map[string]interface{}{} - } - return &UpdateDatasetDocumentResponse{ - ID: doc.ID, - Thumbnail: doc.Thumbnail, - DatasetID: doc.KbID, - ParserID: doc.ParserID, - PipelineID: doc.PipelineID, - ParserConfig: map[string]interface{}(doc.ParserConfig), - SourceType: doc.SourceType, - Type: doc.Type, - CreatedBy: doc.CreatedBy, - Name: doc.Name, - Location: doc.Location, - Size: doc.Size, - TokenCount: doc.TokenNum, - ChunkCount: doc.ChunkNum, - Progress: doc.Progress, - ProgressMsg: doc.ProgressMsg, - ProcessBeginAt: doc.ProcessBeginAt, - ProcessDuration: doc.ProcessDuration, - ContentHash: doc.ContentHash, - MetaFields: metaFields, - Suffix: doc.Suffix, - Run: mapDocumentRunStatus(doc.Run), - Status: doc.Status, - CreateTime: doc.CreateTime, - CreateDate: doc.CreateDate, - UpdateTime: doc.UpdateTime, - UpdateDate: doc.UpdateDate, - } -} - -func mapDocumentRunStatus(run *string) string { - if run == nil { - return "UNSTART" - } - switch *run { - case string(entity.TaskStatusRunning): - return "RUNNING" - case string(entity.TaskStatusCancel): - return "CANCEL" - case string(entity.TaskStatusDone): - return "DONE" - case string(entity.TaskStatusFail): - return "FAIL" - default: - return "UNSTART" - } -} - -// UploadLocalDocuments stores each uploaded file in object storage and inserts a -// matching Document row into the dataset. It mirrors Python -// FileService.upload_document: it derives parser_id by filetype, merges the -// optional parser_config override into the dataset config, dedup-renames the -// filename, records size + xxhash content hash, and links each document into the -// file manager (a File row under the dataset folder + a file2document mapping) -// so it surfaces in the dataset's document list. Chunking/embedding happen later -// in the parse step, so nothing here touches the doc store index. -// -// Gaps vs Python (documented, not yet ported): thumbnail generation and -// read_potential_broken_pdf repair. -func (s *DocumentService) UploadLocalDocuments(kb *entity.Knowledgebase, tenantID string, files []*multipart.FileHeader, parentPath string, parserConfigOverride map[string]interface{}) ([]map[string]interface{}, []string) { - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return nil, []string{"storage not initialized"} - } - - // Resolve (and create if needed) the dataset's file-manager folder up front. - // Without the File / file2document linkage the document list (which inner-joins - // file2document + file) would never surface the uploaded files. - kbFolder, err := s.ensureKBFolder(kb, tenantID) - if err != nil { - return nil, []string{err.Error()} - } - - // Merge parser_config override (allow-listed keys only) over the dataset config. - merged := entity.JSONMap{} - for k, v := range kb.ParserConfig { - merged[k] = v - } - for k, v := range parserConfigOverride { - merged[k] = v - } - - safeParent := utility.SanitizeFilename(parentPath) - - // Don't silently disable dedupe protection: a transient lookup failure means - // the existing-name set is unknown, so fail rather than risk duplicates. - names, err := s.documentDAO.ListNamesByKbID(kb.ID) - if err != nil { - return nil, []string{err.Error()} - } - taken := map[string]bool{} - for _, n := range names { - taken[n] = true - } - - var results []map[string]interface{} - var errMsgs []string - - for _, fh := range files { - blob, err := readFileHeaderBytes(fh) - if err != nil { - errMsgs = append(errMsgs, fh.Filename+": "+err.Error()) - continue - } - - filename := uniqueUploadName(fh.Filename, taken) - - filetype := utility.FilenameType(filename) - if filetype == utility.FileTypeOTHER { - errMsgs = append(errMsgs, fh.Filename+": This type of file has not been supported yet!") - continue - } - - location := filename - if safeParent != "" { - location = safeParent + "/" + filename - } - for storageImpl.ObjExist(kb.ID, location) { - location += "_" - } - if err := storageImpl.Put(kb.ID, location, blob); err != nil { - errMsgs = append(errMsgs, fh.Filename+": "+err.Error()) - continue - } - - 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) - errMsgs = append(errMsgs, fh.Filename+": "+err.Error()) - continue - } - if err := s.addFileFromKB(doc, kbFolder.ID, kb.TenantID); err != nil { - // Linkage failed: roll back the document row and blob so the partial - // state doesn't leave an invisible (unlisted) document behind. - err = s.rollbackAddFileFromKBError(doc, kb.ID, err) - _ = storageImpl.Remove(kb.ID, location) - errMsgs = append(errMsgs, fh.Filename+": "+err.Error()) - continue - } - // Only reserve the name once the write fully succeeds. - taken[filename] = true - results = append(results, docToRawMap(doc)) - } - - return results, errMsgs -} - -// UploadEmptyDocument inserts a zero-byte "virtual" document into the dataset. -func (s *DocumentService) UploadEmptyDocument(kb *entity.Knowledgebase, tenantID, name string) (map[string]interface{}, common.ErrorCode, error) { - // A transient lookup failure means the existing-name set is unknown; fail - // rather than write blind and risk a duplicate. - names, err := s.documentDAO.ListNamesByKbID(kb.ID) - if err != nil { - return nil, common.CodeServerError, err - } - for _, n := range names { - if n == name { - return nil, common.CodeDataError, fmt.Errorf("Duplicated document name in the same dataset.") - } - } - - kbFolder, err := s.ensureKBFolder(kb, tenantID) - if err != nil { - return nil, common.CodeServerError, err - } - - doc := s.newDatasetDocument(kb, tenantID, name, "", "virtual", kb.ParserConfig, "local", 0, nil) - if err := s.InsertDocument(doc); err != nil { - return nil, common.CodeServerError, err - } - if err := s.addFileFromKB(doc, kbFolder.ID, kb.TenantID); err != nil { - return nil, common.CodeServerError, s.rollbackAddFileFromKBError(doc, kb.ID, err) - } - return docToRawMap(doc), common.CodeSuccess, nil -} - -// knowledgebaseFolderName is the file-manager folder under each tenant's root -// that holds per-dataset subfolders, mirroring Python KNOWLEDGEBASE_FOLDER_NAME. -const knowledgebaseFolderName = ".knowledgebase" - -// ensureKBFolder resolves (creating as needed) the per-dataset file-manager -// folder: root -> .knowledgebase -> . Mirrors Python -// get_root_folder + get_kb_folder + new_a_file_from_kb. -func (s *DocumentService) ensureKBFolder(kb *entity.Knowledgebase, tenantID string) (*entity.File, error) { - root, err := s.fileDAO.GetRootFolder(tenantID) - if err != nil { - return nil, err - } - kbRoot, err := s.newAFileFromKB(tenantID, knowledgebaseFolderName, root.ID) - if err != nil { - return nil, err - } - return s.newAFileFromKB(kb.TenantID, kb.Name, kbRoot.ID) -} - -// newAFileFromKB returns the existing folder named name under parentID, or -// creates it. Mirrors Python FileService.new_a_file_from_kb. -func (s *DocumentService) newAFileFromKB(tenantID, name, parentID string) (*entity.File, error) { - for _, f := range s.fileDAO.Query(name, parentID, tenantID) { - if f.TenantID == tenantID { - return f, nil - } - } - loc := "" - folder := &entity.File{ - ID: utility.GenerateToken(), - ParentID: parentID, - TenantID: tenantID, - CreatedBy: tenantID, - Name: name, - Type: "folder", - Size: 0, - Location: &loc, - SourceType: string(entity.FileSourceKnowledgebase), - } - if err := s.fileDAO.Create(folder); err != nil { - return nil, err - } - return folder, nil -} - -// addFileFromKB links a document into the file manager: a File row under the -// dataset folder plus a file2document mapping. Mirrors Python -// FileService.add_file_from_kb (idempotent on the document mapping). -func (s *DocumentService) addFileFromKB(doc *entity.Document, kbFolderID, tenantID string) error { - if existing, err := s.file2DocumentDAO.GetByDocumentID(doc.ID); err == nil && len(existing) > 0 { - return nil - } - name := "" - if doc.Name != nil { - name = *doc.Name - } - loc := "" - if doc.Location != nil { - loc = *doc.Location - } - fileID := utility.GenerateToken() - file := &entity.File{ - ID: fileID, - ParentID: kbFolderID, - TenantID: tenantID, - CreatedBy: tenantID, - Name: name, - Type: doc.Type, - Size: doc.Size, - Location: &loc, - SourceType: string(entity.FileSourceKnowledgebase), - } - if err := s.fileDAO.Create(file); err != nil { - return err - } - docID := doc.ID - if err := s.file2DocumentDAO.Create(&entity.File2Document{ - ID: utility.GenerateToken(), - FileID: &fileID, - DocumentID: &docID, - }); err != nil { - _ = s.fileDAO.Delete(fileID) - return err - } - return nil -} - -func (s *DocumentService) UploadWebDocument(kb *entity.Knowledgebase, tenantID, name, url string) (map[string]interface{}, common.ErrorCode, error) { - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return nil, common.CodeServerError, fmt.Errorf("storage not initialized") - } - - kbFolder, err := s.ensureKBFolder(kb, tenantID) - if err != nil { - return nil, common.CodeServerError, err - } - - names, err := s.documentDAO.ListNamesByKbID(kb.ID) - if err != nil { - return nil, common.CodeServerError, err - } - taken := map[string]bool{} - for _, n := range names { - taken[n] = true - } - - blob, headers, _, err := fetchRemoteFileSafely(url, maxUploadDocSize) - if err != nil { - return nil, common.CodeDataError, err - } - contentType := "" - if headers != nil { - contentType = headers.Get("Content-Type") - } - filename := normalizeWebDocumentName(name, contentType, blob) - filename, _, blob = normalizeUploadInfoContent(filename, contentType, blob) - filename = uniqueUploadName(filename, taken) - - filetype := utility.FilenameType(filename) - if filetype == utility.FileTypeOTHER { - return nil, common.CodeDataError, fmt.Errorf("This type of file has not been supported yet!") - } - - location := filename - for storageImpl.ObjExist(kb.ID, location) { - location += "_" - } - if err := storageImpl.Put(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) - return nil, common.CodeServerError, err - } - if err := s.addFileFromKB(doc, kbFolder.ID, kb.TenantID); err != nil { - err = s.rollbackAddFileFromKBError(doc, kb.ID, err) - _ = storageImpl.Remove(kb.ID, location) - return nil, common.CodeServerError, err - } - return docToRawMap(doc), common.CodeSuccess, nil -} - -func normalizeWebDocumentName(name, contentType string, blob []byte) string { - filename := utility.SanitizeFilename(name) - if filepath.Ext(filename) != "" { - return filename - } - lowerCT := strings.ToLower(strings.TrimSpace(strings.Split(contentType, ";")[0])) - switch { - case lowerCT == "application/pdf" || http.DetectContentType(blob) == "application/pdf" || bytesLooksLikePDF(blob): - return filename + ".pdf" - case lowerCT == "text/html" || lowerCT == "application/xhtml+xml" || looksLikeHTML(blob): - return filename + ".html" - default: - return filename - } -} - -// newDatasetDocument builds a Document row for an upload, deriving parser_id, -// suffix and content hash. blob may be nil for the empty/virtual document. -func (s *DocumentService) newDatasetDocument(kb *entity.Knowledgebase, tenantID, filename, location, filetype string, parserConfig entity.JSONMap, src string, size int64, blob []byte) *entity.Document { - docID := utility.GenerateToken() - run := "0" - status := "1" - suffix := "" - if i := strings.LastIndex(filename, "."); i >= 0 { - suffix = filename[i+1:] - } - parserID := selectUploadParser(utility.FileType(filetype), filename, kb.ParserID) - if kb.PipelineID != nil { - parserID = "" // canvas pipeline mode — parser_id not applicable - } - loc := location - doc := &entity.Document{ - ID: docID, - KbID: kb.ID, - ParserID: parserID, - PipelineID: kb.PipelineID, - ParserConfig: parserConfig, - CreatedBy: tenantID, - Type: filetype, - SourceType: src, - Name: &filename, - Location: &loc, - Size: size, - Suffix: suffix, - Run: &run, - Status: &status, - } - if blob != nil { - hash := contentHashHex(blob) - doc.ContentHash = &hash - } - - // When the document's builtin parser_id differs from the KB's (e.g. visual→picture, - // aural→audio), re-resolve component_params defaults from the document's own DSL - // template so cpnIDs and param keys match the pipeline that will actually execute. - if kb.PipelineID == nil && parserID != kb.ParserID { - if cp, err := resolveComponentParamsDefaults(parserID, nil); err != nil { - common.Warn("newDatasetDocument: resolve component_params defaults", - zap.String("parserID", parserID), zap.Error(err)) - } else if cp != nil { - doc.ParserConfig = cp - } - } - - return doc -} - -// docToRawMap serialises a freshly created Document into the raw key shape the -// handler remaps (chunk_num→chunk_count, kb_id→dataset_id). -func docToRawMap(doc *entity.Document) map[string]interface{} { - m := map[string]interface{}{ - "id": doc.ID, - "kb_id": doc.KbID, - "parser_id": doc.ParserID, - "parser_config": map[string]interface{}(doc.ParserConfig), - "created_by": doc.CreatedBy, - "type": doc.Type, - "source_type": doc.SourceType, - "size": doc.Size, - "chunk_num": doc.ChunkNum, - "token_num": doc.TokenNum, - "suffix": doc.Suffix, - "run": "0", - } - if doc.Name != nil { - m["name"] = *doc.Name - } - if doc.Location != nil { - m["location"] = *doc.Location - } - if doc.PipelineID != nil { - m["pipeline_id"] = *doc.PipelineID - } - if doc.ContentHash != nil { - m["content_hash"] = *doc.ContentHash - } - return m -} - -// uniqueUploadName appends a numeric suffix until the name is free, mirroring -// Python duplicate_name. -func uniqueUploadName(name string, taken map[string]bool) string { - if !taken[name] { - return name - } - base, ext := name, "" - if i := strings.LastIndex(name, "."); i >= 0 { - base, ext = name[:i], name[i:] - } - for i := 1; ; i++ { - candidate := fmt.Sprintf("%s(%d)%s", base, i, ext) - if !taken[candidate] { - return candidate - } - } -} - -// maxUploadDocSize bounds a single uploaded file held in memory, mirroring the -// Python DOC_MAXIMUM_SIZE default (128 MiB; overridable there via MAX_CONTENT_LENGTH). -const maxUploadDocSize = 128 * 1024 * 1024 - -func readFileHeaderBytes(fh *multipart.FileHeader) ([]byte, error) { - if fh.Size > maxUploadDocSize { - return nil, fmt.Errorf("file exceeds the maximum allowed size of %d bytes", maxUploadDocSize) - } - src, err := fh.Open() - if err != nil { - return nil, err - } - defer src.Close() - blob, err := io.ReadAll(io.LimitReader(src, maxUploadDocSize+1)) - if err != nil { - return nil, err - } - if len(blob) > maxUploadDocSize { - return nil, fmt.Errorf("file exceeds the maximum allowed size of %d bytes", maxUploadDocSize) - } - return blob, nil -} - -// MetadataUpdate is one update item: set key to value. -type DocumentMetadataUpdate struct { - Key string `json:"key"` - Value interface{} `json:"value"` - Match interface{} `json:"match,omitempty"` - ValueType string `json:"valueType,omitempty"` -} - -// MetadataDelete removes a whole key, or a specific value from a list field. -type DocumentMetadataDelete struct { - Key string `json:"key"` - Value interface{} `json:"value,omitempty"` -} - -// MetadataSelector selects which documents to target. -type DocumentMetadataSelector struct { - DocumentIDs []string `json:"document_ids"` - MetadataCondition map[string]interface{} `json:"metadata_condition"` -} - -// BatchUpdateDocumentMetadatasResponse summarises the operation. -type BatchUpdateDocumentMetadatasResponse struct { - Updated int `json:"updated"` - MatchedDocs int `json:"matched_docs"` -} - -// BatchUpdateDocumentMetadatas implements the shared logic for -// PATCH /datasets/:dataset_id/documents/metadatas and -// POST /datasets/:dataset_id/metadata/update. -func (s *DocumentService) BatchUpdateDocumentMetadatas( - datasetID string, - selector *DocumentMetadataSelector, - updates []DocumentMetadataUpdate, - deletes []DocumentMetadataDelete, -) (*BatchUpdateDocumentMetadatasResponse, common.ErrorCode, error) { - if selector == nil { - selector = &DocumentMetadataSelector{} - } - if code, err := validateBatchUpdateDocumentMetadatasRequest(selector, updates, deletes); err != nil { - return nil, code, err - } - - // Resolve which document IDs to target. - targetDocIDs := make(map[string]struct{}) - - if len(selector.DocumentIDs) > 0 { - // Validate that supplied IDs actually belong to this dataset. - allRows, err := s.documentDAO.GetAllDocIDsByKBIDs([]string{datasetID}) - if err != nil { - return nil, common.CodeServerError, fmt.Errorf("failed to list dataset documents: %w", err) - } - kbDocIDSet := make(map[string]struct{}, len(allRows)) - for _, row := range allRows { - kbDocIDSet[row["id"]] = struct{}{} - } - var invalidIDs []string - for _, id := range selector.DocumentIDs { - if _, ok := kbDocIDSet[id]; !ok { - invalidIDs = append(invalidIDs, id) - } - } - if len(invalidIDs) > 0 { - return nil, common.CodeDataError, fmt.Errorf("these documents do not belong to dataset %s: %s", - datasetID, strings.Join(invalidIDs, ", ")) - } - for _, id := range selector.DocumentIDs { - targetDocIDs[id] = struct{}{} - } - } - - // Apply metadata_condition filter. - if len(selector.MetadataCondition) > 0 { - flattedMeta, err := s.metadataSvc.GetFlattedMetaByKBs([]string{datasetID}) - if err != nil { - return nil, common.CodeServerError, fmt.Errorf("failed to get flattened metadata: %w", err) - } - - // ParseAndConvert mirrors Python convert_conditions: conditions arrive as - // {name, comparison_operator, value}, the operator is normalised, and the - // (possibly non-string) value is preserved. MetaFilter then matches against - // the common.MetaData returned by GetFlattedMetaByKBs. - filterInput := common.ParseAndConvert(selector.MetadataCondition) - filteredIDs := common.MetaFilter(flattedMeta, filterInput) - - filteredSet := make(map[string]struct{}, len(filteredIDs)) - for _, id := range filteredIDs { - filteredSet[id] = struct{}{} - } - - if len(targetDocIDs) > 0 { - // Intersect with the document_ids restriction. - for id := range targetDocIDs { - if _, ok := filteredSet[id]; !ok { - delete(targetDocIDs, id) - } - } - } else { - targetDocIDs = filteredSet - } - - // 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 - } - } - - ids := make([]string, 0, len(targetDocIDs)) - for id := range targetDocIDs { - ids = append(ids, id) - } - - // Apply updates and deletes per document using Python's batch_update_metadata - // semantics instead of a simple merge-then-delete. - updated := 0 - for _, docID := range ids { - currentMeta, err := s.GetDocumentMetadataByID(docID) - if err != nil { - common.Warn("BatchUpdateDocumentMetadata: get metadata failed", - zap.String("docID", docID), zap.Error(err)) - continue - } - - meta := cloneDocumentMetadata(currentMeta) - originalMeta := cloneDocumentMetadata(meta) - - changed := applyDocumentMetadataUpdates(meta, updates) - if applyDocumentMetadataDeletes(meta, deletes) { - changed = true - } - - if !changed || reflect.DeepEqual(originalMeta, meta) { - continue - } - - if err := s.patchDocumentMetadata(docID, originalMeta, meta); err != nil { - common.Warn("BatchUpdateDocumentMetadata: patch metadata failed", - zap.String("docID", docID), zap.Error(err)) - continue - } - updated++ - } - - return &BatchUpdateDocumentMetadatasResponse{Updated: updated, MatchedDocs: len(ids)}, common.CodeSuccess, nil -} - -func (s *DocumentService) UploadDocumentInfos(userID string, files []*multipart.FileHeader) ([]map[string]interface{}, common.ErrorCode, error) { - fileSvc := &FileService{ - fileDAO: s.fileDAO, - file2DocumentDAO: s.file2DocumentDAO, - documentService: s, - } - data, err := fileSvc.UploadInfos(userID, files) - if err != nil { - return nil, common.CodeDataError, err - } - return data, common.CodeSuccess, nil -} - -func (s *DocumentService) UploadDocumentInfoByURL(userID, rawURL string) (map[string]interface{}, common.ErrorCode, error) { - fileSvc := &FileService{ - fileDAO: s.fileDAO, - file2DocumentDAO: s.file2DocumentDAO, - documentService: s, - } - data, err := fileSvc.UploadFromURL(userID, rawURL) - if err != nil { - return nil, common.CodeDataError, err - } - return data, common.CodeSuccess, nil -} - -func validateBatchUpdateDocumentMetadatasRequest( - selector *DocumentMetadataSelector, - updates []DocumentMetadataUpdate, - deletes []DocumentMetadataDelete, -) (common.ErrorCode, error) { - for _, upd := range updates { - if strings.TrimSpace(upd.Key) == "" || upd.Value == nil { - return common.CodeDataError, errors.New("Each update requires key and value.") - } - } - for _, del := range deletes { - if strings.TrimSpace(del.Key) == "" { - return common.CodeDataError, errors.New("Each delete requires key.") - } - } - if selector != nil && selector.MetadataCondition != nil { - if _, ok := selector.MetadataCondition["conditions"]; !ok && len(selector.MetadataCondition) > 0 { - return common.CodeDataError, errors.New("metadata_condition must be an object.") - } - } - return common.CodeSuccess, nil -} - -func cloneDocumentMetadata(meta map[string]interface{}) map[string]interface{} { - if meta == nil { - return map[string]interface{}{} - } - cloned := make(map[string]interface{}, len(meta)) - for k, v := range meta { - cloned[k] = cloneDocumentMetadataValue(v) - } - return cloned -} - -func cloneDocumentMetadataValue(v interface{}) interface{} { - switch typed := v.(type) { - case []interface{}: - cp := make([]interface{}, len(typed)) - copy(cp, typed) - return cp - case []string: - cp := make([]interface{}, 0, len(typed)) - for _, item := range typed { - cp = append(cp, item) - } - return cp - default: - return typed - } -} - -func applyDocumentMetadataUpdates(meta map[string]interface{}, updates []DocumentMetadataUpdate) bool { - changed := false - for _, upd := range updates { - key := strings.TrimSpace(upd.Key) - if key == "" { - continue - } - normalizedValue := normalizeDocumentMetadataUpdateValue(upd.Value, upd.ValueType) - matchProvided := upd.Match != nil && !(fmt.Sprintf("%v", upd.Match) == "") - current, exists := meta[key] - if !exists { - if matchProvided { - continue - } - if listVal, ok := toMetadataInterfaceSlice(normalizedValue); ok { - meta[key] = dedupeDocumentMetadataList(listVal) - } else { - meta[key] = normalizedValue - } - changed = true - continue - } - - if curList, ok := toMetadataInterfaceSlice(current); ok { - if !matchProvided { - newList := append([]interface{}{}, curList...) - if appendList, ok := toMetadataInterfaceSlice(normalizedValue); ok { - newList = append(newList, appendList...) - } else { - newList = append(newList, normalizedValue) - } - newList = dedupeDocumentMetadataList(newList) - if !reflect.DeepEqual(curList, newList) { - meta[key] = newList - changed = true - } - continue - } - - replaced := false - newList := make([]interface{}, 0, len(curList)) - for _, item := range curList { - if documentMetadataValuesEqual(item, upd.Match) { - if replacementList, ok := toMetadataInterfaceSlice(normalizedValue); ok { - newList = append(newList, replacementList...) - } else { - newList = append(newList, normalizedValue) - } - replaced = true - } else { - newList = append(newList, item) - } - } - newList = dedupeDocumentMetadataList(newList) - if replaced && !reflect.DeepEqual(curList, newList) { - meta[key] = newList - changed = true - } - continue - } - - if !matchProvided { - if !reflect.DeepEqual(current, normalizedValue) { - meta[key] = normalizedValue - changed = true - } - continue - } - if documentMetadataValuesEqual(current, upd.Match) && !reflect.DeepEqual(current, normalizedValue) { - meta[key] = normalizedValue - changed = true - } - } - return changed -} - -func applyDocumentMetadataDeletes(meta map[string]interface{}, deletes []DocumentMetadataDelete) bool { - changed := false - for _, del := range deletes { - key := strings.TrimSpace(del.Key) - current, exists := meta[key] - if key == "" || !exists { - continue - } - - if curList, ok := toMetadataInterfaceSlice(current); ok { - if del.Value == nil { - delete(meta, key) - changed = true - continue - } - newList := make([]interface{}, 0, len(curList)) - for _, item := range curList { - if !documentMetadataValuesEqual(item, del.Value) { - newList = append(newList, item) - } - } - if len(newList) != len(curList) { - if len(newList) == 0 { - delete(meta, key) - } else { - meta[key] = newList - } - changed = true - } - continue - } - - if del.Value == nil || documentMetadataValuesEqual(current, del.Value) { - delete(meta, key) - changed = true - } - } - return changed -} - -func toMetadataInterfaceSlice(v interface{}) ([]interface{}, bool) { - switch typed := v.(type) { - case []interface{}: - cp := make([]interface{}, len(typed)) - copy(cp, typed) - return cp, true - case []string: - cp := make([]interface{}, 0, len(typed)) - for _, item := range typed { - cp = append(cp, item) - } - return cp, true - default: - return nil, false - } -} - -func dedupeDocumentMetadataList(items []interface{}) []interface{} { - result := make([]interface{}, 0, len(items)) - seen := make(map[string]struct{}, len(items)) - for _, item := range items { - key := fmt.Sprintf("%T:%v", item, item) - if _, ok := seen[key]; ok { - continue - } - seen[key] = struct{}{} - result = append(result, item) - } - return result -} - -func documentMetadataValuesEqual(a, b interface{}) bool { - return fmt.Sprintf("%v", a) == fmt.Sprintf("%v", b) -} - -func normalizeDocumentMetadataUpdateValue(value interface{}, valueType string) interface{} { - switch strings.ToLower(strings.TrimSpace(valueType)) { - case "list": - if list, ok := normalizeMetadataListValue(value); ok { - return list - } - return []interface{}{} - case "number": - scalar, ok := firstScalarMetadataValue(value) - if !ok { - return value - } - switch typed := scalar.(type) { - case float64, float32, int, int8, int16, int32, int64: - return typed - case json.Number: - if i, err := typed.Int64(); err == nil { - return i - } - if f, err := typed.Float64(); err == nil { - return f - } - case string: - trimmed := strings.TrimSpace(typed) - if trimmed == "" { - return "" - } - if i, err := strconv.ParseInt(trimmed, 10, 64); err == nil { - return i - } - if f, err := strconv.ParseFloat(trimmed, 64); err == nil { - return f - } - return trimmed - } - return scalar - case "string", "time": - if scalar, ok := firstScalarMetadataValue(value); ok { - return fmt.Sprintf("%v", scalar) - } - return "" - default: - return value - } -} - -func normalizeMetadataListValue(value interface{}) ([]interface{}, bool) { - switch typed := value.(type) { - case []interface{}: - result := make([]interface{}, 0, len(typed)) - for _, item := range typed { - if nested, ok := normalizeMetadataListValue(item); ok { - result = append(result, nested...) - continue - } - if item != nil { - result = append(result, item) - } - } - return result, true - case []string: - result := make([]interface{}, 0, len(typed)) - for _, item := range typed { - result = append(result, item) - } - return result, true - default: - return nil, false - } -} - -func firstScalarMetadataValue(value interface{}) (interface{}, bool) { - if list, ok := normalizeMetadataListValue(value); ok { - for _, item := range list { - if item != nil { - return item, true - } - } - return nil, false - } - if value == nil { - return nil, false - } - return value, true -} diff --git a/internal/service/document/document.go b/internal/service/document/document.go new file mode 100644 index 0000000000..be266f7193 --- /dev/null +++ b/internal/service/document/document.go @@ -0,0 +1,274 @@ +// +// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// + +package document + +import ( + "errors" + "ragflow/internal/service" + "regexp" + "time" + + "ragflow/internal/dao" + "ragflow/internal/engine" +) + +// DocumentService document service +type DocumentService struct { + documentDAO *dao.DocumentDAO + kbDAO *dao.KnowledgebaseDAO + ingestionTaskDAO *dao.IngestionTaskDAO + ingestionTaskLogDAO *dao.IngestionTaskLogDAO + ingestionTaskSvc *service.IngestionTaskService + docEngine engine.DocEngine + metadataSvc *service.MetadataService + taskDAO *dao.TaskDAO + file2DocumentDAO *dao.File2DocumentDAO + fileDAO *dao.FileDAO + canvasDAO *dao.UserCanvasDAO + api4ConvDAO *dao.API4ConversationDAO +} + +// NewDocumentService create document service +func NewDocumentService() *DocumentService { + publisher := service.NewMessageQueueTaskPublisher() + ingestionTaskSvc := service.NewIngestionTaskService() + ingestionTaskSvc.SetTaskPublisher(publisher) + return &DocumentService{ + documentDAO: dao.NewDocumentDAO(), + ingestionTaskDAO: dao.NewIngestionTaskDAO(), + ingestionTaskLogDAO: dao.NewIngestionTaskLogDAO(), + ingestionTaskSvc: ingestionTaskSvc, + kbDAO: dao.NewKnowledgebaseDAO(), + docEngine: engine.Get(), + metadataSvc: service.NewMetadataService(), + taskDAO: dao.NewTaskDAO(), + file2DocumentDAO: dao.NewFile2DocumentDAO(), + fileDAO: dao.NewFileDAO(), + canvasDAO: dao.NewUserCanvasDAO(), + api4ConvDAO: dao.NewAPI4ConversationDAO(), + } +} + +// CreateDocumentRequest create document request +// UpdateDocumentRequest update document request +type UpdateDocumentRequest struct { + Name *string `json:"name"` + Run *string `json:"run"` + TokenNum *int64 `json:"token_num"` + ChunkNum *int64 `json:"chunk_num"` + Progress *float64 `json:"progress"` + ProgressMsg *string `json:"progress_msg"` +} + +// DocumentResponse document response +type DocumentResponse struct { + ID string `json:"id"` + Name *string `json:"name,omitempty"` + KbID string `json:"kb_id"` + ParserID string `json:"parser_id"` + PipelineID *string `json:"pipeline_id,omitempty"` + Type string `json:"type"` + SourceType string `json:"source_type"` + CreatedBy string `json:"created_by"` + Location *string `json:"location,omitempty"` + Size int64 `json:"size"` + TokenNum int64 `json:"token_num"` + ChunkNum int64 `json:"chunk_num"` + Progress float64 `json:"progress"` + ProgressMsg *string `json:"progress_msg,omitempty"` + ProcessDuration float64 `json:"process_duration"` + Suffix string `json:"suffix"` + Run *string `json:"run,omitempty"` + Status *string `json:"status,omitempty"` + CreatedAt string `json:"created_at"` + UpdatedAt string `json:"updated_at"` +} + +type ThumbnailResponse struct { + ID string `json:"id"` + Thumbnail *string `json:"thumbnail,omitempty"` + KbID string `json:"kb_id"` +} + +const imgBase64Prefix = "data:image/png;base64," + +type ArtifactResponse struct { + Data []byte + ContentType string + SafeFilename string + ForceAttachment bool +} + +type UpdateDatasetDocumentRequest struct { + Name *string `json:"name"` + ParserID *string `json:"parser_id"` + ChunkCount *int64 `json:"chunk_count"` + TokenCount *int64 `json:"token_count"` + PipelineID *string `json:"pipeline_id"` + Enabled *int `json:"enabled"` + Progress *float64 `json:"progress"` + ParserConfig map[string]any `json:"parser_config"` + MetaFields map[string]any `json:"meta_fields"` +} + +// PATCH /api/v1/datasets/:dataset_id/documents/:document_id. +type UpdateDatasetDocumentResponse struct { + ID string `json:"id"` + Thumbnail *string `json:"thumbnail,omitempty"` + DatasetID string `json:"dataset_id"` + ParserID string `json:"parser_id"` + PipelineID *string `json:"pipeline_id,omitempty"` + ParserConfig map[string]interface{} `json:"parser_config"` + SourceType string `json:"source_type"` + Type string `json:"type"` + CreatedBy string `json:"created_by"` + Name *string `json:"name,omitempty"` + Location *string `json:"location,omitempty"` + Size int64 `json:"size"` + TokenCount int64 `json:"token_count"` + ChunkCount int64 `json:"chunk_count"` + Progress float64 `json:"progress"` + ProgressMsg *string `json:"progress_msg,omitempty"` + ProcessBeginAt *time.Time `json:"process_begin_at,omitempty"` + ProcessDuration float64 `json:"process_duration"` + ContentHash *string `json:"content_hash,omitempty"` + MetaFields map[string]interface{} `json:"meta_fields,omitempty"` + Suffix string `json:"suffix"` + Run string `json:"run"` + Status *string `json:"status,omitempty"` + CreateTime *int64 `json:"create_time,omitempty"` + CreateDate *time.Time `json:"create_date,omitempty"` + UpdateTime *int64 `json:"update_time,omitempty"` + UpdateDate *time.Time `json:"update_date,omitempty"` +} + +var ( + ErrArtifactInvalidFilename = errors.New("Invalid filename.") + ErrArtifactInvalidFileType = errors.New("Invalid file type.") + ErrArtifactNotFound = errors.New("Artifact not found.") +) + +var artifactContentTypes = map[string]string{ + ".png": "image/png", + ".jpg": "image/jpeg", + ".jpeg": "image/jpeg", + ".svg": "image/svg+xml", + ".pdf": "application/pdf", + ".csv": "text/csv", + ".json": "application/json", + ".html": "text/html", +} + +var artifactForceAttachmentExtensions = map[string]struct{}{ + ".htm": {}, + ".html": {}, + ".shtml": {}, + ".xht": {}, + ".xhtml": {}, + ".xml": {}, + ".mhtml": {}, + ".svg": {}, +} +var artifactForceAttachmentContentTypes = map[string]struct{}{ + "text/html": {}, + "image/svg+xml": {}, + "application/xhtml+xml": {}, + "text/xml": {}, + "application/xml": {}, + "multipart/related": {}, +} + +var artifactUnsafeFilenameChars = regexp.MustCompile(`[^\pL\pN_.-]`) + +type DocumentPreview struct { + Data []byte + ContentType string + FileName string +} + +type DownloadDocumentResp struct { + Data []byte + FileName string + ContentType string +} + +type IngestDocumentRequest struct { + DocIDs []string `json:"doc_ids" binding:"required"` + Run interface{} `json:"run" binding:"required"` + Delete bool `json:"delete"` + ApplyKB bool `json:"apply_kb"` +} + +// StartParseOptions controls StartParseDocuments behavior. +type StartParseOptions struct { + // ApplyKB merges the knowledgebase's parser_config (llm_id, metadata) + // into the document before parsing. + ApplyKB bool + // RerunWithDelete clears prior chunks/tasks/counters before re-parsing. + RerunWithDelete bool +} + +// GetMetadataSummaryRequest request for metadata summary +type GetMetadataSummaryRequest struct { + KBID string `json:"kb_id" binding:"required"` + DocIDs []string `json:"doc_ids"` +} + +// GetMetadataSummaryResponse response for metadata summary +type GetMetadataSummaryResponse struct { + Summary map[string]interface{} `json:"summary"` +} + +// valueInfo holds count and order of first appearance +type valueInfo struct { + count int + firstOrder int +} + +// knowledgebaseFolderName is the file-manager folder under each tenant's root +// that holds per-dataset subfolders, mirroring Python KNOWLEDGEBASE_FOLDER_NAME. +const knowledgebaseFolderName = ".knowledgebase" + +// maxUploadDocSize bounds a single uploaded file held in memory, mirroring the +// Python DOC_MAXIMUM_SIZE default (128 MiB; overridable there via MAX_CONTENT_LENGTH). +const maxUploadDocSize = 128 * 1024 * 1024 + +// MetadataUpdate is one update item: set key to value. +type DocumentMetadataUpdate struct { + Key string `json:"key"` + Value interface{} `json:"value"` + Match interface{} `json:"match,omitempty"` + ValueType string `json:"valueType,omitempty"` +} + +// MetadataDelete removes a whole key, or a specific value from a list field. +type DocumentMetadataDelete struct { + Key string `json:"key"` + Value interface{} `json:"value,omitempty"` +} + +// MetadataSelector selects which documents to target. +type DocumentMetadataSelector struct { + DocumentIDs []string `json:"document_ids"` + MetadataCondition map[string]interface{} `json:"metadata_condition"` +} + +// BatchUpdateDocumentMetadatasResponse summarises the operation. +type BatchUpdateDocumentMetadatasResponse struct { + Updated int `json:"updated"` + MatchedDocs int `json:"matched_docs"` +} diff --git a/internal/service/document/document_artifact.go b/internal/service/document/document_artifact.go new file mode 100644 index 0000000000..6a2707c881 --- /dev/null +++ b/internal/service/document/document_artifact.go @@ -0,0 +1,232 @@ +package document + +import ( + "fmt" + "path/filepath" + "strings" + + "ragflow/internal/common" + "ragflow/internal/dao" + "ragflow/internal/entity" + "ragflow/internal/storage" + "ragflow/internal/utility" +) + +// GetDocumentImage retrieves an image object from storage. +func (s *DocumentService) GetDocumentImage(imageID string) ([]byte, error) { + parts := strings.SplitN(imageID, "-", 2) + if len(parts) != 2 || parts[0] == "" || parts[1] == "" { + return nil, fmt.Errorf("Image not found.") + } + + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return nil, fmt.Errorf("storage not initialized") + } + + return storageImpl.Get(parts[0], parts[1]) +} + +// GetDocumentArtifact retrieves a sandbox artifact from object storage. +// +// userID scopes the lookup: a CodeExec sandbox artifact is only +// returned when the caller owns (or has team access to) at least +// one agent session whose `message` references this filename (or +// its `documents/artifact/` URL form). The authorization +// gate runs BEFORE the storage read so a probe of an unknown +// filename cannot distinguish "you cannot see it" from "it +// exists" — both return ErrArtifactNotFound. Mirrors PR #16169. +func (s *DocumentService) GetDocumentArtifact(filename, userID string) (*ArtifactResponse, error) { + basename := filepath.Base(filename) + if basename != filename || strings.Contains(filename, "/") || strings.Contains(filename, "\\") { + return nil, ErrArtifactInvalidFilename + } + + ext := strings.ToLower(filepath.Ext(basename)) + contentType, ok := artifactContentTypes[ext] + if !ok { + return nil, ErrArtifactInvalidFileType + } + + if !s.sandboxArtifactAccessible(basename, userID) { + // Same error as "object does not exist" to avoid leaking + // whether the artifact exists for a different user/agent. + return nil, ErrArtifactNotFound + } + + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return nil, fmt.Errorf("storage not initialized") + } + + bucket := sandboxArtifactBucket() + if !storageImpl.ObjExist(bucket, basename) { + return nil, ErrArtifactNotFound + } + + data, err := storageImpl.Get(bucket, basename) + if err != nil { + return nil, err + } + if len(data) == 0 { + return nil, ErrArtifactNotFound + } + + return &ArtifactResponse{ + Data: data, + ContentType: contentType, + SafeFilename: sanitizeArtifactFilename(basename), + ForceAttachment: shouldForceArtifactAttachment(ext, contentType), + }, nil +} + +// sandboxArtifactDialogIDsForUser returns the distinct agent +// (canvas) dialog_ids for sessions owned by userID whose +// `message` blob references filename. A CodeExec artifact URL +// appears in `message` as either a bare filename or the +// `documents/artifact/` form, so the helper matches both. +// +// Implemented as a direct GORM query on the +// API4Conversation table — GORM's `Contains` maps to MySQL +// `LIKE '%...%'` which is fine here because the storage path is +// short and indexed lookup on (user_id, exp_user_id) keeps the +// scan narrow. +func (s *DocumentService) sandboxArtifactDialogIDsForUser(filename, userID string) []string { + if filename == "" || userID == "" { + return nil + } + // Escape SQL LIKE wildcards (%, _) before building the pattern. + // Without escaping, a caller could submit a filename like + // "%.png" or "_" and the LIKE query would match arbitrary + // referenced artifacts in any user's conversation — letting the + // caller pass the authorization check against one filename and + // then GET another artifact by name (PR review round 5, Major #8). + // + // Escape character: '!'. We avoid '\\' because SQL string + // literal parsing of '\\' is driver-specific (SQLite treats + // it as a single backslash, MySQL treats it as one, Postgres + // rejects the unterminated string) — '!' is a benign character + // in real filenames (artifact names rarely contain '!') and + // parses identically in every driver. + filenameSafe := escapeSQLLikePattern(filename) + artifactRefSafe := escapeSQLLikePattern("documents/artifact/" + filename) + filenamePattern := "%" + filenameSafe + "%" + artifactRefPattern := "%" + artifactRefSafe + "%" + dialogIDs := make(map[string]struct{}) + rows, err := dao.DB.Model(&entity.API4Conversation{}). + Select("dialog_id"). + Where("user_id = ? OR exp_user_id = ?", userID, userID). + Where(`message LIKE ? ESCAPE '!' OR message LIKE ? ESCAPE '!'`, + filenamePattern, artifactRefPattern). + Distinct("dialog_id"). + Rows() + if err != nil { + return nil + } + defer rows.Close() + for rows.Next() { + var d string + if err := rows.Scan(&d); err == nil && d != "" { + dialogIDs[d] = struct{}{} + } + } + out := make([]string, 0, len(dialogIDs)) + for d := range dialogIDs { + out = append(out, d) + } + return out +} + +// sandboxArtifactAccessible reports whether userID may reach at +// least one agent canvas whose session references filename. +// Mirrors `UserCanvasService.accessible(dialog_id, user_id)` from +// the Python fix; on the Go side this is the same predicate as +// UserCanvasDAO.Accessible (owner or team permission, with the +// latter scoped to the caller's tenant membership — PR review +// round 5). +func (s *DocumentService) sandboxArtifactAccessible(filename, userID string) bool { + if userID == "" { + return false + } + // Fetch the caller's tenant list once; passing it into + // canvasDAO.Accessible ensures the team-permission branch only + // matches canvases the caller can actually see. An empty list + // (callers without tenant data) is safe — it effectively disables + // the team branch, so the only matches are canvases the caller + // directly owns. + tenantIDs, terr := dao.NewUserTenantDAO().GetTenantIDsByUserID(userID) + if terr != nil { + tenantIDs = nil + } + for _, dialogID := range s.sandboxArtifactDialogIDsForUser(filename, userID) { + if s.canvasDAO.Accessible(dialogID, userID, tenantIDs) { + return true + } + } + return false +} + +func sandboxArtifactBucket() string { + if bucket := common.GetEnv(common.EnvSandboxArtifactBucket); bucket != "" { + return bucket + } + return "sandbox-artifacts" +} + +// sanitizeArtifactFilename scrubs characters that are unsafe inside a storage +// object key for sandbox artifacts. It intentionally only replaces the +// artifact-specific unsafe set (artifactUnsafeFilenameChars) and does NOT strip +// directory components, reject reserved device names, or bound length — those +// concerns belong to the general-purpose sanitizeFilename used for uploaded / +// URL-derived filenames. The two are deliberately separate because their +// safety rules differ; do not merge them. +func sanitizeArtifactFilename(filename string) string { + return artifactUnsafeFilenameChars.ReplaceAllString(filename, "_") +} + +func shouldForceArtifactAttachment(ext, contentType string) bool { + if _, ok := artifactForceAttachmentExtensions[strings.ToLower(ext)]; ok { + return true + } + _, ok := artifactForceAttachmentContentTypes[strings.ToLower(contentType)] + return ok +} + +func (s *DocumentService) GetDocumentPreview(docID string) (*DocumentPreview, error) { + doc, err := s.documentDAO.GetByID(docID) + if err != nil { + return nil, err + } + + bucket, name, err := s.GetDocumentStorageAddress(doc) + if err != nil { + return nil, err + } + + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return nil, fmt.Errorf("storage not initialized") + } + + data, err := storageImpl.Get(bucket, name) + if err != nil { + return nil, err + } + if len(data) == 0 { + return nil, ErrArtifactNotFound + } + + fileName := "" + if doc.Name != nil { + fileName = *doc.Name + } + + ext := utility.GetFileExtension(fileName) + contentType := utility.GetContentType(ext, doc.Type) + + return &DocumentPreview{ + Data: data, + ContentType: contentType, + FileName: fileName, + }, nil +} diff --git a/internal/service/document/document_crud.go b/internal/service/document/document_crud.go new file mode 100644 index 0000000000..fa1813c4b8 --- /dev/null +++ b/internal/service/document/document_crud.go @@ -0,0 +1,436 @@ +package document + +import ( + "context" + "fmt" + + "ragflow/internal/common" + "ragflow/internal/dao" + "ragflow/internal/entity" + "ragflow/internal/storage" + + "gorm.io/gorm" +) + +// Accessible reports whether docID belongs to a knowledge base +// reachable by userID. Used by agent endpoints (e.g. RerunAgent, +// PR #15145) to gate destructive / run-again actions on a document +// the caller has access to. Returns false on any lookup failure or +// empty inputs so callers can treat a denial as a 404-equivalent +// and avoid leaking whether the document exists at all. +func (s *DocumentService) Accessible(docID, userID string) bool { + if docID == "" || userID == "" { + return false + } + doc, err := s.documentDAO.GetByID(docID) + if err != nil || doc == nil { + return false + } + return s.kbDAO.Accessible(doc.KbID, userID) +} + +func (s *DocumentService) GetDocumentStorageAddress(doc *entity.Document) (string, string, error) { + if doc == nil { + return "", "", fmt.Errorf("document is nil") + } + + file2DocumentDAO := dao.NewFile2DocumentDAO() + fileDAO := dao.NewFileDAO() + + mappings, err := file2DocumentDAO.GetByDocumentID(doc.ID) + if err != nil { + return "", "", err + } + + if len(mappings) > 0 && mappings[0].FileID != nil { + file, err := fileDAO.GetByID(*mappings[0].FileID) + if err != nil { + return "", "", err + } + + if file.SourceType == "" || entity.FileSource(file.SourceType) == entity.FileSourceLocal { + if file.Location == nil || *file.Location == "" { + return "", "", fmt.Errorf("file location is empty") + } + return file.ParentID, *file.Location, nil + } + } + + if doc.Location == nil || *doc.Location == "" { + return "", "", fmt.Errorf("document location is empty") + } + return doc.KbID, *doc.Location, nil +} + +func (s *DocumentService) DownloadDocument(datasetID, docID string) (*DownloadDocumentResp, error) { + if docID == "" { + return nil, fmt.Errorf("Specify document_id please.") + } + doc, err := s.documentDAO.GetByID(docID) + if err != nil || doc.KbID != datasetID { + return nil, fmt.Errorf("The dataset not own the document %s.", docID) + } + bucket, name, err := s.GetDocumentStorageAddress(doc) + if err != nil { + return nil, err + } + + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return nil, fmt.Errorf("storage not initialized") + } + + data, err := storageImpl.Get(bucket, name) + if err != nil { + return nil, err + } + if len(data) == 0 { + return nil, fmt.Errorf("This file is empty.") + } + + fileName := "" + if doc.Name != nil { + fileName = *doc.Name + } + + return &DownloadDocumentResp{ + Data: data, + FileName: fileName, + ContentType: "application/octet-stream", + }, nil +} + +// CreateDocument create document +// GetDocumentByID get document by ID +func (s *DocumentService) GetDocumentByID(id string) (*DocumentResponse, error) { + document, err := s.documentDAO.GetByID(id) + if err != nil { + return nil, err + } + + return s.toResponse(document), nil +} + +// UpdateDocument update document +func (s *DocumentService) UpdateDocument(id string, req *UpdateDocumentRequest) error { + document, err := s.documentDAO.GetByID(id) + if err != nil { + return err + } + + if req.Name != nil { + document.Name = req.Name + } + if req.Run != nil { + document.Run = req.Run + } + if req.TokenNum != nil { + document.TokenNum = *req.TokenNum + } + if req.ChunkNum != nil { + document.ChunkNum = *req.ChunkNum + } + if req.Progress != nil { + document.Progress = *req.Progress + } + if req.ProgressMsg != nil { + document.ProgressMsg = req.ProgressMsg + } + + return s.documentDAO.Update(document) +} + +// IncrementChunkNum atomically increments chunk/token counters on the document and its knowledge base in a transaction +func (s *DocumentService) IncrementChunkNum(docID, kbID string, chunkNum, tokenNum int, duration float64) error { + return dao.DB.Transaction(func(tx *gorm.DB) error { + // Update document + if err := tx.Model(&entity.Document{}). + Where("id = ? AND kb_id = ?", docID, kbID). + Updates(map[string]interface{}{ + "chunk_num": gorm.Expr("chunk_num + ?", int64(chunkNum)), + "token_num": gorm.Expr("token_num + ?", int64(tokenNum)), + "process_duration": gorm.Expr("process_duration + ?", duration), + }).Error; err != nil { + return err + } + + // Update knowledgebase + if err := tx.Model(&entity.Knowledgebase{}). + Where("id = ?", kbID). + Updates(map[string]interface{}{ + "chunk_num": gorm.Expr("chunk_num + ?", int64(chunkNum)), + "token_num": gorm.Expr("token_num + ?", int64(tokenNum)), + }).Error; err != nil { + return err + } + + return nil + }) +} + +// UpdateRunProgress mirrors a pipeline run's live progress into the document +// row so the document-list endpoint (which reads document.progress/run/ +// progress_msg) reflects in-flight Go pipeline progress. Best-effort by +// design; callers log and continue on error. +func (s *DocumentService) UpdateRunProgress(docID string, progress float64, run, progressMsg string) error { + return s.documentDAO.UpdateByID(docID, map[string]interface{}{ + "progress": progress, + "run": run, + "progress_msg": progressMsg, + }) +} + +// DeleteDocument delete document — delegates to full cleanup logic. +func (s *DocumentService) DeleteDocument(id string) error { + return s.deleteDocumentFull(id) +} + +// DeleteDocuments deletes multiple documents under a dataset. +// +// ids: specific document IDs; deleteAll: delete all docs in the dataset. +// Returns the number of successfully deleted documents. +func (s *DocumentService) DeleteDocuments(ids []string, deleteAll bool, datasetID, userID string) (int, error) { + // 1. Check dataset is accessible by the user + if !s.kbDAO.Accessible(datasetID, userID) { + return 0, fmt.Errorf("You don't own the dataset %s.", datasetID) + } + + // 2. Resolve document IDs + if deleteAll { + if err := dao.DB.Model(&entity.Document{}). + Where("kb_id = ?", datasetID). + Pluck("id", &ids).Error; err != nil { + return 0, fmt.Errorf("failed to query documents: %w", err) + } + } + if len(ids) == 0 { + return 0, nil + } + + // 3. Deduplicate (before validation so dup count doesn't matter) + ids = common.Deduplicate(ids) + + // 4. Validate IDs belong to this dataset (only for explicit ids; deleteAll is already scoped) + if !deleteAll { + if _, err := s.validateDocsInDataset(ids, datasetID); err != nil { + return 0, err + } + } + + // 5. Delete each document (non-critical failures are tolerated per doc) + deleted := 0 + for _, docID := range ids { + if err := s.deleteDocumentFull(docID); err != nil { + common.Warn(fmt.Sprintf("DeleteDocuments: failed to delete %s: %v", docID, err)) + continue + } + deleted++ + } + + return deleted, nil +} + +// deleteDocumentFull performs full document cleanup. Non-critical failures +// are tolerated (logged and continue). Critical failures (e.g. document or +// KB not found) return an error immediately. +func (s *DocumentService) deleteDocumentFull(docID string) error { + doc, kb, err := s.resolveDocAndKB(docID) + if err != nil { + return err + } + + // Delete tasks from DB + ingestionTask, err := s.ingestionTaskDAO.GetByDocumentID(docID) + if err != nil { + common.Error(fmt.Sprintf("failed to get ingestion task by doc:%s", doc.ID), err) + return err + } + if ingestionTask != nil { + taskInfo, err := s.ingestionTaskSvc.Remove(ingestionTask.ID, &ingestionTask.UserID) + if err != nil { + return err + } + // FIXME: need to add logic to delete files in taskInfo + common.Warn(fmt.Sprintf("need to delete files from taskInfo: %v", taskInfo)) + } + + s.deleteDocEngineData(docID, kb.TenantID, doc.KbID) + if err := s.deleteDocRecordWithCounters(doc, kb.ID); err != nil { + return err + } + s.cleanupFileReferences(docID) + + return nil +} + +// RemoveDocumentKeepFile removes a document's chunks/metadata and the document +// row, decrementing the KB counters (doc_num/chunk_num/token_num), WITHOUT +// deleting the underlying file record, its storage blob, or its file2document +// mappings. Mirrors Python DocumentService.remove_document — the caller is +// responsible for cleaning up the file2document mappings separately. +func (s *DocumentService) RemoveDocumentKeepFile(docID string) error { + doc, kb, err := s.resolveDocAndKB(docID) + if err != nil { + return err + } + if _, delErr := s.taskDAO.DeleteByDocIDs([]string{docID}); delErr != nil { + common.Logger.Warn(fmt.Sprintf("RemoveDocumentKeepFile: failed to delete tasks for %s: %v", docID, delErr)) + } + s.deleteDocEngineData(docID, kb.TenantID, doc.KbID) + return s.deleteDocRecordWithCounters(doc, kb.ID) +} + +// InsertDocument creates a document row and increments the owning KB's doc_num +// counter in a single transaction. Mirrors Python DocumentService.insert, which +// updates dataset/document counters on insert. The document's ID and timestamps +// are populated by the caller / model hooks before insertion. +func (s *DocumentService) InsertDocument(doc *entity.Document) error { + return dao.DB.Transaction(func(tx *gorm.DB) error { + if err := tx.Create(doc).Error; err != nil { + return fmt.Errorf("failed to create document: %w", err) + } + // Guard the counter bump with RowsAffected: documents.kb_id has no DB-level + // FK, so Create can succeed against a non-existent KB and the Update would + // then report a nil error with 0 rows touched, silently desyncing doc_num. + // Roll the whole transaction back in that case (mirrors the counter checks + // in deleteDocRecordWithCounters). + result := tx.Model(&entity.Knowledgebase{}). + Where("id = ?", doc.KbID). + Update("doc_num", gorm.Expr("doc_num + 1")) + if result.Error != nil { + return fmt.Errorf("failed to increment doc_num for KB %s: %w", doc.KbID, result.Error) + } + if result.RowsAffected == 0 { + return fmt.Errorf("knowledgebase %s not found", doc.KbID) + } + return nil + }) +} + +// resolveDocAndKB loads the document and its knowledgebase, returning both or +// an error. +func (s *DocumentService) resolveDocAndKB(docID string) (*entity.Document, *entity.Knowledgebase, error) { + doc, err := s.documentDAO.GetByID(docID) + if err != nil { + return nil, nil, fmt.Errorf("document not found: %w", err) + } + kb, err := s.kbDAO.GetByID(doc.KbID) + if err != nil { + return nil, nil, fmt.Errorf("knowledgebase not found: %w", err) + } + return doc, kb, nil +} + +// deleteDocEngineData removes chunks and metadata from the document engine. +// No-op when the engine is nil. +func (s *DocumentService) deleteDocEngineData(docID, tenantID, kbID string) { + if s.docEngine == nil { + return + } + 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)) + } + if s.metadataSvc != nil { + _ = s.DeleteDocumentAllMetadata(docID) // logs internally + } +} + +// deleteDocRecordWithCounters hard-deletes the document row and decrements the +// KB counters in a single transaction. Counters are only decremented when a +// document row was actually removed (RowsAffected > 0), guarding against +// double-decrement on retries or concurrent deletes. +func (s *DocumentService) deleteDocRecordWithCounters(doc *entity.Document, kbID string) error { + return dao.DB.Transaction(func(tx *gorm.DB) error { + result := tx.Where("id = ?", doc.ID).Delete(&entity.Document{}) + if result.Error != nil { + return fmt.Errorf("failed to delete document %s: %w", doc.ID, result.Error) + } + if result.RowsAffected == 0 { + return nil // already deleted by a concurrent request — skip counters + } + + result = tx.Model(&entity.Knowledgebase{}). + Where("id = ?", kbID). + Updates(map[string]interface{}{ + "doc_num": gorm.Expr("doc_num - 1"), + "chunk_num": gorm.Expr("chunk_num - ?", doc.ChunkNum), + "token_num": gorm.Expr("token_num - ?", doc.TokenNum), + }) + if result.Error != nil { + return fmt.Errorf("failed to decrement counters for KB %s: %w", kbID, result.Error) + } + if result.RowsAffected == 0 { + return fmt.Errorf("knowledgebase %s not found", kbID) + } + return nil + }) +} + +func (s *DocumentService) rollbackAddFileFromKBError(doc *entity.Document, kbID string, err error) error { + if cleanupErr := s.deleteDocRecordWithCounters(doc, kbID); cleanupErr != nil { + return fmt.Errorf("%w; rollback cleanup failed: %w", err, cleanupErr) + } + return err +} + +// cleanupFileReferences deletes file2document mappings for docID, and for each +// referenced file, only hard-deletes the file record and its storage blob when +// no other document still references the same file_id. +func (s *DocumentService) cleanupFileReferences(docID string) { + mappings, mapErr := s.file2DocumentDAO.GetByDocumentID(docID) + if mapErr != nil { + common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: failed to get f2d mappings for %s: %v", docID, mapErr)) + } + if len(mappings) == 0 { + return + } + + // Collect unique file_ids + seen := make(map[string]bool) + var fileIDs []string + for _, m := range mappings { + if m.FileID == nil || seen[*m.FileID] { + continue + } + seen[*m.FileID] = true + fileIDs = append(fileIDs, *m.FileID) + } + + // Delete all file2document rows for this document + if delErr := s.file2DocumentDAO.DeleteByDocumentID(docID); delErr != nil { + common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: failed to delete f2d for %s: %v", docID, delErr)) + } + + // For each file, only delete the record and blob when no other doc references it + for _, fileID := range fileIDs { + remaining, remErr := s.file2DocumentDAO.GetByFileID(fileID) + if remErr != nil { + common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: failed to check remaining f2d for %s: %v", fileID, remErr)) + continue + } + if len(remaining) > 0 { + continue + } + + fileDAO := dao.NewFileDAO() + file, fErr := fileDAO.GetByID(fileID) + if fErr != nil || file == nil { + common.Logger.Warn(fmt.Sprintf("cleanupFileReferences: file not found %s: %v", fileID, fErr)) + continue + } + if _, delErr := fileDAO.DeleteByIDs([]string{fileID}); delErr != nil { + common.Logger.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)) + } + } + } + } +} diff --git a/internal/service/document/document_dataset_update.go b/internal/service/document/document_dataset_update.go new file mode 100644 index 0000000000..5f92c649d9 --- /dev/null +++ b/internal/service/document/document_dataset_update.go @@ -0,0 +1,431 @@ +package document + +import ( + "context" + "errors" + "fmt" + "math" + "path/filepath" + "ragflow/internal/service" + "strconv" + "strings" + + "ragflow/internal/common" + "ragflow/internal/dao" + "ragflow/internal/entity" + pipelinepkg "ragflow/internal/ingestion/pipeline" + "ragflow/internal/tokenizer" + + "go.uber.org/zap" +) + +func (s *DocumentService) BatchUpdateDocumentStatus(userID, datasetID, status string, documentIDs []string) (map[string]interface{}, common.ErrorCode, error) { + kb, err := s.kbDAO.GetByIDAndTenantID(datasetID, userID) + if err != nil { + return nil, common.CodeDataError, fmt.Errorf("You don't own the dataset.") + } + statusInt, convErr := strconv.Atoi(status) + if convErr != nil { + return nil, common.CodeArgumentError, fmt.Errorf("invalid status: %s", status) + } + + result := make(map[string]interface{}, len(documentIDs)) + hasError := false + + documents, err := s.documentDAO.GetByIDs(documentIDs) + if err != nil { + return nil, common.CodeServerError, fmt.Errorf("failed to fetch documents: %w", err) + } + documentByID := make(map[string]*entity.Document, len(documents)) + for _, doc := range documents { + documentByID[doc.ID] = doc + } + + for _, docID := range documentIDs { + doc, ok := documentByID[docID] + if !ok { + result[docID] = map[string]string{"error": "Document not found"} + hasError = true + continue + } + + if doc.KbID != datasetID { + result[docID] = map[string]string{"error": "Document not found in this dataset."} + hasError = true + continue + } + + currentStatus := "" + if doc.Status != nil { + currentStatus = *doc.Status + } + if currentStatus == status { + result[docID] = map[string]string{"status": status} + continue + } + previousStatus := interface{}(nil) + if doc.Status != nil { + previousStatus = *doc.Status + } + if err := s.documentDAO.UpdateByID(docID, map[string]interface{}{"status": status}); err != nil { + result[docID] = map[string]string{"error": "Database error (Document update)!"} + hasError = true + continue + } + + if doc.ChunkNum > 0 { + if s.docEngine == nil { + _ = s.documentDAO.UpdateByID(docID, map[string]interface{}{"status": previousStatus}) + result[docID] = map[string]string{"error": "Document store update failed: document engine not initialized"} + hasError = true + continue + } + err := s.docEngine.UpdateChunks( + context.Background(), + map[string]interface{}{"doc_id": docID}, + map[string]interface{}{"available_int": statusInt}, + fmt.Sprintf("ragflow_%s", kb.TenantID), + doc.KbID, + ) + if err != nil { + _ = s.documentDAO.UpdateByID(docID, map[string]interface{}{"status": previousStatus}) + msg := err.Error() + if strings.Contains(msg, "3022") { + result[docID] = map[string]string{"error": "Document store table missing."} + } else { + result[docID] = map[string]string{"error": "Document store update failed: " + msg} + } + hasError = true + continue + } + } + result[docID] = map[string]string{"status": status} + } + + if hasError { + return result, common.CodeServerError, fmt.Errorf("Partial failure") + } + return result, common.CodeSuccess, nil +} + +func (s *DocumentService) UpdateDatasetDocument(userID, datasetID, documentID string, req *UpdateDatasetDocumentRequest, present map[string]bool) (*UpdateDatasetDocumentResponse, common.ErrorCode, error) { + tenantID := userID + kb, err := s.kbDAO.GetByIDAndTenantID(datasetID, tenantID) + if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, errors.New("You don't own the dataset.") + } + return nil, common.CodeDataError, errors.New("Can't find this dataset!") + } + + doc, err := s.documentDAO.GetByDocumentIDAndDatasetID(documentID, datasetID) + if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, errors.New("The dataset doesn't own the document.") + } + return nil, common.CodeServerError, err + } + + if code, err := s.validateDatasetDocumentUpdate(datasetID, documentID, userID, doc, req, present); err != nil { + return nil, code, err + } + + if present["meta_fields"] { + if err := s.replaceDocumentMetadata(documentID, req.MetaFields); err != nil { + return nil, common.CodeDataError, err + } + } + + if present["name"] && req.Name != nil && (doc.Name == nil || *req.Name != *doc.Name) { + if err := s.updateDocumentNameOnly(doc, kb.TenantID, *req.Name); err != nil { + return nil, common.CodeDataError, err + } + } + + if present["parser_config"] && req.ParserConfig != nil { + // Resolve effective pipeline to load the DSL for cleaning. + isCanvas := kb.PipelineID != nil && strings.TrimSpace(*kb.PipelineID) != "" + if req.PipelineID != nil { + isCanvas = strings.TrimSpace(*req.PipelineID) != "" + } + if req.ParserID != nil { + isCanvas = false + } + effParserID := kb.ParserID + if req.ParserID != nil { + effParserID = strings.TrimSpace(*req.ParserID) + } + effPipelineID := kb.PipelineID + if req.PipelineID != nil { + effPipelineID = req.PipelineID + } + if req.ParserID != nil && req.PipelineID == nil && kb.PipelineID != nil { + effPipelineID = nil + } + + dslJSON, err := service.LoadPipelineDSL(isCanvas, effParserID, effPipelineID) + if err != nil { + common.Warn("cleanAndUpdateDocumentParserConfig: failed to load DSL, falling back to merge", + zap.Error(err)) + if err := s.updateDocumentParserConfig(doc.ID, req.ParserConfig); err != nil { + return nil, common.CodeDataError, err + } + } else { + cleaned := pipelinepkg.BuildParserConfig(dslJSON, map[string]interface{}(req.ParserConfig)) + if err := s.documentDAO.UpdateByID(doc.ID, map[string]interface{}{ + "parser_config": cleaned, + }); err != nil { + return nil, common.CodeDataError, err + } + } + } + + if present["pipeline_id"] { + if req.PipelineID != nil && strings.TrimSpace(*req.PipelineID) != "" { + if err := s.resetDocumentForReparse(doc, kb.TenantID, nil, req.PipelineID); err != nil { + return nil, common.CodeDataError, err + } + } else { + // Explicitly cleared: drop the custom canvas so the worker falls + // back to the built-in template, matching validation. + empty := "" + if err := s.resetDocumentForReparse(doc, kb.TenantID, nil, &empty); err != nil { + return nil, common.CodeDataError, err + } + } + } else if present["parser_id"] && req.ParserID != nil && strings.TrimSpace(*req.ParserID) != "" { + parserID := strings.TrimSpace(*req.ParserID) + if err := s.resetDocumentForReparse(doc, kb.TenantID, &parserID, nil); err != nil { + return nil, common.CodeDataError, err + } + } + + if present["enabled"] && req.Enabled != nil { + if err := s.updateDocumentStatusOnly(doc, kb, *req.Enabled); err != nil { + return nil, common.CodeServerError, err + } + } + + updatedDoc, err := s.documentDAO.GetByID(doc.ID) + if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, fmt.Errorf("Can not get document by id:%s", doc.ID) + } + return nil, common.CodeDataError, errors.New("Database operation failed") + } + + metaFields := map[string]interface{}{} + if s.docEngine != nil && s.metadataSvc != nil { + metaFields, _ = s.GetDocumentMetadataByID(updatedDoc.ID) + } + + return s.toUpdateDatasetDocumentResponse(updatedDoc, metaFields), common.CodeSuccess, nil +} + +func (s *DocumentService) validateDatasetDocumentUpdate(datasetID, documentID, userID string, doc *entity.Document, req *UpdateDatasetDocumentRequest, present map[string]bool) (common.ErrorCode, error) { + if req == nil { + return common.CodeDataError, errors.New("Invalid request payload") + } + if present["chunk_count"] && req.ChunkCount != nil && *req.ChunkCount != 0 && *req.ChunkCount != doc.ChunkNum { + return common.CodeDataError, errors.New("Can't change `chunk_count`.") + } + if present["token_count"] && req.TokenCount != nil && *req.TokenCount != 0 && *req.TokenCount != doc.TokenNum { + return common.CodeDataError, errors.New("Can't change `token_count`.") + } + if present["progress"] && req.Progress != nil && *req.Progress != 0 && math.Abs(*req.Progress-doc.Progress) > 1e-9 { + return common.CodeDataError, errors.New("Can't change `progress`.") + } + + if present["enabled"] { + if req.Enabled == nil || (*req.Enabled != 0 && *req.Enabled != 1) { + return common.CodeDataError, errors.New("`enabled` value invalid, only accept 0 or 1") + } + } + + if present["parser_id"] && req.ParserID != nil { + parserID := strings.TrimSpace(*req.ParserID) + if (doc.Type == "visual" && parserID != "picture") || (isPresentationFile(doc.Name) && parserID != "presentation") { + return common.CodeDataError, errors.New("Not supported yet!") + } + } + if present["name"] && req.Name != nil { + if err := s.validateDocumentName(doc, *req.Name); err != nil { + return common.CodeDataError, err + } + } + + if present["meta_fields"] { + if err := validateMetaFields(req.MetaFields); err != nil { + return common.CodeDataError, err + } + } + + return common.CodeSuccess, nil +} + +func (s *DocumentService) validateDocumentName(doc *entity.Document, newName string) error { + if strings.TrimSpace(newName) == "" { + return errors.New("File name can't be empty.") + } + if len([]byte(newName)) > 255 { + return errors.New("File name must be 255 bytes or less.") + } + + oldName := "" + if doc.Name != nil { + oldName = *doc.Name + } + + if strings.ToLower(filepath.Ext(newName)) != strings.ToLower(filepath.Ext(oldName)) { + return errors.New("The extension of file can't be changed") + } + + docs, err := s.documentDAO.GetByNameAndKBID(newName, doc.KbID) + if err != nil { + return err + } + for _, d := range docs { + if d.ID != doc.ID && d.Name != nil && *d.Name == newName { + return errors.New("Duplicated document name in the same dataset.") + } + } + + return nil +} + +func isPresentationFile(name *string) bool { + if name == nil { + return false + } + ext := strings.ToLower(filepath.Ext(*name)) + return ext == ".ppt" || ext == ".pptx" || ext == ".pages" +} + +func validateMetaFields(meta map[string]any) error { + if meta == nil { + return nil + } + + for _, v := range meta { + switch typed := v.(type) { + case string, float64, int, int64, float32: + continue + case []any: + for _, item := range typed { + switch item.(type) { + case string, float64, int, int64, float32: + continue + default: + return fmt.Errorf("The type is not supported in list: %v", typed) + } + } + default: + return fmt.Errorf("The type is not supported: %v", v) + } + } + + return nil +} + +func (s *DocumentService) updateDocumentNameOnly(doc *entity.Document, tenantID, newName string) error { + if err := s.documentDAO.UpdateByID(doc.ID, map[string]interface{}{"name": newName}); err != nil { + return errors.New("Database error (Document rename)!") + } + + mappings, err := s.file2DocumentDAO.GetByDocumentID(doc.ID) + if err == nil && len(mappings) > 0 && mappings[0].FileID != nil && s.fileDAO != nil { + _ = s.fileDAO.UpdateByID(*mappings[0].FileID, map[string]interface{}{"name": newName}) + } + + if s.docEngine == nil { + return nil + } + + titleTks, _ := tokenizer.Tokenize(newName) + titleSmTks, _ := tokenizer.FineGrainedTokenize(titleTks) + indexName := fmt.Sprintf("ragflow_%s", tenantID) + return s.docEngine.UpdateChunks( + context.Background(), + map[string]interface{}{"doc_id": doc.ID}, + map[string]interface{}{ + "docnm_kwd": newName, + "title_tks": titleTks, + "title_sm_tks": titleSmTks, + }, + indexName, + doc.KbID, + ) +} + +func (s *DocumentService) updateDocumentParserConfig(documentID string, config map[string]any) error { + if len(config) == 0 { + return nil + } + + doc, err := s.documentDAO.GetByID(documentID) + if err != nil { + return fmt.Errorf("Document(%s) not found.", documentID) + } + + merged := common.DeepMergeMaps(map[string]interface{}(doc.ParserConfig), map[string]interface{}(config)) + if _, ok := config["raptor"]; !ok { + delete(merged, "raptor") + } + + return s.documentDAO.UpdateByID(documentID, map[string]interface{}{ + "parser_config": entity.JSONMap(merged), + }) +} + +func (s *DocumentService) toUpdateDatasetDocumentResponse(doc *entity.Document, metaFields map[string]interface{}) *UpdateDatasetDocumentResponse { + if metaFields == nil { + metaFields = map[string]interface{}{} + } + return &UpdateDatasetDocumentResponse{ + ID: doc.ID, + Thumbnail: doc.Thumbnail, + DatasetID: doc.KbID, + ParserID: doc.ParserID, + PipelineID: doc.PipelineID, + ParserConfig: map[string]interface{}(doc.ParserConfig), + SourceType: doc.SourceType, + Type: doc.Type, + CreatedBy: doc.CreatedBy, + Name: doc.Name, + Location: doc.Location, + Size: doc.Size, + TokenCount: doc.TokenNum, + ChunkCount: doc.ChunkNum, + Progress: doc.Progress, + ProgressMsg: doc.ProgressMsg, + ProcessBeginAt: doc.ProcessBeginAt, + ProcessDuration: doc.ProcessDuration, + ContentHash: doc.ContentHash, + MetaFields: metaFields, + Suffix: doc.Suffix, + Run: mapDocumentRunStatus(doc.Run), + Status: doc.Status, + CreateTime: doc.CreateTime, + CreateDate: doc.CreateDate, + UpdateTime: doc.UpdateTime, + UpdateDate: doc.UpdateDate, + } +} + +func mapDocumentRunStatus(run *string) string { + if run == nil { + return "UNSTART" + } + switch *run { + case string(entity.TaskStatusRunning): + return "RUNNING" + case string(entity.TaskStatusCancel): + return "CANCEL" + case string(entity.TaskStatusDone): + return "DONE" + case string(entity.TaskStatusFail): + return "FAIL" + default: + return "UNSTART" + } +} diff --git a/internal/service/document/document_ingest.go b/internal/service/document/document_ingest.go new file mode 100644 index 0000000000..92bed3ed02 --- /dev/null +++ b/internal/service/document/document_ingest.go @@ -0,0 +1,148 @@ +package document + +import ( + "context" + "fmt" + "ragflow/internal/service" + + "ragflow/internal/common" + "ragflow/internal/entity" +) + +func (s *DocumentService) ListIngestionTasks(userID string, datasetID *string, page, pageSize int) ([]*entity.IngestionTask, error) { + return s.ingestionTaskSvc.ListByUser(userID, datasetID, page, pageSize) +} + +func (s *DocumentService) IngestDocuments(datasetID, userID string, docIDs []string) ([]*service.ParseDocumentResponse, error) { + responses, err := s.ingestionTaskSvc.CreateForDocuments(datasetID, userID, docIDs) + if err != nil { + return nil, err + } + common.Info(fmt.Sprintf("parse documents, dataset: %s, documents: %v", datasetID, docIDs)) + return responses, nil +} + +func (s *DocumentService) StopIngestionTasks(tasks []string, userID string) ([]*entity.IngestionTask, error) { + return s.ingestionTaskSvc.RequestStopMany(tasks, &userID) +} + +func (s *DocumentService) RemoveIngestionTasks(tasks []string, userID string) ([]map[string]string, error) { + return s.ingestionTaskSvc.RemoveMany(tasks, &userID) +} + +func (s *DocumentService) Ingest(userID string, req *IngestDocumentRequest) (common.ErrorCode, error) { + run := fmt.Sprint(req.Run) + + docs, err := s.documentDAO.GetByIDs(req.DocIDs) + if err != nil { + return common.CodeExceptionError, fmt.Errorf("fail to get documents: %s", err.Error()) + } + + docsByID := make(map[string]*entity.Document, len(docs)) + for _, doc := range docs { + if doc != nil { + docsByID[doc.ID] = doc + } + } + + // First pass: validate every document exists and is accessible before + // mutating any state, so a single invalid doc rejects the whole request. + type validatedDoc struct { + doc *entity.Document + kb *entity.Knowledgebase + } + validated := make([]validatedDoc, 0, len(req.DocIDs)) + validatedIDs := make([]string, 0, len(req.DocIDs)) + for _, docID := range req.DocIDs { + doc := docsByID[docID] + if doc == nil { + return common.CodeDataError, fmt.Errorf("document not found") + } + kb, err := s.kbDAO.GetByID(doc.KbID) + if err != nil { + return common.CodeDataError, fmt.Errorf("dataset not found") + } + if !s.kbDAO.Accessible(kb.ID, userID) { + return common.CodeAuthenticationError, fmt.Errorf("no authorization") + } + validated = append(validated, validatedDoc{doc, kb}) + validatedIDs = append(validatedIDs, docID) + } + + // Batch pre-check for re-parse with delete: use the validated doc IDs + // so we don't silently skip non-existent or unauthorized documents. + if run == string(entity.TaskStatusRunning) && req.Delete { + if err := s.AssertIngestionTasksTerminal(validatedIDs); err != nil { + return common.CodeDataError, err + } + } + + for _, vd := range validated { + doc := vd.doc + kb := vd.kb + + // Start parsing: delegates to the shared start-parse flow. The + // document run status is set by service.IngestionTaskService.StartRunning + // when the task transitions from CREATED, not here. + if run == string(entity.TaskStatusRunning) { + if err := s.StartParseDocuments(doc, kb, userID, StartParseOptions{ + ApplyKB: req.ApplyKB, + RerunWithDelete: req.Delete, + }); err != nil { + common.Error(fmt.Sprintf("go side, doc %s, start parse", doc.ID), err) + return common.CodeExceptionError, err + } + continue + } + + // Cancel: RequestStop (STOPPING) and update doc state. Do NOT + // delete the ingestion task or chunks here — deletion races with + // the worker's async markStopped/settleToTerminal flow. Once the + // worker detects STOPPING and transitions to STOPPED, the task + // is terminal and can be safely cleaned up. + if run == string(entity.TaskStatusCancel) { + if err := s.CancelDocParse(doc); err != nil { + common.Error(fmt.Sprintf("go side, start to process %s, run is cancel", doc.ID), err) + return common.CodeDataError, err + } + if err := s.documentDAO.UpdateByID(doc.ID, map[string]interface{}{ + "run": string(entity.TaskStatusCancel), + "progress": 0, + }); err != nil { + common.Error(fmt.Sprintf("go side, doc %s, UpdateByID failed", doc.ID), err) + return common.CodeExceptionError, err + } + continue + } + + // Delete-only: user asked to remove prior parse results without + // starting a new parse. RUNNING already continued above. + if err := s.documentDAO.UpdateByID(doc.ID, map[string]interface{}{ + "run": run, + "progress": 0, + }); err != nil { + common.Error(fmt.Sprintf("go side, doc %s, UpdateByID failed", doc.ID), err) + return common.CodeExceptionError, err + } + + if req.Delete { + _, _ = s.taskDAO.DeleteIngestionTasksByDocIDs([]string{doc.ID}) + indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) + if s.docEngine != nil { + exists, err := s.docEngine.ChunkStoreExists(context.Background(), indexName, doc.KbID) + if err != nil { + common.Error(fmt.Sprintf("go side, doc %s, ChunkStoreExists failed", doc.ID), err) + return common.CodeExceptionError, err + } + if exists { + if _, err := s.docEngine.DeleteChunks(context.Background(), map[string]interface{}{"doc_id": doc.ID}, indexName, doc.KbID); err != nil { + common.Error(fmt.Sprintf("go side, doc %s, DeleteChunks failed", doc.ID), err) + return common.CodeExceptionError, err + } + } + } + } + } + + return common.CodeSuccess, nil +} diff --git a/internal/service/document/document_list.go b/internal/service/document/document_list.go new file mode 100644 index 0000000000..580095fefc --- /dev/null +++ b/internal/service/document/document_list.go @@ -0,0 +1,232 @@ +package document + +import ( + "fmt" + "strings" + "time" + + "ragflow/internal/dao" + "ragflow/internal/entity" +) + +// escapeSQLLikePattern escapes the SQL LIKE wildcards ('%', '_') and +// the escape character itself ('!') so a literal user-supplied +// filename can be safely interpolated into a `LIKE ? ESCAPE '!'` +// pattern. Without this, "%.png" would match any string ending in +// ".png" and "_" would match a single character — bypassing the +// filename-specific authorization check. PR review round 5, Major #8. +func escapeSQLLikePattern(s string) string { + r := strings.NewReplacer(`!`, `!!`, `%`, `!%`, `_`, `!_`) + return r.Replace(s) +} + +// ListDocuments list documents +func (s *DocumentService) ListDocuments(page, pageSize int) ([]*DocumentResponse, int64, error) { + offset := (page - 1) * pageSize + documents, total, err := s.documentDAO.List(offset, pageSize) + if err != nil { + return nil, 0, err + } + + responses := make([]*DocumentResponse, len(documents)) + for i, doc := range documents { + responses[i] = s.toResponse(doc) + } + + return responses, total, nil +} + +func (s *DocumentService) GetThumbnails(userID string, docIDs []string) (map[string]string, error) { + if len(docIDs) == 0 { + return map[string]string{}, nil + } + + tenantIDs := []string{userID} + if userID != "" { + ids, err := dao.NewUserTenantDAO().GetTenantIDsByUserID(userID) + if err != nil { + return nil, fmt.Errorf("failed to fetch user tenants: %w", err) + } + tenantIDs = append(tenantIDs, ids...) + } + + documents, err := s.documentDAO.GetByIDsAndTenantIDs(docIDs, tenantIDs) + if err != nil { + return nil, fmt.Errorf("failed to fetch document thumbnails: %w", err) + } + + result := make(map[string]string, len(documents)) + for _, document := range documents { + if document == nil { + continue + } + + thumbnail := "" + if document.Thumbnail != nil && *document.Thumbnail != "" { + if strings.HasPrefix(*document.Thumbnail, imgBase64Prefix) { + thumbnail = *document.Thumbnail + } else { + thumbnail = fmt.Sprintf( + "/api/v1/documents/images/%s-%s", + document.KbID, + *document.Thumbnail, + ) + } + } + + result[document.ID] = thumbnail + } + + return result, nil +} + +// ListDocumentsByDatasetID list documents by knowledge base ID +func (s *DocumentService) ListDocumentsByDatasetID(kbID, keywords string, page, pageSize int) ([]*entity.DocumentListItem, int64, error) { + return s.ListDocumentsByDatasetIDWithOptions(dao.DocumentListOptions{ + KbID: kbID, + Keywords: keywords, + OrderBy: "create_time", + Desc: true, + }, page, pageSize) +} + +// ListDocumentsByDatasetIDWithOptions lists documents by knowledge base ID with filters. +func (s *DocumentService) ListDocumentsByDatasetIDWithOptions(opts dao.DocumentListOptions, page, pageSize int) ([]*entity.DocumentListItem, int64, error) { + opts.Offset = (page - 1) * pageSize + opts.Limit = pageSize + if opts.OrderBy == "" { + opts.OrderBy = "create_time" + } + documents, total, err := s.documentDAO.ListByKBIDWithOptions(opts) + if err != nil { + return nil, 0, err + } + + responses := make([]*entity.DocumentListItem, len(documents)) + for i, doc := range documents { + responses[i] = doc + } + + return responses, total, nil +} + +// GetDocumentFiltersByDatasetID returns aggregate filter values for documents in a dataset. +func (s *DocumentService) GetDocumentFiltersByDatasetID(opts dao.DocumentListOptions) (map[string]interface{}, int64, error) { + filters, total, err := s.documentDAO.GetFilterByKBID(opts) + if err != nil { + return nil, 0, err + } + docIDs, err := s.documentDAO.ListIDsByKBIDWithOptions(opts) + if err != nil { + return nil, 0, err + } + metadataFilter, err := s.getDocumentMetadataFilter(opts.KbID, docIDs) + if err != nil { + return nil, 0, err + } + filters["metadata"] = metadataFilter + return filters, total, nil +} + +func (s *DocumentService) getDocumentMetadataFilter(kbID string, docIDs []string) (map[string]interface{}, error) { + metadataByKey, err := s.GetMetadataByKBs([]string{kbID}) + if err != nil { + return nil, err + } + candidateSet := make(map[string]bool, len(docIDs)) + for _, docID := range docIDs { + candidateSet[docID] = true + } + + metadataCounter := map[string]interface{}{} + docIDsWithMetadata := map[string]bool{} + for key, rawValues := range metadataByKey { + values, ok := rawValues.(map[string][]string) + if !ok { + continue + } + valueCounter := map[string]int64{} + for value, valueDocIDs := range values { + for _, docID := range valueDocIDs { + if !candidateSet[docID] { + continue + } + valueCounter[value]++ + docIDsWithMetadata[docID] = true + } + } + if len(valueCounter) > 0 { + metadataCounter[key] = valueCounter + } + } + metadataCounter["empty_metadata"] = map[string]int64{"true": int64(len(docIDs) - len(docIDsWithMetadata))} + return metadataCounter, nil +} + +// ListDocumentIDsByDatasetIDWithOptions lists matching document IDs without pagination. +func (s *DocumentService) ListDocumentIDsByDatasetIDWithOptions(opts dao.DocumentListOptions) ([]string, error) { + return s.documentDAO.ListIDsByKBIDWithOptions(opts) +} + +// GetDocumentsByAuthorID get documents by author ID +func (s *DocumentService) GetDocumentsByAuthorID(authorID, page, pageSize int) ([]*DocumentResponse, int64, error) { + offset := (page - 1) * pageSize + documents, total, err := s.documentDAO.GetByAuthorID(fmt.Sprintf("%d", authorID), offset, pageSize) + if err != nil { + return nil, 0, err + } + + responses := make([]*DocumentResponse, len(documents)) + for i, doc := range documents { + responses[i] = s.toResponse(doc) + } + + return responses, total, nil +} + +// toResponse convert model.Document to DocumentResponse +func (s *DocumentService) toResponse(doc *entity.Document) *DocumentResponse { + createdAt := "" + if doc.CreateTime != nil { + // Check if timestamp is in milliseconds (13 digits) or seconds (10 digits) + var ts int64 + if *doc.CreateTime > 1000000000000 { + // Milliseconds - convert to seconds + ts = *doc.CreateTime / 1000 + } else { + ts = *doc.CreateTime + } + createdAt = time.Unix(ts, 0).Format("2006-01-02 15:04:05") + } + updatedAt := "" + if doc.UpdateTime != nil { + // Accept both historical second-based values and current millisecond-based values. + ts := *doc.UpdateTime + if ts > 1000000000000 { + ts /= 1000 + } + updatedAt = time.Unix(ts, 0).Format("2006-01-02 15:04:05") + } + return &DocumentResponse{ + ID: doc.ID, + Name: doc.Name, + KbID: doc.KbID, + ParserID: doc.ParserID, + PipelineID: doc.PipelineID, + Type: doc.Type, + SourceType: doc.SourceType, + CreatedBy: doc.CreatedBy, + Location: doc.Location, + Size: doc.Size, + TokenNum: doc.TokenNum, + ChunkNum: doc.ChunkNum, + Progress: doc.Progress, + ProgressMsg: doc.ProgressMsg, + ProcessDuration: doc.ProcessDuration, + Suffix: doc.Suffix, + Run: doc.Run, + Status: doc.Status, + CreatedAt: createdAt, + UpdatedAt: updatedAt, + } +} diff --git a/internal/service/document/document_metadata.go b/internal/service/document/document_metadata.go new file mode 100644 index 0000000000..8c70d4e658 --- /dev/null +++ b/internal/service/document/document_metadata.go @@ -0,0 +1,947 @@ +package document + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "ragflow/internal/service" + "reflect" + "regexp" + "sort" + "strconv" + "strings" + + "ragflow/internal/common" + + "go.uber.org/zap" +) + +// GetMetadataSummary get metadata summary for documents +func (s *DocumentService) GetMetadataSummary(kbID string, docIDs []string) (map[string]interface{}, error) { + tenantID, err := s.metadataSvc.GetTenantIDByKBID(kbID) + if err != nil { + return nil, err + } + + searchResult, err := s.metadataSvc.SearchMetadata(kbID, tenantID, docIDs, 1000) + if err != nil { + return nil, err + } + + // Aggregate metadata from results + return aggregateMetadata(searchResult.MetadataRecords), nil +} + +// SetDocumentMetadata sets metadata for a document in the document engine +func (s *DocumentService) SetDocumentMetadata(docID string, meta map[string]interface{}) error { + // Get document to find kb_id + doc, err := s.documentDAO.GetByID(docID) + if err != nil { + return fmt.Errorf("document not found: %w", err) + } + + // Get tenant ID + tenantID, err := s.metadataSvc.GetTenantIDByKBID(doc.KbID) + if err != nil { + return fmt.Errorf("failed to get tenant ID: %w", err) + } + + if err := s.docEngine.UpdateMetadata(context.Background(), docID, doc.KbID, meta, tenantID); err != nil { + return fmt.Errorf("failed to update metadata: %w", err) + } + + return nil +} + +// DeleteDocumentMetadata deletes metadata keys for a document in the document engine +func (s *DocumentService) DeleteDocumentMetadata(docID string, keys []string) error { + // Get document to find kb_id + doc, err := s.documentDAO.GetByID(docID) + if err != nil { + return fmt.Errorf("document not found: %w", err) + } + + // Get tenant ID + tenantID, err := s.metadataSvc.GetTenantIDByKBID(doc.KbID) + if err != nil { + return fmt.Errorf("failed to get tenant ID: %w", err) + } + + // Delete metadata using the document engine + err = s.docEngine.DeleteMetadataKeys(nil, docID, doc.KbID, keys, tenantID) + if err != nil { + return fmt.Errorf("failed to delete metadata: %w", err) + } + + return nil +} + +// DeleteDocumentAllMetadata deletes all metadata for a document in the document engine +func (s *DocumentService) DeleteDocumentAllMetadata(docID string) error { + // Get document to find kb_id + doc, err := s.documentDAO.GetByID(docID) + if err != nil { + return fmt.Errorf("document not found: %w", err) + } + + // Get tenant ID + tenantID, err := s.metadataSvc.GetTenantIDByKBID(doc.KbID) + if err != nil { + return fmt.Errorf("failed to get tenant ID: %w", err) + } + + // Build condition to match the document + condition := map[string]interface{}{ + "id": docID, + "kb_id": doc.KbID, + } + + // Delete entire document metadata + _, err = s.docEngine.DeleteMetadata(nil, condition, tenantID) + if err != nil { + return fmt.Errorf("failed to delete document metadata: %w", err) + } + + return nil +} + +// GetDocumentMetadataByID get metadata for a specific document +func (s *DocumentService) GetDocumentMetadataByID(docID string) (map[string]interface{}, error) { + // Get document to find kb_id + doc, err := s.documentDAO.GetByID(docID) + if err != nil { + return nil, fmt.Errorf("document not found: %w", err) + } + + tenantID, err := s.metadataSvc.GetTenantIDByKBID(doc.KbID) + if err != nil { + return nil, err + } + + searchResult, err := s.metadataSvc.SearchMetadata(doc.KbID, tenantID, []string{docID}, 1) + if err != nil { + return nil, err + } + + // Return metadata if found + if len(searchResult.MetadataRecords) > 0 { + metadata := searchResult.MetadataRecords[0] + return service.ExtractMetaFields(metadata) + } + + return make(map[string]interface{}), nil +} + +// GetMetadataByKBs get metadata for knowledge bases +func (s *DocumentService) GetMetadataByKBs(kbIDs []string) (map[string]interface{}, error) { + if len(kbIDs) == 0 { + return make(map[string]interface{}), nil + } + + searchResult, err := s.metadataSvc.SearchMetadataByKBs(kbIDs, 10000) + if err != nil { + return nil, err + } + + flattenedMeta := make(map[string]map[string][]string) + numMetadata := len(searchResult.MetadataRecords) + + var allMetaFields []map[string]interface{} + if numMetadata > 1 && len(searchResult.MetadataRecords) > 0 { + firstMetadata := searchResult.MetadataRecords[0] + if metaFieldsVal := firstMetadata["meta_fields"]; metaFieldsVal != nil { + if v, ok := metaFieldsVal.([]byte); ok { + allMetaFields = service.ParseAllLengthPrefixedJSON(v) + } + } + } + + for idx, metadata := range searchResult.MetadataRecords { + docID, ok := service.ExtractDocumentID(metadata) + if !ok { + continue + } + + var metaFields map[string]interface{} + var metaFieldsVal interface{} + + if len(allMetaFields) > 0 && idx < len(allMetaFields) { + // Use pre-parsed meta_fields from concatenated data + metaFields = allMetaFields[idx] + } else { + // Normal case - get from chunk + metaFieldsVal = metadata["meta_fields"] + if metaFieldsVal != nil { + switch v := metaFieldsVal.(type) { + case string: + if err := json.Unmarshal([]byte(v), &metaFields); err != nil { + continue + } + case []byte: + // Try direct JSON parse first + if err := json.Unmarshal(v, &metaFields); err != nil { + // Try to parse as concatenated JSON objects + metaFields = service.ParseLengthPrefixedJSON(v) + } + case map[string]interface{}: + metaFields = v + default: + continue + } + } + } + + if metaFields == nil { + continue + } + + // Process each metadata field + for fieldName, fieldValue := range metaFields { + if fieldName == "kb_id" || fieldName == "id" { + continue + } + + if _, ok := flattenedMeta[fieldName]; !ok { + flattenedMeta[fieldName] = make(map[string][]string) + } + + // Handle list and single values + var values []interface{} + switch v := fieldValue.(type) { + case []interface{}: + values = v + default: + values = []interface{}{v} + } + + for _, val := range values { + if val == nil { + continue + } + strVal := fmt.Sprintf("%v", val) + flattenedMeta[fieldName][strVal] = append(flattenedMeta[fieldName][strVal], docID) + } + } + } + + // Convert to map[string]interface{} for return + var metaResult map[string]interface{} = make(map[string]interface{}) + for k, v := range flattenedMeta { + metaResult[k] = v + } + + return metaResult, nil +} + +// aggregateMetadata aggregates metadata from search results +func aggregateMetadata(chunks []map[string]interface{}) map[string]interface{} { + // summary: map[fieldName]map[value]valueInfo + summary := make(map[string]map[string]valueInfo) + typeCounter := make(map[string]map[string]int) + orderCounter := 0 + + for _, chunk := range chunks { + // For metadata table, the actual metadata is in the "meta_fields" JSON field + // Extract it first + metaFieldsVal := chunk["meta_fields"] + if metaFieldsVal == nil { + continue + } + + // Parse meta_fields - could be a string (JSON) or a map + var metaFields map[string]interface{} + switch v := metaFieldsVal.(type) { + case string: + // Parse JSON string + if err := json.Unmarshal([]byte(v), &metaFields); err != nil { + continue + } + case []byte: + // Handle byte slice - Infinity returns concatenated JSON objects with length prefixes + rawBytes := v + + // Try to detect and handle length-prefixed format + // Format: [4-byte length][JSON][4-byte length][JSON]... + parsedMetaFields := make(map[string]interface{}) + offset := 0 + for offset < len(rawBytes) { + // Need at least 4 bytes for length prefix + if offset+4 > len(rawBytes) { + break + } + + // Read 4-byte length (little-endian, not big-endian!) + length := uint32(rawBytes[offset]) | uint32(rawBytes[offset+1])<<8 | + uint32(rawBytes[offset+2])<<16 | uint32(rawBytes[offset+3])<<24 + + // Check if length looks valid (not too large) + if length > 10000 || length == 0 { + // Try to find next '{' from current position + nextBrace := -1 + for i := offset; i < len(rawBytes) && i < offset+100; i++ { + if rawBytes[i] == '{' { + nextBrace = i + break + } + } + if nextBrace > offset { + // Skip to the next '{' + offset = nextBrace + continue + } + break + } + + // Extract JSON data + jsonStart := offset + 4 + jsonEnd := jsonStart + int(length) + if jsonEnd > len(rawBytes) { + jsonEnd = len(rawBytes) + } + + jsonBytes := rawBytes[jsonStart:jsonEnd] + + // Try to parse this JSON + var singleMeta map[string]interface{} + if err := json.Unmarshal(jsonBytes, &singleMeta); err == nil { + // Merge metadata from this document + for k, vv := range singleMeta { + if existing, ok := parsedMetaFields[k]; ok { + // Combine values + if existList, ok := existing.([]interface{}); ok { + if newList, ok := vv.([]interface{}); ok { + parsedMetaFields[k] = append(existList, newList...) + } else { + parsedMetaFields[k] = append(existList, vv) + } + } else { + parsedMetaFields[k] = []interface{}{existing, vv} + } + } else { + parsedMetaFields[k] = vv + } + } + } + + offset = jsonEnd + } + + // If we successfully parsed multiple JSON objects, use the merged result + if len(parsedMetaFields) > 0 { + metaFields = parsedMetaFields + } else { + // Fallback: try the original parsing method + startIdx := -1 + for i, b := range rawBytes { + if b == '{' { + startIdx = i + break + } + } + if startIdx > 0 { + strVal := string(rawBytes[startIdx:]) + if err := json.Unmarshal([]byte(strVal), &metaFields); err != nil { + metaFields = map[string]interface{}{"raw": strVal} + } + } else if err := json.Unmarshal(rawBytes, &metaFields); err != nil { + metaFields = map[string]interface{}{"raw": string(rawBytes)} + } + } + case map[string]interface{}: + metaFields = v + default: + continue + } + + // Now iterate over the extracted metadata fields + for k, v := range metaFields { + // Skip nil values + if v == nil { + continue + } + + // Determine value type + valueType := getMetaValueType(v) + + // Track type counts + if valueType != "" { + if _, ok := typeCounter[k]; !ok { + typeCounter[k] = make(map[string]int) + } + typeCounter[k][valueType] = typeCounter[k][valueType] + 1 + } + + // Aggregate value counts. Flatten nested arrays so malformed values do + // not surface in the UI as the literal string "[]". + values := flattenMetadataSummaryValues(v) + for _, vv := range values { + if vv == nil { + continue + } + sv := fmt.Sprintf("%v", vv) + + if _, ok := summary[k]; !ok { + summary[k] = make(map[string]valueInfo) + } + + if existing, ok := summary[k][sv]; ok { + // Already exists, just increment count + existing.count++ + summary[k][sv] = existing + } else { + // First time seeing this value - record order + summary[k][sv] = valueInfo{count: 1, firstOrder: orderCounter} + orderCounter++ + } + } + } + } + + // Build result with type information and sorted values + result := make(map[string]interface{}) + for k, v := range summary { + // Sort by count descending, then by firstOrder ascending (to match Python stable sort) + // values: [value, count, firstOrder] + values := make([][3]interface{}, 0, len(v)) + for val, info := range v { + values = append(values, [3]interface{}{val, info.count, info.firstOrder}) + } + // Use stable sort - sort by count descending, then by firstOrder + sort.SliceStable(values, func(i, j int) bool { + cntI := values[i][1].(int) + cntJ := values[j][1].(int) + if cntI != cntJ { + return cntI > cntJ // count descending + } + // If counts equal, use firstOrder ascending (earlier appearance first) + return values[i][2].(int) < values[j][2].(int) + }) + + // Determine dominant type + valueType := "string" + if typeCounts, ok := typeCounter[k]; ok { + maxCount := 0 + for t, c := range typeCounts { + if c > maxCount { + maxCount = c + valueType = t + } + } + } + + // Convert from [value, count, firstOrder] to [value, count] for output + outputValues := make([][2]interface{}, len(values)) + for i, val := range values { + outputValues[i] = [2]interface{}{val[0], val[1]} + } + + result[k] = map[string]interface{}{ + "type": valueType, + "values": outputValues, + } + } + + return result +} + +// getMetaValueType determines the type of a metadata value +func getMetaValueType(value interface{}) string { + if value == nil { + return "" + } + + switch v := value.(type) { + case []interface{}: + if len(v) > 0 { + return "list" + } + return "" + case bool: + return "string" + case int, int8, int16, int32, int64: + return "number" + case float32, float64: + return "number" + case string: + if isTimeString(v) { + return "time" + } + return "string" + } + return "string" +} + +func flattenMetadataSummaryValues(value interface{}) []interface{} { + switch typed := value.(type) { + case []interface{}: + result := make([]interface{}, 0, len(typed)) + for _, item := range typed { + result = append(result, flattenMetadataSummaryValues(item)...) + } + return result + case []string: + result := make([]interface{}, 0, len(typed)) + for _, item := range typed { + result = append(result, item) + } + return result + case nil: + return nil + default: + return []interface{}{typed} + } +} + +// isTimeString checks if a string is an ISO 8601 datetime +func isTimeString(s string) bool { + matched, _ := regexp.MatchString(`^\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}$`, s) + return matched +} + +func (s *DocumentService) replaceDocumentMetadata(docID string, meta map[string]any) error { + if s.docEngine == nil || s.metadataSvc == nil { + return nil + } + if err := s.DeleteDocumentAllMetadata(docID); err != nil { + return err + } + return s.SetDocumentMetadata(docID, map[string]interface{}(meta)) +} + +func (s *DocumentService) patchDocumentMetadata(docID string, before, after map[string]interface{}) error { + if s.docEngine == nil || s.metadataSvc == nil { + return nil + } + + deleteKeys := make([]string, 0) + for key := range before { + if _, ok := after[key]; !ok { + deleteKeys = append(deleteKeys, key) + } + } + if len(deleteKeys) > 0 { + if err := s.DeleteDocumentMetadata(docID, deleteKeys); err != nil { + return err + } + } + + updateFields := make(map[string]interface{}) + for key, value := range after { + if !reflect.DeepEqual(before[key], value) { + updateFields[key] = value + } + } + if len(updateFields) == 0 { + return nil + } + return s.SetDocumentMetadata(docID, updateFields) +} + +// BatchUpdateDocumentMetadatas implements the shared logic for +// PATCH /datasets/:dataset_id/documents/metadatas and +// POST /datasets/:dataset_id/metadata/update. +func (s *DocumentService) BatchUpdateDocumentMetadatas( + datasetID string, + selector *DocumentMetadataSelector, + updates []DocumentMetadataUpdate, + deletes []DocumentMetadataDelete, +) (*BatchUpdateDocumentMetadatasResponse, common.ErrorCode, error) { + if selector == nil { + selector = &DocumentMetadataSelector{} + } + if code, err := validateBatchUpdateDocumentMetadatasRequest(selector, updates, deletes); err != nil { + return nil, code, err + } + + // Resolve which document IDs to target. + targetDocIDs := make(map[string]struct{}) + + if len(selector.DocumentIDs) > 0 { + // Validate that supplied IDs actually belong to this dataset. + allRows, err := s.documentDAO.GetAllDocIDsByKBIDs([]string{datasetID}) + if err != nil { + return nil, common.CodeServerError, fmt.Errorf("failed to list dataset documents: %w", err) + } + kbDocIDSet := make(map[string]struct{}, len(allRows)) + for _, row := range allRows { + kbDocIDSet[row["id"]] = struct{}{} + } + var invalidIDs []string + for _, id := range selector.DocumentIDs { + if _, ok := kbDocIDSet[id]; !ok { + invalidIDs = append(invalidIDs, id) + } + } + if len(invalidIDs) > 0 { + return nil, common.CodeDataError, fmt.Errorf("these documents do not belong to dataset %s: %s", + datasetID, strings.Join(invalidIDs, ", ")) + } + for _, id := range selector.DocumentIDs { + targetDocIDs[id] = struct{}{} + } + } + + // Apply metadata_condition filter. + if len(selector.MetadataCondition) > 0 { + flattedMeta, err := s.metadataSvc.GetFlattedMetaByKBs([]string{datasetID}) + if err != nil { + return nil, common.CodeServerError, fmt.Errorf("failed to get flattened metadata: %w", err) + } + + // ParseAndConvert mirrors Python convert_conditions: conditions arrive as + // {name, comparison_operator, value}, the operator is normalised, and the + // (possibly non-string) value is preserved. MetaFilter then matches against + // the common.MetaData returned by GetFlattedMetaByKBs. + filterInput := common.ParseAndConvert(selector.MetadataCondition) + filteredIDs := common.MetaFilter(flattedMeta, filterInput) + + filteredSet := make(map[string]struct{}, len(filteredIDs)) + for _, id := range filteredIDs { + filteredSet[id] = struct{}{} + } + + if len(targetDocIDs) > 0 { + // Intersect with the document_ids restriction. + for id := range targetDocIDs { + if _, ok := filteredSet[id]; !ok { + delete(targetDocIDs, id) + } + } + } else { + targetDocIDs = filteredSet + } + + // 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 + } + } + + ids := make([]string, 0, len(targetDocIDs)) + for id := range targetDocIDs { + ids = append(ids, id) + } + + // Apply updates and deletes per document using Python's batch_update_metadata + // semantics instead of a simple merge-then-delete. + updated := 0 + for _, docID := range ids { + currentMeta, err := s.GetDocumentMetadataByID(docID) + if err != nil { + common.Warn("BatchUpdateDocumentMetadata: get metadata failed", + zap.String("docID", docID), zap.Error(err)) + continue + } + + meta := cloneDocumentMetadata(currentMeta) + originalMeta := cloneDocumentMetadata(meta) + + changed := applyDocumentMetadataUpdates(meta, updates) + if applyDocumentMetadataDeletes(meta, deletes) { + changed = true + } + + if !changed || reflect.DeepEqual(originalMeta, meta) { + continue + } + + if err := s.patchDocumentMetadata(docID, originalMeta, meta); err != nil { + common.Warn("BatchUpdateDocumentMetadata: patch metadata failed", + zap.String("docID", docID), zap.Error(err)) + continue + } + updated++ + } + + return &BatchUpdateDocumentMetadatasResponse{Updated: updated, MatchedDocs: len(ids)}, common.CodeSuccess, nil +} + +func validateBatchUpdateDocumentMetadatasRequest( + selector *DocumentMetadataSelector, + updates []DocumentMetadataUpdate, + deletes []DocumentMetadataDelete, +) (common.ErrorCode, error) { + for _, upd := range updates { + if strings.TrimSpace(upd.Key) == "" || upd.Value == nil { + return common.CodeDataError, errors.New("Each update requires key and value.") + } + } + for _, del := range deletes { + if strings.TrimSpace(del.Key) == "" { + return common.CodeDataError, errors.New("Each delete requires key.") + } + } + if selector != nil && selector.MetadataCondition != nil { + if _, ok := selector.MetadataCondition["conditions"]; !ok && len(selector.MetadataCondition) > 0 { + return common.CodeDataError, errors.New("metadata_condition must be an object.") + } + } + return common.CodeSuccess, nil +} + +func cloneDocumentMetadata(meta map[string]interface{}) map[string]interface{} { + if meta == nil { + return map[string]interface{}{} + } + cloned := make(map[string]interface{}, len(meta)) + for k, v := range meta { + cloned[k] = cloneDocumentMetadataValue(v) + } + return cloned +} + +func cloneDocumentMetadataValue(v interface{}) interface{} { + switch typed := v.(type) { + case []interface{}: + cp := make([]interface{}, len(typed)) + copy(cp, typed) + return cp + case []string: + cp := make([]interface{}, 0, len(typed)) + for _, item := range typed { + cp = append(cp, item) + } + return cp + default: + return typed + } +} + +func applyDocumentMetadataUpdates(meta map[string]interface{}, updates []DocumentMetadataUpdate) bool { + changed := false + for _, upd := range updates { + key := strings.TrimSpace(upd.Key) + if key == "" { + continue + } + normalizedValue := normalizeDocumentMetadataUpdateValue(upd.Value, upd.ValueType) + matchProvided := upd.Match != nil && !(fmt.Sprintf("%v", upd.Match) == "") + current, exists := meta[key] + if !exists { + if matchProvided { + continue + } + if listVal, ok := toMetadataInterfaceSlice(normalizedValue); ok { + meta[key] = dedupeDocumentMetadataList(listVal) + } else { + meta[key] = normalizedValue + } + changed = true + continue + } + + if curList, ok := toMetadataInterfaceSlice(current); ok { + if !matchProvided { + newList := append([]interface{}{}, curList...) + if appendList, ok := toMetadataInterfaceSlice(normalizedValue); ok { + newList = append(newList, appendList...) + } else { + newList = append(newList, normalizedValue) + } + newList = dedupeDocumentMetadataList(newList) + if !reflect.DeepEqual(curList, newList) { + meta[key] = newList + changed = true + } + continue + } + + replaced := false + newList := make([]interface{}, 0, len(curList)) + for _, item := range curList { + if documentMetadataValuesEqual(item, upd.Match) { + if replacementList, ok := toMetadataInterfaceSlice(normalizedValue); ok { + newList = append(newList, replacementList...) + } else { + newList = append(newList, normalizedValue) + } + replaced = true + } else { + newList = append(newList, item) + } + } + newList = dedupeDocumentMetadataList(newList) + if replaced && !reflect.DeepEqual(curList, newList) { + meta[key] = newList + changed = true + } + continue + } + + if !matchProvided { + if !reflect.DeepEqual(current, normalizedValue) { + meta[key] = normalizedValue + changed = true + } + continue + } + if documentMetadataValuesEqual(current, upd.Match) && !reflect.DeepEqual(current, normalizedValue) { + meta[key] = normalizedValue + changed = true + } + } + return changed +} + +func applyDocumentMetadataDeletes(meta map[string]interface{}, deletes []DocumentMetadataDelete) bool { + changed := false + for _, del := range deletes { + key := strings.TrimSpace(del.Key) + current, exists := meta[key] + if key == "" || !exists { + continue + } + + if curList, ok := toMetadataInterfaceSlice(current); ok { + if del.Value == nil { + delete(meta, key) + changed = true + continue + } + newList := make([]interface{}, 0, len(curList)) + for _, item := range curList { + if !documentMetadataValuesEqual(item, del.Value) { + newList = append(newList, item) + } + } + if len(newList) != len(curList) { + if len(newList) == 0 { + delete(meta, key) + } else { + meta[key] = newList + } + changed = true + } + continue + } + + if del.Value == nil || documentMetadataValuesEqual(current, del.Value) { + delete(meta, key) + changed = true + } + } + return changed +} + +func toMetadataInterfaceSlice(v interface{}) ([]interface{}, bool) { + switch typed := v.(type) { + case []interface{}: + cp := make([]interface{}, len(typed)) + copy(cp, typed) + return cp, true + case []string: + cp := make([]interface{}, 0, len(typed)) + for _, item := range typed { + cp = append(cp, item) + } + return cp, true + default: + return nil, false + } +} + +func dedupeDocumentMetadataList(items []interface{}) []interface{} { + result := make([]interface{}, 0, len(items)) + seen := make(map[string]struct{}, len(items)) + for _, item := range items { + key := fmt.Sprintf("%T:%v", item, item) + if _, ok := seen[key]; ok { + continue + } + seen[key] = struct{}{} + result = append(result, item) + } + return result +} + +func documentMetadataValuesEqual(a, b interface{}) bool { + return fmt.Sprintf("%v", a) == fmt.Sprintf("%v", b) +} + +func normalizeDocumentMetadataUpdateValue(value interface{}, valueType string) interface{} { + switch strings.ToLower(strings.TrimSpace(valueType)) { + case "list": + if list, ok := normalizeMetadataListValue(value); ok { + return list + } + return []interface{}{} + case "number": + scalar, ok := firstScalarMetadataValue(value) + if !ok { + return value + } + switch typed := scalar.(type) { + case float64, float32, int, int8, int16, int32, int64: + return typed + case json.Number: + if i, err := typed.Int64(); err == nil { + return i + } + if f, err := typed.Float64(); err == nil { + return f + } + case string: + trimmed := strings.TrimSpace(typed) + if trimmed == "" { + return "" + } + if i, err := strconv.ParseInt(trimmed, 10, 64); err == nil { + return i + } + if f, err := strconv.ParseFloat(trimmed, 64); err == nil { + return f + } + return trimmed + } + return scalar + case "string", "time": + if scalar, ok := firstScalarMetadataValue(value); ok { + return fmt.Sprintf("%v", scalar) + } + return "" + default: + return value + } +} + +func normalizeMetadataListValue(value interface{}) ([]interface{}, bool) { + switch typed := value.(type) { + case []interface{}: + result := make([]interface{}, 0, len(typed)) + for _, item := range typed { + if nested, ok := normalizeMetadataListValue(item); ok { + result = append(result, nested...) + continue + } + if item != nil { + result = append(result, item) + } + } + return result, true + case []string: + result := make([]interface{}, 0, len(typed)) + for _, item := range typed { + result = append(result, item) + } + return result, true + default: + return nil, false + } +} + +func firstScalarMetadataValue(value interface{}) (interface{}, bool) { + if list, ok := normalizeMetadataListValue(value); ok { + for _, item := range list { + if item != nil { + return item, true + } + } + return nil, false + } + if value == nil { + return nil, false + } + return value, true +} diff --git a/internal/service/document/document_metadata_test.go b/internal/service/document/document_metadata_test.go new file mode 100644 index 0000000000..49013e6abe --- /dev/null +++ b/internal/service/document/document_metadata_test.go @@ -0,0 +1,107 @@ +package document + +import "testing" + +func TestCloneDocumentMetadata(t *testing.T) { + if got := cloneDocumentMetadata(nil); got == nil { + t.Fatal("nil input should return a non-nil empty map") + } + + orig := map[string]interface{}{ + "list": []interface{}{"a", "b"}, + "str": "x", + } + clone := cloneDocumentMetadata(orig) + + // Mutating the clone must not affect the original. + clone["str"] = "changed" + if orig["str"] != "x" { + t.Error("clone mutation leaked into original scalar") + } + cloneList := clone["list"].([]interface{}) + cloneList[0] = "mutated" + if origList, ok := orig["list"].([]interface{}); ok && origList[0] == "mutated" { + t.Error("nested slice was not deep-copied") + } +} + +func TestDocumentMetadataValuesEqual(t *testing.T) { + if !documentMetadataValuesEqual(1, "1") { + t.Error("1 and \"1\" should compare equal (formatted form)") + } + if !documentMetadataValuesEqual(int64(1), "1") { + t.Error("int64(1) and \"1\" should compare equal") + } + if documentMetadataValuesEqual("a", "b") { + t.Error("distinct values should differ") + } +} + +func TestNormalizeMetadataListValue(t *testing.T) { + if _, ok := normalizeMetadataListValue("scalar"); ok { + t.Error("a scalar should not be reported as a list") + } + list, ok := normalizeMetadataListValue([]interface{}{"a", "b"}) + if !ok || len(list) != 2 { + t.Errorf("[]interface{} should normalize: %v (ok=%v)", list, ok) + } + list2, ok2 := normalizeMetadataListValue([]string{"a", "b"}) + if !ok2 || len(list2) != 2 { + t.Errorf("[]string should normalize: %v (ok=%v)", list2, ok2) + } +} + +func TestFirstScalarMetadataValue(t *testing.T) { + if v, ok := firstScalarMetadataValue([]interface{}{"a", "b"}); !ok || v != "a" { + t.Errorf("should return first non-nil scalar: %v (ok=%v)", v, ok) + } + if _, ok := firstScalarMetadataValue(nil); ok { + t.Error("nil should report not-found") + } +} + +func TestNormalizeDocumentMetadataUpdateValue(t *testing.T) { + if v := normalizeDocumentMetadataUpdateValue("42", "number"); v != int64(42) { + t.Errorf("number string → int64(42), got %v (%T)", v, v) + } + list, ok := normalizeDocumentMetadataUpdateValue([]interface{}{"a"}, "list").([]interface{}) + if !ok || len(list) != 1 { + t.Errorf("list normalize failed: %v (ok=%v)", list, ok) + } + if v := normalizeDocumentMetadataUpdateValue(123, "string"); v != "123" { + t.Errorf("string normalize: got %v", v) + } + if v := normalizeDocumentMetadataUpdateValue("bar", "unknown"); v != "bar" { + t.Errorf("unknown valueType should pass through: got %v", v) + } +} + +func TestAggregateMetadata(t *testing.T) { + chunks := []map[string]interface{}{ + {"meta_fields": map[string]interface{}{"author": "alice"}}, + {"meta_fields": map[string]interface{}{"author": "bob"}}, + {"meta_fields": map[string]interface{}{"author": "alice"}}, + } + result := aggregateMetadata(chunks) + + field, ok := result["author"].(map[string]interface{}) + if !ok { + t.Fatalf("author field missing from summary: %v", result) + } + values, ok := field["values"].([][2]interface{}) + if !ok { + t.Fatalf("values has unexpected shape: %v", field["values"]) + } + counts := map[string]int{} + for _, pair := range values { + if s, ok := pair[0].(string); ok { + counts[s] = pair[1].(int) + } + } + if counts["alice"] != 2 { + t.Errorf("alice count = %d, want 2", counts["alice"]) + } + if counts["bob"] != 1 { + t.Errorf("bob count = %d, want 1", counts["bob"]) + } +} diff --git a/internal/service/document/document_parse.go b/internal/service/document/document_parse.go new file mode 100644 index 0000000000..9168b7ff73 --- /dev/null +++ b/internal/service/document/document_parse.go @@ -0,0 +1,471 @@ +package document + +import ( + "context" + "errors" + "fmt" + "ragflow/internal/service" + "strconv" + "strings" + + "ragflow/internal/common" + "ragflow/internal/dao" + enginetypes "ragflow/internal/engine/types" + "ragflow/internal/entity" + "ragflow/internal/storage" + + "gorm.io/gorm" + "gorm.io/gorm/clause" +) + +// StartParseDocuments starts parsing a document via the DSL ingestion +// pipeline. It optionally clears prior results (RerunWithDelete), applies +// KB config (ApplyKB), validates storage, and enqueues an ingestion task. +// The document run status is NOT set here; service.IngestionTaskService.StartRunning +// sets it to RUNNING when the worker picks up the task and transitions it from +// CREATED. Extracted from Ingest so other entry points (e.g. ChunkService.Parse) +// can reuse the same start-parse flow. +func (s *DocumentService) StartParseDocuments(doc *entity.Document, kb *entity.Knowledgebase, userID string, opts StartParseOptions) error { + // Validate storage first so we don't clear prior results and then fail + // because the document can't be read, leaving the document with neither + // old nor new parse results. + if _, _, err := s.GetDocumentStorageAddress(doc); err != nil { + return err + } + + if opts.RerunWithDelete { + if err := s.clearDocumentParseResults(doc, kb.TenantID); err != nil { + return err + } + } + + if _, err := s.IngestDocuments(doc.KbID, userID, []string{doc.ID}); err != nil { + return err + } + return nil +} + +// AssertIngestionTasksTerminal verifies none of the documents has an +// in-flight (RUNNING/STOPPING) ingestion task. Used as a batch pre-check +// before re-parsing so a single non-terminal doc rejects the whole request +// up front instead of partially cleaning some docs then failing. +func (s *DocumentService) AssertIngestionTasksTerminal(docIDs []string) error { + for _, docID := range docIDs { + task, err := s.ingestionTaskDAO.GetByDocumentID(docID) + if err != nil { + return fmt.Errorf("check ingestion task for %s: %w", docID, err) + } + if task == nil { + continue + } + if task.Status == common.RUNNING || task.Status == common.STOPPING { + return fmt.Errorf("document %s ingestion task is %s; stop it and wait for a terminal state before re-parsing", docID, task.Status) + } + } + return nil +} + +func (s *DocumentService) clearDocumentParseResults(doc *entity.Document, tenantID string) error { + if doc == nil { + return fmt.Errorf("document is nil") + } + + // Refuse to clear a non-terminal ingestion task. An in-flight worker + // (RUNNING) or one mid-stop (STOPPING) would keep writing chunks and + // corrupt the new run's results. The caller must stop the task first + // and wait for a terminal state (COMPLETED/STOPPED/FAILED) or CREATED. + if task, _ := s.ingestionTaskDAO.GetByDocumentID(doc.ID); task != nil { + if task.Status == common.RUNNING || task.Status == common.STOPPING { + return fmt.Errorf("document %s ingestion task is %s; stop it and wait for a terminal state before re-parsing", doc.ID, task.Status) + } + } + + // Delete terminal and CREATED ingestion tasks atomically, leaving + // RUNNING/STOPPING tasks untouched so the check-then-delete window + // between GetByDocumentID and the delete above cannot delete a task + // that just transitioned to RUNNING. + if _, err := s.ingestionTaskDAO.DeleteIfTerminal(doc.ID); err != nil { + return err + } + + if err := s.clearDocumentAndKBCountersForRerun(doc.ID, doc.KbID); err != nil { + return err + } + + if s.docEngine == nil { + return nil + } + + indexName := fmt.Sprintf("ragflow_%s", tenantID) + exists, err := s.docEngine.ChunkStoreExists(context.Background(), 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 { + return err + } + return nil +} + +func (s *DocumentService) clearDocumentAndKBCountersForRerun(docID, kbID string) error { + return dao.DB.Transaction(func(tx *gorm.DB) error { + var current entity.Document + if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). + Where("id = ? AND kb_id = ?", docID, kbID). + First(¤t).Error; err != nil { + return err + } + + if current.TokenNum == 0 && current.ChunkNum == 0 && current.ProcessDuration == 0 { + return nil + } + + result := tx.Model(&entity.Document{}). + Where("id = ? AND kb_id = ?", docID, kbID). + Updates(map[string]interface{}{ + "token_num": 0, + "chunk_num": 0, + "process_duration": 0, + }) + if result.Error != nil { + return result.Error + } + if current.TokenNum == 0 && current.ChunkNum == 0 { + return nil + } + + result = tx.Model(&entity.Knowledgebase{}). + Where("id = ?", kbID). + Updates(map[string]interface{}{ + "token_num": gorm.Expr("token_num - ?", current.TokenNum), + "chunk_num": gorm.Expr("chunk_num - ?", current.ChunkNum), + }) + if result.Error != nil { + return result.Error + } + if result.RowsAffected == 0 { + return fmt.Errorf("knowledgebase not found") + } + return nil + }) +} + +func (s *DocumentService) countDoneDocuments(datasetID string) (int64, error) { + var count int64 + err := dao.GetDB().Model(&entity.Document{}). + Where("kb_id = ? AND run = ?", datasetID, string(entity.TaskStatusDone)). + Count(&count).Error + return count, err +} + +func (s *DocumentService) clearKBChunkNumWhenRerun(doc *entity.Document) error { + if doc == nil { + return fmt.Errorf("document is nil") + } + return dao.GetDB().Model(&entity.Knowledgebase{}).Where("id = ?", doc.KbID).Updates(map[string]interface{}{ + "token_num": gorm.Expr("token_num - ?", doc.TokenNum), + "chunk_num": gorm.Expr("chunk_num - ?", doc.ChunkNum), + }).Error +} + +func (s *DocumentService) ParseDocuments(datasetID, userID string, docIDs []string) ([]*service.ParseDocumentResponse, error) { + // create document parse id + // save to task table + // send to message queue + + // deduplicate the document id + uniqueDocIDs := common.Deduplicate(docIDs) + if uniqueDocIDs == nil || len(uniqueDocIDs) == 0 { + return nil, fmt.Errorf("no documents to parse") + } + + var responses []*service.ParseDocumentResponse + + // query database, if the document ids are valid + for _, docID := range uniqueDocIDs { + doc, err := s.documentDAO.GetByID(docID) + if err != nil { + errorMessage := err.Error() + responses = append(responses, &service.ParseDocumentResponse{ + DocumentID: docID, + Result: errorMessage, + }) + continue + } + if doc == nil { + errorMessage := "no such document" + responses = append(responses, &service.ParseDocumentResponse{ + DocumentID: docID, + Result: errorMessage, + }) + continue + } + + if doc.Status != nil && *doc.Status != "0" { + errorMessage := fmt.Sprintf("document %s is already parsed", docID) + responses = append(responses, &service.ParseDocumentResponse{ + DocumentID: docID, + Result: errorMessage, + }) + continue + } + + // create task for each document + //task := &entity.IngestionTask{ + // ID: utility.GenerateToken(), + // DocumentID: docID, + // UserID: userID, + //} + + // save the task to database + //err = s.ingestionTaskDAO.Create(task) + //if err != nil { + // errorMessage := err.Error() + // responses = append(responses, &service.ParseDocumentResponse{ + // DocumentID: docID, + // Result: &errorMessage, + // }) + // continue + //} + + // Send task to message queue + + } + + common.Info(fmt.Sprintf("parse documents, dataset: %s, documents: %v", datasetID, docIDs)) + return responses, nil +} + +// StopParseDocuments stops parsing for the given documents in a dataset. +// It sets Redis cancel signals for associated tasks and updates doc.run to CANCEL. +// Returns a map with success_count and optionally errors. +func (s *DocumentService) StopParseDocuments(datasetID string, docIDs []string) (map[string]interface{}, error) { + deduped := common.Deduplicate(docIDs) + if len(deduped) == 0 { + return nil, fmt.Errorf("no document IDs provided") + } + + docs, err := s.validateDocsInDataset(deduped, datasetID) + if err != nil { + return nil, err + } + + var errors []string + successCount := 0 + for _, doc := range docs { + if cancelErr := s.CancelDocParse(doc); cancelErr != nil { + errors = append(errors, cancelErr.Error()) + continue + } + successCount++ + } + + result := map[string]interface{}{"success_count": successCount} + if len(errors) > 0 { + result["errors"] = errors + } + return result, nil +} + +// validateDocsInDataset deduplicates IDs, fetches the documents, and ensures +// every document exists and belongs to the given dataset. Returns the resolved +// documents. +func (s *DocumentService) validateDocsInDataset(docIDs []string, datasetID string) ([]*entity.Document, error) { + docs, err := s.documentDAO.GetByIDs(docIDs) + if err != nil { + return nil, fmt.Errorf("failed to fetch documents: %w", err) + } + if len(docs) != len(docIDs) { + return nil, fmt.Errorf("some document IDs not found in dataset %s", datasetID) + } + var invalid []string + for _, d := range docs { + if d.KbID != datasetID { + invalid = append(invalid, d.ID) + } + } + if len(invalid) > 0 { + return nil, fmt.Errorf("these documents do not belong to dataset %s: %v", datasetID, invalid) + } + return docs, nil +} + +// CancelDocParse stops the ingestion task for the document by calling +// RequestStop (STOPPING), then marks the document run status as CANCEL. +func (s *DocumentService) CancelDocParse(doc *entity.Document) error { + task, err := s.ingestionTaskDAO.GetByDocumentID(doc.ID) + if err != nil { + return fmt.Errorf("failed to get ingestion task for %s: %v", doc.ID, err) + } + if task == nil { + return fmt.Errorf("no ingestion task found for document %s", doc.ID) + } + + if _, err := s.ingestionTaskSvc.RequestStop(task.ID); err != nil { + return fmt.Errorf("failed to stop ingestion task %s: %v", task.ID, err) + } + + if upErr := s.documentDAO.UpdateByID(doc.ID, map[string]interface{}{"run": string(entity.TaskStatusCancel)}); upErr != nil { + return fmt.Errorf("failed to update document %s: %v", doc.ID, upErr) + } + return nil +} + +func (s *DocumentService) resetDocumentForReparse(doc *entity.Document, tenantID string, parserID *string, pipelineID *string) error { + progressMsg := "" + run := string(entity.TaskStatusUnstart) + updates := map[string]interface{}{ + "progress": 0, + "progress_msg": progressMsg, + "run": run, + } + if parserID != nil { + updates["parser_id"] = *parserID + } + if pipelineID != nil { + updates["pipeline_id"] = *pipelineID + } + + if err := s.documentDAO.UpdateByID(doc.ID, updates); err != nil { + return errors.New("Document not found!") + } + + if doc.TokenNum > 0 { + decremented, err := s.decrementDocumentAndKBCountersForReparse(doc) + if err != nil { + return errors.New("Document not found!") + } + if !decremented { + return nil + } + 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 { + return err + } + } + } + + return nil +} + +func (s *DocumentService) deleteChunkImages(doc *entity.Document, indexName string) { + if s.docEngine == nil { + return + } + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return + } + + const pageSize = 1000 + for offset := 0; ; offset += pageSize { + result, err := s.docEngine.Search(context.Background(), &enginetypes.SearchRequest{ + IndexNames: []string{indexName}, + KbIDs: []string{doc.KbID}, + Offset: offset, + Limit: pageSize, + SelectFields: []string{"id", "img_id"}, + Filter: map[string]interface{}{"doc_id": doc.ID}, + MatchExprs: nil, + OrderBy: nil, + RankFeature: nil, + }) + if err != nil || result == nil || len(result.Chunks) == 0 { + return + } + for _, chunk := range result.Chunks { + imageKey, ok := chunkImageStorageKey(doc.KbID, chunk) + if !ok { + continue + } + if storageImpl.ObjExist(doc.KbID, imageKey) { + _ = storageImpl.Remove(doc.KbID, imageKey) + } + } + } +} + +func chunkImageStorageKey(defaultBucket string, chunk map[string]interface{}) (string, bool) { + imgID := firstStringField(chunk, "img_id") + if imgID != "" { + prefix := defaultBucket + "-" + if strings.HasPrefix(imgID, prefix) && len(imgID) > len(prefix) { + return strings.TrimPrefix(imgID, prefix), true + } + return imgID, true + } + + chunkID := firstStringField(chunk, "id", "_id") + if chunkID == "" { + return "", false + } + return chunkID, true +} + +func firstStringField(m map[string]interface{}, keys ...string) string { + for _, key := range keys { + if value, ok := m[key]; ok { + if s, ok := value.(string); ok { + return s + } + } + } + return "" +} + +func (s *DocumentService) decrementDocumentAndKBCountersForReparse(doc *entity.Document) (bool, error) { + decremented := false + err := dao.DB.Transaction(func(tx *gorm.DB) error { + result := tx.Model(&entity.Document{}). + Where("id = ? AND kb_id = ? AND token_num = ? AND chunk_num = ?", doc.ID, doc.KbID, doc.TokenNum, doc.ChunkNum). + Updates(map[string]interface{}{ + "token_num": gorm.Expr("token_num - ?", doc.TokenNum), + "chunk_num": gorm.Expr("chunk_num - ?", doc.ChunkNum), + "process_duration": gorm.Expr("process_duration - ?", doc.ProcessDuration), + }) + if result.Error != nil { + return result.Error + } + if result.RowsAffected == 0 { + return nil + } + decremented = true + + return tx.Model(&entity.Knowledgebase{}). + Where("id = ?", doc.KbID). + Updates(map[string]interface{}{ + "token_num": gorm.Expr("token_num - ?", doc.TokenNum), + "chunk_num": gorm.Expr("chunk_num - ?", doc.ChunkNum), + }).Error + }) + return decremented, err +} + +func (s *DocumentService) updateDocumentStatusOnly(doc *entity.Document, kb *entity.Knowledgebase, status int) error { + statusStr := strconv.Itoa(status) + if doc.Status != nil && *doc.Status == statusStr { + return nil + } + + if err := s.documentDAO.UpdateByID(doc.ID, map[string]interface{}{"status": statusStr}); err != nil { + return errors.New("Database error (Document update)!") + } + + if s.docEngine == nil { + return nil + } + + indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) + return s.docEngine.UpdateChunks( + context.Background(), + map[string]interface{}{"doc_id": doc.ID}, + map[string]interface{}{"available_int": status}, + indexName, + doc.KbID, + ) +} diff --git a/internal/service/document_test.go b/internal/service/document/document_test.go similarity index 96% rename from internal/service/document_test.go rename to internal/service/document/document_test.go index 0320b50376..44fae39faa 100644 --- a/internal/service/document_test.go +++ b/internal/service/document/document_test.go @@ -14,7 +14,7 @@ // limitations under the License. // -package service +package document import ( "bytes" @@ -37,10 +37,24 @@ import ( "ragflow/internal/dao" "ragflow/internal/engine/types" "ragflow/internal/entity" + "ragflow/internal/service" + "ragflow/internal/service/file" "ragflow/internal/storage" "ragflow/internal/utility" ) +// recordingTaskPublisher implements service.TaskPublisher and records published messages. +type recordingTaskPublisher struct { + subject string + messages []common.TaskMessage +} + +func (r *recordingTaskPublisher) PublishTaskMessage(subject string, msg common.TaskMessage) error { + r.subject = subject + r.messages = append(r.messages, msg) + return nil +} + type fakeUploadStorage struct { objects map[string][]byte } @@ -424,7 +438,7 @@ func testDocumentService(t *testing.T) *DocumentService { file2DocumentDAO: dao.NewFile2DocumentDAO(), fileDAO: dao.NewFileDAO(), ingestionTaskDAO: dao.NewIngestionTaskDAO(), - ingestionTaskSvc: NewIngestionTaskService(), + ingestionTaskSvc: service.NewIngestionTaskService(), docEngine: nil, metadataSvc: nil, // nil engine → metadata ops skipped } @@ -1960,7 +1974,7 @@ func TestUpdateDatasetDocumentPropagatesMetadataDeleteFailure(t *testing.T) { engine := &failingDeleteMetadataEngine{deleteErr: errors.New("delete failed")} svc := testDocumentService(t) svc.docEngine = engine - svc.metadataSvc = &MetadataService{} + svc.metadataSvc = service.NewMetadataServiceForTest(nil, nil) _, code, err := svc.UpdateDatasetDocument("tenant-1", "kb-1", "doc-1", &UpdateDatasetDocumentRequest{ MetaFields: map[string]any{"new": "value"}, @@ -1993,7 +2007,7 @@ func TestSetDocumentMetadataMergesMetadataRow(t *testing.T) { }, map[string]string{"doc-1": "kb-1"}) svc := testDocumentService(t) svc.docEngine = engine - svc.metadataSvc = &MetadataService{kbDAO: dao.NewKnowledgebaseDAO(), docEngine: engine} + svc.metadataSvc = service.NewMetadataServiceForTest(dao.NewKnowledgebaseDAO(), engine) if err := svc.SetDocumentMetadata("doc-1", map[string]interface{}{"category": "tech", "year": 2026}); err != nil { t.Fatalf("SetDocumentMetadata failed: %v", err) @@ -2065,7 +2079,7 @@ func TestBatchUpdateDocumentMetadatasMatchesPythonSemantics(t *testing.T) { svc := testDocumentService(t) svc.docEngine = engine - svc.metadataSvc = &MetadataService{kbDAO: dao.NewKnowledgebaseDAO(), docEngine: engine} + svc.metadataSvc = service.NewMetadataServiceForTest(dao.NewKnowledgebaseDAO(), engine) resp, code, err := svc.BatchUpdateDocumentMetadatas("kb-1", &DocumentMetadataSelector{ DocumentIDs: []string{"doc-1", "doc-2", "doc-3"}, @@ -2126,7 +2140,7 @@ func TestBatchUpdateDocumentMetadatasDoesNotReplaceWhenCurrentSearchIsStale(t *t svc := testDocumentService(t) svc.docEngine = engine - svc.metadataSvc = &MetadataService{kbDAO: dao.NewKnowledgebaseDAO(), docEngine: engine} + svc.metadataSvc = service.NewMetadataServiceForTest(dao.NewKnowledgebaseDAO(), engine) resp, code, err := svc.BatchUpdateDocumentMetadatas("kb-1", &DocumentMetadataSelector{ DocumentIDs: []string{"doc-1"}, @@ -2163,7 +2177,7 @@ func TestBatchUpdateDocumentMetadatasDeletesEmptyMetadataAndNoOps(t *testing.T) svc := testDocumentService(t) svc.docEngine = engine - svc.metadataSvc = &MetadataService{kbDAO: dao.NewKnowledgebaseDAO(), docEngine: engine} + svc.metadataSvc = service.NewMetadataServiceForTest(dao.NewKnowledgebaseDAO(), engine) resp, code, err := svc.BatchUpdateDocumentMetadatas("kb-1", &DocumentMetadataSelector{ DocumentIDs: []string{"doc-1", "doc-2"}, @@ -2192,7 +2206,7 @@ func TestBatchUpdateDocumentMetadatasNormalizesNumberValues(t *testing.T) { svc := testDocumentService(t) svc.docEngine = engine - svc.metadataSvc = &MetadataService{kbDAO: dao.NewKnowledgebaseDAO(), docEngine: engine} + svc.metadataSvc = service.NewMetadataServiceForTest(dao.NewKnowledgebaseDAO(), engine) resp, code, err := svc.BatchUpdateDocumentMetadatas("kb-1", &DocumentMetadataSelector{ DocumentIDs: []string{"doc-1"}, @@ -2266,9 +2280,9 @@ func TestAggregateMetadataIgnoresNestedEmptyLists(t *testing.T) { } func TestMergeFieldValuesKeepsNumericValues(t *testing.T) { - got := mergeFieldValues(1.0, 2.0) + got := service.MergeFieldValues(1.0, 2.0) if len(got) != 2 || got[0] != 1.0 || got[1] != 2.0 { - t.Fatalf("mergeFieldValues = %#v, want [1 2]", got) + t.Fatalf("MergeFieldValues = %#v, want [1 2]", got) } } @@ -2646,3 +2660,69 @@ func TestIngest_CancelDoesNotDeleteIngestionTask(t *testing.T) { t.Fatal("ingestion task must NOT be deleted by cancel") } } + +func TestUpdateRunProgressMirrorsFields(t *testing.T) { + db := setupServiceTestDB(t) + pushServiceDB(t, db) + insertTestDoc(t, "doc-1", "kb-1", 0, 0) + + svc := testDocumentService(t) + if err := svc.UpdateRunProgress("doc-1", 0.5, "1", "halfway"); err != nil { + t.Fatalf("UpdateRunProgress failed: %v", err) + } + doc, err := dao.NewDocumentDAO().GetByID("doc-1") + if err != nil { + t.Fatalf("load document: %v", err) + } + if doc.Progress != 0.5 { + t.Fatalf("progress = %v, want 0.5", doc.Progress) + } + if doc.Run == nil || *doc.Run != "1" { + t.Fatalf("run = %v, want 1", doc.Run) + } + if doc.ProgressMsg == nil || *doc.ProgressMsg != "halfway" { + t.Fatalf("progress_msg = %v, want halfway", doc.ProgressMsg) + } +} + +func TestFileDeleteRemovesLinkedDocument(t *testing.T) { + db := setupServiceTestDB(t) + pushServiceDB(t, db) + + insertTestKB(t, "kb-1", "tenant-1", 1, 30, 10) + insertTestDoc(t, "doc-1", "kb-1", 30, 10) + insertTestTask(t, "task-1", "doc-1") + // insertTestFile creates a knowledgebase-source file, which DeleteFiles skips. + // We need a local-source file so the deletion path is exercised. + loc := "doc.pdf" + testFile := &entity.File{ + ID: "file-1", ParentID: "kb-1", TenantID: "tenant-1", CreatedBy: "user-1", + Name: "test.pdf", Location: &loc, SourceType: "", Type: "pdf", + } + if err := dao.DB.Create(testFile).Error; err != nil { + t.Fatalf("insert test file: %v", err) + } + insertTestFile2Document(t, "f2d-1", "file-1", "doc-1") + + mockStorage := newFakeUploadStorage() + factory := storage.GetStorageFactory() + originalStorage := factory.GetStorage() + factory.SetStorage(mockStorage) + t.Cleanup(func() { factory.SetStorage(originalStorage) }) + + docSvc := testDocumentService(t) + fileSvc := file.NewFileService( + func(_ *dao.FileDAO, _ *entity.File, _ string) bool { return true }, + docSvc, + ) + + success, msg := fileSvc.DeleteFiles(context.Background(), "tenant-1", []string{"file-1"}) + if !success { + t.Fatalf("DeleteFiles failed: %s", msg) + } + + _, err := dao.NewDocumentDAO().GetByID("doc-1") + if err == nil { + t.Fatal("document should have been deleted but still exists") + } +} diff --git a/internal/service/document/document_upload.go b/internal/service/document/document_upload.go new file mode 100644 index 0000000000..a653cc9b88 --- /dev/null +++ b/internal/service/document/document_upload.go @@ -0,0 +1,424 @@ +package document + +import ( + "fmt" + "io" + "mime/multipart" + "net/http" + "path/filepath" + "ragflow/internal/service" + "strings" + + "ragflow/internal/common" + "ragflow/internal/entity" + "ragflow/internal/storage" + "ragflow/internal/utility" + + "go.uber.org/zap" +) + +// UploadLocalDocuments stores each uploaded file in object storage and inserts a +// matching Document row into the dataset. It mirrors Python +// FileService.upload_document: it derives parser_id by filetype, merges the +// optional parser_config override into the dataset config, dedup-renames the +// filename, records size + xxhash content hash, and links each document into the +// file manager (a File row under the dataset folder + a file2document mapping) +// so it surfaces in the dataset's document list. Chunking/embedding happen later +// in the parse step, so nothing here touches the doc store index. +// +// Gaps vs Python (documented, not yet ported): thumbnail generation and +// read_potential_broken_pdf repair. +func (s *DocumentService) UploadLocalDocuments(kb *entity.Knowledgebase, tenantID string, files []*multipart.FileHeader, parentPath string, parserConfigOverride map[string]interface{}) ([]map[string]interface{}, []string) { + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return nil, []string{"storage not initialized"} + } + + // Resolve (and create if needed) the dataset's file-manager folder up front. + // Without the File / file2document linkage the document list (which inner-joins + // file2document + file) would never surface the uploaded files. + kbFolder, err := s.ensureKBFolder(kb, tenantID) + if err != nil { + return nil, []string{err.Error()} + } + + // Merge parser_config override (allow-listed keys only) over the dataset config. + merged := entity.JSONMap{} + for k, v := range kb.ParserConfig { + merged[k] = v + } + for k, v := range parserConfigOverride { + merged[k] = v + } + + safeParent := utility.SanitizeFilename(parentPath) + + // Don't silently disable dedupe protection: a transient lookup failure means + // the existing-name set is unknown, so fail rather than risk duplicates. + names, err := s.documentDAO.ListNamesByKbID(kb.ID) + if err != nil { + return nil, []string{err.Error()} + } + taken := map[string]bool{} + for _, n := range names { + taken[n] = true + } + + var results []map[string]interface{} + var errMsgs []string + + for _, fh := range files { + blob, err := readFileHeaderBytes(fh) + if err != nil { + errMsgs = append(errMsgs, fh.Filename+": "+err.Error()) + continue + } + + filename := uniqueUploadName(fh.Filename, taken) + + filetype := utility.FilenameType(filename) + if filetype == utility.FileTypeOTHER { + errMsgs = append(errMsgs, fh.Filename+": This type of file has not been supported yet!") + continue + } + + location := filename + if safeParent != "" { + location = safeParent + "/" + filename + } + for storageImpl.ObjExist(kb.ID, location) { + location += "_" + } + if err := storageImpl.Put(kb.ID, location, blob); err != nil { + errMsgs = append(errMsgs, fh.Filename+": "+err.Error()) + continue + } + + 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) + errMsgs = append(errMsgs, fh.Filename+": "+err.Error()) + continue + } + if err := s.addFileFromKB(doc, kbFolder.ID, kb.TenantID); err != nil { + // Linkage failed: roll back the document row and blob so the partial + // state doesn't leave an invisible (unlisted) document behind. + err = s.rollbackAddFileFromKBError(doc, kb.ID, err) + _ = storageImpl.Remove(kb.ID, location) + errMsgs = append(errMsgs, fh.Filename+": "+err.Error()) + continue + } + // Only reserve the name once the write fully succeeds. + taken[filename] = true + results = append(results, docToRawMap(doc)) + } + + return results, errMsgs +} + +// UploadEmptyDocument inserts a zero-byte "virtual" document into the dataset. +func (s *DocumentService) UploadEmptyDocument(kb *entity.Knowledgebase, tenantID, name string) (map[string]interface{}, common.ErrorCode, error) { + // A transient lookup failure means the existing-name set is unknown; fail + // rather than write blind and risk a duplicate. + names, err := s.documentDAO.ListNamesByKbID(kb.ID) + if err != nil { + return nil, common.CodeServerError, err + } + for _, n := range names { + if n == name { + return nil, common.CodeDataError, fmt.Errorf("Duplicated document name in the same dataset.") + } + } + + kbFolder, err := s.ensureKBFolder(kb, tenantID) + if err != nil { + return nil, common.CodeServerError, err + } + + doc := s.newDatasetDocument(kb, tenantID, name, "", "virtual", kb.ParserConfig, "local", 0, nil) + if err := s.InsertDocument(doc); err != nil { + return nil, common.CodeServerError, err + } + if err := s.addFileFromKB(doc, kbFolder.ID, kb.TenantID); err != nil { + return nil, common.CodeServerError, s.rollbackAddFileFromKBError(doc, kb.ID, err) + } + return docToRawMap(doc), common.CodeSuccess, nil +} + +// ensureKBFolder resolves (creating as needed) the per-dataset file-manager +// folder: root -> .knowledgebase -> . Mirrors Python +// get_root_folder + get_kb_folder + new_a_file_from_kb. +func (s *DocumentService) ensureKBFolder(kb *entity.Knowledgebase, tenantID string) (*entity.File, error) { + root, err := s.fileDAO.GetRootFolder(tenantID) + if err != nil { + return nil, err + } + kbRoot, err := s.newAFileFromKB(tenantID, knowledgebaseFolderName, root.ID) + if err != nil { + return nil, err + } + return s.newAFileFromKB(kb.TenantID, kb.Name, kbRoot.ID) +} + +// newAFileFromKB returns the existing folder named name under parentID, or +// creates it. Mirrors Python FileService.new_a_file_from_kb. +func (s *DocumentService) newAFileFromKB(tenantID, name, parentID string) (*entity.File, error) { + for _, f := range s.fileDAO.Query(name, parentID, tenantID) { + if f.TenantID == tenantID { + return f, nil + } + } + loc := "" + folder := &entity.File{ + ID: utility.GenerateToken(), + ParentID: parentID, + TenantID: tenantID, + CreatedBy: tenantID, + Name: name, + Type: "folder", + Size: 0, + Location: &loc, + SourceType: string(entity.FileSourceKnowledgebase), + } + if err := s.fileDAO.Create(folder); err != nil { + return nil, err + } + return folder, nil +} + +// addFileFromKB links a document into the file manager: a File row under the +// dataset folder plus a file2document mapping. Mirrors Python +// FileService.add_file_from_kb (idempotent on the document mapping). +func (s *DocumentService) addFileFromKB(doc *entity.Document, kbFolderID, tenantID string) error { + if existing, err := s.file2DocumentDAO.GetByDocumentID(doc.ID); err == nil && len(existing) > 0 { + return nil + } + name := "" + if doc.Name != nil { + name = *doc.Name + } + loc := "" + if doc.Location != nil { + loc = *doc.Location + } + fileID := utility.GenerateToken() + file := &entity.File{ + ID: fileID, + ParentID: kbFolderID, + TenantID: tenantID, + CreatedBy: tenantID, + Name: name, + Type: doc.Type, + Size: doc.Size, + Location: &loc, + SourceType: string(entity.FileSourceKnowledgebase), + } + if err := s.fileDAO.Create(file); err != nil { + return err + } + docID := doc.ID + if err := s.file2DocumentDAO.Create(&entity.File2Document{ + ID: utility.GenerateToken(), + FileID: &fileID, + DocumentID: &docID, + }); err != nil { + _ = s.fileDAO.Delete(fileID) + return err + } + return nil +} + +func (s *DocumentService) UploadWebDocument(kb *entity.Knowledgebase, tenantID, name, url string) (map[string]interface{}, common.ErrorCode, error) { + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return nil, common.CodeServerError, fmt.Errorf("storage not initialized") + } + + kbFolder, err := s.ensureKBFolder(kb, tenantID) + if err != nil { + return nil, common.CodeServerError, err + } + + names, err := s.documentDAO.ListNamesByKbID(kb.ID) + if err != nil { + return nil, common.CodeServerError, err + } + taken := map[string]bool{} + for _, n := range names { + taken[n] = true + } + + blob, headers, _, err := utility.FetchRemoteFileSafely(url, maxUploadDocSize) + if err != nil { + return nil, common.CodeDataError, err + } + contentType := "" + if headers != nil { + contentType = headers.Get("Content-Type") + } + filename := normalizeWebDocumentName(name, contentType, blob) + filename, _, blob = utility.NormalizeUploadInfoContent(filename, contentType, blob) + filename = uniqueUploadName(filename, taken) + + filetype := utility.FilenameType(filename) + if filetype == utility.FileTypeOTHER { + return nil, common.CodeDataError, fmt.Errorf("This type of file has not been supported yet!") + } + + location := filename + for storageImpl.ObjExist(kb.ID, location) { + location += "_" + } + if err := storageImpl.Put(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) + return nil, common.CodeServerError, err + } + if err := s.addFileFromKB(doc, kbFolder.ID, kb.TenantID); err != nil { + err = s.rollbackAddFileFromKBError(doc, kb.ID, err) + _ = storageImpl.Remove(kb.ID, location) + return nil, common.CodeServerError, err + } + return docToRawMap(doc), common.CodeSuccess, nil +} + +func normalizeWebDocumentName(name, contentType string, blob []byte) string { + filename := utility.SanitizeFilename(name) + if filepath.Ext(filename) != "" { + return filename + } + lowerCT := strings.ToLower(strings.TrimSpace(strings.Split(contentType, ";")[0])) + switch { + case lowerCT == "application/pdf" || http.DetectContentType(blob) == "application/pdf" || utility.BytesLooksLikePDF(blob): + return filename + ".pdf" + case lowerCT == "text/html" || lowerCT == "application/xhtml+xml" || utility.LooksLikeHTML(blob): + return filename + ".html" + default: + return filename + } +} + +// newDatasetDocument builds a Document row for an upload, deriving parser_id, +// suffix and content hash. blob may be nil for the empty/virtual document. +func (s *DocumentService) newDatasetDocument(kb *entity.Knowledgebase, tenantID, filename, location, filetype string, parserConfig entity.JSONMap, src string, size int64, blob []byte) *entity.Document { + docID := utility.GenerateToken() + run := "0" + status := "1" + suffix := "" + if i := strings.LastIndex(filename, "."); i >= 0 { + suffix = filename[i+1:] + } + parserID := selectUploadParser(utility.FileType(filetype), filename, kb.ParserID) + if kb.PipelineID != nil { + parserID = "" // canvas pipeline mode — parser_id not applicable + } + loc := location + doc := &entity.Document{ + ID: docID, + KbID: kb.ID, + ParserID: parserID, + PipelineID: kb.PipelineID, + ParserConfig: parserConfig, + CreatedBy: tenantID, + Type: filetype, + SourceType: src, + Name: &filename, + Location: &loc, + Size: size, + Suffix: suffix, + Run: &run, + Status: &status, + } + if blob != nil { + hash := contentHashHex(blob) + doc.ContentHash = &hash + } + + // When the document's builtin parser_id differs from the KB's (e.g. visual→picture, + // aural→audio), re-resolve component_params defaults from the document's own DSL + // template so cpnIDs and param keys match the pipeline that will actually execute. + if kb.PipelineID == nil && parserID != kb.ParserID { + if cp, err := service.ResolveComponentParamsDefaults(parserID, nil); err != nil { + common.Warn("newDatasetDocument: resolve component_params defaults", + zap.String("parserID", parserID), zap.Error(err)) + } else if cp != nil { + doc.ParserConfig = cp + } + } + + return doc +} + +// docToRawMap serialises a freshly created Document into the raw key shape the +// handler remaps (chunk_num→chunk_count, kb_id→dataset_id). +func docToRawMap(doc *entity.Document) map[string]interface{} { + m := map[string]interface{}{ + "id": doc.ID, + "kb_id": doc.KbID, + "parser_id": doc.ParserID, + "parser_config": map[string]interface{}(doc.ParserConfig), + "created_by": doc.CreatedBy, + "type": doc.Type, + "source_type": doc.SourceType, + "size": doc.Size, + "chunk_num": doc.ChunkNum, + "token_num": doc.TokenNum, + "suffix": doc.Suffix, + "run": "0", + } + if doc.Name != nil { + m["name"] = *doc.Name + } + if doc.Location != nil { + m["location"] = *doc.Location + } + if doc.PipelineID != nil { + m["pipeline_id"] = *doc.PipelineID + } + if doc.ContentHash != nil { + m["content_hash"] = *doc.ContentHash + } + return m +} + +// uniqueUploadName appends a numeric suffix until the name is free, mirroring +// Python duplicate_name. +func uniqueUploadName(name string, taken map[string]bool) string { + if !taken[name] { + return name + } + base, ext := name, "" + if i := strings.LastIndex(name, "."); i >= 0 { + base, ext = name[:i], name[i:] + } + for i := 1; ; i++ { + candidate := fmt.Sprintf("%s(%d)%s", base, i, ext) + if !taken[candidate] { + return candidate + } + } +} + +func readFileHeaderBytes(fh *multipart.FileHeader) ([]byte, error) { + if fh.Size > maxUploadDocSize { + return nil, fmt.Errorf("file exceeds the maximum allowed size of %d bytes", maxUploadDocSize) + } + src, err := fh.Open() + if err != nil { + return nil, err + } + defer src.Close() + blob, err := io.ReadAll(io.LimitReader(src, maxUploadDocSize+1)) + if err != nil { + return nil, err + } + if len(blob) > maxUploadDocSize { + return nil, fmt.Errorf("file exceeds the maximum allowed size of %d bytes", maxUploadDocSize) + } + return blob, nil +} diff --git a/internal/service/document_upload_helpers.go b/internal/service/document/document_upload_helpers.go similarity index 97% rename from internal/service/document_upload_helpers.go rename to internal/service/document/document_upload_helpers.go index f277f84913..96c4fcf853 100644 --- a/internal/service/document_upload_helpers.go +++ b/internal/service/document/document_upload_helpers.go @@ -1,4 +1,4 @@ -package service +package document import ( "encoding/hex" diff --git a/internal/service/document/document_upload_test.go b/internal/service/document/document_upload_test.go new file mode 100644 index 0000000000..66639f497a --- /dev/null +++ b/internal/service/document/document_upload_test.go @@ -0,0 +1,38 @@ +package document + +import "testing" + +func TestNormalizeWebDocumentName(t *testing.T) { + pdfBlob := []byte("%PDF-1.4 fake") + htmlBlob := []byte("hi") + cases := []struct { + name, filename, ct string + blob []byte + want string + }{ + {"pdf detected by blob", "report", "application/octet-stream", pdfBlob, "report.pdf"}, + {"html detected by blob", "page", "application/octet-stream", htmlBlob, "page.html"}, + {"dot stripped by utility, no type hint", "image.png", "application/octet-stream", []byte("plain"), "imagepng"}, + {"no hint, no extension", "doc", "application/json", []byte("{}"), "doc"}, + } + for _, c := range cases { + if got := normalizeWebDocumentName(c.filename, c.ct, c.blob); got != c.want { + t.Errorf("%s: normalizeWebDocumentName = %q, want %q", c.name, got, c.want) + } + } +} + +func TestUniqueUploadName(t *testing.T) { + if got := uniqueUploadName("a.txt", map[string]bool{}); got != "a.txt" { + t.Errorf("free name: got %q", got) + } + if got := uniqueUploadName("a.txt", map[string]bool{"a.txt": true}); got != "a(1).txt" { + t.Errorf("single clash: got %q", got) + } + if got := uniqueUploadName("a.txt", map[string]bool{"a.txt": true, "a(1).txt": true}); got != "a(2).txt" { + t.Errorf("double clash: got %q", got) + } + if got := uniqueUploadName("noext", map[string]bool{"noext": true}); got != "noext(1)" { + t.Errorf("no-extension clash: got %q", got) + } +} diff --git a/internal/service/file2document.go b/internal/service/document/file2document.go similarity index 89% rename from internal/service/file2document.go rename to internal/service/document/file2document.go index efe8b15270..a4fe3116fb 100644 --- a/internal/service/file2document.go +++ b/internal/service/document/file2document.go @@ -14,11 +14,12 @@ // limitations under the License. // -package service +package document import ( "errors" "path/filepath" + "ragflow/internal/service" "strings" "go.uber.org/zap" @@ -110,7 +111,7 @@ func (s *File2DocumentService) LinkToDatasets(userID string, req *LinkToDatasets expanded := make([]string, 0, len(req.FileIDs)) for _, id := range req.FileIDs { file := filesSet[id] - if file.Type == FileTypeFolder { + if file.Type == "folder" { inner, err := s.getAllInnermostFileIDs(id) if err != nil { common.Warn("LinkToDatasets: folder expansion failed", zap.String("fileID", id), zap.Error(err)) @@ -129,14 +130,14 @@ func (s *File2DocumentService) LinkToDatasets(userID string, req *LinkToDatasets if err != nil || file == nil { return ErrLinkFileNotFound } - if !s.checkFileTeamPermission(file, userID) { + if !service.CheckFileTeamPermission(s.fileDAO, file, userID) { return ErrLinkNoAuthorization } } // ── 5. Validate KB permissions ──────────────────────────────────────────── for _, kb := range kbMap { - if !s.checkKBTeamPermission(kb, userID) { + if !service.HasKBTeamPermission(kb, userID, dao.NewTenantDAO()) { return ErrLinkNoAuthorization } } @@ -223,7 +224,7 @@ func (s *File2DocumentService) convertFiles(fileIDs, kbIDs []string, userID stri // defaults from the document's own DSL template so cpnIDs and // param keys match the pipeline that will actually execute. if kb.PipelineID == nil && parserID != kb.ParserID { - if cp, err := resolveComponentParamsDefaults(parserID, nil); err != nil { + if cp, err := service.ResolveComponentParamsDefaults(parserID, nil); err != nil { common.Warn("convertFiles: resolve component_params defaults", zap.String("parserID", parserID), zap.Error(err)) } else if cp != nil { @@ -262,7 +263,7 @@ func (s *File2DocumentService) getAllInnermostFileIDs(folderID string) ([]string } var ids []string for _, child := range children { - if child.Type == FileTypeFolder { + if child.Type == "folder" { sub, err := s.getAllInnermostFileIDs(child.ID) if err != nil { return nil, err @@ -275,36 +276,6 @@ func (s *File2DocumentService) getAllInnermostFileIDs(folderID string) ([]string return ids, nil } -// checkFileTeamPermission mirrors Python check_file_team_permission: -// true when file.TenantID == userID or user is in the file tenant's team. -func (s *File2DocumentService) checkFileTeamPermission(file *entity.File, userID string) bool { - if file.TenantID == userID { - return true - } - - datasetIDs, err := s.fileDAO.GetDatasetIDByFileID(file.ID) - if err != nil || len(datasetIDs) == 0 { - return false - } - - for _, datasetID := range datasetIDs { - kb, err := s.kbDAO.GetByID(datasetID) - if err != nil || kb == nil { - continue - } - if s.checkKBTeamPermission(kb, userID) { - return true - } - } - return false -} - -// checkKBTeamPermission mirrors Python check_kb_team_permission: -// true when kb.TenantID == userID or user is in the KB tenant's team. -func (s *File2DocumentService) checkKBTeamPermission(kb *entity.Knowledgebase, userID string) bool { - return hasKBTeamPermission(kb, userID, dao.NewTenantDAO()) -} - // dedupeStrings returns the input slice with duplicates removed, preserving the // first-seen order. func dedupeStrings(in []string) []string { diff --git a/internal/service/file.go b/internal/service/file.go deleted file mode 100644 index a7954d77f1..0000000000 --- a/internal/service/file.go +++ /dev/null @@ -1,1545 +0,0 @@ -// -// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. -// - -package service - -import ( - "context" - "encoding/base64" - "encoding/json" - "fmt" - "html" - "io" - "mime/multipart" - "net" - "net/http" - "net/url" - "path/filepath" - "ragflow/internal/common" - "ragflow/internal/dao" - "ragflow/internal/entity" - "ragflow/internal/parser/parser" - "ragflow/internal/storage" - "ragflow/internal/utility" - "regexp" - "strings" - "time" -) - -// FileService file service -type FileService struct { - fileDAO *dao.FileDAO - file2DocumentDAO *dao.File2DocumentDAO - documentService *DocumentService -} - -// NewFileService create file service -func NewFileService() *FileService { - return &FileService{ - fileDAO: dao.NewFileDAO(), - file2DocumentDAO: dao.NewFile2DocumentDAO(), - documentService: NewDocumentService(), - } -} - -// FileInfo file info with additional fields -type FileInfo struct { - *entity.File - Size int64 `json:"size"` - KbsInfo []map[string]interface{} `json:"kbs_info"` - HasChildFolder bool `json:"has_child_folder,omitempty"` -} - -// ListFilesResponse list files response -type ListFilesResponse struct { - Total int64 `json:"total"` - Files []map[string]interface{} `json:"files"` - ParentFolder map[string]interface{} `json:"parent_folder"` -} - -// GetRootFolder gets or creates root folder for tenant -func (s *FileService) GetRootFolder(tenantID string) (map[string]interface{}, error) { - file, err := s.fileDAO.GetRootFolder(tenantID) - if err != nil { - return nil, err - } - return s.toFileResponse(file), nil -} - -// ListFiles lists files by parent folder ID (matching Python /files endpoint) -// This method includes init_dataset_docs initialization when parent_id is empty -func (s *FileService) ListFiles(tenantID, pfID string, page, pageSize int, orderby string, desc bool, keywords string) (*ListFilesResponse, error) { - // If pfID is empty, get root folder and initialize dataset docs - if pfID == "" { - rootFolder, err := s.fileDAO.GetRootFolder(tenantID) - if err != nil { - return nil, fmt.Errorf("failed to get root folder: %w", err) - } - pfID = rootFolder.ID - - // Initialize dataset docs (matching Python init_knowledgebase_docs logic) - if err := s.initDatasetDocs(pfID, tenantID); err != nil { - return nil, fmt.Errorf("failed to initialize dataset docs: %w", err) - } - - // Initialize skills folder (matching Python init_skills_folder logic) - if err := s.initSkillsFolder(pfID, tenantID); err != nil { - return nil, fmt.Errorf("failed to initialize skills folder: %w", err) - } - } - - // Check if parent folder exists - if _, err := s.fileDAO.GetByID(pfID); err != nil { - return nil, fmt.Errorf("Folder not found!") - } - - // Get files by parent folder ID - files, total, err := s.fileDAO.GetByPfID(tenantID, pfID, page, pageSize, orderby, desc, keywords) - if err != nil { - return nil, err - } - - // Get parent folder - parentFolder, err := s.fileDAO.GetParentFolder(pfID) - if err != nil { - return nil, fmt.Errorf("File not found!") - } - - // Process files to add additional info, deduplicating by ID as a safety net - // against any leftover duplicate rows (e.g. duplicate 'skills' or '.knowledgebase' folders). - fileResponses := make([]map[string]interface{}, 0, len(files)) - seenIDs := make(map[string]struct{}) - for _, file := range files { - if _, ok := seenIDs[file.ID]; ok { - continue - } - seenIDs[file.ID] = struct{}{} - fileInfo := s.toFileInfo(file) - - // If folder, calculate size and check for child folders - if file.Type == FileTypeFolder { - folderSize, err := s.fileDAO.GetFolderSize(file.ID) - if err == nil { - fileInfo.Size = folderSize - } - hasChild, err := s.fileDAO.HasChildFolder(file.ID) - if err == nil { - fileInfo.HasChildFolder = hasChild - } - fileInfo.KbsInfo = []map[string]interface{}{} - } else { - // Get KB info for non-folder files - kbsInfo, err := s.file2DocumentDAO.GetKBInfoByFileID(file.ID) - if err != nil { - kbsInfo = []map[string]interface{}{} - } - fileInfo.KbsInfo = kbsInfo - } - - fileResponses = append(fileResponses, s.fileInfoToResponse(fileInfo)) - } - - return &ListFilesResponse{ - Total: total, - Files: fileResponses, - ParentFolder: s.toFileResponse(parentFolder), - }, nil -} - -// initDatasetDocs initializes dataset documents for tenant -// This matches Python's FileService.init_dataset_docs method -func (s *FileService) initDatasetDocs(rootID, tenantID string) error { - return s.fileDAO.InitDatasetDocs(rootID, tenantID, s.file2DocumentDAO) -} - -// DatasetFolderName is the folder name for dataset -const DatasetFolderName = ".knowledgebase" - -// SkillsFolderName is the folder name for skills -const SkillsFolderName = "skills" - -// initSkillsFolder initializes the skills folder under the root folder. -// Deduplicates duplicate entries that may have been created by -// concurrent race conditions (TOCTOU). -func (s *FileService) initSkillsFolder(rootID, tenantID string) error { - existing := s.fileDAO.Query(SkillsFolderName, rootID, tenantID) - if len(existing) > 0 { - if len(existing) > 1 { - common.Logger.Warn(fmt.Sprintf( - "Found %d duplicate '%s' folders under root %s, keeping only the first", - len(existing), SkillsFolderName, rootID, - )) - keepID := existing[0].ID - for _, dup := range existing[1:] { - children, _ := s.fileDAO.ListAllFilesByParentID(dup.ID) - for _, child := range children { - s.fileDAO.UpdateByID(child.ID, map[string]interface{}{"parent_id": keepID}) - } - if delErr := s.fileDAO.Delete(dup.ID); delErr != nil { - common.Logger.Warn(fmt.Sprintf("Failed to delete duplicate skills folder %s: %v", dup.ID, delErr)) - } - } - } - return nil - } - - folder := &entity.File{ - ID: utility.GenerateToken(), - ParentID: rootID, - TenantID: tenantID, - CreatedBy: tenantID, - Name: SkillsFolderName, - Type: FileTypeFolder, - Size: 0, - SourceType: "", - } - return s.fileDAO.Insert(folder) -} - -// FileSourceDataset represents dataset as file source -const FileSourceDataset = "knowledgebase" - -var ( - assertURLSafe = utility.AssertURLSafe - pinnedHTTPClient = utility.PinnedHTTPClient -) - -// toFileResponse converts file model to response format -func (s *FileService) toFileResponse(file *entity.File) map[string]interface{} { - result := map[string]interface{}{ - "id": file.ID, - "parent_id": file.ParentID, - "tenant_id": file.TenantID, - "created_by": file.CreatedBy, - "name": file.Name, - "size": file.Size, - "type": file.Type, - "create_time": file.CreateTime, - "update_time": file.UpdateTime, - } - - if file.Location != nil { - result["location"] = *file.Location - } - result["source_type"] = file.SourceType - - return result -} - -// toFileInfo converts file model to FileInfo -func (s *FileService) toFileInfo(file *entity.File) *FileInfo { - return &FileInfo{ - File: file, - Size: file.Size, - KbsInfo: []map[string]interface{}{}, - HasChildFolder: false, - } -} - -// fileInfoToResponse converts FileInfo to response map -func (s *FileService) fileInfoToResponse(info *FileInfo) map[string]interface{} { - result := map[string]interface{}{ - "id": info.File.ID, - "parent_id": info.File.ParentID, - "tenant_id": info.File.TenantID, - "created_by": info.File.CreatedBy, - "name": info.File.Name, - "size": info.Size, - "type": info.File.Type, - "create_time": info.File.CreateTime, - "update_time": info.File.UpdateTime, - "kbs_info": info.KbsInfo, - } - - if info.File.Location != nil { - result["location"] = *info.File.Location - } - result["source_type"] = info.File.SourceType - - if info.File.Type == "folder" { - result["has_child_folder"] = info.HasChildFolder - } - - return result -} - -// GetParentFolder gets parent folder of a file with permission check -func (s *FileService) GetParentFolder(userID, fileID string) (map[string]interface{}, error) { - // Get file - file, err := s.fileDAO.GetByID(fileID) - if err != nil { - return nil, err - } - - // Permission check - if !s.checkFileTeamPermission(file, userID) { - return nil, fmt.Errorf("No authorization.") - } - - // Get parent folder - parentFolder, err := s.fileDAO.GetParentFolder(fileID) - if err != nil { - return nil, err - } - - return s.toFileResponse(parentFolder), nil -} - -// GetAllParentFolders gets all parent folders in path with permission check -func (s *FileService) GetAllParentFolders(userID, fileID string) ([]map[string]interface{}, error) { - // Get file - file, err := s.fileDAO.GetByID(fileID) - if err != nil { - return nil, err - } - - // Permission check - if !s.checkFileTeamPermission(file, userID) { - return nil, fmt.Errorf("No authorization.") - } - - // Get all parent folders - parentFolders, err := s.fileDAO.GetAllParentFolders(fileID) - if err != nil { - return nil, err - } - - // Convert to response format - result := make([]map[string]interface{}, len(parentFolders)) - for i, folder := range parentFolders { - result[i] = s.toFileResponse(folder) - } - - return result, nil -} - -const ( - FileTypeFolder = "folder" - FileTypeVirtual = "virtual" -) - -// GetDocCount gets document count for a tenant -func (s *FileService) GetDocCount(tenantID string) (int64, error) { - documentDAO := dao.NewDocumentDAO() - return documentDAO.CountByTenantID(tenantID) -} - -// UploadFile uploads files to a folder -func (s *FileService) UploadFile(tenantID, parentID string, files []*multipart.FileHeader) ([]map[string]interface{}, error) { - if parentID == "" { - rootFolder, err := s.fileDAO.GetRootFolder(tenantID) - if err != nil { - return nil, fmt.Errorf("failed to get root folder: %w", err) - } - parentID = rootFolder.ID - } - - _, err := s.fileDAO.GetByID(parentID) - if err != nil { - return nil, fmt.Errorf("Can't find this folder!") - } - - maxFileNumPerUser := common.GetEnv(common.EnvMaxFileNumPerUser) - if maxFileNumPerUser != "" { - var maxNum int64 - if _, err = fmt.Sscanf(maxFileNumPerUser, "%d", &maxNum); err == nil && maxNum > 0 { - var docCount int64 - docCount, err = s.GetDocCount(tenantID) - if err != nil { - return nil, fmt.Errorf("failed to get document count: %w", err) - } - if docCount >= maxNum { - return nil, fmt.Errorf("Exceed the maximum file number of a free user!") - } - } - } - - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return nil, fmt.Errorf("storage not initialized") - } - - var result []map[string]interface{} - - for _, fileHeader := range files { - filename := fileHeader.Filename - if filename == "" { - return nil, fmt.Errorf("No file selected!") - } - - fileType := utility.FilenameType(filename) - - fileObjNames := s.parseFilePath(filename) - - var idList []string - idList, err = s.fileDAO.GetIDListByID(parentID, fileObjNames, 1, []string{parentID}) - if err != nil { - return nil, fmt.Errorf("failed to get file ID list: %w", err) - } - - var lastFolder *entity.File - if len(fileObjNames) != len(idList)-1 { - lastID := idList[len(idList)-1] - lastFolder, err = s.fileDAO.GetByID(lastID) - if err != nil { - return nil, fmt.Errorf("Folder not found!") - } - var createdFolder *entity.File - createdFolder, err = s.createFolderRecursive(lastFolder, fileObjNames, len(idList), tenantID) - if err != nil { - return nil, fmt.Errorf("failed to create folder: %w", err) - } - lastFolder = createdFolder - } else { - lastID := idList[len(idList)-2] - lastFolder, err = s.fileDAO.GetByID(lastID) - if err != nil { - return nil, fmt.Errorf("Folder not found!") - } - } - - location := fileObjNames[len(fileObjNames)-1] - for storageImpl.ObjExist(lastFolder.ID, location) { - location += "_" - } - - src, err := fileHeader.Open() - if err != nil { - return nil, fmt.Errorf("failed to open uploaded file: %w", err) - } - defer src.Close() - - data, err := io.ReadAll(src) - if err != nil { - return nil, fmt.Errorf("failed to read file data: %w", err) - } - - if err = storageImpl.Put(lastFolder.ID, location, data); err != nil { - return nil, fmt.Errorf("failed to store file: %w", err) - } - - uniqueName := s.getUniqueFilename(fileObjNames[len(fileObjNames)-1], lastFolder.ID, tenantID) - - fileRecord := &entity.File{ - ID: utility.GenerateToken(), - ParentID: lastFolder.ID, - TenantID: tenantID, - CreatedBy: tenantID, - Name: uniqueName, - Location: &location, - Size: int64(len(data)), - Type: string(fileType), - SourceType: "", - } - - if err = s.fileDAO.Insert(fileRecord); err != nil { - return nil, fmt.Errorf("failed to insert file record: %w", err) - } - - result = append(result, s.toFileResponse(fileRecord)) - } - - return result, nil -} - -// UploadInfos mirrors Python's upload_info file branch: store raw bytes in the -// per-user downloads bucket and return lightweight upload descriptors instead -// of creating full File rows in the file-management tree. -func (s *FileService) UploadInfos(userID string, files []*multipart.FileHeader) ([]map[string]interface{}, error) { - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return nil, fmt.Errorf("storage not initialized") - } - - results := make([]map[string]interface{}, 0, len(files)) - for _, fileHeader := range files { - filename := fileHeader.Filename - if err := s.checkUploadInfoHealth(userID, filename); err != nil { - return nil, err - } - src, err := fileHeader.Open() - if err != nil { - return nil, fmt.Errorf("failed to open uploaded file: %w", err) - } - data, readErr := readUploadInfoData(src) - src.Close() - if readErr != nil { - return nil, fmt.Errorf("failed to read file data: %w", readErr) - } - - contentType := fileHeader.Header.Get("Content-Type") - if contentType == "" { - contentType = http.DetectContentType(data) - } - filename, contentType, data = normalizeUploadInfoContent(filename, contentType, data) - resp, err := s.storeUploadInfoBlob(storageImpl, userID, filename, contentType, data) - if err != nil { - return nil, err - } - results = append(results, resp) - } - return results, nil -} - -func readUploadInfoData(r io.Reader) ([]byte, error) { - limited := &io.LimitedReader{R: r, N: maxRemoteFileSize + 1} - data, err := io.ReadAll(limited) - if err != nil { - return nil, err - } - if int64(len(data)) > maxRemoteFileSize { - return nil, fmt.Errorf("file size exceeds %d bytes", maxRemoteFileSize) - } - return data, nil -} - -func (s *FileService) parseFilePath(filename string) []string { - filename = strings.TrimPrefix(filename, "/") - parts := strings.Split(filename, "/") - var result []string - for _, part := range parts { - if part != "" { - result = append(result, part) - } - } - return result -} - -func (s *FileService) createFolderRecursive(parentFolder *entity.File, names []string, count int, tenantID string) (*entity.File, error) { - if count > len(names)-2 { - return parentFolder, nil - } - - newFolder, err := s.fileDAO.CreateFolder(parentFolder.ID, tenantID, names[count], FileTypeFolder) - if err != nil { - return nil, err - } - - return s.createFolderRecursive(newFolder, names, count+1, tenantID) -} - -func (s *FileService) getUniqueFilename(name, parentID, tenantID string) string { - existingFiles := s.fileDAO.Query(name, parentID, tenantID) - if len(existingFiles) == 0 { - return name - } - - base := filepath.Base(name) - ext := filepath.Ext(name) - nameWithoutExt := strings.TrimSuffix(base, ext) - - counter := 1 - for { - newName := fmt.Sprintf("%s_%d%s", nameWithoutExt, counter, ext) - existingFiles = s.fileDAO.Query(newName, parentID, tenantID) - if len(existingFiles) == 0 { - return newName - } - counter++ - } -} - -// CreateFolder creates a new folder or virtual file -func (s *FileService) CreateFolder(tenantID, name, parentID, fileType string) (map[string]interface{}, error) { - if parentID == "" { - rootFolder, err := s.fileDAO.GetRootFolder(tenantID) - if err != nil { - return nil, fmt.Errorf("failed to get root folder: %w", err) - } - parentID = rootFolder.ID - } - - if !s.fileDAO.IsParentFolderExist(parentID) { - return nil, fmt.Errorf("Parent Folder Doesn't Exist!") - } - - existingFiles := s.fileDAO.Query(name, parentID, tenantID) - if len(existingFiles) > 0 { - return nil, fmt.Errorf("Duplicated folder name in the same folder.") - } - - if fileType == "" { - fileType = FileTypeVirtual - } - - if fileType == FileTypeFolder { - fileType = FileTypeFolder - } else { - fileType = FileTypeVirtual - } - - folder, err := s.fileDAO.CreateFolder(parentID, tenantID, name, fileType) - if err != nil { - return nil, fmt.Errorf("failed to create folder: %w", err) - } - - return s.toFileResponse(folder), nil -} - -// DeleteFiles deletes files by IDs -// Returns (success, message) where success is true if all files were deleted -func (s *FileService) DeleteFiles(ctx context.Context, uid string, fileIDs []string) (bool, string) { - for _, fileID := range fileIDs { - // 1. Get file - file, err := s.fileDAO.GetByID(fileID) - if err != nil || file == nil { - return false, "File or Folder not found!" - } - - // 2. Check tenant_id - if file.TenantID == "" { - return false, "Tenant not found!" - } - - // Block root-folder deletion (root folders have parent_id == id) - if file.ParentID == file.ID { - return false, "Root folder cannot be deleted." - } - - // 3. Permission check - if !s.checkFileTeamPermission(file, uid) { - return false, "No authorization." - } - - // 4. Skip dataset source files - if file.SourceType == FileSourceDataset { - continue - } - - // 5. Delete based on type - if file.Type == FileTypeFolder { - if err := s.deleteFolderRecursive(ctx, file, uid); err != nil { - return false, fmt.Sprintf("Failed to delete folder: %v", err) - } - } else { - if err := s.deleteSingleFile(ctx, file); err != nil { - return false, fmt.Sprintf("Failed to delete file: %v", err) - } - } - } - - return true, "" -} - -// checkFileTeamPermission checks if user has permission to access the file -// Matches Python's check_file_team_permission function -func (s *FileService) checkFileTeamPermission(file *entity.File, uid string) bool { - // File's tenant directly authorized - if file.TenantID == uid { - return true - } - - // Check KB permissions - datasetIDs, err := s.fileDAO.GetDatasetIDByFileID(file.ID) - if err != nil || len(datasetIDs) == 0 { - return false - } - - kbDAO := dao.NewKnowledgebaseDAO() - for _, datasetID := range datasetIDs { - ds, err := kbDAO.GetByID(datasetID) - if err != nil || ds == nil { - continue - } - - // Check KB tenant permission - if s.checkDatasetTeamPermission(ds, uid) { - return true - } - } - - return false -} - -// checkDatasetTeamPermission checks if user has permission to access the dataset -// Matches Python's check_kb_team_permission function -func (s *FileService) checkDatasetTeamPermission(ds *entity.Knowledgebase, uid string) bool { - return hasKBTeamPermission(ds, uid, dao.NewTenantDAO()) -} - -// deleteSingleFile deletes a single file (not folder) -// Matches Python's _delete_single_file function -func (s *FileService) deleteSingleFile(ctx context.Context, file *entity.File) error { - // 1. Delete storage object - if file.Location != nil && *file.Location != "" { - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl != nil { - if err := storageImpl.Remove(file.ParentID, *file.Location); err != nil { - common.Logger.Error(fmt.Sprintf("Fail to remove object: %s/%s, error: %v", file.ParentID, *file.Location, err)) - } - } - } - - // 2. Handle associated documents - informs, err := s.file2DocumentDAO.GetByFileID(file.ID) - if err != nil { - return fmt.Errorf("failed to get file2document mappings: %w", err) - } - if len(informs) > 0 { - for _, inform := range informs { - if inform.DocumentID == nil { - continue - } - docID := *inform.DocumentID - if err := s.documentService.RemoveDocumentKeepFile(docID); err != nil { - common.Logger.Error(fmt.Sprintf("Fail to remove document: %s, error: %v", docID, err)) - } - } - - // Delete file2document mapping (outside the loop, called once - matching Python behavior) - if err := s.file2DocumentDAO.DeleteByFileID(file.ID); err != nil { - return fmt.Errorf("failed to delete file2document mapping: %w", err) - } - } - - // 3. Delete file record - if err := s.fileDAO.Delete(file.ID); err != nil { - return err - } - - return nil -} - -// deleteFolderRecursive recursively deletes a folder and its contents -// Matches Python's _delete_folder_recursive function -func (s *FileService) deleteFolderRecursive(ctx context.Context, folder *entity.File, uid string) error { - // Get all sub-files - subFiles, err := s.fileDAO.ListByParentID(folder.ID) - if err != nil { - return err - } - - for _, subFile := range subFiles { - if subFile.Type == FileTypeFolder { - // Recursively delete subfolder - if err := s.deleteFolderRecursive(ctx, subFile, uid); err != nil { - return err - } - } else { - // Delete single file - if err := s.deleteSingleFile(ctx, subFile); err != nil { - return err - } - } - } - - // Delete the folder itself - if err := s.fileDAO.Delete(folder.ID); err != nil { - return err - } - - return nil -} - -// MoveFileReq represents the request body for move files operation -type MoveFileReq struct { - SrcFileIDs []string `json:"src_file_ids" binding:"required,min=1"` - DestFileID string `json:"dest_file_id"` - NewName string `json:"new_name"` -} - -// MoveFiles moves and/or renames files -// Follows Linux mv semantics: -// - new_name only: rename in place (no storage operation) -// - dest_file_id only: move to new folder (keep names) -// - both: move and rename simultaneously -func (s *FileService) MoveFiles(uid string, srcFileIDs []string, destFileID string, newName string) (bool, string) { - // 1. Get all source files - files, err := s.fileDAO.GetByIDs(srcFileIDs) - if err != nil || len(files) == 0 { - return false, "Source files not found!" - } - - // Create a map for quick lookup - filesMap := make(map[string]*entity.File) - for _, f := range files { - filesMap[f.ID] = f - } - - // 2. Validate all source files - for _, fileID := range srcFileIDs { - file, ok := filesMap[fileID] - if !ok { - return false, "File or folder not found!" - } - if file.TenantID == "" { - return false, "Tenant not found!" - } - // 3. Permission check - if !s.checkFileTeamPermission(file, uid) { - return false, "No authorization." - } - } - - // 4. Validate destination folder if provided - var destFolder *entity.File - if destFileID != "" { - destFolder, err = s.fileDAO.GetByID(destFileID) - if err != nil || destFolder == nil { - return false, "Parent folder not found!" - } - // Check destination folder permission - if !s.checkFileTeamPermission(destFolder, uid) { - return false, "No authorization to write to destination folder." - } - - if destFolder.Type != FileTypeFolder { - return false, "Destination is not a folder." - } - - destAncestors, err := s.fileDAO.GetAllParentFolders(destFolder.ID) - if err != nil { - return false, "Parent folder not found!" - } - - destAncestorIDs := make(map[string]struct{}, len(destAncestors)) - for _, folder := range destAncestors { - destAncestorIDs[folder.ID] = struct{}{} - } - - for _, file := range files { - if file.Type != FileTypeFolder { - continue - } - - if file.ID == destFolder.ID { - return false, "Cannot move a folder to itself." - } - - if _, ok := destAncestorIDs[file.ID]; ok { - return false, "Cannot move a folder into its own subfolder." - } - } - } - - // 5. Validate new_name if provided - if newName != "" { - if len(srcFileIDs) > 1 { - return false, "new_name can only be used with a single file" - } - - file := filesMap[srcFileIDs[0]] - // Check extension for non-folder files - if file.Type != FileTypeFolder { - oldExt := utility.GetFileExtension(file.Name) - newExt := utility.GetFileExtension(newName) - if oldExt != newExt { - return false, "The extension of file can't be changed" - } - } - - // Check for duplicate names in target folder - targetParentID := file.ParentID - if destFolder != nil { - targetParentID = destFolder.ID - } - existingFiles := s.fileDAO.Query(newName, targetParentID, file.TenantID) - for _, f := range existingFiles { - if f.Name == newName { - return false, "Duplicated file name in the same folder." - } - } - } else if destFolder != nil { - // Plain move (no rename): check for duplicate names in destination folder - for _, file := range files { - existingFiles := s.fileDAO.Query(file.Name, destFolder.ID, file.TenantID) - for _, f := range existingFiles { - // Ignore the source file itself - if f.ID != file.ID { - return false, "Duplicated file name in the same folder." - } - } - } - } - - // 6. Perform the move operation - if destFolder != nil { - // Move to destination folder - for _, file := range files { - if err := s.moveEntryRecursive(file, destFolder, newName); err != nil { - return false, err.Error() - } - } - } else { - // Pure rename: no storage operation needed - if newName == "" { - return false, "new_name is required for rename" - } - if len(srcFileIDs) == 0 { - return false, "Source files not found!" - } - file := filesMap[srcFileIDs[0]] - if err := s.fileDAO.UpdateByID(file.ID, map[string]interface{}{"name": newName}); err != nil { - return false, "Database error (File rename)!" - } - - // Update associated document name if exists - informs, err := s.file2DocumentDAO.GetByFileID(file.ID) - if err == nil && len(informs) > 0 && informs[0].DocumentID != nil { - docID := *informs[0].DocumentID - documentDAO := dao.NewDocumentDAO() - if err := documentDAO.UpdateByID(docID, map[string]interface{}{"name": newName}); err != nil { - return false, "Database error (Document rename)!" - } - } - } - - return true, "" -} - -// moveEntryRecursive recursively moves a file or folder entry -func (s *FileService) moveEntryRecursive(sourceFile *entity.File, destFolder *entity.File, overrideName string) error { - effectiveName := overrideName - if effectiveName == "" { - effectiveName = sourceFile.Name - } - - if sourceFile.Type == FileTypeFolder { - // Handle folder move - existingFolders := s.fileDAO.Query(effectiveName, destFolder.ID, sourceFile.TenantID) - var newFolder *entity.File - if len(existingFolders) > 0 { - // Prevent moving a folder into itself (self-target merge) - if existingFolders[0].ID == sourceFile.ID { - return fmt.Errorf("cannot move folder into itself") - } - newFolder = existingFolders[0] - } else { - // Create new folder - var err error - newFolder, err = s.fileDAO.CreateFolder(destFolder.ID, sourceFile.TenantID, effectiveName, FileTypeFolder) - if err != nil { - return fmt.Errorf("failed to create destination folder: %w", err) - } - } - - // Recursively move sub-files - subFiles, err := s.fileDAO.ListAllFilesByParentID(sourceFile.ID) - if err != nil { - return err - } - for _, subFile := range subFiles { - if err := s.moveEntryRecursive(subFile, newFolder, ""); err != nil { - return err - } - } - - // Delete the source folder - return s.fileDAO.Delete(sourceFile.ID) - } - - // Handle non-folder file move - needStorageMove := destFolder.ID != sourceFile.ParentID - updates := map[string]interface{}{} - - if needStorageMove { - // Get storage - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return fmt.Errorf("storage not initialized") - } - - // Calculate new location - newLocation := effectiveName - for storageImpl.ObjExist(destFolder.ID, newLocation) { - newLocation += "_" - } - - // Perform storage move (copy + delete) - if sourceFile.Location == nil || *sourceFile.Location == "" { - return fmt.Errorf("file location is empty") - } - - if !storageImpl.Move(sourceFile.ParentID, *sourceFile.Location, destFolder.ID, newLocation) { - return fmt.Errorf("move file failed at storage layer") - } - - updates["parent_id"] = destFolder.ID - updates["location"] = newLocation - } - - if overrideName != "" { - updates["name"] = overrideName - } - - if len(updates) > 0 { - if err := s.fileDAO.UpdateByID(sourceFile.ID, updates); err != nil { - return fmt.Errorf("database error (File update): %w", err) - } - } - - // Update associated document name if renamed - if overrideName != "" { - informs, err := s.file2DocumentDAO.GetByFileID(sourceFile.ID) - if err == nil && len(informs) > 0 && informs[0].DocumentID != nil { - docID := *informs[0].DocumentID - documentDAO := dao.NewDocumentDAO() - if err := documentDAO.UpdateByID(docID, map[string]interface{}{"name": overrideName}); err != nil { - return fmt.Errorf("database error (Document rename): %w", err) - } - } - } - - return nil -} - -// GetFileContent gets file metadata and checks permission for download -// Matches Python's file_api_service.get_file_content function -func (s *FileService) GetFileContent(uid, fileID string) (*entity.File, error) { - file, err := s.fileDAO.GetByID(fileID) - if err != nil || file == nil { - return nil, fmt.Errorf("Document not found!") - } - if !s.checkFileTeamPermission(file, uid) { - return nil, fmt.Errorf("No authorization.") - } - return file, nil -} - -// StorageAddress represents bucket and object name for storage -type StorageAddress struct { - Bucket string - Name string -} - -// GetStorageAddress gets storage address for a file (fallback for when direct blob is empty) -// Matches Python's File2DocumentService.get_storage_address function -func (s *FileService) GetStorageAddress(fileID string) (*StorageAddress, error) { - // Get file2document mapping - f2d, err := s.file2DocumentDAO.GetByFileID(fileID) - if err != nil || len(f2d) == 0 { - return nil, fmt.Errorf("file2document mapping not found") - } - - // Get the file - if f2d[0].FileID == nil { - return nil, fmt.Errorf("file_id is nil in file2document mapping") - } - file, err := s.fileDAO.GetByID(*f2d[0].FileID) - if err != nil || file == nil { - return nil, fmt.Errorf("file not found") - } - - // If source_type is empty or local, return file's parent_id and location - if file.SourceType == "" || entity.FileSource(file.SourceType) == entity.FileSourceLocal { - if file.Location == nil || *file.Location == "" { - return nil, fmt.Errorf("file location is empty") - } - return &StorageAddress{ - Bucket: file.ParentID, - Name: *file.Location, - }, nil - } - - // Otherwise, use document's kb_id and location - if f2d[0].DocumentID == nil { - return nil, fmt.Errorf("document_id is required") - } - - documentDAO := dao.NewDocumentDAO() - doc, err := documentDAO.GetByID(*f2d[0].DocumentID) - if err != nil || doc == nil { - return nil, fmt.Errorf("document not found") - } - - if doc.Location == nil || *doc.Location == "" { - return nil, fmt.Errorf("document location is empty") - } - - return &StorageAddress{ - Bucket: doc.KbID, - Name: *doc.Location, - }, nil -} - -// DownloadAgentFile downloads an agent-generated file directly from MinIO without querying the database. -func (s *FileService) DownloadAgentFile(tenantID, location string) ([]byte, error) { - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return nil, fmt.Errorf("storage not initialized") - } - - bucketName := fmt.Sprintf("%s-downloads", tenantID) - - blob, err := storageImpl.Get(bucketName, location) - if err != nil { - return nil, fmt.Errorf("failed to read file from storage: %w", err) - } - - return blob, nil -} - -// GetFileContents fetches file contents (text + image) from storage -// for the given file dicts. -// - raw=false: images returned as base64 data URIs in images; non-images parsed and returned as text. -// - raw=true: images returned as raw bytes in images; non-images parsed and returned as text. -func (s *FileService) GetFileContents(uid string, fileDicts []map[string]interface{}, raw bool) (texts []string, images []string, err error) { - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return nil, nil, fmt.Errorf("storage not initialized") - } - - for _, fd := range fileDicts { - id, _ := fd["id"].(string) - if id == "" { - continue - } - file, ferr := s.fileDAO.GetByID(id) - if ferr != nil || file == nil || file.Location == nil || *file.Location == "" { - continue - } - if !s.checkFileTeamPermission(file, uid) { - return nil, nil, fmt.Errorf("No authorization.") - } - data, derr := storageImpl.Get(file.ParentID, *file.Location) - if derr != nil || len(data) == 0 { - continue - } - ft := utility.FilenameType(file.Name) - if ft == utility.FileTypeVISUAL { - if raw { - images = append(images, string(data)) - } else { - ext := utility.GetFileExtension(file.Name) - mime := utility.GetContentType(ext, string(ft)) - images = append(images, "data:"+mime+";base64,"+base64.StdEncoding.EncodeToString(data)) - } - } else { - texts = append(texts, parseFileContent(file.Name, data)) - } - } - return texts, images, nil -} - -// parseAgentUploads resolves descriptors returned by upload_info from the -// caller's downloads bucket and converts them to sys.files values. -func (s *FileService) parseAgentUploads(userID string, fileDicts []map[string]interface{}, layoutRecognize string) ([]string, error) { - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return nil, fmt.Errorf("storage not initialized") - } - - contents := make([]string, 0, len(fileDicts)) - for i, fd := range fileDicts { - id, _ := fd["id"].(string) - name, _ := fd["name"].(string) - mimeType, _ := fd["mime_type"].(string) - createdBy, _ := fd["created_by"].(string) - if id == "" || name == "" || mimeType == "" || createdBy == "" { - return nil, fmt.Errorf("file %d: id, name, mime_type, and created_by are required", i) - } - if createdBy != userID { - return nil, fmt.Errorf("file %q: created_by does not match the current user", name) - } - - data, err := storageImpl.Get(createdBy+"-downloads", id) - if err != nil { - return nil, fmt.Errorf("file %q: read upload: %w", name, err) - } - if len(data) == 0 { - return nil, fmt.Errorf("file %q: upload is empty", name) - } - - mediaType := strings.ToLower(strings.TrimSpace(strings.Split(mimeType, ";")[0])) - if strings.HasPrefix(mediaType, "image/") { - contents = append(contents, "data:"+mediaType+";base64,"+base64.StdEncoding.EncodeToString(data)) - continue - } - - content, err := parseAgentUploadContent(name, data, layoutRecognize) - if err != nil { - return nil, fmt.Errorf("file %q: parse upload: %w", name, err) - } - contents = append(contents, content) - } - return contents, nil -} - -func parseAgentUploadContent(filename string, data []byte, layoutRecognize string) (string, error) { - content := string(data) - fileType := utility.GetFileType(filename) - if fileType != utility.FileTypeOTHER { - fp, err := parser.GetParser(fileType) - if err != nil { - return "", err - } - if configurable, ok := fp.(interface{ ConfigureFromSetup(map[string]any) }); ok { - configurable.ConfigureFromSetup(map[string]any{"layout_recognize": layoutRecognize}) - } - res := fp.ParseWithResult(filename, data) - if res.Err != nil { - return "", res.Err - } - switch res.OutputFormat { - case "text": - content = res.Text - case "markdown": - content = res.Markdown - case "html": - content = res.HTML - case "json": - parts := make([]string, 0, len(res.JSON)) - for _, item := range res.JSON { - if text, ok := item["text"].(string); ok { - parts = append(parts, text) - continue - } - raw, err := json.Marshal(item) - if err != nil { - return "", err - } - parts = append(parts, string(raw)) - } - content = strings.Join(parts, "\n") - } - } - return fmt.Sprintf("\n -----------------\nFile: %s\nContent as following: \n%s", filename, content), nil -} - -// parseFileContent tries to parse a file's contents using the appropriate parser. -// Falls back to returning raw text if no parser is available. -func parseFileContent(filename string, data []byte) string { - fileType := utility.GetFileType(filename) - if fileType == utility.FileTypeOTHER { - return string(data) - } - fp, err := parser.GetParser(fileType) - if err != nil { - return string(data) - } - res := fp.ParseWithResult(filename, data) - if res.Err != nil { - return string(data) - } - switch res.OutputFormat { - case "text": - return res.Text - case "markdown": - return res.Markdown - case "html": - return res.HTML - case "json": - return string(data) - default: - return string(data) - } -} - -// toUploadInfoResponse converts a newly-uploaded file record to the shape -// Python's upload_info endpoint returns. -func (s *FileService) toUploadInfoResponse(file *entity.File, mimeType string) map[string]interface{} { - ext := "" - if idx := strings.LastIndex(file.Name, "."); idx >= 0 { - ext = strings.ToLower(file.Name[idx+1:]) - } - return map[string]interface{}{ - "id": file.ID, - "name": file.Name, - "size": file.Size, - "extension": ext, - "mime_type": mimeType, - "created_by": file.CreatedBy, - "created_at": float64(time.Now().UnixMilli()) / 1000.0, - "preview_url": nil, - } -} - -// maxRemoteFileSize bounds the body of a ?url= upload (100 MB). -const maxRemoteFileSize = 100 << 20 - -// UploadFromURL fetches a remote URL, saves the content to the tenant's root -// folder, and returns the file metadata map — mirroring Python -// FileService.upload_info(tenant_id, None, url). -// -// The remote fetch is SSRF-guarded (mirrors Python's assert_url_is_safe): the -// scheme must be http/https and every address the host resolves to must be -// globally routable; the validated IP is pinned for the actual connection — and -// re-validated on each redirect hop — to defeat DNS-rebinding. The HTTP client -// carries connect and overall timeouts, and the response body is bounded with -// truncation detection so an oversized file is rejected rather than silently -// clipped. -func (s *FileService) UploadFromURL(tenantID, rawURL string) (map[string]interface{}, error) { - if rawURL == "" { - return nil, fmt.Errorf("url is required") - } - parsed, err := url.Parse(rawURL) - if err != nil || (parsed.Scheme != "http" && parsed.Scheme != "https") || parsed.Hostname() == "" { - return nil, fmt.Errorf("invalid or unsafe URL") - } - - data, headers, finalURL, err := fetchRemoteFileSafely(rawURL, maxRemoteFileSize) - if err != nil { - return nil, err - } - - storageImpl := storage.GetStorageFactory().GetStorage() - if storageImpl == nil { - return nil, fmt.Errorf("storage not initialized") - } - - contentType := headers.Get("Content-Type") - filename := normalizeRemoteUploadFilename(finalURL, contentType, data) - if err := s.checkUploadInfoHealth(tenantID, filename); err != nil { - return nil, err - } - filename, contentType, data = normalizeUploadInfoContent(filename, contentType, data) - return s.storeUploadInfoBlob(storageImpl, tenantID, filename, contentType, data) -} - -// fetchRemoteFileSafely downloads rawURL with SSRF protection, connect/overall -// timeouts, and a hard size cap that rejects (rather than truncates) oversized -// bodies. -func fetchRemoteFileSafely(rawURL string, maxSize int64) ([]byte, http.Header, string, error) { - currentURL := rawURL - for redirects := 0; redirects < 10; redirects++ { - hostname, resolvedIP, err := assertURLSafe(currentURL) - if err != nil { - return nil, nil, "", err - } - client := pinnedHTTPClient(hostname, resolvedIP, 10*time.Second) - client.CheckRedirect = func(req *http.Request, via []*http.Request) error { - return http.ErrUseLastResponse - } - - // runs assertURLSafe(currentURL) on every iteration (including - // redirects), which rejects private/loopback IPs and other - // SSRF targets. The "nosec G107" comment is for gosec; - // CodeQL needs an explicit suppression. - // codeql[go/request-forgery] False positive: the loop above - resp, err := client.Get(currentURL) // #nosec G107 - if err != nil { - return nil, nil, "", fmt.Errorf("failed to fetch URL: %w", err) - } - - if resp.StatusCode == http.StatusMovedPermanently || - resp.StatusCode == http.StatusFound || - resp.StatusCode == http.StatusSeeOther || - resp.StatusCode == http.StatusTemporaryRedirect || - resp.StatusCode == http.StatusPermanentRedirect { - location := resp.Header.Get("Location") - resp.Body.Close() - if location == "" { - return nil, nil, "", fmt.Errorf("redirect response missing Location header") - } - baseURL, parseErr := url.Parse(currentURL) - if parseErr != nil { - return nil, nil, "", parseErr - } - nextURL, resolveErr := baseURL.Parse(location) - if resolveErr != nil { - return nil, nil, "", resolveErr - } - currentURL = nextURL.String() - continue - } - - if resp.StatusCode >= 400 { - resp.Body.Close() - return nil, nil, "", fmt.Errorf("remote URL returned HTTP %d", resp.StatusCode) - } - - data, readErr := io.ReadAll(io.LimitReader(resp.Body, maxSize+1)) - resp.Body.Close() - if readErr != nil { - return nil, nil, "", fmt.Errorf("failed to read remote content: %w", readErr) - } - if int64(len(data)) > maxSize { - return nil, nil, "", fmt.Errorf("remote file exceeds the maximum allowed size of %d bytes", maxSize) - } - return data, resp.Header.Clone(), currentURL, nil - } - return nil, nil, "", fmt.Errorf("stopped after too many redirects") -} - -// isPublicIP reports whether ip is a globally routable address. It mirrors the -// allowlist intent of Python's assert_url_is_safe (which requires ip.is_global) -// by rejecting loopback, private, link-local, multicast, unspecified, and -// carrier-grade NAT ranges. IPv4-mapped IPv6 addresses are handled by the -// stdlib predicates. -func isPublicIP(ip net.IP) bool { - if ip == nil { - return false - } - if ip.IsLoopback() || ip.IsPrivate() || ip.IsUnspecified() || - ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || - ip.IsMulticast() || ip.IsInterfaceLocalMulticast() { - return false - } - // Carrier-grade NAT 100.64.0.0/10 (RFC 6598) — not covered by IsPrivate. - if ip4 := ip.To4(); ip4 != nil && ip4[0] == 100 && ip4[1]&0xc0 == 0x40 { - return false - } - return true -} - -func (s *FileService) checkUploadInfoHealth(userID, filename string) error { - if filename == "" { - return fmt.Errorf("No file selected!") - } - maxFileNumPerUser := common.GetEnv(common.EnvMaxFileNumPerUser) - if maxFileNumPerUser != "" { - var maxNum int64 - if _, err := fmt.Sscanf(maxFileNumPerUser, "%d", &maxNum); err == nil && maxNum > 0 { - var docCount int64 - docCount, err = s.GetDocCount(userID) - if err != nil { - return fmt.Errorf("failed to get document count: %w", err) - } - if docCount >= maxNum { - return fmt.Errorf("Exceed the maximum file number of a free user!") - } - } - } - if len([]byte(filename)) > 255 { - return fmt.Errorf("Exceed the maximum length of file name!") - } - return nil -} - -func (s *FileService) storeUploadInfoBlob(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 { - return nil, fmt.Errorf("failed to store file: %w", err) - } - ext := "" - if idx := strings.LastIndex(filename, "."); idx >= 0 { - ext = strings.ToLower(filename[idx+1:]) - } - return map[string]interface{}{ - "id": location, - "name": filename, - "size": int64(len(data)), - "extension": ext, - "mime_type": contentType, - "created_by": userID, - "created_at": float64(time.Now().UnixMilli()) / 1000.0, - "preview_url": nil, - }, nil -} - -func normalizeRemoteUploadFilename(rawURL, contentType string, data []byte) string { - parsed, err := url.Parse(rawURL) - filename := "download" - if err == nil { - filename = sanitizeFilename(filepath.Base(parsed.Path)) - } - ct := strings.ToLower(strings.TrimSpace(strings.Split(contentType, ";")[0])) - if ct == "application/pdf" || bytesLooksLikePDF(data) { - if !strings.HasSuffix(strings.ToLower(filename), ".pdf") { - filename += ".pdf" - } - } - return filename -} - -func normalizeUploadInfoContent(filename, contentType string, data []byte) (string, string, []byte) { - lowerCT := strings.ToLower(strings.TrimSpace(strings.Split(contentType, ";")[0])) - if lowerCT == "" { - lowerCT = http.DetectContentType(data) - } - - if lowerCT == "application/pdf" || bytesLooksLikePDF(data) { - if !strings.HasSuffix(strings.ToLower(filename), ".pdf") { - filename += ".pdf" - } - lowerCT = "application/pdf" - } - if lowerCT == "text/html" || lowerCT == "application/xhtml+xml" || looksLikeHTML(data) { - data = htmlToReadableMarkdown(data) - if lowerCT == "" { - lowerCT = "text/html" - } - } - return filename, lowerCT, data -} - -func bytesLooksLikePDF(data []byte) bool { - return len(data) >= 4 && string(data[:4]) == "%PDF" -} - -func looksLikeHTML(data []byte) bool { - snippet := strings.ToLower(string(data)) - return strings.Contains(snippet, "]*>.*?`) - htmlTagRE = regexp.MustCompile(`(?s)<[^>]+>`) - multiSpaceRE = regexp.MustCompile(`[ \t]+`) - multiNewlineRE = regexp.MustCompile(`\n{3,}`) -) - -func htmlToReadableMarkdown(data []byte) []byte { - text := string(data) - text = htmlScriptStyleRE.ReplaceAllString(text, " ") - text = strings.ReplaceAll(text, "
", "\n") - text = strings.ReplaceAll(text, "
", "\n") - text = strings.ReplaceAll(text, "
", "\n") - text = strings.ReplaceAll(text, "

", "\n\n") - text = strings.ReplaceAll(text, "", "\n") - text = strings.ReplaceAll(text, "", "\n") - text = htmlTagRE.ReplaceAllString(text, " ") - text = html.UnescapeString(text) - text = strings.ReplaceAll(text, "\r", "\n") - text = multiSpaceRE.ReplaceAllString(text, " ") - text = multiNewlineRE.ReplaceAllString(text, "\n\n") - text = strings.TrimSpace(text) - return []byte(text) -} - -// reservedDeviceNames are Windows reserved filenames that must never be used. -var reservedDeviceNames = map[string]bool{ - "CON": true, "PRN": true, "AUX": true, "NUL": true, - "COM1": true, "COM2": true, "COM3": true, "COM4": true, "COM5": true, - "COM6": true, "COM7": true, "COM8": true, "COM9": true, - "LPT1": true, "LPT2": true, "LPT3": true, "LPT4": true, "LPT5": true, - "LPT6": true, "LPT7": true, "LPT8": true, "LPT9": true, -} - -// sanitizeFilename produces a safe, filesystem-friendly filename from an -// arbitrary URL path segment: it strips directory components, replaces unsafe / -// control characters, rejects reserved names, bounds the length, and falls back -// to "download". -func sanitizeFilename(name string) string { - name = filepath.Base(name) - name = strings.TrimSpace(name) - - name = strings.Map(func(r rune) rune { - switch r { - case '/', '\\', ':', '*', '?', '"', '<', '>', '|', 0: - return '_' - } - if r < 0x20 { // control characters - return '_' - } - return r - }, name) - - // Strip leading/trailing dots and spaces to avoid hidden or reserved forms. - name = strings.Trim(name, ". ") - - if name == "" || name == "." || name == ".." { - return "download" - } - if stem := strings.SplitN(strings.ToUpper(name), ".", 2)[0]; reservedDeviceNames[stem] { - return "download" - } - if len(name) > 255 { - name = name[:255] - } - return name -} diff --git a/internal/service/file/file.go b/internal/service/file/file.go new file mode 100644 index 0000000000..32c9cc8bc1 --- /dev/null +++ b/internal/service/file/file.go @@ -0,0 +1,114 @@ +// +// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// + +package file + +import ( + "ragflow/internal/dao" + "ragflow/internal/entity" + "ragflow/internal/utility" +) + +var ( + // assertURLSafe and pinnedHTTPClient are aliased from utility so tests + // can override them for mock injection. + assertURLSafe = utility.AssertURLSafe + pinnedHTTPClient = utility.PinnedHTTPClient +) + +// DocRemover is the narrow interface FileService needs from the document domain. +type DocRemover interface { + RemoveDocumentKeepFile(docID string) error +} + +// CheckFilePermFunc is the function signature for file-team permission checks, +// injected by the parent adapter so the file subpackage does not need to import +// the parent service package. +type CheckFilePermFunc func(fileDAO *dao.FileDAO, file *entity.File, userID string) bool + +// FileService file service +type FileService struct { + fileDAO *dao.FileDAO + file2DocumentDAO *dao.File2DocumentDAO + documentService DocRemover + checkFilePerm CheckFilePermFunc +} + +// NewFileService create file service. checkFilePerm is always required; +// dr may be nil when the caller only uses read/parse methods and never +// calls DeleteFiles. +func NewFileService(checkFilePerm CheckFilePermFunc, dr DocRemover) *FileService { + return &FileService{ + fileDAO: dao.NewFileDAO(), + file2DocumentDAO: dao.NewFile2DocumentDAO(), + documentService: dr, + checkFilePerm: checkFilePerm, + } +} + +// FileInfo file info with additional fields +type FileInfo struct { + *entity.File + Size int64 `json:"size"` + KbsInfo []map[string]interface{} `json:"kbs_info"` + HasChildFolder bool `json:"has_child_folder,omitempty"` +} + +// ListFilesResponse list files response +type ListFilesResponse struct { + Total int64 `json:"total"` + Files []map[string]interface{} `json:"files"` + ParentFolder map[string]interface{} `json:"parent_folder"` +} + +// DatasetFolderName is the folder name for dataset +const DatasetFolderName = ".knowledgebase" + +// SkillsFolderName is the folder name for skills +const SkillsFolderName = "skills" + +// FileSourceDataset represents dataset as file source +const FileSourceDataset = "knowledgebase" + +const ( + FileTypeFolder = "folder" + FileTypeVirtual = "virtual" +) + +// MoveFileReq represents the request body for move files operation +type MoveFileReq struct { + SrcFileIDs []string `json:"src_file_ids" binding:"required,min=1"` + DestFileID string `json:"dest_file_id"` + NewName string `json:"new_name"` +} + +// StorageAddress represents bucket and object name for storage +type StorageAddress struct { + Bucket string + Name string +} + +// maxRemoteFileSize bounds the body of a ?url= upload (100 MB). +const maxRemoteFileSize = 100 << 20 + +// reservedDeviceNames are Windows reserved filenames that must never be used. +var reservedDeviceNames = map[string]bool{ + "CON": true, "PRN": true, "AUX": true, "NUL": true, + "COM1": true, "COM2": true, "COM3": true, "COM4": true, "COM5": true, + "COM6": true, "COM7": true, "COM8": true, "COM9": true, + "LPT1": true, "LPT2": true, "LPT3": true, "LPT4": true, "LPT5": true, + "LPT6": true, "LPT7": true, "LPT8": true, "LPT9": true, +} diff --git a/internal/service/file_commit.go b/internal/service/file/file_commit.go similarity index 99% rename from internal/service/file_commit.go rename to internal/service/file/file_commit.go index d2f74c5558..14932b1d43 100644 --- a/internal/service/file_commit.go +++ b/internal/service/file/file_commit.go @@ -14,7 +14,7 @@ // limitations under the License. // -package service +package file import ( "crypto/sha256" diff --git a/internal/service/file/file_content.go b/internal/service/file/file_content.go new file mode 100644 index 0000000000..ccb4504cf1 --- /dev/null +++ b/internal/service/file/file_content.go @@ -0,0 +1,249 @@ +package file + +import ( + "encoding/base64" + "encoding/json" + "fmt" + "ragflow/internal/dao" + "ragflow/internal/entity" + "ragflow/internal/parser/parser" + "ragflow/internal/storage" + "ragflow/internal/utility" + "strings" +) + +// GetFileContent gets file metadata and checks permission for download +// Matches Python's file_api_service.get_file_content function +func (s *FileService) GetFileContent(uid, fileID string) (*entity.File, error) { + file, err := s.fileDAO.GetByID(fileID) + if err != nil || file == nil { + return nil, fmt.Errorf("Document not found!") + } + if !s.checkFilePerm(s.fileDAO, file, uid) { + return nil, fmt.Errorf("No authorization.") + } + return file, nil +} + +// GetStorageAddress gets storage address for a file (fallback for when direct blob is empty) +// Matches Python's File2DocumentService.get_storage_address function +func (s *FileService) GetStorageAddress(fileID string) (*StorageAddress, error) { + // Get file2document mapping + f2d, err := s.file2DocumentDAO.GetByFileID(fileID) + if err != nil || len(f2d) == 0 { + return nil, fmt.Errorf("file2document mapping not found") + } + + // Get the file + if f2d[0].FileID == nil { + return nil, fmt.Errorf("file_id is nil in file2document mapping") + } + file, err := s.fileDAO.GetByID(*f2d[0].FileID) + if err != nil || file == nil { + return nil, fmt.Errorf("file not found") + } + + // If source_type is empty or local, return file's parent_id and location + if file.SourceType == "" || entity.FileSource(file.SourceType) == entity.FileSourceLocal { + if file.Location == nil || *file.Location == "" { + return nil, fmt.Errorf("file location is empty") + } + return &StorageAddress{ + Bucket: file.ParentID, + Name: *file.Location, + }, nil + } + + // Otherwise, use document's kb_id and location + if f2d[0].DocumentID == nil { + return nil, fmt.Errorf("document_id is required") + } + + documentDAO := dao.NewDocumentDAO() + doc, err := documentDAO.GetByID(*f2d[0].DocumentID) + if err != nil || doc == nil { + return nil, fmt.Errorf("document not found") + } + + if doc.Location == nil || *doc.Location == "" { + return nil, fmt.Errorf("document location is empty") + } + + return &StorageAddress{ + Bucket: doc.KbID, + Name: *doc.Location, + }, nil +} + +// DownloadAgentFile downloads an agent-generated file directly from MinIO without querying the database. +func (s *FileService) DownloadAgentFile(tenantID, location string) ([]byte, error) { + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return nil, fmt.Errorf("storage not initialized") + } + + bucketName := fmt.Sprintf("%s-downloads", tenantID) + + blob, err := storageImpl.Get(bucketName, location) + if err != nil { + return nil, fmt.Errorf("failed to read file from storage: %w", err) + } + + return blob, nil +} + +// GetFileContents fetches file contents (text + image) from storage +// for the given file dicts. +// - raw=false: images returned as base64 data URIs in images; non-images parsed and returned as text. +// - raw=true: images returned as raw bytes in images; non-images parsed and returned as text. +func (s *FileService) GetFileContents(uid string, fileDicts []map[string]interface{}, raw bool) (texts []string, images []string, err error) { + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return nil, nil, fmt.Errorf("storage not initialized") + } + + for _, fd := range fileDicts { + id, _ := fd["id"].(string) + if id == "" { + continue + } + file, ferr := s.fileDAO.GetByID(id) + if ferr != nil || file == nil || file.Location == nil || *file.Location == "" { + continue + } + if !s.checkFilePerm(s.fileDAO, file, uid) { + return nil, nil, fmt.Errorf("No authorization.") + } + data, derr := storageImpl.Get(file.ParentID, *file.Location) + if derr != nil || len(data) == 0 { + continue + } + ft := utility.FilenameType(file.Name) + if ft == utility.FileTypeVISUAL { + if raw { + images = append(images, string(data)) + } else { + ext := utility.GetFileExtension(file.Name) + mime := utility.GetContentType(ext, string(ft)) + images = append(images, "data:"+mime+";base64,"+base64.StdEncoding.EncodeToString(data)) + } + } else { + texts = append(texts, parseFileContent(file.Name, data)) + } + } + return texts, images, nil +} + +// parseAgentUploads resolves descriptors returned by upload_info from the +// caller's downloads bucket and converts them to sys.files values. +func (s *FileService) ParseAgentUploads(userID string, fileDicts []map[string]interface{}, layoutRecognize string) ([]string, error) { + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return nil, fmt.Errorf("storage not initialized") + } + + contents := make([]string, 0, len(fileDicts)) + for i, fd := range fileDicts { + id, _ := fd["id"].(string) + name, _ := fd["name"].(string) + mimeType, _ := fd["mime_type"].(string) + createdBy, _ := fd["created_by"].(string) + if id == "" || name == "" || mimeType == "" || createdBy == "" { + return nil, fmt.Errorf("file %d: id, name, mime_type, and created_by are required", i) + } + if createdBy != userID { + return nil, fmt.Errorf("file %q: created_by does not match the current user", name) + } + + data, err := storageImpl.Get(createdBy+"-downloads", id) + if err != nil { + return nil, fmt.Errorf("file %q: read upload: %w", name, err) + } + if len(data) == 0 { + return nil, fmt.Errorf("file %q: upload is empty", name) + } + + mediaType := strings.ToLower(strings.TrimSpace(strings.Split(mimeType, ";")[0])) + if strings.HasPrefix(mediaType, "image/") { + contents = append(contents, "data:"+mediaType+";base64,"+base64.StdEncoding.EncodeToString(data)) + continue + } + + content, err := parseAgentUploadContent(name, data, layoutRecognize) + if err != nil { + return nil, fmt.Errorf("file %q: parse upload: %w", name, err) + } + contents = append(contents, content) + } + return contents, nil +} + +func parseAgentUploadContent(filename string, data []byte, layoutRecognize string) (string, error) { + content := string(data) + fileType := utility.GetFileType(filename) + if fileType != utility.FileTypeOTHER { + fp, err := parser.GetParser(fileType) + if err != nil { + return "", err + } + if configurable, ok := fp.(interface{ ConfigureFromSetup(map[string]any) }); ok { + configurable.ConfigureFromSetup(map[string]any{"layout_recognize": layoutRecognize}) + } + res := fp.ParseWithResult(filename, data) + if res.Err != nil { + return "", res.Err + } + switch res.OutputFormat { + case "text": + content = res.Text + case "markdown": + content = res.Markdown + case "html": + content = res.HTML + case "json": + parts := make([]string, 0, len(res.JSON)) + for _, item := range res.JSON { + if text, ok := item["text"].(string); ok { + parts = append(parts, text) + continue + } + raw, err := json.Marshal(item) + if err != nil { + return "", err + } + parts = append(parts, string(raw)) + } + content = strings.Join(parts, "\n") + } + } + return fmt.Sprintf("\n -----------------\nFile: %s\nContent as following: \n%s", filename, content), nil +} + +// parseFileContent tries to parse a file's contents using the appropriate parser. +// Falls back to returning raw text if no parser is available. +func parseFileContent(filename string, data []byte) string { + fileType := utility.GetFileType(filename) + if fileType == utility.FileTypeOTHER { + return string(data) + } + fp, err := parser.GetParser(fileType) + if err != nil { + return string(data) + } + res := fp.ParseWithResult(filename, data) + if res.Err != nil { + return string(data) + } + switch res.OutputFormat { + case "text": + return res.Text + case "markdown": + return res.Markdown + case "html": + return res.HTML + case "json": + return string(data) + default: + return string(data) + } +} diff --git a/internal/service/file/file_content_test.go b/internal/service/file/file_content_test.go new file mode 100644 index 0000000000..e50eef9cc2 --- /dev/null +++ b/internal/service/file/file_content_test.go @@ -0,0 +1,60 @@ +package file + +import ( + "ragflow/internal/utility" + "testing" +) + +func TestBytesLooksLikePDF(t *testing.T) { + cases := []struct { + name string + data []byte + want bool + }{ + {"valid header", []byte("%PDF-1.4 content"), true}, + {"too short", []byte("%PD"), false}, + {"plain text", []byte("hello"), false}, + {"nil", nil, false}, + } + for _, c := range cases { + if got := utility.BytesLooksLikePDF(c.data); got != c.want { + t.Errorf("%s: utility.BytesLooksLikePDF = %v, want %v", c.name, got, c.want) + } + } +} + +func TestLooksLikeHTML(t *testing.T) { + if !utility.LooksLikeHTML([]byte("x")) { + t.Error("should detect ") + } + if !utility.LooksLikeHTML([]byte("
x
")) { + t.Error("should detect
case-insensitively") + } + if !utility.LooksLikeHTML([]byte("x")) { + t.Error("should detect ") + } + if utility.LooksLikeHTML([]byte("just text")) { + t.Error("should not detect plain text") + } +} + +func TestParseFileContent_HTMLOutputFormat(t *testing.T) { + // CSV parser produces OutputFormat "html". + result := parseFileContent("data.csv", []byte("a,b,c\n1,2,3\n")) + if result == "" || result == string([]byte("a,b,c\n1,2,3\n")) { + t.Skip("CSV parser not available or returned raw text; integration-only test") + } + // CSV parser emits an HTML table; must not contain raw CSV comma-separated rows. + if result == "a,b,c\n1,2,3\n" { + t.Errorf("CSV should produce HTML output, got raw CSV: %q", result) + } +} + +func TestParseFileContent_MarkdownOutputFormat(t *testing.T) { + // .md files fall through to default parser or produce text/markdown. + result := parseFileContent("doc.md", []byte("# Title\n\nBody")) + // This is integration-dependent; verify it doesn't crash and returns something. + if result == "" { + t.Skip("markdown parser returned empty (integration-only)") + } +} diff --git a/internal/service/file/file_delete.go b/internal/service/file/file_delete.go new file mode 100644 index 0000000000..5a609a351a --- /dev/null +++ b/internal/service/file/file_delete.go @@ -0,0 +1,130 @@ +package file + +import ( + "context" + "fmt" + "ragflow/internal/common" + "ragflow/internal/entity" + "ragflow/internal/storage" +) + +// DeleteFiles deletes files by IDs +// Returns (success, message) where success is true if all files were deleted +func (s *FileService) DeleteFiles(ctx context.Context, uid string, fileIDs []string) (bool, string) { + for _, fileID := range fileIDs { + // 1. Get file + file, err := s.fileDAO.GetByID(fileID) + if err != nil || file == nil { + return false, "File or Folder not found!" + } + + // 2. Check tenant_id + if file.TenantID == "" { + return false, "Tenant not found!" + } + + // Block root-folder deletion (root folders have parent_id == id) + if file.ParentID == file.ID { + return false, "Root folder cannot be deleted." + } + + // 3. Permission check + if !s.checkFilePerm(s.fileDAO, file, uid) { + return false, "No authorization." + } + + // 4. Skip dataset source files + if file.SourceType == FileSourceDataset { + continue + } + + // 5. Delete based on type + if file.Type == FileTypeFolder { + if err := s.deleteFolderRecursive(ctx, file, uid); err != nil { + return false, fmt.Sprintf("Failed to delete folder: %v", err) + } + } else { + if err := s.deleteSingleFile(ctx, file); err != nil { + return false, fmt.Sprintf("Failed to delete file: %v", err) + } + } + } + + return true, "" +} + +// deleteSingleFile deletes a single file (not folder) +// Matches Python's _delete_single_file function +func (s *FileService) deleteSingleFile(ctx context.Context, file *entity.File) error { + // 1. Delete storage object + if file.Location != nil && *file.Location != "" { + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl != nil { + if err := storageImpl.Remove(file.ParentID, *file.Location); err != nil { + common.Logger.Error(fmt.Sprintf("Fail to remove object: %s/%s, error: %v", file.ParentID, *file.Location, err)) + } + } + } + + // 2. Handle associated documents + informs, err := s.file2DocumentDAO.GetByFileID(file.ID) + if err != nil { + return fmt.Errorf("failed to get file2document mappings: %w", err) + } + if len(informs) > 0 { + for _, inform := range informs { + if inform.DocumentID == nil { + continue + } + docID := *inform.DocumentID + if s.documentService != nil { + if err := s.documentService.RemoveDocumentKeepFile(docID); err != nil { + common.Logger.Error(fmt.Sprintf("Fail to remove document: %s, error: %v", docID, err)) + } + } + } + + // Delete file2document mapping (outside the loop, called once - matching Python behavior) + if err := s.file2DocumentDAO.DeleteByFileID(file.ID); err != nil { + return fmt.Errorf("failed to delete file2document mapping: %w", err) + } + } + + // 3. Delete file record + if err := s.fileDAO.Delete(file.ID); err != nil { + return err + } + + return nil +} + +// deleteFolderRecursive recursively deletes a folder and its contents +// Matches Python's _delete_folder_recursive function +func (s *FileService) deleteFolderRecursive(ctx context.Context, folder *entity.File, uid string) error { + // Get all sub-files + subFiles, err := s.fileDAO.ListByParentID(folder.ID) + if err != nil { + return err + } + + for _, subFile := range subFiles { + if subFile.Type == FileTypeFolder { + // Recursively delete subfolder + if err := s.deleteFolderRecursive(ctx, subFile, uid); err != nil { + return err + } + } else { + // Delete single file + if err := s.deleteSingleFile(ctx, subFile); err != nil { + return err + } + } + } + + // Delete the folder itself + if err := s.fileDAO.Delete(folder.ID); err != nil { + return err + } + + return nil +} diff --git a/internal/service/file/file_folder.go b/internal/service/file/file_folder.go new file mode 100644 index 0000000000..4f2c0f638f --- /dev/null +++ b/internal/service/file/file_folder.go @@ -0,0 +1,576 @@ +package file + +import ( + "fmt" + "path/filepath" + "ragflow/internal/common" + "ragflow/internal/dao" + "ragflow/internal/entity" + "ragflow/internal/storage" + "ragflow/internal/utility" + "strings" +) + +// GetRootFolder gets or creates root folder for tenant +func (s *FileService) GetRootFolder(tenantID string) (map[string]interface{}, error) { + file, err := s.fileDAO.GetRootFolder(tenantID) + if err != nil { + return nil, err + } + return s.toFileResponse(file), nil +} + +// ListFiles lists files by parent folder ID (matching Python /files endpoint) +// This method includes init_dataset_docs initialization when parent_id is empty +func (s *FileService) ListFiles(tenantID, pfID string, page, pageSize int, orderby string, desc bool, keywords string) (*ListFilesResponse, error) { + // If pfID is empty, get root folder and initialize dataset docs + if pfID == "" { + rootFolder, err := s.fileDAO.GetRootFolder(tenantID) + if err != nil { + return nil, fmt.Errorf("failed to get root folder: %w", err) + } + pfID = rootFolder.ID + + // Initialize dataset docs (matching Python init_knowledgebase_docs logic) + if err := s.initDatasetDocs(pfID, tenantID); err != nil { + return nil, fmt.Errorf("failed to initialize dataset docs: %w", err) + } + + // Initialize skills folder (matching Python init_skills_folder logic) + if err := s.initSkillsFolder(pfID, tenantID); err != nil { + return nil, fmt.Errorf("failed to initialize skills folder: %w", err) + } + } + + // Check if parent folder exists + if _, err := s.fileDAO.GetByID(pfID); err != nil { + return nil, fmt.Errorf("Folder not found!") + } + + // Get files by parent folder ID + files, total, err := s.fileDAO.GetByPfID(tenantID, pfID, page, pageSize, orderby, desc, keywords) + if err != nil { + return nil, err + } + + // Get parent folder + parentFolder, err := s.fileDAO.GetParentFolder(pfID) + if err != nil { + return nil, fmt.Errorf("File not found!") + } + + // Process files to add additional info, deduplicating by ID as a safety net + // against any leftover duplicate rows (e.g. duplicate 'skills' or '.knowledgebase' folders). + fileResponses := make([]map[string]interface{}, 0, len(files)) + seenIDs := make(map[string]struct{}) + for _, file := range files { + if _, ok := seenIDs[file.ID]; ok { + continue + } + seenIDs[file.ID] = struct{}{} + fileInfo := s.toFileInfo(file) + + // If folder, calculate size and check for child folders + if file.Type == FileTypeFolder { + folderSize, err := s.fileDAO.GetFolderSize(file.ID) + if err == nil { + fileInfo.Size = folderSize + } + hasChild, err := s.fileDAO.HasChildFolder(file.ID) + if err == nil { + fileInfo.HasChildFolder = hasChild + } + fileInfo.KbsInfo = []map[string]interface{}{} + } else { + // Get KB info for non-folder files + kbsInfo, err := s.file2DocumentDAO.GetKBInfoByFileID(file.ID) + if err != nil { + kbsInfo = []map[string]interface{}{} + } + fileInfo.KbsInfo = kbsInfo + } + + fileResponses = append(fileResponses, s.fileInfoToResponse(fileInfo)) + } + + return &ListFilesResponse{ + Total: total, + Files: fileResponses, + ParentFolder: s.toFileResponse(parentFolder), + }, nil +} + +// initDatasetDocs initializes dataset documents for tenant +// This matches Python's FileService.init_dataset_docs method +func (s *FileService) initDatasetDocs(rootID, tenantID string) error { + return s.fileDAO.InitDatasetDocs(rootID, tenantID, s.file2DocumentDAO) +} + +// initSkillsFolder initializes the skills folder under the root folder. +// Deduplicates duplicate entries that may have been created by +// concurrent race conditions (TOCTOU). +func (s *FileService) initSkillsFolder(rootID, tenantID string) error { + existing := s.fileDAO.Query(SkillsFolderName, rootID, tenantID) + if len(existing) > 0 { + if len(existing) > 1 { + common.Logger.Warn(fmt.Sprintf( + "Found %d duplicate '%s' folders under root %s, keeping only the first", + len(existing), SkillsFolderName, rootID, + )) + keepID := existing[0].ID + for _, dup := range existing[1:] { + children, _ := s.fileDAO.ListAllFilesByParentID(dup.ID) + for _, child := range children { + s.fileDAO.UpdateByID(child.ID, map[string]interface{}{"parent_id": keepID}) + } + if delErr := s.fileDAO.Delete(dup.ID); delErr != nil { + common.Logger.Warn(fmt.Sprintf("Failed to delete duplicate skills folder %s: %v", dup.ID, delErr)) + } + } + } + return nil + } + + folder := &entity.File{ + ID: utility.GenerateToken(), + ParentID: rootID, + TenantID: tenantID, + CreatedBy: tenantID, + Name: SkillsFolderName, + Type: FileTypeFolder, + Size: 0, + SourceType: "", + } + return s.fileDAO.Insert(folder) +} + +// toFileResponse converts file model to response format +func (s *FileService) toFileResponse(file *entity.File) map[string]interface{} { + result := map[string]interface{}{ + "id": file.ID, + "parent_id": file.ParentID, + "tenant_id": file.TenantID, + "created_by": file.CreatedBy, + "name": file.Name, + "size": file.Size, + "type": file.Type, + "create_time": file.CreateTime, + "update_time": file.UpdateTime, + } + + if file.Location != nil { + result["location"] = *file.Location + } + result["source_type"] = file.SourceType + + return result +} + +// toFileInfo converts file model to FileInfo +func (s *FileService) toFileInfo(file *entity.File) *FileInfo { + return &FileInfo{ + File: file, + Size: file.Size, + KbsInfo: []map[string]interface{}{}, + HasChildFolder: false, + } +} + +// fileInfoToResponse converts FileInfo to response map +func (s *FileService) fileInfoToResponse(info *FileInfo) map[string]interface{} { + result := map[string]interface{}{ + "id": info.File.ID, + "parent_id": info.File.ParentID, + "tenant_id": info.File.TenantID, + "created_by": info.File.CreatedBy, + "name": info.File.Name, + "size": info.Size, + "type": info.File.Type, + "create_time": info.File.CreateTime, + "update_time": info.File.UpdateTime, + "kbs_info": info.KbsInfo, + } + + if info.File.Location != nil { + result["location"] = *info.File.Location + } + result["source_type"] = info.File.SourceType + + if info.File.Type == "folder" { + result["has_child_folder"] = info.HasChildFolder + } + + return result +} + +// GetParentFolder gets parent folder of a file with permission check +func (s *FileService) GetParentFolder(userID, fileID string) (map[string]interface{}, error) { + // Get file + file, err := s.fileDAO.GetByID(fileID) + if err != nil { + return nil, err + } + + // Permission check + if !s.checkFilePerm(s.fileDAO, file, userID) { + return nil, fmt.Errorf("No authorization.") + } + + // Get parent folder + parentFolder, err := s.fileDAO.GetParentFolder(fileID) + if err != nil { + return nil, err + } + + return s.toFileResponse(parentFolder), nil +} + +// GetAllParentFolders gets all parent folders in path with permission check +func (s *FileService) GetAllParentFolders(userID, fileID string) ([]map[string]interface{}, error) { + // Get file + file, err := s.fileDAO.GetByID(fileID) + if err != nil { + return nil, err + } + + // Permission check + if !s.checkFilePerm(s.fileDAO, file, userID) { + return nil, fmt.Errorf("No authorization.") + } + + // Get all parent folders + parentFolders, err := s.fileDAO.GetAllParentFolders(fileID) + if err != nil { + return nil, err + } + + // Convert to response format + result := make([]map[string]interface{}, len(parentFolders)) + for i, folder := range parentFolders { + result[i] = s.toFileResponse(folder) + } + + return result, nil +} + +// GetDocCount gets document count for a tenant +func (s *FileService) GetDocCount(tenantID string) (int64, error) { + documentDAO := dao.NewDocumentDAO() + return documentDAO.CountByTenantID(tenantID) +} + +func (s *FileService) createFolderRecursive(parentFolder *entity.File, names []string, count int, tenantID string) (*entity.File, error) { + if count > len(names)-2 { + return parentFolder, nil + } + + newFolder, err := s.fileDAO.CreateFolder(parentFolder.ID, tenantID, names[count], FileTypeFolder) + if err != nil { + return nil, err + } + + return s.createFolderRecursive(newFolder, names, count+1, tenantID) +} + +func (s *FileService) getUniqueFilename(name, parentID, tenantID string) string { + existingFiles := s.fileDAO.Query(name, parentID, tenantID) + if len(existingFiles) == 0 { + return name + } + + base := filepath.Base(name) + ext := filepath.Ext(name) + nameWithoutExt := strings.TrimSuffix(base, ext) + + counter := 1 + for { + newName := fmt.Sprintf("%s_%d%s", nameWithoutExt, counter, ext) + existingFiles = s.fileDAO.Query(newName, parentID, tenantID) + if len(existingFiles) == 0 { + return newName + } + counter++ + } +} + +// CreateFolder creates a new folder or virtual file +func (s *FileService) CreateFolder(tenantID, name, parentID, fileType string) (map[string]interface{}, error) { + if parentID == "" { + rootFolder, err := s.fileDAO.GetRootFolder(tenantID) + if err != nil { + return nil, fmt.Errorf("failed to get root folder: %w", err) + } + parentID = rootFolder.ID + } + + if !s.fileDAO.IsParentFolderExist(parentID) { + return nil, fmt.Errorf("Parent Folder Doesn't Exist!") + } + + existingFiles := s.fileDAO.Query(name, parentID, tenantID) + if len(existingFiles) > 0 { + return nil, fmt.Errorf("Duplicated folder name in the same folder.") + } + + if fileType == "" { + fileType = FileTypeVirtual + } + + if fileType == FileTypeFolder { + fileType = FileTypeFolder + } else { + fileType = FileTypeVirtual + } + + folder, err := s.fileDAO.CreateFolder(parentID, tenantID, name, fileType) + if err != nil { + return nil, fmt.Errorf("failed to create folder: %w", err) + } + + return s.toFileResponse(folder), nil +} + +// MoveFiles moves and/or renames files +// Follows Linux mv semantics: +// - new_name only: rename in place (no storage operation) +// - dest_file_id only: move to new folder (keep names) +// - both: move and rename simultaneously +func (s *FileService) MoveFiles(uid string, srcFileIDs []string, destFileID string, newName string) (bool, string) { + // 1. Get all source files + files, err := s.fileDAO.GetByIDs(srcFileIDs) + if err != nil || len(files) == 0 { + return false, "Source files not found!" + } + + // Create a map for quick lookup + filesMap := make(map[string]*entity.File) + for _, f := range files { + filesMap[f.ID] = f + } + + // 2. Validate all source files + for _, fileID := range srcFileIDs { + file, ok := filesMap[fileID] + if !ok { + return false, "File or folder not found!" + } + if file.TenantID == "" { + return false, "Tenant not found!" + } + // 3. Permission check + if !s.checkFilePerm(s.fileDAO, file, uid) { + return false, "No authorization." + } + } + + // 4. Validate destination folder if provided + var destFolder *entity.File + if destFileID != "" { + destFolder, err = s.fileDAO.GetByID(destFileID) + if err != nil || destFolder == nil { + return false, "Parent folder not found!" + } + // Check destination folder permission + if !s.checkFilePerm(s.fileDAO, destFolder, uid) { + return false, "No authorization to write to destination folder." + } + + if destFolder.Type != FileTypeFolder { + return false, "Destination is not a folder." + } + + destAncestors, err := s.fileDAO.GetAllParentFolders(destFolder.ID) + if err != nil { + return false, "Parent folder not found!" + } + + destAncestorIDs := make(map[string]struct{}, len(destAncestors)) + for _, folder := range destAncestors { + destAncestorIDs[folder.ID] = struct{}{} + } + + for _, file := range files { + if file.Type != FileTypeFolder { + continue + } + + if file.ID == destFolder.ID { + return false, "Cannot move a folder to itself." + } + + if _, ok := destAncestorIDs[file.ID]; ok { + return false, "Cannot move a folder into its own subfolder." + } + } + } + + // 5. Validate new_name if provided + if newName != "" { + if len(srcFileIDs) > 1 { + return false, "new_name can only be used with a single file" + } + + file := filesMap[srcFileIDs[0]] + // Check extension for non-folder files + if file.Type != FileTypeFolder { + oldExt := utility.GetFileExtension(file.Name) + newExt := utility.GetFileExtension(newName) + if oldExt != newExt { + return false, "The extension of file can't be changed" + } + } + + // Check for duplicate names in target folder + targetParentID := file.ParentID + if destFolder != nil { + targetParentID = destFolder.ID + } + existingFiles := s.fileDAO.Query(newName, targetParentID, file.TenantID) + for _, f := range existingFiles { + if f.Name == newName { + return false, "Duplicated file name in the same folder." + } + } + } else if destFolder != nil { + // Plain move (no rename): check for duplicate names in destination folder + for _, file := range files { + existingFiles := s.fileDAO.Query(file.Name, destFolder.ID, file.TenantID) + for _, f := range existingFiles { + // Ignore the source file itself + if f.ID != file.ID { + return false, "Duplicated file name in the same folder." + } + } + } + } + + // 6. Perform the move operation + if destFolder != nil { + // Move to destination folder + for _, file := range files { + if err := s.moveEntryRecursive(file, destFolder, newName); err != nil { + return false, err.Error() + } + } + } else { + // Pure rename: no storage operation needed + if newName == "" { + return false, "new_name is required for rename" + } + if len(srcFileIDs) == 0 { + return false, "Source files not found!" + } + file := filesMap[srcFileIDs[0]] + if err := s.fileDAO.UpdateByID(file.ID, map[string]interface{}{"name": newName}); err != nil { + return false, "Database error (File rename)!" + } + + // Update associated document name if exists + informs, err := s.file2DocumentDAO.GetByFileID(file.ID) + if err == nil && len(informs) > 0 && informs[0].DocumentID != nil { + docID := *informs[0].DocumentID + documentDAO := dao.NewDocumentDAO() + if err := documentDAO.UpdateByID(docID, map[string]interface{}{"name": newName}); err != nil { + return false, "Database error (Document rename)!" + } + } + } + + return true, "" +} + +// moveEntryRecursive recursively moves a file or folder entry +func (s *FileService) moveEntryRecursive(sourceFile *entity.File, destFolder *entity.File, overrideName string) error { + effectiveName := overrideName + if effectiveName == "" { + effectiveName = sourceFile.Name + } + + if sourceFile.Type == FileTypeFolder { + // Handle folder move + existingFolders := s.fileDAO.Query(effectiveName, destFolder.ID, sourceFile.TenantID) + var newFolder *entity.File + if len(existingFolders) > 0 { + // Prevent moving a folder into itself (self-target merge) + if existingFolders[0].ID == sourceFile.ID { + return fmt.Errorf("cannot move folder into itself") + } + newFolder = existingFolders[0] + } else { + // Create new folder + var err error + newFolder, err = s.fileDAO.CreateFolder(destFolder.ID, sourceFile.TenantID, effectiveName, FileTypeFolder) + if err != nil { + return fmt.Errorf("failed to create destination folder: %w", err) + } + } + + // Recursively move sub-files + subFiles, err := s.fileDAO.ListAllFilesByParentID(sourceFile.ID) + if err != nil { + return err + } + for _, subFile := range subFiles { + if err := s.moveEntryRecursive(subFile, newFolder, ""); err != nil { + return err + } + } + + // Delete the source folder + return s.fileDAO.Delete(sourceFile.ID) + } + + // Handle non-folder file move + needStorageMove := destFolder.ID != sourceFile.ParentID + updates := map[string]interface{}{} + + if needStorageMove { + // Get storage + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return fmt.Errorf("storage not initialized") + } + + // Calculate new location + newLocation := effectiveName + for storageImpl.ObjExist(destFolder.ID, newLocation) { + newLocation += "_" + } + + // Perform storage move (copy + delete) + if sourceFile.Location == nil || *sourceFile.Location == "" { + return fmt.Errorf("file location is empty") + } + + if !storageImpl.Move(sourceFile.ParentID, *sourceFile.Location, destFolder.ID, newLocation) { + return fmt.Errorf("move file failed at storage layer") + } + + updates["parent_id"] = destFolder.ID + updates["location"] = newLocation + } + + if overrideName != "" { + updates["name"] = overrideName + } + + if len(updates) > 0 { + if err := s.fileDAO.UpdateByID(sourceFile.ID, updates); err != nil { + return fmt.Errorf("database error (File update): %w", err) + } + } + + // Update associated document name if renamed + if overrideName != "" { + informs, err := s.file2DocumentDAO.GetByFileID(sourceFile.ID) + if err == nil && len(informs) > 0 && informs[0].DocumentID != nil { + docID := *informs[0].DocumentID + documentDAO := dao.NewDocumentDAO() + if err := documentDAO.UpdateByID(docID, map[string]interface{}{"name": overrideName}); err != nil { + return fmt.Errorf("database error (Document rename): %w", err) + } + } + } + + return nil +} diff --git a/internal/service/file/file_permission_test.go b/internal/service/file/file_permission_test.go new file mode 100644 index 0000000000..14e48266c2 --- /dev/null +++ b/internal/service/file/file_permission_test.go @@ -0,0 +1,51 @@ +package file + +import ( + "testing" + + "github.com/glebarez/sqlite" + "gorm.io/gorm" + + "ragflow/internal/dao" + "ragflow/internal/entity" +) + +func setupPermissionDB(t *testing.T) { + t.Helper() + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{TranslateError: true}) + if err != nil { + t.Fatalf("open sqlite: %v", err) + } + if err := db.AutoMigrate( + &entity.File{}, + &entity.File2Document{}, + &entity.Document{}, + &entity.Knowledgebase{}, + ); err != nil { + t.Fatalf("migrate: %v", err) + } + old := dao.DB + dao.DB = db + t.Cleanup(func() { dao.DB = old }) +} + +// checkFileTeamPermissionStub is a minimal permission checker for tests +// that avoids importing the parent service package (which would create a +// cycle in the test binary). +func checkFileTeamPermissionStub(_ *dao.FileDAO, file *entity.File, userID string) bool { + return file.TenantID == userID +} + +func TestCheckFileTeamPermissionStub(t *testing.T) { + setupPermissionDB(t) + + // Direct tenant match short-circuits before any DB lookup. + if !checkFileTeamPermissionStub(nil, &entity.File{TenantID: "u1"}, "u1") { + t.Error("file tenant match should be authorized") + } + + // A file owned by another tenant is denied for a different user. + if checkFileTeamPermissionStub(nil, &entity.File{TenantID: "other"}, "u1") { + t.Error("file with different tenant should be denied") + } +} diff --git a/internal/service/file_test.go b/internal/service/file/file_test.go similarity index 76% rename from internal/service/file_test.go rename to internal/service/file/file_test.go index 58799548a7..06e7257f59 100644 --- a/internal/service/file_test.go +++ b/internal/service/file/file_test.go @@ -1,8 +1,8 @@ -package service +package file import ( "bytes" - "context" + "errors" "io" "net/http" @@ -17,6 +17,7 @@ import ( "ragflow/internal/dao" "ragflow/internal/entity" "ragflow/internal/storage" + "ragflow/internal/utility" ) // fakeStorage mocks storage.Storage for testing DownloadAgentFile. @@ -29,11 +30,20 @@ type fakeStorage struct { getCalls int } +// sptr returns a pointer to the given string. +func sptr(s string) *string { return &s } + +// testFilePerm controls the permission check returned by testFileService. +// Tests that need to simulate denied access can set it to a function that +// returns false. +var testFilePerm CheckFilePermFunc = func(_ *dao.FileDAO, _ *entity.File, _ string) bool { return true } + func testFileService() *FileService { return &FileService{ fileDAO: dao.NewFileDAO(), file2DocumentDAO: dao.NewFile2DocumentDAO(), - documentService: &DocumentService{}, + documentService: nil, + checkFilePerm: testFilePerm, } } @@ -179,6 +189,10 @@ func setupFileContentPermissionDB(t *testing.T, accessible bool) { func TestFileService_GetFileContents_NotAccessible(t *testing.T) { setupFileContentPermissionDB(t, false) + orig := testFilePerm + testFilePerm = func(_ *dao.FileDAO, _ *entity.File, _ string) bool { return false } + t.Cleanup(func() { testFilePerm = orig }) + mockStorage := &fakeStorage{blob: []byte("secret")} factory := storage.GetStorageFactory() originalStorage := factory.GetStorage() @@ -242,7 +256,7 @@ func TestFileService_ParseAgentUploads_TextAndImageInRequestOrder(t *testing.T) factory.SetStorage(memory) t.Cleanup(func() { factory.SetStorage(originalStorage) }) - contents, err := testFileService().parseAgentUploads("user-1", []map[string]interface{}{ + contents, err := testFileService().ParseAgentUploads("user-1", []map[string]interface{}{ {"id": "text-id", "name": "notes.txt", "mime_type": "text/plain", "created_by": "user-1"}, {"id": "image-id", "name": "photo.bin", "mime_type": "image/png", "created_by": "user-1"}, }, "Plain Text") @@ -267,7 +281,7 @@ func TestFileService_ParseAgentUploads_RejectsForeignOwner(t *testing.T) { factory.SetStorage(memory) t.Cleanup(func() { factory.SetStorage(originalStorage) }) - _, err := testFileService().parseAgentUploads("user-1", []map[string]interface{}{ + _, err := testFileService().ParseAgentUploads("user-1", []map[string]interface{}{ {"id": "file-id", "name": "secret.txt", "mime_type": "text/plain", "created_by": "user-2"}, }, "") if err == nil || !strings.Contains(err.Error(), "created_by does not match") { @@ -282,7 +296,7 @@ func TestFileService_ParseAgentUploads_MissingObjectFails(t *testing.T) { factory.SetStorage(memory) t.Cleanup(func() { factory.SetStorage(originalStorage) }) - _, err := testFileService().parseAgentUploads("user-1", []map[string]interface{}{ + _, err := testFileService().ParseAgentUploads("user-1", []map[string]interface{}{ {"id": "missing", "name": "missing.txt", "mime_type": "text/plain", "created_by": "user-1"}, }, "") if err == nil || !strings.Contains(err.Error(), "read upload") { @@ -363,17 +377,17 @@ func TestFileService_UploadFromURL_PDFAddsExtensionAndStoresToDownloads(t *testi })) defer server.Close() - origAssert := assertURLSafe - origPinned := pinnedHTTPClient - assertURLSafe = func(rawURL string) (string, string, error) { + origAssert := utility.AssertURLSafe + origPinned := utility.PinnedHTTPClient + utility.AssertURLSafe = func(rawURL string) (string, string, error) { return "127.0.0.1", "127.0.0.1", nil } - pinnedHTTPClient = func(hostname, resolvedIP string, timeout time.Duration) *http.Client { + utility.PinnedHTTPClient = func(hostname, resolvedIP string, timeout time.Duration) *http.Client { return server.Client() } t.Cleanup(func() { - assertURLSafe = origAssert - pinnedHTTPClient = origPinned + utility.AssertURLSafe = origAssert + utility.PinnedHTTPClient = origPinned }) mockStorage := &fakeStorage{} @@ -409,17 +423,17 @@ func TestFileService_UploadFromURL_HTMLNormalizesReadableContent(t *testing.T) { })) defer server.Close() - origAssert := assertURLSafe - origPinned := pinnedHTTPClient - assertURLSafe = func(rawURL string) (string, string, error) { + origAssert := utility.AssertURLSafe + origPinned := utility.PinnedHTTPClient + utility.AssertURLSafe = func(rawURL string) (string, string, error) { return "127.0.0.1", "127.0.0.1", nil } - pinnedHTTPClient = func(hostname, resolvedIP string, timeout time.Duration) *http.Client { + utility.PinnedHTTPClient = func(hostname, resolvedIP string, timeout time.Duration) *http.Client { return server.Client() } t.Cleanup(func() { - assertURLSafe = origAssert - pinnedHTTPClient = origPinned + utility.AssertURLSafe = origAssert + utility.PinnedHTTPClient = origPinned }) mockStorage := &fakeStorage{} @@ -447,7 +461,7 @@ func TestFileService_UploadFromURL_HTMLNormalizesReadableContent(t *testing.T) { } func TestNormalizeUploadInfoContent_PDFTakesPrecedenceOverHTML(t *testing.T) { - filename, contentType, data := normalizeUploadInfoContent( + filename, contentType, data := utility.NormalizeUploadInfoContent( "report", "text/html", []byte("%PDF-1.7 fake pdf"), @@ -474,92 +488,7 @@ func TestReadUploadInfoData_RejectsOversizedInput(t *testing.T) { } } -func TestFileService_DeleteSingleFile_RemovesLinkedDocumentThroughDocumentService(t *testing.T) { - db := setupServiceTestDB(t) - pushServiceDB(t, db) - - mockStorage := newFakeUploadStorage() - factory := storage.GetStorageFactory() - originalStorage := factory.GetStorage() - factory.SetStorage(mockStorage) - t.Cleanup(func() { factory.SetStorage(originalStorage) }) - - insertTestKB(t, "kb-file-delete", "tenant-1", 1, 30, 10) - insertTestDoc(t, "doc-file-delete", "kb-file-delete", 30, 10) - insertTestTask(t, "task-file-delete", "doc-file-delete") - - location := "doc.pdf" - insertTestFile(t, "file-delete", "folder-delete", "doc.pdf", &location) - insertTestFile2Document(t, "f2d-file-delete", "file-delete", "doc-file-delete") - if err := mockStorage.Put("folder-delete", location, []byte("blob")); err != nil { - t.Fatalf("put test blob: %v", err) - } - - file, err := dao.NewFileDAO().GetByID("file-delete") - if err != nil { - t.Fatalf("get file: %v", err) - } - docSvc := testDocumentService(t) - docEngine := &rerunDeleteDocEngine{} - docSvc.docEngine = docEngine - svc := &FileService{ - fileDAO: dao.NewFileDAO(), - file2DocumentDAO: dao.NewFile2DocumentDAO(), - documentService: docSvc, - } - - if err := svc.deleteSingleFile(context.Background(), file); err != nil { - t.Fatalf("deleteSingleFile failed: %v", err) - } - - if mockStorage.ObjExist("folder-delete", location) { - t.Fatal("expected storage object to be removed") - } - if _, err := dao.NewDocumentDAO().GetByID("doc-file-delete"); err == nil { - t.Fatal("expected linked document to be deleted") - } - var taskCount int64 - if err := dao.DB.Model(&entity.Task{}).Where("doc_id = ?", "doc-file-delete").Count(&taskCount).Error; err != nil { - t.Fatalf("count tasks: %v", err) - } - if taskCount != 0 { - t.Fatalf("expected linked tasks to be deleted, got %d", taskCount) - } - mappings, err := dao.NewFile2DocumentDAO().GetByFileID("file-delete") - if err != nil { - t.Fatalf("get file2document mappings: %v", err) - } - if len(mappings) != 0 { - t.Fatalf("expected file2document mappings to be deleted, got %d", len(mappings)) - } - files, err := dao.NewFileDAO().GetByIDs([]string{"file-delete"}) - if err != nil { - t.Fatalf("get deleted file: %v", err) - } - if len(files) != 0 { - t.Fatalf("expected file record to be deleted, got %d", len(files)) - } - if docEngine.deleteCalls != 1 { - t.Fatalf("expected document engine cleanup to be called once, got %d", docEngine.deleteCalls) - } - if docEngine.indexName != "ragflow_tenant-1" { - t.Fatalf("expected document engine index ragflow_tenant-1, got %q", docEngine.indexName) - } - if docEngine.datasetID != "kb-file-delete" { - t.Fatalf("expected document engine dataset kb-file-delete, got %q", docEngine.datasetID) - } - if docEngine.condition["doc_id"] != "doc-file-delete" { - t.Fatalf("expected document engine doc_id condition doc-file-delete, got %v", docEngine.condition["doc_id"]) - } - kb, err := dao.NewKnowledgebaseDAO().GetByID("kb-file-delete") - if err != nil { - t.Fatalf("get kb: %v", err) - } - if kb.DocNum != 0 || kb.TokenNum != 0 || kb.ChunkNum != 0 { - t.Fatalf("expected KB counters to be decremented to zero, got doc=%d token=%d chunk=%d", kb.DocNum, kb.TokenNum, kb.ChunkNum) - } -} - +// zeroReader is an io.Reader that returns an infinite stream of zero bytes. type zeroReader struct{} func (zeroReader) Read(p []byte) (int, error) { diff --git a/internal/service/file/file_upload.go b/internal/service/file/file_upload.go new file mode 100644 index 0000000000..48f9118f89 --- /dev/null +++ b/internal/service/file/file_upload.go @@ -0,0 +1,281 @@ +package file + +import ( + "fmt" + "io" + "mime/multipart" + "net/http" + "ragflow/internal/common" + "ragflow/internal/entity" + "ragflow/internal/storage" + "ragflow/internal/utility" + "strings" + "time" +) + +// UploadFile uploads files to a folder +func (s *FileService) UploadFile(tenantID, parentID string, files []*multipart.FileHeader) ([]map[string]interface{}, error) { + if parentID == "" { + rootFolder, err := s.fileDAO.GetRootFolder(tenantID) + if err != nil { + return nil, fmt.Errorf("failed to get root folder: %w", err) + } + parentID = rootFolder.ID + } + + _, err := s.fileDAO.GetByID(parentID) + if err != nil { + return nil, fmt.Errorf("Can't find this folder!") + } + + maxFileNumPerUser := common.GetEnv(common.EnvMaxFileNumPerUser) + if maxFileNumPerUser != "" { + var maxNum int64 + if _, err = fmt.Sscanf(maxFileNumPerUser, "%d", &maxNum); err == nil && maxNum > 0 { + var docCount int64 + docCount, err = s.GetDocCount(tenantID) + if err != nil { + return nil, fmt.Errorf("failed to get document count: %w", err) + } + if docCount >= maxNum { + return nil, fmt.Errorf("Exceed the maximum file number of a free user!") + } + } + } + + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return nil, fmt.Errorf("storage not initialized") + } + + var result []map[string]interface{} + + for _, fileHeader := range files { + filename := fileHeader.Filename + if filename == "" { + return nil, fmt.Errorf("No file selected!") + } + + fileType := utility.FilenameType(filename) + + fileObjNames := s.parseFilePath(filename) + + var idList []string + idList, err = s.fileDAO.GetIDListByID(parentID, fileObjNames, 1, []string{parentID}) + if err != nil { + return nil, fmt.Errorf("failed to get file ID list: %w", err) + } + + var lastFolder *entity.File + if len(fileObjNames) != len(idList)-1 { + lastID := idList[len(idList)-1] + lastFolder, err = s.fileDAO.GetByID(lastID) + if err != nil { + return nil, fmt.Errorf("Folder not found!") + } + var createdFolder *entity.File + createdFolder, err = s.createFolderRecursive(lastFolder, fileObjNames, len(idList), tenantID) + if err != nil { + return nil, fmt.Errorf("failed to create folder: %w", err) + } + lastFolder = createdFolder + } else { + lastID := idList[len(idList)-2] + lastFolder, err = s.fileDAO.GetByID(lastID) + if err != nil { + return nil, fmt.Errorf("Folder not found!") + } + } + + location := fileObjNames[len(fileObjNames)-1] + for storageImpl.ObjExist(lastFolder.ID, location) { + location += "_" + } + + src, err := fileHeader.Open() + if err != nil { + return nil, fmt.Errorf("failed to open uploaded file: %w", err) + } + defer src.Close() + + data, err := io.ReadAll(src) + if err != nil { + return nil, fmt.Errorf("failed to read file data: %w", err) + } + + if err = storageImpl.Put(lastFolder.ID, location, data); err != nil { + return nil, fmt.Errorf("failed to store file: %w", err) + } + + uniqueName := s.getUniqueFilename(fileObjNames[len(fileObjNames)-1], lastFolder.ID, tenantID) + + fileRecord := &entity.File{ + ID: utility.GenerateToken(), + ParentID: lastFolder.ID, + TenantID: tenantID, + CreatedBy: tenantID, + Name: uniqueName, + Location: &location, + Size: int64(len(data)), + Type: string(fileType), + SourceType: "", + } + + if err = s.fileDAO.Insert(fileRecord); err != nil { + return nil, fmt.Errorf("failed to insert file record: %w", err) + } + + result = append(result, s.toFileResponse(fileRecord)) + } + + return result, nil +} + +// UploadInfos mirrors Python's upload_info file branch: store raw bytes in the +// per-user downloads bucket and return lightweight upload descriptors instead +// of creating full File rows in the file-management tree. +func (s *FileService) UploadInfos(userID string, files []*multipart.FileHeader) ([]map[string]interface{}, error) { + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return nil, fmt.Errorf("storage not initialized") + } + + results := make([]map[string]interface{}, 0, len(files)) + for _, fileHeader := range files { + filename := fileHeader.Filename + if err := s.checkUploadInfoHealth(userID, filename); err != nil { + return nil, err + } + src, err := fileHeader.Open() + if err != nil { + return nil, fmt.Errorf("failed to open uploaded file: %w", err) + } + data, readErr := readUploadInfoData(src) + src.Close() + if readErr != nil { + return nil, fmt.Errorf("failed to read file data: %w", readErr) + } + + contentType := fileHeader.Header.Get("Content-Type") + if contentType == "" { + contentType = http.DetectContentType(data) + } + filename, contentType, data = utility.NormalizeUploadInfoContent(filename, contentType, data) + resp, err := s.storeUploadInfoBlob(storageImpl, userID, filename, contentType, data) + if err != nil { + return nil, err + } + results = append(results, resp) + } + return results, nil +} + +func readUploadInfoData(r io.Reader) ([]byte, error) { + limited := &io.LimitedReader{R: r, N: maxRemoteFileSize + 1} + data, err := io.ReadAll(limited) + if err != nil { + return nil, err + } + if int64(len(data)) > maxRemoteFileSize { + return nil, fmt.Errorf("file size exceeds %d bytes", maxRemoteFileSize) + } + return data, nil +} + +func (s *FileService) parseFilePath(filename string) []string { + filename = strings.TrimPrefix(filename, "/") + parts := strings.Split(filename, "/") + var result []string + for _, part := range parts { + if part != "" { + result = append(result, part) + } + } + return result +} + +// toUploadInfoResponse converts a newly-uploaded file record to the shape +// Python's upload_info endpoint returns. +func (s *FileService) toUploadInfoResponse(file *entity.File, mimeType string) map[string]interface{} { + ext := "" + if idx := strings.LastIndex(file.Name, "."); idx >= 0 { + ext = strings.ToLower(file.Name[idx+1:]) + } + return map[string]interface{}{ + "id": file.ID, + "name": file.Name, + "size": file.Size, + "extension": ext, + "mime_type": mimeType, + "created_by": file.CreatedBy, + "created_at": float64(time.Now().UnixMilli()) / 1000.0, + "preview_url": nil, + } +} + +func (s *FileService) checkUploadInfoHealth(userID, filename string) error { + if filename == "" { + return fmt.Errorf("No file selected!") + } + maxFileNumPerUser := common.GetEnv(common.EnvMaxFileNumPerUser) + if maxFileNumPerUser != "" { + var maxNum int64 + if _, err := fmt.Sscanf(maxFileNumPerUser, "%d", &maxNum); err == nil && maxNum > 0 { + var docCount int64 + docCount, err = s.GetDocCount(userID) + if err != nil { + return fmt.Errorf("failed to get document count: %w", err) + } + if docCount >= maxNum { + return fmt.Errorf("Exceed the maximum file number of a free user!") + } + } + } + if len([]byte(filename)) > 255 { + return fmt.Errorf("Exceed the maximum length of file name!") + } + return nil +} + +func (s *FileService) storeUploadInfoBlob(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 { + return nil, fmt.Errorf("failed to store file: %w", err) + } + ext := "" + if idx := strings.LastIndex(filename, "."); idx >= 0 { + ext = strings.ToLower(filename[idx+1:]) + } + return map[string]interface{}{ + "id": location, + "name": filename, + "size": int64(len(data)), + "extension": ext, + "mime_type": contentType, + "created_by": userID, + "created_at": float64(time.Now().UnixMilli()) / 1000.0, + "preview_url": nil, + }, nil +} + +// UploadDocumentInfos is the document-level wrapper that stores uploaded blobs +// without creating Document rows, then returns the file metadata including +// size/mime-type/extension. +func (s *FileService) UploadDocumentInfos(userID string, files []*multipart.FileHeader) ([]map[string]interface{}, common.ErrorCode, error) { + data, err := s.UploadInfos(userID, files) + if err != nil { + return nil, common.CodeDataError, err + } + return data, common.CodeSuccess, nil +} + +// UploadDocumentInfoByURL fetches a remote URL, stores the content without +// creating a Document row, then returns file metadata. +func (s *FileService) UploadDocumentInfoByURL(userID, rawURL string) (map[string]interface{}, common.ErrorCode, error) { + data, err := s.UploadFromURL(userID, rawURL) + if err != nil { + return nil, common.CodeDataError, err + } + return data, common.CodeSuccess, nil +} diff --git a/internal/service/file/file_url.go b/internal/service/file/file_url.go new file mode 100644 index 0000000000..8e19965682 --- /dev/null +++ b/internal/service/file/file_url.go @@ -0,0 +1,93 @@ +package file + +import ( + "fmt" + "net/url" + "path/filepath" + "ragflow/internal/storage" + "ragflow/internal/utility" + "strings" +) + +// UploadFromURL fetches a remote URL, saves the content to the tenant's root +// folder, and returns the file metadata map — mirroring Python +// FileService.upload_info(tenant_id, None, url). +// +// The remote fetch is SSRF-guarded (mirrors Python's assert_url_is_safe): the +// scheme must be http/https and every address the host resolves to must be +// globally routable; the validated IP is pinned for the actual connection — and +// re-validated on each redirect hop — to defeat DNS-rebinding. The HTTP client +// carries connect and overall timeouts, and the response body is bounded with +// truncation detection so an oversized file is rejected rather than silently +// clipped. +func (s *FileService) UploadFromURL(tenantID, rawURL string) (map[string]interface{}, error) { + if rawURL == "" { + return nil, fmt.Errorf("url is required") + } + parsed, err := url.Parse(rawURL) + if err != nil || (parsed.Scheme != "http" && parsed.Scheme != "https") || parsed.Hostname() == "" { + return nil, fmt.Errorf("invalid or unsafe URL") + } + + data, headers, finalURL, err := utility.FetchRemoteFileSafely(rawURL, maxRemoteFileSize) + if err != nil { + return nil, err + } + + storageImpl := storage.GetStorageFactory().GetStorage() + if storageImpl == nil { + return nil, fmt.Errorf("storage not initialized") + } + + contentType := headers.Get("Content-Type") + filename := normalizeRemoteUploadFilename(finalURL, contentType, data) + if err := s.checkUploadInfoHealth(tenantID, filename); err != nil { + return nil, err + } + filename, contentType, data = utility.NormalizeUploadInfoContent(filename, contentType, data) + return s.storeUploadInfoBlob(storageImpl, tenantID, filename, contentType, data) +} + +func normalizeRemoteUploadFilename(rawURL, contentType string, data []byte) string { + parsed, err := url.Parse(rawURL) + filename := "download" + if err == nil { + filename = sanitizeFilename(filepath.Base(parsed.Path)) + } + ct := strings.ToLower(strings.TrimSpace(strings.Split(contentType, ";")[0])) + if ct == "application/pdf" || utility.BytesLooksLikePDF(data) { + if !strings.HasSuffix(strings.ToLower(filename), ".pdf") { + filename += ".pdf" + } + } + return filename +} + +func sanitizeFilename(name string) string { + name = filepath.Base(name) + name = strings.TrimSpace(name) + + name = strings.Map(func(r rune) rune { + switch r { + case '/', '\\', ':', '*', '?', '"', '<', '>', '|', 0: + return '_' + } + if r < 0x20 { + return '_' + } + return r + }, name) + + name = strings.Trim(name, ". ") + + if name == "" || name == "." || name == ".." { + return "download" + } + if stem := strings.SplitN(strings.ToUpper(name), ".", 2)[0]; reservedDeviceNames[stem] { + return "download" + } + if len(name) > 255 { + name = name[:255] + } + return name +} diff --git a/internal/service/file/file_url_test.go b/internal/service/file/file_url_test.go new file mode 100644 index 0000000000..2e4d71fa18 --- /dev/null +++ b/internal/service/file/file_url_test.go @@ -0,0 +1,37 @@ +package file + +import "testing" + +func TestSanitizeFilename(t *testing.T) { + cases := []struct { + name string + in string + want string + }{ + {"no special", "report.txt", "report.txt"}, + {"path stripped", "dir/sub/report.txt", "report.txt"}, + {"forbidden chars", "a:b*c?", "a_b_c_"}, + {"reserved device", "CON", "download"}, + {"only dots", "...", "download"}, + {"leading dots", ".hidden", "hidden"}, + } + for _, c := range cases { + got := sanitizeFilename(c.in) + if c.want != "" && got != c.want { + t.Errorf("%s: sanitizeFilename(%q) = %q, want %q", c.name, c.in, got, c.want) + } + // Safety invariant: never emit characters that are unsafe in a path. + for _, r := range got { + switch r { + case '/', '\\', ':', '*', '?', '"', '<', '>', '|', 0: + t.Errorf("%s: forbidden char %q in %q", c.name, r, got) + } + if r < 0x20 { + t.Errorf("%s: control char in %q", c.name, got) + } + } + if got == "" { + t.Errorf("%s: empty result for %q", c.name, c.in) + } + } +} diff --git a/internal/service/file_permission.go b/internal/service/file_permission.go new file mode 100644 index 0000000000..a0b5d17f08 --- /dev/null +++ b/internal/service/file_permission.go @@ -0,0 +1,34 @@ +package service + +import ( + "ragflow/internal/dao" + "ragflow/internal/entity" +) + +// CheckFileTeamPermission reports whether userID may access the given file: either +// the file's owning tenant matches userID, or userID is a team member of any +// dataset the file is linked to. It mirrors Python's check_file_team_permission +// and is shared by FileService and File2DocumentService, which previously each +// carried an identical copy of this logic. +func CheckFileTeamPermission(fileDAO *dao.FileDAO, file *entity.File, userID string) bool { + if file.TenantID == userID { + return true + } + + datasetIDs, err := fileDAO.GetDatasetIDByFileID(file.ID) + if err != nil || len(datasetIDs) == 0 { + return false + } + + kbDAO := dao.NewKnowledgebaseDAO() + for _, datasetID := range datasetIDs { + kb, err := kbDAO.GetByID(datasetID) + if err != nil || kb == nil { + continue + } + if HasKBTeamPermission(kb, userID, dao.NewTenantDAO()) { + return true + } + } + return false +} diff --git a/internal/service/ingestion_task_service_test.go b/internal/service/ingestion_task_service_test.go index c4d1b88563..9c65ca63c2 100644 --- a/internal/service/ingestion_task_service_test.go +++ b/internal/service/ingestion_task_service_test.go @@ -850,27 +850,3 @@ func TestIngestionTaskServiceMarkCompletedIdempotentOnAlreadyTerminal(t *testing t.Fatalf("MarkCompleted on already FAILED task should be idempotent, got: %v", err) } } - -func TestDocumentServiceUpdateRunProgressMirrorsFields(t *testing.T) { - db := setupServiceTestDB(t) - pushServiceDB(t, db) - insertTestDoc(t, "doc-1", "kb-1", 0, 0) - - svc := testDocumentService(t) - if err := svc.UpdateRunProgress("doc-1", 0.5, "1", "halfway"); err != nil { - t.Fatalf("UpdateRunProgress failed: %v", err) - } - doc, err := dao.NewDocumentDAO().GetByID("doc-1") - if err != nil { - t.Fatalf("load document: %v", err) - } - if doc.Progress != 0.5 { - t.Fatalf("progress = %v, want 0.5", doc.Progress) - } - if doc.Run == nil || *doc.Run != "1" { - t.Fatalf("run = %v, want 1", doc.Run) - } - if doc.ProgressMsg == nil || *doc.ProgressMsg != "halfway" { - t.Fatalf("progress_msg = %v, want halfway", doc.ProgressMsg) - } -} diff --git a/internal/service/memory_message_test.go b/internal/service/memory_message_test.go index 40bf26c193..9a5b40c751 100644 --- a/internal/service/memory_message_test.go +++ b/internal/service/memory_message_test.go @@ -62,6 +62,12 @@ func (e *memoryMessageDocEngine) UpdateChunks(ctx context.Context, condition map return nil } +func (e *memoryMessageDocEngine) FilterDocIdsByMetaPushdown(_ context.Context, _ []string, _ []map[string]interface{}, _ string) []string { + return nil +} + +func (e *memoryMessageDocEngine) GetType() string { return "memory" } + func setupMemoryMessageTestDB(t *testing.T) { t.Helper() diff --git a/internal/service/metadata.go b/internal/service/metadata.go index 98f9a1e586..dee0dae91c 100644 --- a/internal/service/metadata.go +++ b/internal/service/metadata.go @@ -50,6 +50,15 @@ func NewMetadataService() *MetadataService { } } +// NewMetadataServiceForTest creates a MetadataService with injected dependencies +// for tests that need to control the DAO and engine. +func NewMetadataServiceForTest(kbDAO *dao.KnowledgebaseDAO, docEngine engine.DocEngine) *MetadataService { + return &MetadataService{ + kbDAO: kbDAO, + docEngine: docEngine, + } +} + // BuildMetadataIndexName constructs the metadata index name for a tenant func BuildMetadataIndexName(tenantID string) string { return fmt.Sprintf("ragflow_doc_meta_%s", tenantID) @@ -390,7 +399,7 @@ func ExtractMetaFields(chunk map[string]interface{}) (map[string]interface{}, er for k, val := range result { if existing, exists := metaFields[k]; exists { // Key already exists - merge values - metaFields[k] = mergeFieldValues(existing, val) + metaFields[k] = MergeFieldValues(existing, val) } else { metaFields[k] = val } @@ -409,7 +418,7 @@ func ExtractMetaFields(chunk map[string]interface{}) (map[string]interface{}, er // mergeFieldValues merges two field values when the same key appears multiple times // If both are arrays, append all elements. If one is array and other is string, append string to array. // Returns []interface{} with all merged values (flattened). -func mergeFieldValues(existing, new interface{}) []interface{} { +func MergeFieldValues(existing, new interface{}) []interface{} { result := []interface{}{} var addValue func(v interface{}) diff --git a/internal/service/parse_types.go b/internal/service/parse_types.go new file mode 100644 index 0000000000..6f01b6ab77 --- /dev/null +++ b/internal/service/parse_types.go @@ -0,0 +1,23 @@ +// +// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// + +package service + +// ParseDocumentResponse is returned by document parse/ingest operations. +type ParseDocumentResponse struct { + DocumentID string `json:"document_id"` + Result string `json:"result"` +} diff --git a/internal/service/pipeline_params.go b/internal/service/pipeline_params.go new file mode 100644 index 0000000000..fbdc002f49 --- /dev/null +++ b/internal/service/pipeline_params.go @@ -0,0 +1,151 @@ +package service + +import ( + "encoding/json" + "errors" + "fmt" + "strings" + + "ragflow/internal/dao" + "ragflow/internal/entity" + pipelinepkg "ragflow/internal/ingestion/pipeline" +) + +// loadCanvasDSLJSON returns the DSL JSON for a custom canvas pipeline. The +// canvas's dsl column holds the same component-graph structure that built-in +// templates use, so it can be validated by the same schema extractor. It is a +// package-level function so both document and knowledge-base updates reuse it. +func loadCanvasDSLJSON(canvasID string) ([]byte, error) { + if strings.TrimSpace(canvasID) == "" { + return nil, fmt.Errorf("empty canvas id") + } + canvas, err := dao.NewUserCanvasDAO().GetByID(canvasID) + if err != nil { + if errors.Is(err, dao.ErrUserCanvasNotFound) { + return nil, fmt.Errorf("canvas %s not found", canvasID) + } + return nil, fmt.Errorf("load canvas %s: %w", canvasID, err) + } + if len(canvas.DSL) == 0 { + return nil, fmt.Errorf("canvas %s has no DSL", canvasID) + } + raw, err := json.Marshal(canvas.DSL) + if err != nil { + return nil, fmt.Errorf("marshal canvas %s DSL: %w", canvasID, err) + } + return raw, nil +} + +// LoadPipelineDSL loads the DSL JSON for a pipeline identified by parserID +// (built-in) or pipelineID (custom canvas). When both are provided, isCanvas +// selects which one to use. +func LoadPipelineDSL(isCanvas bool, parserID string, pipelineID *string) ([]byte, error) { + if isCanvas { + return loadCanvasDSLJSON(strings.TrimSpace(*pipelineID)) + } + registry, err := pipelinepkg.DefaultRegistry() + if err != nil { + return nil, fmt.Errorf("builtin pipeline registry: %w", err) + } + if !registry.IsValid(parserID) { + return nil, fmt.Errorf("unknown builtin parser_id: %s", parserID) + } + dslStr, err := pipelinepkg.LoadBuiltinDSL(parserID) + if err != nil { + return nil, fmt.Errorf("load builtin DSL for %q: %w", parserID, err) + } + return []byte(dslStr), nil +} + +// ResolveComponentParamsDefaults loads the DSL for the target pipeline and +// returns the component params defaults as an entity.JSONMap {cpnID: {param: value}}. +// For builtin templates the DSL is loaded from the embedded registry; for custom +// canvas pipelines it is loaded from the canvas row in the database. +func ResolveComponentParamsDefaults(parserID string, pipelineID *string) (entity.JSONMap, error) { + isCanvas := pipelineID != nil && strings.TrimSpace(*pipelineID) != "" + var cp map[string]map[string]any + var err error + if isCanvas { + dslJSON, lerr := loadCanvasDSLJSON(strings.TrimSpace(*pipelineID)) + if lerr != nil { + return nil, fmt.Errorf("load canvas DSL: %w", lerr) + } + cp, err = pipelinepkg.ComponentParamsDefaults(dslJSON) + } else { + registry, regErr := pipelinepkg.DefaultRegistry() + if regErr != nil { + return nil, fmt.Errorf("builtin registry: %w", regErr) + } + if !registry.IsValid(parserID) { + return nil, fmt.Errorf("unknown builtin parser_id: %q", parserID) + } + dslStr, dslErr := pipelinepkg.LoadBuiltinDSL(parserID) + if dslErr != nil { + return nil, fmt.Errorf("load builtin DSL: %w", dslErr) + } + cp, err = pipelinepkg.ComponentParamsDefaults([]byte(dslStr)) + } + if err != nil { + return nil, err + } + out := make(entity.JSONMap, len(cp)) + for k, v := range cp { + out[k] = v + } + return out, nil +} + +// ValidateDatasetEmbeddingModels checks that all knowledge bases in the list +// either have an embedding model or none do, and that they all use the same model. +func ValidateDatasetEmbeddingModels(kbs []*entity.Knowledgebase) error { + embdIDs := make(map[string]struct{}) + hasEmbd := false + noEmbd := false + for _, kb := range kbs { + if kb.EmbdID != "" { + hasEmbd = true + baseName := kb.EmbdID + if idx := strings.LastIndex(kb.EmbdID, "@"); idx > 0 { + baseName = kb.EmbdID[:idx] + // Strip the second-to-last @-segment too (instance name), + // matching Python's _base_model_name which uses rsplit("@", 2). + if idx2 := strings.LastIndex(baseName, "@"); idx2 > 0 { + baseName = baseName[:idx2] + } + } + embdIDs[baseName] = struct{}{} + } else { + noEmbd = true + } + } + if hasEmbd && noEmbd { + return fmt.Errorf("Cannot search across datasets where some have embedding models and others do not.") + } + if len(embdIDs) > 1 { + return fmt.Errorf("Datasets use different embedding models: %v", getEmbdIDs(kbs)) + } + return nil +} + +func getEmbdIDs(kbs []*entity.Knowledgebase) []string { + ids := make([]string, len(kbs)) + for i, kb := range kbs { + ids[i] = kb.EmbdID + } + return ids +} + +// Backward-compat lowercase aliases for callers within the service package. +// These will be removed when callers are updated to use the exported names. + +func loadPipelineDSL(isCanvas bool, parserID string, pipelineID *string) ([]byte, error) { + return LoadPipelineDSL(isCanvas, parserID, pipelineID) +} + +func resolveComponentParamsDefaults(parserID string, pipelineID *string) (entity.JSONMap, error) { + return ResolveComponentParamsDefaults(parserID, pipelineID) +} + +func validateDatasetEmbeddingModels(kbs []*entity.Knowledgebase) error { + return ValidateDatasetEmbeddingModels(kbs) +} diff --git a/internal/service/pipeline_params_test.go b/internal/service/pipeline_params_test.go new file mode 100644 index 0000000000..38ea632132 --- /dev/null +++ b/internal/service/pipeline_params_test.go @@ -0,0 +1,81 @@ +package service + +import ( + "testing" + + "ragflow/internal/entity" +) + +func TestValidateDatasetEmbeddingModels_AllHaveEmbeddingModel(t *testing.T) { + kbs := []*entity.Knowledgebase{ + {EmbdID: "BAAI/bge-large-zh-v1.5@Builtin"}, + {EmbdID: "BAAI/bge-large-zh-v1.5@Builtin"}, + } + if err := ValidateDatasetEmbeddingModels(kbs); err != nil { + t.Fatalf("expected nil, got %v", err) + } +} + +func TestValidateDatasetEmbeddingModels_NoneHasEmbeddingModel(t *testing.T) { + kbs := []*entity.Knowledgebase{ + {EmbdID: ""}, + {EmbdID: ""}, + } + if err := ValidateDatasetEmbeddingModels(kbs); err != nil { + t.Fatalf("expected nil, got %v", err) + } +} + +func TestValidateDatasetEmbeddingModels_MixedErrors(t *testing.T) { + kbs := []*entity.Knowledgebase{ + {EmbdID: "BAAI/bge-large-zh-v1.5@Builtin"}, + {EmbdID: ""}, + } + err := ValidateDatasetEmbeddingModels(kbs) + if err == nil { + t.Fatal("expected error for mixed embedding") + } + if err.Error() != "Cannot search across datasets where some have embedding models and others do not." { + t.Errorf("unexpected error: %v", err) + } +} + +func TestValidateDatasetEmbeddingModels_DifferentEmbeddingsErrors(t *testing.T) { + kbs := []*entity.Knowledgebase{ + {EmbdID: "model-a@provider-1"}, + {EmbdID: "model-b@provider-2"}, + } + err := ValidateDatasetEmbeddingModels(kbs) + if err == nil { + t.Fatal("expected error for different embeddings") + } +} + +func TestValidateDatasetEmbeddingModels_SameBaseDifferentInstanceOK(t *testing.T) { + // Two KBs using the same base model through different provider instances + // (the rsplit("@", 2) logic should treat them as the same base). + kbs := []*entity.Knowledgebase{ + {EmbdID: "BAAI/bge-large-zh-v1.5@instance1@provider1"}, + {EmbdID: "BAAI/bge-large-zh-v1.5@instance2@provider2"}, + } + if err := ValidateDatasetEmbeddingModels(kbs); err != nil { + t.Fatalf("expected nil, got %v", err) + } +} + +func TestValidateDatasetEmbeddingModels_DifferentBasesErrors(t *testing.T) { + kbs := []*entity.Knowledgebase{ + {EmbdID: "model-a@instance1@provider1"}, + {EmbdID: "model-b@instance2@provider2"}, + } + err := ValidateDatasetEmbeddingModels(kbs) + if err == nil { + t.Fatal("expected error for different base models") + } +} + +func TestValidateDatasetEmbeddingModels_EmptyList(t *testing.T) { + if err := ValidateDatasetEmbeddingModels(nil); err != nil { + t.Fatalf("expected nil for empty list, got %v", err) + } +} diff --git a/internal/service/skill_space.go b/internal/service/skill_space.go index f25446941c..2b73c7a0b3 100644 --- a/internal/service/skill_space.go +++ b/internal/service/skill_space.go @@ -23,6 +23,7 @@ import ( "ragflow/internal/dao" "ragflow/internal/engine" "ragflow/internal/entity" + "ragflow/internal/service/file" "ragflow/internal/utility" "sync" "time" @@ -34,7 +35,7 @@ import ( type SkillSpaceService struct { spaceDAO *dao.SkillSpaceDAO fileDAO *dao.FileDAO - fileService *FileService + fileService *file.FileService configDAO *dao.SkillSearchConfigDAO tenantDAO *dao.TenantDAO skillsFolderCache map[string]string // tenant-keyed cache for skills folder ID @@ -43,12 +44,13 @@ type SkillSpaceService struct { spaceCreateMu sync.Map // tenant-scoped locks for space creation (prevents TOCTOU races) } -// NewSkillSpaceService creates a new SkillSpaceService instance -func NewSkillSpaceService() *SkillSpaceService { +// NewSkillSpaceService creates a new SkillSpaceService instance. +// dr is the document remover used when deleting files; it must be non-nil. +func NewSkillSpaceService(dr file.DocRemover) *SkillSpaceService { return &SkillSpaceService{ spaceDAO: dao.NewSkillSpaceDAO(), fileDAO: dao.NewFileDAO(), - fileService: NewFileService(), + fileService: file.NewFileService(CheckFileTeamPermission, dr), configDAO: dao.NewSkillSearchConfigDAO(), tenantDAO: dao.NewTenantDAO(), skillsFolderCache: make(map[string]string), diff --git a/internal/service/team_permission.go b/internal/service/team_permission.go index 73a0882138..57b932b477 100644 --- a/internal/service/team_permission.go +++ b/internal/service/team_permission.go @@ -5,10 +5,10 @@ import ( "ragflow/internal/entity" ) -// hasKBTeamPermission mirrors Python check_kb_team_permission: +// 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(kb *entity.Knowledgebase, userID string, tenantDAO *dao.TenantDAO) bool { if kb == nil { return false } diff --git a/internal/service/test_helpers_test.go b/internal/service/test_helpers_test.go new file mode 100644 index 0000000000..099d3ec03b --- /dev/null +++ b/internal/service/test_helpers_test.go @@ -0,0 +1,219 @@ +// +// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// + +package service + +import ( + "context" + "testing" + + "github.com/glebarez/sqlite" + "gorm.io/gorm" + + "ragflow/internal/common" + "ragflow/internal/dao" + "ragflow/internal/engine/types" + "ragflow/internal/entity" +) + +// sptr returns a pointer to the given string. +func sptr(s string) *string { return &s } + +// setupServiceTestDB initializes an in-memory SQLite database for service tests. +func setupServiceTestDB(t *testing.T) *gorm.DB { + t.Helper() + + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + TranslateError: true, + }) + if err != nil { + t.Fatalf("failed to open sqlite: %v", err) + } + + if err = db.AutoMigrate( + &entity.Document{}, + &entity.Knowledgebase{}, + &entity.Task{}, + &entity.IngestionTask{}, + &entity.IngestionTaskLog{}, + &entity.File2Document{}, + &entity.File{}, + &entity.User{}, + &entity.Tenant{}, + &entity.UserTenant{}, + &entity.API4Conversation{}, + ); err != nil { + t.Fatalf("failed to migrate: %v", err) + } + + return db +} + +// pushServiceDB swaps dao.DB for the test and restores after. +func pushServiceDB(t *testing.T, testDB *gorm.DB) { + t.Helper() + orig := dao.DB + dao.DB = testDB + t.Cleanup(func() { + dao.DB = orig + }) +} + +// fakeChatDocEngine is a stub engine.DocEngine used by parent-package tests. +type fakeChatDocEngine struct{} + +func (fakeChatDocEngine) CreateChunkStore(context.Context, string, string, int, string) error { + return nil +} +func (fakeChatDocEngine) InsertChunks(context.Context, []map[string]interface{}, string, string) ([]string, error) { + return nil, nil +} +func (fakeChatDocEngine) UpdateChunks(context.Context, map[string]interface{}, map[string]interface{}, string, string) error { + return nil +} +func (fakeChatDocEngine) DeleteChunks(context.Context, map[string]interface{}, string, string) (int64, error) { + return 0, nil +} +func (fakeChatDocEngine) Search(context.Context, *types.SearchRequest) (*types.SearchResult, error) { + return nil, nil +} +func (fakeChatDocEngine) GetChunk(context.Context, string, string, []string) (interface{}, error) { + return nil, nil +} +func (fakeChatDocEngine) DropChunkStore(context.Context, string, string) error { return nil } +func (fakeChatDocEngine) ChunkStoreExists(context.Context, string, string) (bool, error) { + return false, nil +} +func (fakeChatDocEngine) CreateMetadataStore(context.Context, string) error { return nil } +func (fakeChatDocEngine) InsertMetadata(context.Context, []map[string]interface{}, string) ([]string, error) { + return nil, nil +} +func (fakeChatDocEngine) UpdateMetadata(context.Context, string, string, map[string]interface{}, string) error { + return nil +} +func (fakeChatDocEngine) DeleteMetadata(context.Context, map[string]interface{}, string) (int64, error) { + return 0, nil +} +func (fakeChatDocEngine) DeleteMetadataKeys(context.Context, string, string, []string, string) error { + return nil +} +func (fakeChatDocEngine) DropMetadataStore(context.Context, string) error { return nil } +func (fakeChatDocEngine) MetadataStoreExists(context.Context, string) (bool, error) { + return false, nil +} +func (fakeChatDocEngine) SearchMetadata(context.Context, *types.SearchMetadataRequest) (*types.SearchMetadataResult, error) { + return nil, nil +} +func (fakeChatDocEngine) IndexDocument(context.Context, string, string, interface{}) error { + return nil +} +func (fakeChatDocEngine) DeleteDocument(context.Context, string, string) error { return nil } +func (fakeChatDocEngine) BulkIndex(context.Context, string, []interface{}) (interface{}, error) { + return nil, nil +} +func (fakeChatDocEngine) GetFields([]map[string]interface{}, []string) map[string]map[string]interface{} { + return nil +} +func (fakeChatDocEngine) GetAggregation([]map[string]interface{}, string) []map[string]interface{} { + return nil +} +func (fakeChatDocEngine) GetHighlight([]map[string]interface{}, []string, string) map[string]string { + return nil +} +func (fakeChatDocEngine) RunSQL(context.Context, string, string, []string, string) ([]map[string]interface{}, error) { + return nil, nil +} +func (fakeChatDocEngine) GetChunkIDs([]map[string]interface{}) []string { return nil } +func (fakeChatDocEngine) KNNScores(context.Context, []map[string]interface{}, []float64, int) (map[string]interface{}, error) { + return nil, nil +} +func (fakeChatDocEngine) GetScores(map[string]interface{}) map[string]float64 { return nil } +func (fakeChatDocEngine) Ping(context.Context) error { return nil } +func (fakeChatDocEngine) Close() error { return nil } +func (fakeChatDocEngine) CheckStatus() error { return nil } +func (fakeChatDocEngine) FilterDocIdsByMetaPushdown(context.Context, []string, []map[string]interface{}, string) []string { + return nil +} +func (fakeChatDocEngine) GetMessages(context.Context, string, int, int) ([]interface{}, error) { + return nil, nil +} +func (fakeChatDocEngine) GetType() string { return "fake" } +func (fakeChatDocEngine) Init() error { return nil } +func (fakeChatDocEngine) InitConsumer() error { return nil } +func (fakeChatDocEngine) ListMessages(context.Context, string, int, int) ([]interface{}, error) { + return nil, nil +} +func (fakeChatDocEngine) PublishTask(map[string]interface{}) error { return nil } +func (fakeChatDocEngine) ShowMessageQueue() string { return "" } +func (fakeChatDocEngine) SupportsPageRank() bool { return false } + +// insertTestKB inserts a test Knowledgebase row. +func insertTestKB(t *testing.T, id, tenantID string, docNum, tokenNum, chunkNum int64) { + t.Helper() + kb := &entity.Knowledgebase{ + ID: id, + TenantID: tenantID, + Name: "test-kb", + EmbdID: "embd-1", + CreatedBy: "user-1", + Permission: string(entity.TenantPermissionTeam), + DocNum: docNum, + TokenNum: tokenNum, + ChunkNum: chunkNum, + Status: sptr(string(entity.StatusValid)), + } + if err := dao.DB.Create(kb).Error; err != nil { + t.Fatalf("insert test kb: %v", err) + } +} + +// insertTestDoc inserts a test Document row. +func insertTestDoc(t *testing.T, id, kbID string, tokenNum, chunkNum int64) { + t.Helper() + doc := &entity.Document{ + ID: id, + KbID: kbID, + ParserID: "naive", + ParserConfig: entity.JSONMap{}, + TokenNum: tokenNum, + ChunkNum: chunkNum, + Suffix: ".txt", + Status: sptr("1"), + } + if err := dao.DB.Create(doc).Error; err != nil { + t.Fatalf("insert test doc: %v", err) + } +} + +// insertTestIngestionTask inserts a test IngestionTask row with CREATED status. +func insertTestIngestionTask(t *testing.T, id, userID, docID, datasetID string) { + insertTestIngestionTaskWithStatus(t, id, userID, docID, datasetID, common.CREATED) +} + +// insertTestIngestionTaskWithStatus inserts a test IngestionTask with a specific status. +func insertTestIngestionTaskWithStatus(t *testing.T, id, userID, docID, datasetID, status string) { + t.Helper() + task := &entity.IngestionTask{ + ID: id, + UserID: userID, + DocumentID: docID, + DatasetID: datasetID, + Status: status, + } + if err := dao.DB.Create(task).Error; err != nil { + t.Fatalf("insert test ingestion task: %v", err) + } +} diff --git a/internal/utility/ssrf.go b/internal/utility/ssrf.go index 74a39c3d08..c271dda88c 100644 --- a/internal/utility/ssrf.go +++ b/internal/utility/ssrf.go @@ -63,7 +63,7 @@ func allowAnyHost() bool { // prevent rebinding between validation and the actual TCP connection. // // Mirrors common/ssrf_guard.py:assert_url_is_safe. -func AssertURLSafe(rawURL string) (hostname, resolvedIP string, err error) { +var AssertURLSafe = func(rawURL string) (hostname, resolvedIP string, err error) { parsed, err := url.Parse(strings.TrimSpace(rawURL)) if err != nil { return "", "", fmt.Errorf("invalid url") @@ -176,7 +176,7 @@ func allZero(b []byte) bool { // outbound dial for hostname:port to resolvedIP:port, closing the TOCTOU // window between AssertURLSafe and the actual TCP connection. Pins are // scoped to this client only. -func PinnedHTTPClient(hostname, resolvedIP string, timeout time.Duration) *http.Client { +var PinnedHTTPClient = func(hostname, resolvedIP string, timeout time.Duration) *http.Client { dialer := &net.Dialer{ Timeout: timeout, KeepAlive: 30 * time.Second, diff --git a/internal/utility/upload.go b/internal/utility/upload.go new file mode 100644 index 0000000000..4d357c121b --- /dev/null +++ b/internal/utility/upload.go @@ -0,0 +1,149 @@ +// +// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// + +package utility + +import ( + "fmt" + "html" + "io" + "net/http" + "net/url" + "regexp" + "strings" + "time" +) + +var ( + htmlScriptStyleRE = regexp.MustCompile(`(?is)<(script|style)[^>]*>.*?`) + htmlTagRE = regexp.MustCompile(`(?s)<[^>]+>`) + multiSpaceRE = regexp.MustCompile(`[ \t]+`) + multiNewlineRE = regexp.MustCompile(`\n{3,}`) +) + +// FetchRemoteFileSafely downloads rawURL with SSRF protection, connect/overall +// timeouts, and a hard size cap that rejects (rather than truncates) oversized +// bodies. +func FetchRemoteFileSafely(rawURL string, maxSize int64) ([]byte, http.Header, string, error) { + currentURL := rawURL + for redirects := 0; redirects < 10; redirects++ { + hostname, resolvedIP, err := AssertURLSafe(currentURL) + if err != nil { + return nil, nil, "", err + } + client := PinnedHTTPClient(hostname, resolvedIP, 10*time.Second) + client.CheckRedirect = func(req *http.Request, via []*http.Request) error { + return http.ErrUseLastResponse + } + + // codeql[go/request-forgery] False positive: the loop above + resp, err := client.Get(currentURL) // #nosec G107 + if err != nil { + return nil, nil, "", fmt.Errorf("failed to fetch URL: %w", err) + } + + if resp.StatusCode == http.StatusMovedPermanently || + resp.StatusCode == http.StatusFound || + resp.StatusCode == http.StatusSeeOther || + resp.StatusCode == http.StatusTemporaryRedirect || + resp.StatusCode == http.StatusPermanentRedirect { + location := resp.Header.Get("Location") + resp.Body.Close() + if location == "" { + return nil, nil, "", fmt.Errorf("redirect response missing Location header") + } + baseURL, parseErr := url.Parse(currentURL) + if parseErr != nil { + return nil, nil, "", parseErr + } + nextURL, resolveErr := baseURL.Parse(location) + if resolveErr != nil { + return nil, nil, "", resolveErr + } + currentURL = nextURL.String() + continue + } + + if resp.StatusCode >= 400 { + resp.Body.Close() + return nil, nil, "", fmt.Errorf("remote URL returned HTTP %d", resp.StatusCode) + } + + data, readErr := io.ReadAll(io.LimitReader(resp.Body, maxSize+1)) + resp.Body.Close() + if readErr != nil { + return nil, nil, "", fmt.Errorf("failed to read remote content: %w", readErr) + } + if int64(len(data)) > maxSize { + return nil, nil, "", fmt.Errorf("remote file exceeds the maximum allowed size of %d bytes", maxSize) + } + return data, resp.Header.Clone(), currentURL, nil + } + return nil, nil, "", fmt.Errorf("stopped after too many redirects") +} + +// NormalizeUploadInfoContent normalizes an uploaded file's filename, content +// type, and content bytes: detects PDF by magic bytes, converts HTML to +// readable markdown, and fixes the filename extension. +func NormalizeUploadInfoContent(filename, contentType string, data []byte) (string, string, []byte) { + lowerCT := strings.ToLower(strings.TrimSpace(strings.Split(contentType, ";")[0])) + if lowerCT == "" { + lowerCT = http.DetectContentType(data) + } + + if lowerCT == "application/pdf" || BytesLooksLikePDF(data) { + if !strings.HasSuffix(strings.ToLower(filename), ".pdf") { + filename += ".pdf" + } + lowerCT = "application/pdf" + } + if lowerCT == "text/html" || lowerCT == "application/xhtml+xml" || LooksLikeHTML(data) { + data = htmlToReadableMarkdown(data) + if lowerCT == "" { + lowerCT = "text/html" + } + } + return filename, lowerCT, data +} + +// BytesLooksLikePDF reports whether data starts with the PDF magic bytes. +func BytesLooksLikePDF(data []byte) bool { + return len(data) >= 4 && string(data[:4]) == "%PDF" +} + +// LooksLikeHTML reports whether data contains common HTML tag markers. +func LooksLikeHTML(data []byte) bool { + snippet := strings.ToLower(string(data)) + return strings.Contains(snippet, "", "\n") + text = strings.ReplaceAll(text, "
", "\n") + text = strings.ReplaceAll(text, "
", "\n") + text = strings.ReplaceAll(text, "

", "\n\n") + text = strings.ReplaceAll(text, "
", "\n") + text = strings.ReplaceAll(text, "", "\n") + text = htmlTagRE.ReplaceAllString(text, " ") + text = html.UnescapeString(text) + text = strings.ReplaceAll(text, "\r", "\n") + text = multiSpaceRE.ReplaceAllString(text, " ") + text = multiNewlineRE.ReplaceAllString(text, "\n\n") + text = strings.TrimSpace(text) + return []byte(text) +} diff --git a/internal/utility/upload_test.go b/internal/utility/upload_test.go new file mode 100644 index 0000000000..0016c67c54 --- /dev/null +++ b/internal/utility/upload_test.go @@ -0,0 +1,103 @@ +// +// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// + +package utility + +import ( + "bytes" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" +) + +func TestHtmlToReadableMarkdown(t *testing.T) { + out := htmlToReadableMarkdown([]byte("

hello

world
")) + if bytes.Contains(out, []byte("<")) { + t.Errorf("tags not stripped: %q", out) + } + if !bytes.Contains(out, []byte("hello")) || !bytes.Contains(out, []byte("world")) { + t.Errorf("text lost: %q", out) + } +} + +func TestFetchRemoteFileSafely_PDFAddsExtension(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/pdf") + _, _ = w.Write([]byte("%PDF-1.7 fake pdf")) + })) + defer server.Close() + + origAssert := AssertURLSafe + origPinned := PinnedHTTPClient + AssertURLSafe = func(rawURL string) (string, string, error) { + return "127.0.0.1", "127.0.0.1", nil + } + PinnedHTTPClient = func(hostname, resolvedIP string, timeout time.Duration) *http.Client { + return server.Client() + } + t.Cleanup(func() { + AssertURLSafe = origAssert + PinnedHTTPClient = origPinned + }) + + data, headers, _, err := FetchRemoteFileSafely(server.URL+"/report", 100<<20) + if err != nil { + t.Fatalf("FetchRemoteFileSafely failed: %v", err) + } + if ct := headers.Get("Content-Type"); ct != "application/pdf" { + t.Fatalf("Content-Type = %q, want application/pdf", ct) + } + if !bytes.Equal(data, []byte("%PDF-1.7 fake pdf")) { + t.Fatalf("data = %q, want %%PDF-1.7 fake pdf", string(data)) + } +} + +func TestFetchRemoteFileSafely_ReturnsContentAndHeaders(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/html; charset=utf-8") + _, _ = w.Write([]byte(`Hello World`)) + })) + defer server.Close() + + origAssert := AssertURLSafe + origPinned := PinnedHTTPClient + AssertURLSafe = func(rawURL string) (string, string, error) { + return "127.0.0.1", "127.0.0.1", nil + } + PinnedHTTPClient = func(hostname, resolvedIP string, timeout time.Duration) *http.Client { + return server.Client() + } + t.Cleanup(func() { + AssertURLSafe = origAssert + PinnedHTTPClient = origPinned + }) + + data, headers, finalURL, err := FetchRemoteFileSafely(server.URL+"/page", 100<<20) + if err != nil { + t.Fatalf("FetchRemoteFileSafely failed: %v", err) + } + if ct := headers.Get("Content-Type"); !strings.Contains(ct, "text/html") { + t.Fatalf("Content-Type = %q, want text/html", ct) + } + if finalURL != server.URL+"/page" { + t.Fatalf("finalURL = %q, want %q", finalURL, server.URL+"/page") + } + if !bytes.Equal(data, []byte(`Hello World`)) { + t.Fatalf("data = %q", string(data)) + } +}