From 965590ccbee8e4b0d3a41ca0f777d72cba9e534b Mon Sep 17 00:00:00 2001
From: Jack
Date: Mon, 20 Jul 2026 09:48:24 +0800
Subject: [PATCH] Refactor: dataset/document/file service (#17071)
### Summary
Refactor dataset.go document.do file.go file2document.go in
internal/service.
---
cmd/ragflow_server.go | 15 +-
internal/handler/agent.go | 3 +-
internal/handler/dataset.go | 5 +-
internal/handler/document.go | 68 +-
internal/handler/document_test.go | 104 +-
internal/handler/file.go | 22 +-
internal/handler/mcp_server.go | 5 +-
internal/handler/skill_search.go | 8 +-
internal/handler/tenant.go | 5 +-
.../ingestion/pipeline/pipeline_params.go | 151 +
.../pipeline/pipeline_params_test.go | 268 ++
internal/ingestion/service/doc_state.go | 4 +-
.../ingestion/service/ingestion_service.go | 5 +-
internal/ingestion/service/progress_sink.go | 3 +-
.../ingestion/service/progress_sink_test.go | 3 +-
internal/service/agent.go | 7 +-
internal/service/chat.go | 2 +-
internal/service/chat_pipeline.go | 21 +-
internal/service/chunk/chunk.go | 11 +-
internal/service/chunk/chunk_test.go | 7 +-
internal/service/dataset.go | 3190 --------------
internal/service/dataset/create_test.go | 157 +
internal/service/dataset/crud.go | 452 ++
internal/service/dataset/crud_test.go | 101 +
.../service/dataset/fake_doc_engine_test.go | 82 +
internal/service/dataset/helpers.go | 270 ++
internal/service/dataset/helpers_test.go | 389 ++
internal/service/dataset/index.go | 743 ++++
.../index_delete_test.go} | 2 +-
internal/service/dataset/ingestion.go | 167 +
internal/service/dataset/metadata.go | 116 +
.../metadata_config_test.go} | 2 +-
internal/service/dataset/permission.go | 60 +
internal/service/dataset/search.go | 292 ++
.../search_test.go} | 10 +-
internal/service/dataset/service.go | 40 +
internal/service/dataset/setup_test.go | 55 +
internal/service/dataset/tags.go | 233 ++
.../tags_aggregate_test.go} | 2 +-
.../tags_list_test.go} | 2 +-
.../tags_rename_test.go} | 2 +-
.../task_cleanup_test.go} | 2 +-
internal/service/dataset/update.go | 269 ++
.../update_test.go} | 43 +-
internal/service/dataset/utils.go | 322 ++
internal/service/dataset_create_test.go | 207 -
internal/service/dataset_parser_test.go | 64 -
internal/service/dataset_types.go | 159 +
internal/service/document.go | 3652 -----------------
internal/service/document/document.go | 274 ++
.../service/document/document_artifact.go | 232 ++
internal/service/document/document_crud.go | 436 ++
.../document/document_dataset_update.go | 431 ++
internal/service/document/document_ingest.go | 148 +
internal/service/document/document_list.go | 232 ++
.../service/document/document_metadata.go | 947 +++++
.../document/document_metadata_test.go | 107 +
internal/service/document/document_parse.go | 471 +++
.../service/{ => document}/document_test.go | 100 +-
internal/service/document/document_upload.go | 424 ++
.../{ => document}/document_upload_helpers.go | 2 +-
.../service/document/document_upload_test.go | 38 +
.../service/{ => document}/file2document.go | 43 +-
internal/service/file.go | 1545 -------
internal/service/file/file.go | 114 +
internal/service/{ => file}/file_commit.go | 2 +-
internal/service/file/file_content.go | 249 ++
internal/service/file/file_content_test.go | 60 +
internal/service/file/file_delete.go | 130 +
internal/service/file/file_folder.go | 576 +++
internal/service/file/file_permission_test.go | 51 +
internal/service/{ => file}/file_test.go | 139 +-
internal/service/file/file_upload.go | 281 ++
internal/service/file/file_url.go | 93 +
internal/service/file/file_url_test.go | 37 +
internal/service/file_permission.go | 34 +
.../service/ingestion_task_service_test.go | 24 -
internal/service/memory_message_test.go | 6 +
internal/service/metadata.go | 13 +-
internal/service/parse_types.go | 23 +
internal/service/pipeline_params.go | 151 +
internal/service/pipeline_params_test.go | 81 +
internal/service/skill_space.go | 10 +-
internal/service/team_permission.go | 4 +-
internal/service/test_helpers_test.go | 219 +
internal/utility/ssrf.go | 4 +-
internal/utility/upload.go | 149 +
internal/utility/upload_test.go | 103 +
88 files changed, 10776 insertions(+), 9009 deletions(-)
create mode 100644 internal/ingestion/pipeline/pipeline_params.go
create mode 100644 internal/ingestion/pipeline/pipeline_params_test.go
delete mode 100644 internal/service/dataset.go
create mode 100644 internal/service/dataset/create_test.go
create mode 100644 internal/service/dataset/crud.go
create mode 100644 internal/service/dataset/crud_test.go
create mode 100644 internal/service/dataset/fake_doc_engine_test.go
create mode 100644 internal/service/dataset/helpers.go
create mode 100644 internal/service/dataset/helpers_test.go
create mode 100644 internal/service/dataset/index.go
rename internal/service/{dataset_delete_index_test.go => dataset/index_delete_test.go} (99%)
create mode 100644 internal/service/dataset/ingestion.go
create mode 100644 internal/service/dataset/metadata.go
rename internal/service/{dataset_document_metadata_config_test.go => dataset/metadata_config_test.go} (99%)
create mode 100644 internal/service/dataset/permission.go
create mode 100644 internal/service/dataset/search.go
rename internal/service/{dataset_search_test.go => dataset/search_test.go} (94%)
create mode 100644 internal/service/dataset/service.go
create mode 100644 internal/service/dataset/setup_test.go
create mode 100644 internal/service/dataset/tags.go
rename internal/service/{dataset_aggregate_tags_test.go => dataset/tags_aggregate_test.go} (99%)
rename internal/service/{dataset_list_tags_test.go => dataset/tags_list_test.go} (99%)
rename internal/service/{dataset_rename_tag_test.go => dataset/tags_rename_test.go} (99%)
rename internal/service/{dataset_task_cleanup_test.go => dataset/task_cleanup_test.go} (99%)
create mode 100644 internal/service/dataset/update.go
rename internal/service/{dataset_update_test.go => dataset/update_test.go} (95%)
create mode 100644 internal/service/dataset/utils.go
delete mode 100644 internal/service/dataset_create_test.go
delete mode 100644 internal/service/dataset_parser_test.go
create mode 100644 internal/service/dataset_types.go
delete mode 100644 internal/service/document.go
create mode 100644 internal/service/document/document.go
create mode 100644 internal/service/document/document_artifact.go
create mode 100644 internal/service/document/document_crud.go
create mode 100644 internal/service/document/document_dataset_update.go
create mode 100644 internal/service/document/document_ingest.go
create mode 100644 internal/service/document/document_list.go
create mode 100644 internal/service/document/document_metadata.go
create mode 100644 internal/service/document/document_metadata_test.go
create mode 100644 internal/service/document/document_parse.go
rename internal/service/{ => document}/document_test.go (96%)
create mode 100644 internal/service/document/document_upload.go
rename internal/service/{ => document}/document_upload_helpers.go (97%)
create mode 100644 internal/service/document/document_upload_test.go
rename internal/service/{ => document}/file2document.go (89%)
delete mode 100644 internal/service/file.go
create mode 100644 internal/service/file/file.go
rename internal/service/{ => file}/file_commit.go (99%)
create mode 100644 internal/service/file/file_content.go
create mode 100644 internal/service/file/file_content_test.go
create mode 100644 internal/service/file/file_delete.go
create mode 100644 internal/service/file/file_folder.go
create mode 100644 internal/service/file/file_permission_test.go
rename internal/service/{ => file}/file_test.go (76%)
create mode 100644 internal/service/file/file_upload.go
create mode 100644 internal/service/file/file_url.go
create mode 100644 internal/service/file/file_url_test.go
create mode 100644 internal/service/file_permission.go
create mode 100644 internal/service/parse_types.go
create mode 100644 internal/service/pipeline_params.go
create mode 100644 internal/service/pipeline_params_test.go
create mode 100644 internal/service/test_helpers_test.go
create mode 100644 internal/utility/upload.go
create mode 100644 internal/utility/upload_test.go
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(`?(table|td|caption|tr|th)( [^<>]{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, "]*>.*?(script|style)>`)
- 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)[^>]*>.*?(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))
+ }
+}