Go: add context, part5 (#17392)

Signed-off-by: Jin Hai <haijin.chn@gmail.com>
This commit is contained in:
Jin Hai
2026-07-27 10:20:16 +08:00
committed by GitHub
parent 4b39106cd4
commit f53518c110
182 changed files with 1249 additions and 1066 deletions

View File

@@ -218,7 +218,7 @@ func (h *AgentHandler) DebugComponent(c *gin.Context) {
invokeCtx := runtime.WithState(c.Request.Context(), debugState)
outputs, err := runtime.TrackElapsed(name, func() (map[string]any, error) {
return comp.Invoke(invokeCtx, inputs)
return comp.Invoke(invokeCtx, dao.DB, inputs)
})
if err != nil {
common.ResponseWithCodeData(c, common.CodeServerError, nil, "invoke: "+err.Error())

View File

@@ -25,6 +25,7 @@ import (
"testing"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
_ "ragflow/internal/agent/component" // registers the production factory
agentruntime "ragflow/internal/agent/runtime"
@@ -329,7 +330,7 @@ func TestDebugComponent_HappyPath_Begin(t *testing.T) {
type sysEchoComponent struct{}
func (s *sysEchoComponent) Invoke(ctx context.Context, _ map[string]any) (map[string]any, error) {
func (s *sysEchoComponent) Invoke(ctx context.Context, _ *gorm.DB, _ map[string]any) (map[string]any, error) {
state, _, err := agentruntime.GetStateFromContext[*agentruntime.CanvasState](ctx)
if err != nil {
return nil, err

View File

@@ -17,6 +17,7 @@
package handler
import (
"context"
"encoding/json"
"fmt"
"net/http"
@@ -42,11 +43,11 @@ type DatasetsHandler struct {
}
type searchDatasetsService interface {
SearchDatasets(req *service.SearchDatasetsRequest, userID string) (*service.SearchDatasetsResponse, error)
SearchDatasets(ctx context.Context, req *service.SearchDatasetsRequest, userID string) (*service.SearchDatasetsResponse, error)
}
type searchDatasetService interface {
SearchDataset(datasetID, userID string, req *service.SearchDatasetRequest) (*service.SearchDatasetsResponse, error)
SearchDataset(ctx context.Context, datasetID, userID string, req *service.SearchDatasetRequest) (*service.SearchDatasetsResponse, error)
}
type listDatasetsExt struct {
@@ -116,7 +117,9 @@ func (h *DatasetsHandler) ListDatasets(c *gin.Context) {
ownerIDs = ext.OwnerIDs
}
ctx := c.Request.Context()
data, total, code, err := h.datasetsService.ListDatasets(
ctx,
c.Query("id"),
c.Query("name"),
page,
@@ -154,7 +157,9 @@ func (h *DatasetsHandler) CreateDataset(c *gin.Context) {
return
}
result, code, err := h.datasetsService.CreateDataset(&req, user.ID)
ctx := c.Request.Context()
result, code, err := h.datasetsService.CreateDataset(ctx, &req, user.ID)
if err != nil {
common.ErrorWithCode(c, code, err.Error())
return
@@ -244,8 +249,10 @@ func (h *DatasetsHandler) GetMetadataConfig(c *gin.Context) {
return
}
ctx := c.Request.Context()
datasetID := c.Param("dataset_id")
result, code, err := h.datasetsService.GetMetadataConfig(datasetID, user.ID)
result, code, err := h.datasetsService.GetMetadataConfig(ctx, datasetID, user.ID)
if err != nil {
common.ErrorWithCode(c, code, err.Error())
return
@@ -268,7 +275,9 @@ func (h *DatasetsHandler) UpdateMetadataConfig(c *gin.Context) {
return
}
result, code, err := h.datasetsService.UpdateMetadataConfig(datasetID, user.ID, &req)
ctx := c.Request.Context()
result, code, err := h.datasetsService.UpdateMetadataConfig(ctx, datasetID, user.ID, &req)
if err != nil {
common.ErrorWithCode(c, code, err.Error())
return
@@ -336,7 +345,9 @@ func (h *DatasetsHandler) ListIngestionLogs(c *gin.Context) {
logType := c.DefaultQuery("log_type", "dataset")
keywords := c.Query("keywords")
result, code, err := h.datasetsService.ListIngestionLogs(datasetID, user.ID, page, pageSize, orderby, desc, operationStatus, createDateFrom, createDateTo, logType, keywords)
ctx := c.Request.Context()
result, code, err := h.datasetsService.ListIngestionLogs(ctx, datasetID, user.ID, page, pageSize, orderby, desc, operationStatus, createDateFrom, createDateTo, logType, keywords)
if err != nil {
common.ErrorWithCode(c, code, err.Error())
return
@@ -353,9 +364,11 @@ func (h *DatasetsHandler) GetIngestionLog(c *gin.Context) {
return
}
ctx := c.Request.Context()
datasetID := c.Param("dataset_id")
logID := c.Param("log_id")
result, code, err := h.datasetsService.GetIngestionLog(datasetID, user.ID, logID)
result, code, err := h.datasetsService.GetIngestionLog(ctx, datasetID, user.ID, logID)
if err != nil {
common.ErrorWithCode(c, code, err.Error())
return
@@ -388,7 +401,9 @@ func (h *DatasetsHandler) DeleteDatasets(c *gin.Context) {
ids = *req.IDs
}
result, code, err := h.datasetsService.DeleteDatasets(ids, req.DeleteAll, user.ID)
ctx := c.Request.Context()
result, code, err := h.datasetsService.DeleteDatasets(ctx, ids, req.DeleteAll, user.ID)
if err != nil {
common.ErrorWithCode(c, code, err.Error())
return
@@ -513,8 +528,10 @@ func (h *DatasetsHandler) ListTags(c *gin.Context) {
return
}
ctx := c.Request.Context()
datasetID := strings.TrimSpace(c.Param("dataset_id"))
result, code, err := h.datasetsService.ListTags(datasetID, user.ID)
result, code, err := h.datasetsService.ListTags(ctx, datasetID, user.ID)
if err != nil {
common.ErrorWithCode(c, code, err.Error())
return
@@ -559,7 +576,9 @@ func (h *DatasetsHandler) RenameTag(c *gin.Context) {
return
}
result, code, err := h.datasetsService.RenameTag(datasetID, user.ID, req.FromTag, req.ToTag)
ctx := c.Request.Context()
result, code, err := h.datasetsService.RenameTag(ctx, datasetID, user.ID, req.FromTag, req.ToTag)
if err != nil {
common.ErrorWithCode(c, code, err.Error())
return
@@ -756,7 +775,9 @@ func (h *DatasetsHandler) AggregateTags(c *gin.Context) {
return
}
result, code, err := h.datasetsService.AggregateTags(datasetIDs, user.ID)
ctx := c.Request.Context()
result, code, err := h.datasetsService.AggregateTags(ctx, datasetIDs, user.ID)
if err != nil {
common.ErrorWithCode(c, code, err.Error())
return
@@ -815,8 +836,10 @@ func (h *DatasetsHandler) TraceIndex(c *gin.Context) {
return
}
ctx := c.Request.Context()
indexType := strings.ToLower(strings.TrimSpace(c.Query("type")))
result, code, err := h.datasetsService.TraceIndex(datasetID, userID, indexType)
result, code, err := h.datasetsService.TraceIndex(ctx, datasetID, userID, indexType)
if err != nil {
common.ErrorWithCode(c, code, err.Error())
return
@@ -861,7 +884,9 @@ func (h *DatasetsHandler) DeleteIndex(c *gin.Context) {
wipe = false
}
code, err := h.datasetsService.DeleteIndex(userID, datasetID, indexType, wipe)
ctx := c.Request.Context()
code, err := h.datasetsService.DeleteIndex(ctx, userID, datasetID, indexType, wipe)
if err != nil {
common.ErrorWithCode(c, code, err.Error())
return
@@ -905,15 +930,16 @@ func (h *DatasetsHandler) ListMetadataFlattened(c *gin.Context) {
return
}
ctx := c.Request.Context()
// Check access for each dataset
for _, datasetID := range datasetIDs {
if !h.datasetsService.Accessible(datasetID, user.ID) {
if !h.datasetsService.Accessible(ctx, datasetID, user.ID) {
common.ResponseWithCodeData(c, common.CodeAuthenticationError, nil, "No authorization for dataset: "+datasetID)
return
}
}
flattenedMeta, err := h.metadataService.GetFlattedMetaByKBs(datasetIDs)
flattenedMeta, err := h.metadataService.GetFlattedMetaByKBs(ctx, datasetIDs)
if err != nil {
common.ResponseWithCodeData(c, common.CodeServerError, nil, "Failed to get metadata: "+err.Error())
return
@@ -1026,7 +1052,9 @@ func (h *DatasetsHandler) SearchDatasets(c *gin.Context) {
return
}
resp, err := searchService.SearchDatasets(&req, user.ID)
ctx := c.Request.Context()
resp, err := searchService.SearchDatasets(ctx, &req, user.ID)
if err != nil {
common.ResponseWithCodeData(c, common.CodeDataError, nil, err.Error())
return
@@ -1080,7 +1108,9 @@ func (h *DatasetsHandler) SearchDataset(c *gin.Context) {
return
}
resp, err := searchService.SearchDataset(datasetID, user.ID, &req)
ctx := c.Request.Context()
resp, err := searchService.SearchDataset(ctx, datasetID, user.ID, &req)
if err != nil {
common.ResponseWithCodeData(c, common.CodeDataError, nil, err.Error())
return

View File

@@ -1,6 +1,7 @@
package handler
import (
"context"
"encoding/json"
"errors"
"net/http"
@@ -23,7 +24,7 @@ type fakeSearchDatasetService struct {
err error
}
func (f *fakeSearchDatasetService) SearchDataset(datasetID, userID string, req *service.SearchDatasetRequest) (*service.SearchDatasetsResponse, error) {
func (f *fakeSearchDatasetService) SearchDataset(ctx context.Context, datasetID, userID string, req *service.SearchDatasetRequest) (*service.SearchDatasetsResponse, error) {
f.datasetID = datasetID
f.userID = userID
f.req = req
@@ -37,7 +38,7 @@ type fakeSearchDatasetsService struct {
err error
}
func (f *fakeSearchDatasetsService) SearchDatasets(req *service.SearchDatasetsRequest, userID string) (*service.SearchDatasetsResponse, error) {
func (f *fakeSearchDatasetsService) SearchDatasets(ctx context.Context, req *service.SearchDatasetsRequest, userID string) (*service.SearchDatasetsResponse, error) {
f.userID = userID
f.req = req
return f.resp, f.err

View File

@@ -43,21 +43,21 @@ import (
// KBServiceIface abstracts KnowledgebaseService for the Dify handler.
type KBServiceIface interface {
GetByID(kbID string) (*entity.Knowledgebase, error)
Accessible(kbID, userID string) bool
GetByID(ctx context.Context, kbID string) (*entity.Knowledgebase, error)
Accessible(ctx context.Context, kbID, userID string) bool
}
// ModelServiceIface abstracts ModelProviderService for the Dify handler.
type ModelServiceIface interface {
GetEmbeddingModel(tenantID, embdID string) (*modelModule.EmbeddingModel, error)
GetChatModel(tenantID, compositeModelName string) (*modelModule.ChatModel, error)
GetEmbeddingModel(ctx context.Context, tenantID, embdID string) (*modelModule.EmbeddingModel, error)
GetChatModel(ctx context.Context, tenantID, compositeModelName string) (*modelModule.ChatModel, error)
}
// MetadataServiceIface abstracts MetadataService for the Dify handler.
type MetadataServiceIface interface {
GetFlattedMetaByKBs(kbIDs []string) (common.MetaData, error)
SearchMetadataByKBs(kbIDs []string, size int) (*service.SearchMetadataResponse, error)
LabelQuestion(question string, kbs []*entity.Knowledgebase) map[string]float64
GetFlattedMetaByKBs(ctx context.Context, kbIDs []string) (common.MetaData, error)
SearchMetadataByKBs(ctx context.Context, kbIDs []string, size int) (*service.SearchMetadataResponse, error)
LabelQuestion(ctx context.Context, question string, kbs []*entity.Knowledgebase) map[string]float64
}
// RetrievalServiceIface abstracts RetrievalService for the Dify handler.
@@ -204,7 +204,8 @@ func (h *DifyRetrievalHandler) Retrieval(c *gin.Context) {
return
}
kb, err := h.kbSvc.GetByID(req.KnowledgeID)
ctx := c.Request.Context()
kb, err := h.kbSvc.GetByID(ctx, req.KnowledgeID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
common.ResponseWithHttpCodeData(c, http.StatusNotFound, common.CodeNotFound, nil, "Knowledge base not found!")
@@ -214,7 +215,7 @@ func (h *DifyRetrievalHandler) Retrieval(c *gin.Context) {
return
}
if !h.kbSvc.Accessible(req.KnowledgeID, user.ID) {
if !h.kbSvc.Accessible(ctx, req.KnowledgeID, user.ID) {
common.ResponseWithHttpCodeData(c, http.StatusUnauthorized, common.CodeAuthenticationError, nil, "No authorization")
return
}
@@ -234,14 +235,14 @@ func (h *DifyRetrievalHandler) Retrieval(c *gin.Context) {
}
// Get embedding model
embModel, err := h.modelSvc.GetEmbeddingModel(kb.TenantID, kb.EmbdID)
embModel, err := h.modelSvc.GetEmbeddingModel(ctx, kb.TenantID, kb.EmbdID)
if err != nil {
common.ResponseWithHttpCodeData(c, http.StatusInternalServerError, common.CodeServerError, nil, fmt.Sprintf("failed to get embedding model: %v", err))
return
}
// Metadata filter
metas, metaErr := h.metadataSvc.GetFlattedMetaByKBs([]string{req.KnowledgeID})
metas, metaErr := h.metadataSvc.GetFlattedMetaByKBs(ctx, []string{req.KnowledgeID})
docIDs := make([]string, 0)
if metaErr == nil && req.MetadataCondition != nil {
logic := req.MetadataCondition.Logic
@@ -257,7 +258,7 @@ func (h *DifyRetrievalHandler) Retrieval(c *gin.Context) {
// Label question for rank features
kbs := []*entity.Knowledgebase{kb}
rankFeature := h.metadataSvc.LabelQuestion(req.Query, kbs)
rankFeature := h.metadataSvc.LabelQuestion(ctx, req.Query, kbs)
// Chunk retrieval
sr := &nlp.RetrievalRequest{
@@ -290,7 +291,7 @@ func (h *DifyRetrievalHandler) Retrieval(c *gin.Context) {
// KG retrieval (optional)
if req.UseKG {
chatModel, kgErr := h.modelSvc.GetChatModel(kb.TenantID, "")
chatModel, kgErr := h.modelSvc.GetChatModel(ctx, kb.TenantID, "")
if kgErr != nil {
common.Warn("KG retrieval: failed to get chat model", zap.String("kbID", req.KnowledgeID), zap.Error(kgErr))
} else if chatModel != nil {
@@ -322,7 +323,6 @@ func (h *DifyRetrievalHandler) Retrieval(c *gin.Context) {
allDocIDs = append(allDocIDs, id)
}
ctx := c.Request.Context()
docMap := make(map[string]*entity.Document)
if len(allDocIDs) > 0 {
var docs []*entity.Document
@@ -337,7 +337,7 @@ func (h *DifyRetrievalHandler) Retrieval(c *gin.Context) {
}
metaMap := make(map[string]map[string]interface{})
metaResult, err := h.metadataSvc.SearchMetadataByKBs([]string{kb.ID}, 10000)
metaResult, err := h.metadataSvc.SearchMetadataByKBs(ctx, []string{kb.ID}, 10000)
if err == nil {
for _, metadata := range metaResult.MetadataRecords {
docID, ok := service.ExtractDocumentID(metadata)

View File

@@ -42,63 +42,63 @@ import (
type mockKBService struct {
KBServiceIface
getByIDFn func(kbID string) (*entity.Knowledgebase, error)
accessibleFn func(kbID, userID string) bool
getByIDFn func(ctx context.Context, kbID string) (*entity.Knowledgebase, error)
accessibleFn func(ctx context.Context, kbID, userID string) bool
}
func (m *mockKBService) GetByID(kbID string) (*entity.Knowledgebase, error) {
func (m *mockKBService) GetByID(ctx context.Context, kbID string) (*entity.Knowledgebase, error) {
if m.getByIDFn != nil {
return m.getByIDFn(kbID)
return m.getByIDFn(ctx, kbID)
}
return &entity.Knowledgebase{
ID: kbID, TenantID: "tenant1", EmbdID: "text-embedding",
}, nil
}
func (m *mockKBService) Accessible(kbID, userID string) bool {
func (m *mockKBService) Accessible(ctx context.Context, kbID, userID string) bool {
if m.accessibleFn != nil {
return m.accessibleFn(kbID, userID)
return m.accessibleFn(ctx, kbID, userID)
}
return true
}
type mockModelService struct {
ModelServiceIface
getEmbeddingFn func(tenantID, embdID string) (*modelModule.EmbeddingModel, error)
getChatModelFn func(tenantID, llmID string) (*modelModule.ChatModel, error)
getEmbeddingFn func(ctx context.Context, tenantID, embdID string) (*modelModule.EmbeddingModel, error)
getChatModelFn func(ctx context.Context, tenantID, llmID string) (*modelModule.ChatModel, error)
}
func (m *mockModelService) GetEmbeddingModel(tenantID, embdID string) (*modelModule.EmbeddingModel, error) {
func (m *mockModelService) GetEmbeddingModel(ctx context.Context, tenantID, embdID string) (*modelModule.EmbeddingModel, error) {
if m.getEmbeddingFn != nil {
return m.getEmbeddingFn(tenantID, embdID)
return m.getEmbeddingFn(ctx, tenantID, embdID)
}
return &modelModule.EmbeddingModel{}, nil
}
func (m *mockModelService) GetChatModel(tenantID, llmID string) (*modelModule.ChatModel, error) {
func (m *mockModelService) GetChatModel(ctx context.Context, tenantID, llmID string) (*modelModule.ChatModel, error) {
if m.getChatModelFn != nil {
return m.getChatModelFn(tenantID, llmID)
return m.getChatModelFn(ctx, tenantID, llmID)
}
return &modelModule.ChatModel{}, nil
}
type mockMetadataService struct {
MetadataServiceIface
getFlattedMetaFn func(kbIDs []string) (common.MetaData, error)
searchMetadataByKBs func(kbIDs []string, size int) (*service.SearchMetadataResponse, error)
labelQuestionFn func(question string, kbs []*entity.Knowledgebase) map[string]float64
getFlattedMetaFn func(ctx context.Context, kbIDs []string) (common.MetaData, error)
searchMetadataByKBs func(ctx context.Context, kbIDs []string, size int) (*service.SearchMetadataResponse, error)
labelQuestionFn func(ctx context.Context, question string, kbs []*entity.Knowledgebase) map[string]float64
}
func (m *mockMetadataService) GetFlattedMetaByKBs(kbIDs []string) (common.MetaData, error) {
func (m *mockMetadataService) GetFlattedMetaByKBs(ctx context.Context, kbIDs []string) (common.MetaData, error) {
if m.getFlattedMetaFn != nil {
return m.getFlattedMetaFn(kbIDs)
return m.getFlattedMetaFn(ctx, kbIDs)
}
return common.MetaData{}, nil
}
func (m *mockMetadataService) SearchMetadataByKBs(kbIDs []string, size int) (*service.SearchMetadataResponse, error) {
func (m *mockMetadataService) SearchMetadataByKBs(ctx context.Context, kbIDs []string, size int) (*service.SearchMetadataResponse, error) {
if m.searchMetadataByKBs != nil {
return m.searchMetadataByKBs(kbIDs, size)
return m.searchMetadataByKBs(ctx, kbIDs, size)
}
return &service.SearchMetadataResponse{
MetadataRecords: []map[string]interface{}{
@@ -107,9 +107,9 @@ func (m *mockMetadataService) SearchMetadataByKBs(kbIDs []string, size int) (*se
}, nil
}
func (m *mockMetadataService) LabelQuestion(question string, kbs []*entity.Knowledgebase) map[string]float64 {
func (m *mockMetadataService) LabelQuestion(ctx context.Context, question string, kbs []*entity.Knowledgebase) map[string]float64 {
if m.labelQuestionFn != nil {
return m.labelQuestionFn(question, kbs)
return m.labelQuestionFn(ctx, question, kbs)
}
return nil
}
@@ -275,7 +275,7 @@ func TestDifyRetrieval_MissingArgs(t *testing.T) {
func TestDifyRetrieval_KBNotFound(t *testing.T) {
h, r := setupDifyTest("user1")
h.kbSvc = &mockKBService{
getByIDFn: func(kbID string) (*entity.Knowledgebase, error) {
getByIDFn: func(ctx context.Context, kbID string) (*entity.Knowledgebase, error) {
return nil, gorm.ErrRecordNotFound
},
}
@@ -304,7 +304,7 @@ func TestDifyRetrieval_NoAuth(t *testing.T) {
func TestDifyRetrieval_Unauthorized(t *testing.T) {
h, r := setupDifyTest("user1")
h.kbSvc = &mockKBService{
accessibleFn: func(kbID, userID string) bool { return false },
accessibleFn: func(ctx context.Context, kbID, userID string) bool { return false },
}
w := httptest.NewRecorder()
body := `{"knowledge_id": "kb1", "query": "test"}`
@@ -319,7 +319,7 @@ func TestDifyRetrieval_Unauthorized(t *testing.T) {
func TestDifyRetrieval_WithMetadataFilter(t *testing.T) {
h, r := setupDifyTest("user1")
h.metadataSvc = &mockMetadataService{
getFlattedMetaFn: func(kbIDs []string) (common.MetaData, error) {
getFlattedMetaFn: func(ctx context.Context, kbIDs []string) (common.MetaData, error) {
return common.MetaData{}, nil
},
}
@@ -347,7 +347,7 @@ func TestDifyRetrieval_InvalidJSON(t *testing.T) {
func TestDifyRetrieval_UseKG(t *testing.T) {
h, r := setupDifyTest("user1")
h.metadataSvc = &mockMetadataService{
labelQuestionFn: func(question string, kbs []*entity.Knowledgebase) map[string]float64 {
labelQuestionFn: func(ctx context.Context, question string, kbs []*entity.Knowledgebase) map[string]float64 {
return map[string]float64{"tag_1": 0.8}
},
}
@@ -366,7 +366,7 @@ func strPtr(s string) *string { return &s }
func TestDifyRetrieval_KBDBError(t *testing.T) {
h, r := setupDifyTest("user1")
h.kbSvc = &mockKBService{
getByIDFn: func(kbID string) (*entity.Knowledgebase, error) {
getByIDFn: func(ctx context.Context, kbID string) (*entity.Knowledgebase, error) {
return nil, errors.New("connection refused")
},
}

View File

@@ -308,7 +308,7 @@ func (h *DocumentHandler) UpdateDocument(c *gin.Context) {
common.ResponseWithHttpCodeData(c, http.StatusBadRequest, 1, nil, "document not found!")
return
}
if !h.datasetService.Accessible(doc.KbID, user.ID) {
if !h.datasetService.Accessible(ctx, doc.KbID, user.ID) {
common.ResponseWithCodeData(c, common.CodeAuthenticationError, nil, "No authorization.")
return
}
@@ -361,7 +361,7 @@ func (h *DocumentHandler) DeleteDocument(c *gin.Context) {
common.ResponseWithHttpCodeData(c, http.StatusBadRequest, 1, nil, "document not found!")
return
}
if !h.datasetService.Accessible(doc.KbID, user.ID) {
if !h.datasetService.Accessible(ctx, doc.KbID, user.ID) {
common.ResponseWithCodeData(c, common.CodeAuthenticationError, nil, "No authorization.")
return
}
@@ -505,7 +505,8 @@ func (h *DocumentHandler) ListDocuments(c *gin.Context) {
userID := c.GetString("user_id")
if !h.datasetService.Accessible(datasetID, userID) {
ctx := c.Request.Context()
if !h.datasetService.Accessible(ctx, datasetID, userID) {
common.ResponseWithCodeData(c, common.CodeDataError, nil, fmt.Sprintf("You don't own the dataset %s.", datasetID))
return
}
@@ -528,7 +529,6 @@ func (h *DocumentHandler) ListDocuments(c *gin.Context) {
return
}
ctx := c.Request.Context()
if c.Query("type") == "filter" {
filters, total, err := h.documentService.GetDocumentFiltersByDatasetID(ctx, opts)
if err != nil {
@@ -830,7 +830,8 @@ func (h *DocumentHandler) UploadDocuments(c *gin.Context) {
datasetID := c.Param("dataset_id")
uploadType := strings.ToLower(c.DefaultQuery("type", "local"))
kb, err := h.datasetService.GetKnowledgebaseByID(datasetID)
ctx := c.Request.Context()
kb, err := h.datasetService.GetKnowledgebaseByID(ctx, datasetID)
if err != nil || kb == nil {
common.ResponseWithCodeData(c, common.CodeDataError, nil, fmt.Sprintf("Can't find the dataset with ID %s!", datasetID))
return
@@ -1236,7 +1237,7 @@ func (h *DocumentHandler) SetMeta(c *gin.Context) {
common.ResponseWithHttpCodeData(c, http.StatusBadRequest, 1, nil, "document not found")
return
}
if !h.datasetService.Accessible(doc.KbID, user.ID) {
if !h.datasetService.Accessible(ctx, doc.KbID, user.ID) {
common.ResponseWithCodeData(c, common.CodeAuthenticationError, nil, "No authorization.")
return
}
@@ -1331,7 +1332,7 @@ func (h *DocumentHandler) DeleteMeta(c *gin.Context) {
common.ResponseWithHttpCodeData(c, http.StatusBadRequest, 1, nil, "document not found")
return
}
if !h.datasetService.Accessible(doc.KbID, user.ID) {
if !h.datasetService.Accessible(ctx, doc.KbID, user.ID) {
common.ResponseWithCodeData(c, common.CodeAuthenticationError, nil, "No authorization.")
return
}
@@ -1392,13 +1393,13 @@ func (h *DocumentHandler) ListIngestionTasks(c *gin.Context) {
var parseResult []*entity.IngestionTask
var err error
ctx := c.Request.Context()
if req.DatasetID != nil {
if !h.datasetService.Accessible(*req.DatasetID, userID) {
if !h.datasetService.Accessible(ctx, *req.DatasetID, userID) {
common.ResponseWithCodeData(c, common.CodeAuthenticationError, nil, "No authorization to access the dataset.")
return
}
}
ctx := c.Request.Context()
parseResult, err = h.documentService.ListIngestionTasks(ctx, userID, req.DatasetID, 0, 0)
if err != nil {
common.ResponseWithCodeData(c, IngestionTaskErrorCode(err), nil, err.Error())
@@ -1422,12 +1423,12 @@ func (h *DocumentHandler) StartIngestionTask(c *gin.Context) {
}
userID := c.GetString("user_id")
if !h.datasetService.Accessible(datasetID, userID) {
ctx := c.Request.Context()
if !h.datasetService.Accessible(ctx, datasetID, userID) {
common.ResponseWithCodeData(c, common.CodeDataError, nil, fmt.Sprintf("You don't own the dataset %s.", datasetID))
return
}
ctx := c.Request.Context()
parseResult, err := h.documentService.IngestDocuments(ctx, datasetID, userID, req.DocumentIDs)
if err != nil {
common.ResponseWithCodeData(c, common.CodeExceptionError, nil, err.Error())
@@ -1505,12 +1506,12 @@ func (h *DocumentHandler) ParseDocuments(c *gin.Context) {
}
userID := c.GetString("user_id")
if !h.datasetService.Accessible(datasetID, userID) {
ctx := c.Request.Context()
if !h.datasetService.Accessible(ctx, datasetID, userID) {
common.ResponseWithCodeData(c, common.CodeAuthenticationError, nil, "No authorization to access the dataset.")
return
}
ctx := c.Request.Context()
parseResult, err := h.documentService.ParseDocuments(ctx, datasetID, userID, req.Documents)
if err != nil {
common.ResponseWithCodeData(c, common.CodeExceptionError, nil, err.Error())
@@ -1538,12 +1539,12 @@ func (h *DocumentHandler) StopParseDocuments(c *gin.Context) {
}
userID := c.GetString("user_id")
if !h.datasetService.Accessible(datasetID, userID) {
ctx := c.Request.Context()
if !h.datasetService.Accessible(ctx, datasetID, userID) {
common.ResponseWithCodeData(c, common.CodeDataError, nil, fmt.Sprintf("You don't own the dataset %s.", datasetID))
return
}
ctx := c.Request.Context()
result, err := h.documentService.StopParseDocuments(ctx, datasetID, req.DocumentIDs)
if err != nil {
common.ResponseWithCodeData(c, common.CodeExceptionError, nil, err.Error())
@@ -1564,7 +1565,8 @@ func (h *DocumentHandler) MetadataSummaryByDataset(c *gin.Context) {
common.ErrorWithCode(c, common.CodeServerError, "dataset_id is required")
return
}
if !h.datasetService.Accessible(datasetID, user.ID) {
ctx := c.Request.Context()
if !h.datasetService.Accessible(ctx, datasetID, user.ID) {
common.ErrorWithCode(c, common.CodeServerError, "You don't own the dataset "+datasetID)
return
}
@@ -1573,7 +1575,7 @@ func (h *DocumentHandler) MetadataSummaryByDataset(c *gin.Context) {
if docIDsParam := c.Query("doc_ids"); docIDsParam != "" {
docIDS = strings.Split(docIDsParam, ",")
}
ctx := c.Request.Context()
summary, err := h.documentService.GetMetadataSummary(ctx, datasetID, docIDS)
if err != nil {
common.ResponseWithHttpCodeData(c, http.StatusInternalServerError, common.CodeServerError, nil, "Failed to get metadata summary"+err.Error())
@@ -1710,7 +1712,8 @@ func (h *DocumentHandler) handleBatchUpdateDocumentMetadatas(c *gin.Context) {
common.ResponseWithCodeData(c, common.CodeArgumentError, nil, "dataset_id is required")
return
}
if !h.datasetService.Accessible(datasetID, user.ID) {
ctx := c.Request.Context()
if !h.datasetService.Accessible(ctx, datasetID, user.ID) {
common.ResponseWithCodeData(c, common.CodeDataError, nil, "You don't own the dataset "+datasetID+".")
return
}
@@ -1742,7 +1745,7 @@ func (h *DocumentHandler) handleBatchUpdateDocumentMetadatas(c *gin.Context) {
Updates: updates,
Deletes: deletes,
}
ctx := c.Request.Context()
resp, code, err := h.documentService.BatchUpdateDocumentMetadatas(ctx, datasetID, req.Selector, req.Updates, req.Deletes)
if err != nil {
common.ErrorWithCode(c, code, err.Error())

View File

@@ -94,7 +94,7 @@ func CommitFolderResolver(h *FileCommitHandler, entityType, urlParam string) gin
}
func (h *FileCommitHandler) resolveDatasetFolderID(ctx context.Context, datasetID string) (string, error) {
kb, err := h.kbDAO.GetByID(datasetID)
kb, err := h.kbDAO.GetByID(ctx, dao.DB, datasetID)
if err != nil {
return "", err
}

View File

@@ -41,18 +41,18 @@ type MCPRetrievalService interface {
// MCPServerHandler handles MCP protocol requests (JSON-RPC over HTTP).
// It exposes RAGFlow capabilities as MCP tools to external AI clients.
type MCPServerHandler struct {
listDatasetsFunc func(userID string, page, pageSize int, orderby string, desc bool) ([]map[string]interface{}, int64, error)
listChatsFunc func(userID string, page, pageSize int, orderby string, desc bool) ([]map[string]interface{}, int64, error)
retrievalFunc func(userID string, req mcp.RetrievalRequest) (string, error)
listDatasetsFunc func(ctx context.Context, userID string, page, pageSize int, orderby string, desc bool) ([]map[string]interface{}, int64, error)
listChatsFunc func(ctx context.Context, userID string, page, pageSize int, orderby string, desc bool) ([]map[string]interface{}, int64, error)
retrievalFunc func(ctx context.Context, userID string, req mcp.RetrievalRequest) (string, error)
}
// NewMCPServerHandler creates a new MCPServerHandler.
// The service functions are passed as closures to avoid importing the service
// package directly from the handler layer.
func NewMCPServerHandler(
listDatasetsFunc func(userID string, page, pageSize int, orderby string, desc bool) ([]map[string]interface{}, int64, error),
listChatsFunc func(userID string, page, pageSize int, orderby string, desc bool) ([]map[string]interface{}, int64, error),
retrievalFunc func(userID string, req mcp.RetrievalRequest) (string, error),
listDatasetsFunc func(ctx context.Context, userID string, page, pageSize int, orderby string, desc bool) ([]map[string]interface{}, int64, error),
listChatsFunc func(ctx context.Context, userID string, page, pageSize int, orderby string, desc bool) ([]map[string]interface{}, int64, error),
retrievalFunc func(ctx context.Context, userID string, req mcp.RetrievalRequest) (string, error),
) *MCPServerHandler {
return &MCPServerHandler{
listDatasetsFunc: listDatasetsFunc,
@@ -95,8 +95,9 @@ func (h *MCPServerHandler) HandleMCP(c *gin.Context) {
h.retrievalFunc,
)
ctx := c.Request.Context()
server := mcp.NewServer(connector)
respBody, hasResponse, err := server.HandleRequest(body)
respBody, hasResponse, err := server.HandleRequest(ctx, body)
if err != nil {
common.ResponseWithHttpCodeData(c, http.StatusInternalServerError, common.CodeBadRequest, nil, "MCP server error: "+err.Error())
return
@@ -114,8 +115,8 @@ 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 *dataset.DatasetService, userID string, page, pageSize int, orderby string, desc bool) ([]map[string]interface{}, int64, error) {
data, total, _, err := ds.ListDatasets(
func MCPListDatasets(ctx context.Context, ds *dataset.DatasetService, userID string, page, pageSize int, orderby string, desc bool) ([]map[string]interface{}, int64, error) {
data, total, _, err := ds.ListDatasets(ctx,
"", "", page, pageSize, orderby, desc,
"", nil, "", userID,
)
@@ -143,7 +144,7 @@ func MCPListChats(ctx context.Context, chatService *service.ChatService, userID
// 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 *dataset.DatasetService, userID string, req mcp.RetrievalRequest) (string, error) {
func MCPRetrieval(ctx context.Context, 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
@@ -151,7 +152,7 @@ func MCPRetrieval(ds *dataset.DatasetService, userID string, req mcp.RetrievalRe
const maxPageSize = 100
page := 1
for {
data, _, _, err := ds.ListDatasets(
data, _, _, err := ds.ListDatasets(ctx,
"", "", page, maxPageSize, "create_time", true,
"", nil, "", userID,
)
@@ -213,7 +214,7 @@ func MCPRetrieval(ds *dataset.DatasetService, userID string, req mcp.RetrievalRe
searchReq.Keyword = &v
}
resp, err := ds.SearchDatasets(searchReq, userID)
resp, err := ds.SearchDatasets(ctx, searchReq, userID)
if err != nil {
return "", err
}

View File

@@ -382,8 +382,8 @@ func (h *SearchHandler) Completion(c *gin.Context) {
if searchSvc == nil {
searchSvc = service.NewSearchService()
}
plan, code, err := searchSvc.PrepareCompletion(user.ID, c.Param("search_id"), &req)
ctx := c.Request.Context()
plan, code, err := searchSvc.PrepareCompletion(ctx, user.ID, c.Param("search_id"), &req)
if err != nil {
if code == common.CodeAuthenticationError {
common.ResponseWithCodeData(c, code, false, err.Error())

View File

@@ -257,8 +257,9 @@ func (h *TenantHandler) CreateChunkStore(c *gin.Context) {
return
}
ctx := c.Request.Context()
// Check authorization - user must have access to this kb
if !h.datasetService.Accessible(req.KBID, user.ID) {
if !h.datasetService.Accessible(ctx, req.KBID, user.ID) {
common.ResponseWithCodeData(c, common.CodeAuthenticationError, nil, "No authorization.")
return
}
@@ -267,7 +268,7 @@ func (h *TenantHandler) CreateChunkStore(c *gin.Context) {
KBID: req.KBID,
VectorSize: req.VectorSize,
}
result, code, err := h.tenantService.CreateChunkStore(serviceReq)
result, code, err := h.tenantService.CreateChunkStore(ctx, serviceReq)
if err != nil {
common.ErrorWithCode(c, code, err.Error())
return
@@ -298,6 +299,7 @@ func (h *TenantHandler) DeleteChunkStore(c *gin.Context) {
return
}
ctx := c.Request.Context()
var req DeleteChunkTableRequest
if err := c.ShouldBindJSON(&req); err != nil {
common.ResponseWithCodeData(c, common.CodeDataError, nil, err.Error())
@@ -305,12 +307,12 @@ func (h *TenantHandler) DeleteChunkStore(c *gin.Context) {
}
// Check authorization
if !h.datasetService.Accessible(req.KBID, user.ID) {
if !h.datasetService.Accessible(ctx, req.KBID, user.ID) {
common.ResponseWithCodeData(c, common.CodeAuthenticationError, nil, "No authorization.")
return
}
code, err := h.tenantService.DeleteChunkStore(req.KBID)
code, err := h.tenantService.DeleteChunkStore(ctx, req.KBID)
if err != nil {
common.ErrorWithCode(c, code, err.Error())
return