Go: add context, part4 (#17381)

Signed-off-by: Jin Hai <haijin.chn@gmail.com>
This commit is contained in:
Jin Hai
2026-07-26 18:31:56 +08:00
committed by GitHub
parent d2a769c532
commit 53e83dcadf
23 changed files with 370 additions and 315 deletions

View File

@@ -1019,7 +1019,8 @@ func (h *Handler) RemoveIngestionTasks(c *gin.Context) {
}
if req.Email == nil && req.Status == nil {
tasks, err := h.service.RemoveIngestionTasks(req.Tasks)
ctx := c.Request.Context()
tasks, err := h.service.RemoveIngestionTasks(ctx, req.Tasks)
if err != nil {
common.ErrorWithCode(c, handler.IngestionTaskErrorCode(err), err.Error())
return
@@ -1050,7 +1051,8 @@ func (h *Handler) StopIngestionTasks(c *gin.Context) {
}
if req.Email == nil && req.Status == nil {
tasks, err := h.service.StopIngestionTasks(req.Tasks)
ctx := c.Request.Context()
tasks, err := h.service.StopIngestionTasks(ctx, req.Tasks)
if err != nil {
common.ErrorWithCode(c, handler.IngestionTaskErrorCode(err), err.Error())
return
@@ -1085,7 +1087,8 @@ func (h *Handler) ListIngestionTasks(c *gin.Context) {
var tasks []map[string]interface{}
var req ListIngestionTasksRequest
if err = c.ShouldBindJSON(&req); err != nil {
tasks, err = h.service.ListIngestionTasks()
ctx := c.Request.Context()
tasks, err = h.service.ListIngestionTasks(ctx)
} else {
tasks, err = h.service.ListIngestionTasksByCondition(req.Email, req.Status)
}

View File

@@ -104,16 +104,16 @@ func (s *Service) Logout(user interface{}) error {
}
// ListIngestionTasks list all ingestion tasks for admin user
func (s *Service) ListIngestionTasks() ([]map[string]interface{}, error) {
return s.ingestionTaskSvc.ListAllForAdmin()
func (s *Service) ListIngestionTasks(ctx context.Context) ([]map[string]interface{}, error) {
return s.ingestionTaskSvc.ListAllForAdmin(ctx)
}
func (s *Service) RemoveIngestionTasks(tasks []string) ([]map[string]string, error) {
return s.ingestionTaskSvc.RemoveMany(tasks, nil)
func (s *Service) RemoveIngestionTasks(ctx context.Context, tasks []string) ([]map[string]string, error) {
return s.ingestionTaskSvc.RemoveMany(ctx, tasks, nil)
}
func (s *Service) StopIngestionTasks(tasks []string) ([]*entity.IngestionTask, error) {
return s.ingestionTaskSvc.RequestStopMany(tasks, nil)
func (s *Service) StopIngestionTasks(ctx context.Context, tasks []string) ([]*entity.IngestionTask, error) {
return s.ingestionTaskSvc.RequestStopMany(ctx, tasks, nil)
}
// GetUserByToken get user by access token

View File

@@ -17,7 +17,10 @@
package dao
import (
"context"
"ragflow/internal/entity"
"gorm.io/gorm"
)
// FileCommitDAO file commit data access object
@@ -29,9 +32,9 @@ func NewFileCommitDAO() *FileCommitDAO {
}
// GetByID gets a file commit by ID
func (dao *FileCommitDAO) GetByID(id string) (*entity.FileCommit, error) {
func (dao *FileCommitDAO) GetByID(ctx context.Context, db *gorm.DB, id string) (*entity.FileCommit, error) {
var commit entity.FileCommit
err := DB.Where("id = ?", id).First(&commit).Error
err := db.WithContext(ctx).Where("id = ?", id).First(&commit).Error
if err != nil {
return nil, err
}
@@ -39,19 +42,19 @@ func (dao *FileCommitDAO) GetByID(id string) (*entity.FileCommit, error) {
}
// Create creates a new file commit record
func (dao *FileCommitDAO) Create(commit *entity.FileCommit) error {
return DB.Create(commit).Error
func (dao *FileCommitDAO) Create(ctx context.Context, db *gorm.DB, commit *entity.FileCommit) error {
return db.WithContext(ctx).Create(commit).Error
}
// UpdateTreeState updates the tree_state field for a commit
func (dao *FileCommitDAO) UpdateTreeState(id string, treeState string) error {
return DB.Model(&entity.FileCommit{}).Where("id = ?", id).Update("tree_state", treeState).Error
func (dao *FileCommitDAO) UpdateTreeState(ctx context.Context, db *gorm.DB, id string, treeState string) error {
return db.WithContext(ctx).Model(&entity.FileCommit{}).Where("id = ?", id).Update("tree_state", treeState).Error
}
// GetLatestByFolderID gets the latest (most recent) commit for a folder
func (dao *FileCommitDAO) GetLatestByFolderID(folderID string) (*entity.FileCommit, error) {
func (dao *FileCommitDAO) GetLatestByFolderID(ctx context.Context, db *gorm.DB, folderID string) (*entity.FileCommit, error) {
var commit entity.FileCommit
err := DB.Where("folder_id = ?", folderID).
err := db.WithContext(ctx).Where("folder_id = ?", folderID).
Order("create_time DESC").
First(&commit).Error
if err != nil {
@@ -70,11 +73,11 @@ var allowedFileCommitSorts = map[string]string{
}
// ListByFolderID lists commits for a folder with pagination
func (dao *FileCommitDAO) ListByFolderID(folderID string, page, pageSize int, orderBy string, desc bool) ([]*entity.FileCommit, int64, error) {
func (dao *FileCommitDAO) ListByFolderID(ctx context.Context, db *gorm.DB, folderID string, page, pageSize int, orderBy string, desc bool) ([]*entity.FileCommit, int64, error) {
var commits []*entity.FileCommit
var total int64
query := DB.Model(&entity.FileCommit{}).Where("folder_id = ?", folderID)
query := db.WithContext(ctx).Model(&entity.FileCommit{}).Where("folder_id = ?", folderID)
if err := query.Count(&total).Error; err != nil {
return nil, 0, err
@@ -116,28 +119,28 @@ func NewFileCommitItemDAO() *FileCommitItemDAO {
}
// Create creates a new file commit item record
func (dao *FileCommitItemDAO) Create(item *entity.FileCommitItem) error {
return DB.Create(item).Error
func (dao *FileCommitItemDAO) Create(ctx context.Context, db *gorm.DB, item *entity.FileCommitItem) error {
return db.WithContext(ctx).Create(item).Error
}
// ListByCommitID lists all items for a commit
func (dao *FileCommitItemDAO) ListByCommitID(commitID string) ([]*entity.FileCommitItem, error) {
func (dao *FileCommitItemDAO) ListByCommitID(ctx context.Context, db *gorm.DB, commitID string) ([]*entity.FileCommitItem, error) {
var items []*entity.FileCommitItem
err := DB.Where("commit_id = ?", commitID).Order("create_time ASC").Find(&items).Error
err := db.WithContext(ctx).Where("commit_id = ?", commitID).Order("create_time ASC").Find(&items).Error
return items, err
}
// ListByFileID lists all commit items for a specific file (for version history)
func (dao *FileCommitItemDAO) ListByFileID(fileID string) ([]*entity.FileCommitItem, error) {
func (dao *FileCommitItemDAO) ListByFileID(ctx context.Context, db *gorm.DB, fileID string) ([]*entity.FileCommitItem, error) {
var items []*entity.FileCommitItem
err := DB.Where("file_id = ?", fileID).Order("create_time DESC").Find(&items).Error
err := db.WithContext(ctx).Where("file_id = ?", fileID).Order("create_time DESC").Find(&items).Error
return items, err
}
// GetByCommitIDAndFileID gets a single commit item by commit and file ID
func (dao *FileCommitItemDAO) GetByCommitIDAndFileID(commitID, fileID string) (*entity.FileCommitItem, error) {
func (dao *FileCommitItemDAO) GetByCommitIDAndFileID(ctx context.Context, db *gorm.DB, commitID, fileID string) (*entity.FileCommitItem, error) {
var item entity.FileCommitItem
err := DB.Where("commit_id = ? AND file_id = ?", commitID, fileID).First(&item).Error
err := db.WithContext(ctx).Where("commit_id = ? AND file_id = ?", commitID, fileID).First(&item).Error
if err != nil {
return nil, err
}

View File

@@ -17,6 +17,7 @@
package dao
import (
"context"
"errors"
"fmt"
"ragflow/internal/common"
@@ -32,8 +33,8 @@ func NewIngestionTaskDAO() *IngestionTaskDAO {
return &IngestionTaskDAO{}
}
func (dao *IngestionTaskDAO) Create(ingestionTask *entity.IngestionTask) (*entity.IngestionTask, error) {
existing, err := dao.GetByDocumentID(ingestionTask.DocumentID)
func (dao *IngestionTaskDAO) Create(ctx context.Context, db *gorm.DB, ingestionTask *entity.IngestionTask) (*entity.IngestionTask, error) {
existing, err := dao.GetByDocumentID(ctx, db, ingestionTask.DocumentID)
if err != nil {
return nil, err
}
@@ -43,9 +44,9 @@ func (dao *IngestionTaskDAO) Create(ingestionTask *entity.IngestionTask) (*entit
if ingestionTask.ID == "" {
ingestionTask.ID = utility.GenerateUUID()
}
if err := DB.Create(ingestionTask).Error; err != nil {
if err = db.WithContext(ctx).Create(ingestionTask).Error; err != nil {
if errors.Is(err, gorm.ErrDuplicatedKey) {
existing, getErr := dao.GetByDocumentID(ingestionTask.DocumentID)
existing, getErr := dao.GetByDocumentID(ctx, db, ingestionTask.DocumentID)
if getErr != nil {
return nil, getErr
}
@@ -58,8 +59,8 @@ func (dao *IngestionTaskDAO) Create(ingestionTask *entity.IngestionTask) (*entit
return ingestionTask, nil
}
func (dao *IngestionTaskDAO) UpdateStatusIfCurrent(taskID, fromStatus, toStatus string) (bool, error) {
result := DB.Model(&entity.IngestionTask{}).
func (dao *IngestionTaskDAO) UpdateStatusIfCurrent(ctx context.Context, db *gorm.DB, taskID, fromStatus, toStatus string) (bool, error) {
result := db.WithContext(ctx).Model(&entity.IngestionTask{}).
Where("id = ? AND status = ?", taskID, fromStatus).
Update("status", toStatus)
if result.Error != nil {
@@ -70,8 +71,8 @@ func (dao *IngestionTaskDAO) UpdateStatusIfCurrent(taskID, fromStatus, toStatus
// UpdateComponentTotal records the number of components in the task's DSL
// graph. It is the authoritative denominator for progress percentage.
func (dao *IngestionTaskDAO) UpdateComponentTotal(taskID string, total int) error {
return DB.Model(&entity.IngestionTask{}).Where("id = ?", taskID).Update("component_total", total).Error
func (dao *IngestionTaskDAO) UpdateComponentTotal(ctx context.Context, db *gorm.DB, taskID string, total int) error {
return db.WithContext(ctx).Model(&entity.IngestionTask{}).Where("id = ?", taskID).Update("component_total", total).Error
}
type TaskInfo struct {
@@ -79,8 +80,8 @@ type TaskInfo struct {
FilesToDelete []string `json:"files_to_delete"`
}
func (dao *IngestionTaskDAO) Delete(taskID string, userID *string) (*TaskInfo, error) {
tx := DB.Begin()
func (dao *IngestionTaskDAO) Delete(ctx context.Context, db *gorm.DB, taskID string, userID *string) (*TaskInfo, error) {
tx := db.WithContext(ctx).Begin()
if tx.Error != nil {
return nil, tx.Error
}
@@ -141,50 +142,50 @@ func (dao *IngestionTaskDAO) Delete(taskID string, userID *string) (*TaskInfo, e
}
}
func (dao *IngestionTaskDAO) GetAllTasks(page, pageSize int) ([]*entity.IngestionTask, error) {
func (dao *IngestionTaskDAO) GetAllTasks(ctx context.Context, db *gorm.DB, page, pageSize int) ([]*entity.IngestionTask, error) {
var tasks []*entity.IngestionTask
var err error
if pageSize == 0 {
err = DB.Find(&tasks).Error
err = db.WithContext(ctx).Find(&tasks).Error
} else {
err = DB.Order("create_time DESC").Offset((page - 1) * pageSize).Limit(pageSize).Find(&tasks).Error
err = db.WithContext(ctx).Order("create_time DESC").Offset((page - 1) * pageSize).Limit(pageSize).Find(&tasks).Error
}
return tasks, err
}
func (dao *IngestionTaskDAO) ListByUserID(userID string, page, pageSize int) ([]*entity.IngestionTask, error) {
func (dao *IngestionTaskDAO) ListByUserID(ctx context.Context, db *gorm.DB, userID string, page, pageSize int) ([]*entity.IngestionTask, error) {
var tasks []*entity.IngestionTask
var err error
if pageSize == 0 {
err = DB.Where("user_id = ?", userID).Order("create_time DESC").Find(&tasks).Error
err = db.WithContext(ctx).Where("user_id = ?", userID).Order("create_time DESC").Find(&tasks).Error
} else {
err = DB.Where("user_id = ?", userID).Order("create_time DESC").Offset((page - 1) * pageSize).Limit(pageSize).Find(&tasks).Error
err = db.WithContext(ctx).Where("user_id = ?", userID).Order("create_time DESC").Offset((page - 1) * pageSize).Limit(pageSize).Find(&tasks).Error
}
return tasks, err
}
func (dao *IngestionTaskDAO) ListByUserIDAndDatasetID(userID, datasetID string, page, pageSize int) ([]*entity.IngestionTask, error) {
func (dao *IngestionTaskDAO) ListByUserIDAndDatasetID(ctx context.Context, db *gorm.DB, userID, datasetID string, page, pageSize int) ([]*entity.IngestionTask, error) {
var tasks []*entity.IngestionTask
var err error
if pageSize == 0 {
err = DB.Where("user_id = ? AND dataset_id = ?", userID, datasetID).Order("create_time DESC").Find(&tasks).Error
err = db.WithContext(ctx).Where("user_id = ? AND dataset_id = ?", userID, datasetID).Order("create_time DESC").Find(&tasks).Error
} else {
err = DB.Where("user_id = ? AND dataset_id = ?", userID, datasetID).Order("create_time DESC").Offset((page - 1) * pageSize).Limit(pageSize).Find(&tasks).Error
err = db.WithContext(ctx).Where("user_id = ? AND dataset_id = ?", userID, datasetID).Order("create_time DESC").Offset((page - 1) * pageSize).Limit(pageSize).Find(&tasks).Error
}
return tasks, err
}
func (dao *IngestionTaskDAO) GetByID(id string) (*entity.IngestionTask, error) {
func (dao *IngestionTaskDAO) GetByID(ctx context.Context, db *gorm.DB, id string) (*entity.IngestionTask, error) {
var task *entity.IngestionTask
err := DB.Where("id = ?", id).First(&task).Error
err := db.WithContext(ctx).Where("id = ?", id).First(&task).Error
return task, err
}
func (dao *IngestionTaskDAO) GetByDocumentID(documentId string) (*entity.IngestionTask, error) {
func (dao *IngestionTaskDAO) GetByDocumentID(ctx context.Context, db *gorm.DB, documentId string) (*entity.IngestionTask, error) {
var tasks []*entity.IngestionTask
err := DB.Where("document_id = ?", documentId).Limit(1).Find(&tasks).Error
err := db.WithContext(ctx).Where("document_id = ?", documentId).Limit(1).Find(&tasks).Error
if err != nil {
return nil, err
}
@@ -199,8 +200,8 @@ func (dao *IngestionTaskDAO) GetByDocumentID(documentId string) (*entity.Ingesti
// RUNNING and STOPPING tasks are NOT deleted because an in-flight worker
// would keep writing chunks and corrupt a new run's results.
// Returns the number of rows deleted.
func (dao *IngestionTaskDAO) DeleteIfTerminal(documentID string) (int64, error) {
result := DB.Where("document_id = ? AND status NOT IN (?, ?)",
func (dao *IngestionTaskDAO) DeleteIfTerminal(ctx context.Context, db *gorm.DB, documentID string) (int64, error) {
result := db.WithContext(ctx).Where("document_id = ? AND status NOT IN (?, ?)",
documentID, common.RUNNING, common.STOPPING).
Delete(&entity.IngestionTask{})
if result.Error != nil {
@@ -215,12 +216,12 @@ func NewIngestionTaskLogDAO() *IngestionTaskLogDAO {
return &IngestionTaskLogDAO{}
}
func (dao *IngestionTaskLogDAO) Create(ingestionLog *entity.IngestionTaskLog) error {
return DB.Create(ingestionLog).Error
func (dao *IngestionTaskLogDAO) Create(ctx context.Context, db *gorm.DB, ingestionLog *entity.IngestionTaskLog) error {
return db.WithContext(ctx).Create(ingestionLog).Error
}
func (dao *IngestionTaskLogDAO) Update(ingestionLog *entity.IngestionTaskLog) error {
return DB.Save(ingestionLog).Error
func (dao *IngestionTaskLogDAO) Update(ctx context.Context, db *gorm.DB, ingestionLog *entity.IngestionTaskLog) error {
return db.WithContext(ctx).Save(ingestionLog).Error
}
// ListLogsByTaskID returns the task's logs in chronological (write) order.
@@ -229,9 +230,9 @@ func (dao *IngestionTaskLogDAO) Update(ingestionLog *entity.IngestionTaskLog) er
// arbitrarily; `id` is monotonic and always reflects write order. This
// feeds the frontend log stream (GET .../logs), which renders each row by
// phase (0 started / 1 done / -1 failed).
func (dao *IngestionTaskLogDAO) ListLogsByTaskID(taskID string) ([]*entity.IngestionTaskLog, error) {
func (dao *IngestionTaskLogDAO) ListLogsByTaskID(ctx context.Context, db *gorm.DB, taskID string) ([]*entity.IngestionTaskLog, error) {
var tasks []*entity.IngestionTaskLog
err := DB.Where("task_id = ?", taskID).Order("id ASC").Find(&tasks).Error
err := db.WithContext(ctx).Where("task_id = ?", taskID).Order("id ASC").Find(&tasks).Error
return tasks, err
}
@@ -257,9 +258,9 @@ type TaskProgress struct {
// `total` is the authoritative denominator from ingestion_task.component_total.
// The classification is forward-compatible with the §5.1 ProgressPhase
// renumbering (exit=1 stays; error moves -1 -> 2).
func (dao *IngestionTaskLogDAO) AggregateProgress(taskID string, total int) (*TaskProgress, error) {
func (dao *IngestionTaskLogDAO) AggregateProgress(ctx context.Context, db *gorm.DB, taskID string, total int) (*TaskProgress, error) {
// Latest row id per component for this task.
latestIDs := DB.Model(&entity.IngestionTaskLog{}).
latestIDs := db.WithContext(ctx).Model(&entity.IngestionTaskLog{}).
Select("MAX(id)").
Where("task_id = ?", taskID).
Group("component")
@@ -268,7 +269,7 @@ func (dao *IngestionTaskLogDAO) AggregateProgress(taskID string, total int) (*Ta
Phase int
}
var rows []phaseRow
err := DB.Model(&entity.IngestionTaskLog{}).
err := db.WithContext(ctx).Model(&entity.IngestionTaskLog{}).
Select("phase").
Where("id IN (?)", latestIDs).
Scan(&rows).Error
@@ -293,19 +294,19 @@ func (dao *IngestionTaskLogDAO) AggregateProgress(taskID string, total int) (*Ta
return progress, nil
}
func (dao *IngestionTaskLogDAO) LatestLogByTaskID(taskID string) (*entity.IngestionTaskLog, error) {
func (dao *IngestionTaskLogDAO) LatestLogByTaskID(ctx context.Context, db *gorm.DB, taskID string) (*entity.IngestionTaskLog, error) {
var task *entity.IngestionTaskLog
err := DB.Where("task_id = ?", taskID).Order("create_time DESC").First(&task).Error
err := db.WithContext(ctx).Where("task_id = ?", taskID).Order("create_time DESC").First(&task).Error
return task, err
}
func (dao *IngestionTaskLogDAO) GetLogByLogID(logID string) (*entity.IngestionTaskLog, error) {
func (dao *IngestionTaskLogDAO) GetLogByLogID(ctx context.Context, db *gorm.DB, logID string) (*entity.IngestionTaskLog, error) {
var task *entity.IngestionTaskLog
err := DB.Where("id = ?", logID).First(&task).Error
err := db.WithContext(ctx).Where("id = ?", logID).First(&task).Error
return task, err
}
func (dao *IngestionTaskLogDAO) DeleteByTaskID(taskID string) (int64, error) {
result := DB.Unscoped().Where("task_id = ?", taskID).Delete(&entity.IngestionTaskLog{})
func (dao *IngestionTaskLogDAO) DeleteByTaskID(ctx context.Context, db *gorm.DB, taskID string) (int64, error) {
result := db.WithContext(ctx).Unscoped().Where("task_id = ?", taskID).Delete(&entity.IngestionTaskLog{})
return result.RowsAffected, result.Error
}

View File

@@ -28,7 +28,8 @@ func TestIngestionTaskDAOUpdateStatusIfCurrentSucceeds(t *testing.T) {
t.Fatalf("create task: %v", err)
}
updated, err := NewIngestionTaskDAO().UpdateStatusIfCurrent("task-1", common.CREATED, common.RUNNING)
ctx := t.Context()
updated, err := NewIngestionTaskDAO().UpdateStatusIfCurrent(ctx, db, "task-1", common.CREATED, common.RUNNING)
if err != nil {
t.Fatalf("UpdateStatusIfCurrent failed: %v", err)
}
@@ -36,7 +37,7 @@ func TestIngestionTaskDAOUpdateStatusIfCurrentSucceeds(t *testing.T) {
t.Fatal("expected update to succeed")
}
reloaded, err := NewIngestionTaskDAO().GetByID("task-1")
reloaded, err := NewIngestionTaskDAO().GetByID(ctx, db, "task-1")
if err != nil {
t.Fatalf("reload task: %v", err)
}
@@ -59,20 +60,21 @@ func TestIngestionTaskDAOCreateRejectsExistingTerminalTask(t *testing.T) {
{name: "stopped", status: common.STOPPED},
}
ctx := t.Context()
for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
if err := db.Where("id = ?", "task-1").Delete(&entity.IngestionTask{}).Error; err != nil {
if err := db.WithContext(ctx).Where("id = ?", "task-1").Delete(&entity.IngestionTask{}).Error; err != nil {
t.Fatalf("clear task: %v", err)
}
task := &entity.IngestionTask{ID: "task-1", UserID: "user-1", DocumentID: "doc-1", DatasetID: "kb-1", Status: tc.status}
if err := db.Create(task).Error; err != nil {
if err := db.WithContext(ctx).Create(task).Error; err != nil {
t.Fatalf("create task: %v", err)
}
_, err := NewIngestionTaskDAO().Create(&entity.IngestionTask{ID: "task-2", UserID: "user-1", DocumentID: "doc-1", DatasetID: "kb-1", Status: common.CREATED})
_, err := NewIngestionTaskDAO().Create(ctx, db, &entity.IngestionTask{ID: "task-2", UserID: "user-1", DocumentID: "doc-1", DatasetID: "kb-1", Status: common.CREATED})
if err == nil {
t.Fatal("expected Create to reject duplicate document task")
}
reloaded, err := NewIngestionTaskDAO().GetByID("task-1")
reloaded, err := NewIngestionTaskDAO().GetByID(ctx, db, "task-1")
if err != nil {
t.Fatalf("reload task: %v", err)
}
@@ -118,7 +120,8 @@ func TestIngestionTaskDAOUpdateStatusIfCurrentRejectsMismatchedStatus(t *testing
t.Fatalf("create task: %v", err)
}
updated, err := NewIngestionTaskDAO().UpdateStatusIfCurrent("task-1", common.CREATED, common.RUNNING)
ctx := t.Context()
updated, err := NewIngestionTaskDAO().UpdateStatusIfCurrent(ctx, db, "task-1", common.CREATED, common.RUNNING)
if err != nil {
t.Fatalf("UpdateStatusIfCurrent failed: %v", err)
}
@@ -126,7 +129,7 @@ func TestIngestionTaskDAOUpdateStatusIfCurrentRejectsMismatchedStatus(t *testing
t.Fatal("expected update to be rejected")
}
reloaded, err := NewIngestionTaskDAO().GetByID("task-1")
reloaded, err := NewIngestionTaskDAO().GetByID(ctx, db, "task-1")
if err != nil {
t.Fatalf("reload task: %v", err)
}
@@ -157,13 +160,14 @@ func TestIngestionTaskDAODeleteIfTerminal_RemovesOnlyTerminal(t *testing.T) {
}
}
ctx := t.Context()
// DeleteIfTerminal deletes everything except RUNNING and STOPPING.
// CREATED is safe to delete (no worker has claimed it yet);
// COMPLETED/STOPPED/FAILED are terminal.
// Call it for every doc and verify the negative cases survived.
for i := 0; i < len(statuses); i++ {
docID := fmt.Sprintf("doc-%d", i)
_, err := NewIngestionTaskDAO().DeleteIfTerminal(docID)
_, err := NewIngestionTaskDAO().DeleteIfTerminal(ctx, db, docID)
if err != nil {
t.Fatalf("DeleteIfTerminal(doc-%d): %v", i, err)
}
@@ -172,7 +176,7 @@ func TestIngestionTaskDAODeleteIfTerminal_RemovesOnlyTerminal(t *testing.T) {
// RUNNING and STOPPING must survive.
for _, i := range []int{1, 2} {
docID := fmt.Sprintf("doc-%d", i)
task, err := NewIngestionTaskDAO().GetByDocumentID(docID)
task, err := NewIngestionTaskDAO().GetByDocumentID(ctx, db, docID)
if err != nil {
t.Fatalf("GetByDocumentID %s: %v", docID, err)
}
@@ -183,7 +187,7 @@ func TestIngestionTaskDAODeleteIfTerminal_RemovesOnlyTerminal(t *testing.T) {
// CREATED, COMPLETED, STOPPED, FAILED must be gone.
for _, i := range []int{0, 3, 4, 5} {
docID := fmt.Sprintf("doc-%d", i)
task, err := NewIngestionTaskDAO().GetByDocumentID(docID)
task, err := NewIngestionTaskDAO().GetByDocumentID(ctx, db, docID)
if err != nil {
t.Fatalf("GetByDocumentID %s: %v", docID, err)
}

View File

@@ -53,18 +53,19 @@ func TestExecuteTask_AcksMessageOnCompletion(t *testing.T) {
testutil.WithTenantID("tenant-1"),
)
ctx := t.Context()
ingestor := NewIngestor("test", 1, []string{"pdf"})
ingestor.runDocumentTask = func(ctx context.Context, _ *entity.IngestionTask) error {
return nil
}
handle := &fakeTaskHandle{}
ingestor.executeTask(newAckTaskCtx(context.Background(), taskID, docID, handle))
ingestor.executeTask(ctx, newAckTaskCtx(context.Background(), taskID, docID, handle))
if handle.acks.Load() != 1 || handle.nacks.Load() != 0 {
t.Fatalf("expected 1 Ack / 0 Nack on completion, got acks=%d nacks=%d", handle.acks.Load(), handle.nacks.Load())
}
finalTask, err := dao.NewIngestionTaskDAO().GetByID(taskID)
finalTask, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load final task: %v", err)
}
@@ -85,6 +86,7 @@ func TestExecuteTask_AcksMessageOnFailure(t *testing.T) {
testutil.WithPipelineID("flow-1"),
testutil.WithTenantID("tenant-1"),
)
ctx := t.Context()
ingestor := NewIngestor("test", 1, []string{"pdf"})
ingestor.runDocumentTask = func(ctx context.Context, _ *entity.IngestionTask) error {
@@ -92,12 +94,12 @@ func TestExecuteTask_AcksMessageOnFailure(t *testing.T) {
}
handle := &fakeTaskHandle{}
ingestor.executeTask(newAckTaskCtx(context.Background(), taskID, docID, handle))
ingestor.executeTask(ctx, newAckTaskCtx(context.Background(), taskID, docID, handle))
if handle.acks.Load() != 1 || handle.nacks.Load() != 0 {
t.Fatalf("expected 1 Ack / 0 Nack on failure, got acks=%d nacks=%d", handle.acks.Load(), handle.nacks.Load())
}
finalTask, err := dao.NewIngestionTaskDAO().GetByID(taskID)
finalTask, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load final task: %v", err)
}
@@ -131,7 +133,7 @@ func TestExecuteTask_AcksMessageOnContextCancel(t *testing.T) {
cancel()
handle := &fakeTaskHandle{}
ingestor.executeTask(newAckTaskCtx(ctx, taskID, docID, handle))
ingestor.executeTask(t.Context(), newAckTaskCtx(ctx, taskID, docID, handle))
if runCalled {
t.Fatal("expected runDocumentTask to be skipped on cancelled ctx")
@@ -154,6 +156,7 @@ func TestExecuteTask_HeartbeatsInProgressDuringLongTask(t *testing.T) {
testutil.WithPipelineID("flow-1"),
testutil.WithTenantID("tenant-1"),
)
ctx := t.Context()
ingestor := NewIngestor("test", 1, []string{"pdf"})
ingestor.heartbeatInterval = 5 * time.Millisecond
@@ -167,7 +170,7 @@ func TestExecuteTask_HeartbeatsInProgressDuringLongTask(t *testing.T) {
}
handle := &fakeTaskHandle{}
go ingestor.executeTask(newAckTaskCtx(context.Background(), taskID, docID, handle))
go ingestor.executeTask(ctx, newAckTaskCtx(context.Background(), taskID, docID, handle))
<-started
@@ -246,9 +249,10 @@ func TestExecuteTask_ReleasesTaskFromCurrentTasks(t *testing.T) {
return nil
}
ingestor.claimTask(taskID)
ctx := t.Context()
handle := &fakeTaskHandle{}
ingestor.executeTask(newAckTaskCtx(context.Background(), taskID, docID, handle))
ingestor.executeTask(ctx, newAckTaskCtx(context.Background(), taskID, docID, handle))
if _, stillActive := ingestor.currentTasks[taskID]; stillActive {
t.Fatal("expected task released from currentTasks after executeTask finished")
@@ -263,9 +267,10 @@ func TestExecuteTask_ReleasesTaskFromCurrentTasks(t *testing.T) {
func TestSettleMessage_AckOnTerminal(t *testing.T) {
ingestor := NewIngestor("test", 1, []string{"pdf"})
handle := &fakeTaskHandle{}
ctx := t.Context()
taskCtx := newAckTaskCtx(context.Background(), "task-1", "doc-1", handle)
ingestor.settleMessage(taskCtx, func(ctx context.Context) bool { return true })
ingestor.settleMessage(ctx, taskCtx, func(ctx context.Context) bool { return true })
if handle.acks.Load() != 1 || handle.nacks.Load() != 0 {
t.Fatalf("body=true: expected 1 Ack/0 Nack, got acks=%d nacks=%d", handle.acks.Load(), handle.nacks.Load())
@@ -276,9 +281,10 @@ func TestSettleMessage_AckOnTerminal(t *testing.T) {
func TestSettleMessage_NackOnNonTerminal(t *testing.T) {
ingestor := NewIngestor("test", 1, []string{"pdf"})
handle := &fakeTaskHandle{}
ctx := t.Context()
taskCtx := newAckTaskCtx(context.Background(), "task-1", "doc-1", handle)
ingestor.settleMessage(taskCtx, func(ctx context.Context) bool { return false })
ingestor.settleMessage(ctx, taskCtx, func(ctx context.Context) bool { return false })
if handle.nacks.Load() != 1 || handle.acks.Load() != 0 {
t.Fatalf("body=false: expected 1 Nack/0 Ack, got acks=%d nacks=%d", handle.acks.Load(), handle.nacks.Load())
@@ -301,6 +307,7 @@ func TestSettleMessage_RecoversPanicAndAcksWhenTaskTerminal(t *testing.T) {
ingestor := NewIngestor("test", 1, []string{"pdf"})
handle := &fakeTaskHandle{}
ctx := t.Context()
taskCtx := newAckTaskCtx(context.Background(), taskID, docID, handle)
panicked := false
@@ -310,7 +317,7 @@ func TestSettleMessage_RecoversPanicAndAcksWhenTaskTerminal(t *testing.T) {
panicked = true
}
}()
ingestor.settleMessage(taskCtx, func(ctx context.Context) bool {
ingestor.settleMessage(ctx, taskCtx, func(ctx context.Context) bool {
panic("boom")
})
}()
@@ -322,7 +329,7 @@ func TestSettleMessage_RecoversPanicAndAcksWhenTaskTerminal(t *testing.T) {
if handle.acks.Load() != 1 || handle.nacks.Load() != 0 {
t.Fatalf("body=panic: expected 1 Ack/0 Nack (markFailed→FAILED→DB terminal), got acks=%d nacks=%d", handle.acks.Load(), handle.nacks.Load())
}
finalTask, err := dao.NewIngestionTaskDAO().GetByID(taskID)
finalTask, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load final task: %v", err)
}
@@ -393,18 +400,19 @@ func TestSettleMessage_DBTruthOverridesBodyReturn(t *testing.T) {
// recovery where markFailed succeeded: the task is terminal (FAILED)
// but the body signals non-terminal.
body := func(ctx context.Context) bool {
ingestor.markFailed(taskID)
ingestor.markFailed(ctx, taskID)
return false
}
ingestor.settleMessage(taskCtx, body)
ctx := t.Context()
ingestor.settleMessage(ctx, taskCtx, body)
// DB shows FAILED → terminal → Ack, overriding body's false.
if handle.acks.Load() != 1 || handle.nacks.Load() != 0 {
t.Fatalf("expected 1 Ack / 0 Nack (DB FAILED overrides body=false), got acks=%d nacks=%d",
handle.acks.Load(), handle.nacks.Load())
}
finalTask, err := dao.NewIngestionTaskDAO().GetByID(taskID)
finalTask, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load final task: %v", err)
}

View File

@@ -51,7 +51,8 @@ func TestExecuteTask_CheckpointParseFailureDoesNotKillProcess(t *testing.T) {
)
// Execute the task - this should NOT panic or fatal exit (this is our main validation!)
ingestor.executeTask(taskCtx)
ctx := t.Context()
ingestor.executeTask(ctx, taskCtx)
// Corrupted run_count values are skipped by IncrementRunCount, so the task
// proceeds to runDocumentTask and completes normally.
@@ -60,7 +61,7 @@ func TestExecuteTask_CheckpointParseFailureDoesNotKillProcess(t *testing.T) {
}
// Verify task status was set to COMPLETED
finalTask, err := dao.NewIngestionTaskDAO().GetByID(taskID)
finalTask, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load final ingestion task: %v", err)
}
@@ -173,7 +174,8 @@ func TestExecuteTask_RunsDocumentTask(t *testing.T) {
&entity.IngestionTask{ID: taskID, DocumentID: docID, DatasetID: "kb-1", Status: common.RUNNING},
)
ingestor.executeTask(taskCtx)
ctx := t.Context()
ingestor.executeTask(ctx, taskCtx)
if !runDocumentTaskCalled {
t.Fatal("expected executeTask to run runDocumentTask")
@@ -181,7 +183,7 @@ func TestExecuteTask_RunsDocumentTask(t *testing.T) {
if gotTaskID != taskID {
t.Fatalf("runDocumentTask got task ID %q, want %q", gotTaskID, taskID)
}
finalTask, err := dao.NewIngestionTaskDAO().GetByID(taskID)
finalTask, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load final ingestion task: %v", err)
}
@@ -211,7 +213,7 @@ func TestExecuteTask_CancelBeforePipeline(t *testing.T) {
)
ingestor := NewIngestor("test", 1, []string{"pdf"})
ingestor.cancelCheck = func(taskID string) bool { return true }
ingestor.cancelCheck = func(ctx context.Context, taskID string) bool { return true }
var runDocumentTaskCalled bool
ingestor.runDocumentTask = func(ctx context.Context, ingestionTask *entity.IngestionTask) error {
@@ -223,13 +225,13 @@ func TestExecuteTask_CancelBeforePipeline(t *testing.T) {
context.Background(),
&entity.IngestionTask{ID: taskID, DocumentID: docID, DatasetID: "kb-1", Status: common.RUNNING},
)
ingestor.executeTask(taskCtx)
ctx := t.Context()
ingestor.executeTask(ctx, taskCtx)
if runDocumentTaskCalled {
t.Fatal("expected runDocumentTask to NOT be called when cancel is detected before pipeline")
}
ctx := t.Context()
doc, err := dao.NewDocumentDAO().GetByID(ctx, db, docID)
if err != nil {
t.Fatalf("load document: %v", err)

View File

@@ -80,7 +80,7 @@ type Ingestor struct {
// pipeline to stop at the next ctx.Err() check. Defaults to a Redis
// cancel-flag lookup that mirrors Python's has_canceled(). Tests may
// override this to simulate cancel without Redis.
cancelCheck func(taskID string) bool
cancelCheck func(ctx context.Context, taskID string) bool
}
func NewIngestor(name string, maxConcurrency int32, supportedTypes []string) *Ingestor {
@@ -280,12 +280,12 @@ func (e *Ingestor) workerLoop(id int32) {
return
case taskCtx := <-e.taskChan:
common.Info("task context:" + taskCtx.IngestionTask.ID)
e.executeTask(taskCtx)
e.executeTask(e.ctx, taskCtx)
}
}
}
func (e *Ingestor) executeTask(taskCtx *taskpkg.TaskContext) {
func (e *Ingestor) executeTask(ctx context.Context, taskCtx *taskpkg.TaskContext) {
task := taskCtx.IngestionTask
common.Info(fmt.Sprintf("Starting task %s", task.ID))
@@ -312,12 +312,12 @@ func (e *Ingestor) executeTask(taskCtx *taskpkg.TaskContext) {
// Synchronous check: if already cancelled (e.g. flag set between MQ
// delivery and worker claim), stop before the pipeline even starts.
if e.cancelCheck(task.ID) {
if e.cancelCheck(ctx, task.ID) {
common.Info(fmt.Sprintf("Task %s cancel flag detected before pipeline start, cancelling", task.ID))
perTaskCancel()
}
e.settleMessage(taskCtx, func(ctx context.Context) bool {
e.settleMessage(ctx, taskCtx, func(ctx context.Context) bool {
return e.runTask(ctx, task)
})
}
@@ -326,12 +326,12 @@ func (e *Ingestor) executeTask(taskCtx *taskpkg.TaskContext) {
// RequestStop to handle RUNNING → STOPPING, then MarkStopped for the final
// STOPPING → STOPPED transition. Finally it cleans up the Redis cancel flag
// so that a future retry of the same task does not immediately re-cancel.
func (e *Ingestor) markStopped(taskID string) bool {
if _, err := e.ingestionTaskSvc.RequestStop(taskID); err != nil {
func (e *Ingestor) markStopped(ctx context.Context, taskID string) bool {
if _, err := e.ingestionTaskSvc.RequestStop(ctx, taskID); err != nil {
common.Error(fmt.Sprintf("markStopped: RequestStop task %s: %v", taskID, err), err)
return false
}
if err := e.ingestionTaskSvc.MarkStopped(taskID); err != nil {
if err := e.ingestionTaskSvc.MarkStopped(ctx, taskID); err != nil {
common.Error(fmt.Sprintf("markStopped: MarkStopped task %s: %v", taskID, err), err)
return false
}
@@ -346,8 +346,8 @@ func (e *Ingestor) markStopped(taskID string) bool {
// markFailed persists FAILED status for the task and reports whether the
// terminal status was durably written, so the caller can decide Ack vs Nack.
func (e *Ingestor) markFailed(taskID string) bool {
if uErr := e.ingestionTaskSvc.MarkFailed(taskID); uErr != nil {
func (e *Ingestor) markFailed(ctx context.Context, taskID string) bool {
if uErr := e.ingestionTaskSvc.MarkFailed(ctx, taskID); uErr != nil {
common.Error(fmt.Sprintf("Failed to set task %s to FAILED", taskID), uErr)
return false
}
@@ -362,13 +362,13 @@ func (e *Ingestor) runTask(ctx context.Context, task *entity.IngestionTask) bool
case <-ctx.Done():
common.Info(fmt.Sprintf("Task %s cancelled", task.ID))
e.markCancelProgress(task)
return e.markStopped(task.ID)
return e.markStopped(context.Background(), task.ID)
default:
}
if err := e.ingestionTaskSvc.IncrementRunCount(task.ID); err != nil {
if err := e.ingestionTaskSvc.IncrementRunCount(ctx, task.ID); err != nil {
common.Error(fmt.Sprintf("Failed to increment run count for task %s", task.ID), err)
return e.markFailed(task.ID)
return e.markFailed(ctx, task.ID)
}
// This is a new run (IncrementRunCount succeeded). Any Redis cancel flag
@@ -389,15 +389,15 @@ func (e *Ingestor) runTask(ctx context.Context, task *entity.IngestionTask) bool
if errors.Is(err, context.Canceled) {
common.Info(fmt.Sprintf("Task %s cancelled during pipeline", task.ID))
e.markCancelProgress(task)
return e.markStopped(task.ID)
return e.markStopped(ctx, task.ID)
}
if errors.Is(err, context.DeadlineExceeded) {
common.Info(fmt.Sprintf("Task %s timed out during pipeline", task.ID))
e.markTimeoutProgress(task)
return e.markFailed(task.ID)
return e.markFailed(ctx, task.ID)
}
common.Error(fmt.Sprintf("Task %s failed", task.ID), err)
return e.markFailed(task.ID)
return e.markFailed(ctx, task.ID)
}
if err := e.completeTask(ctx, task.ID); err != nil {
@@ -417,7 +417,7 @@ func (e *Ingestor) runTask(ctx context.Context, task *entity.IngestionTask) bool
// terminal state and the caller Acks instead of redelivering.
func (e *Ingestor) completeTask(ctx context.Context, taskID string) error {
_, err := backoff.Retry(ctx, func() (struct{}, error) {
return struct{}{}, e.completeOrSettle(taskID)
return struct{}{}, e.completeOrSettle(ctx, taskID)
}, backoff.WithMaxTries(3))
return err
}
@@ -426,10 +426,10 @@ func (e *Ingestor) completeTask(ctx context.Context, taskID string) error {
// terminally invalid because the task is no longer RUNNING - settles it to its
// actual terminal state. Returns nil once the task is in any terminal state;
// returns a non-terminal (transient) error only for retry-worthy DB failures.
func (e *Ingestor) completeOrSettle(taskID string) error {
if err := e.ingestionTaskSvc.MarkCompleted(taskID); err != nil {
func (e *Ingestor) completeOrSettle(ctx context.Context, taskID string) error {
if err := e.ingestionTaskSvc.MarkCompleted(ctx, taskID); err != nil {
if isTerminalTransitionError(err) {
return e.settleToTerminal(taskID)
return e.settleToTerminal(ctx, taskID)
}
return err
}
@@ -452,14 +452,14 @@ func isTerminalTransitionError(err error) bool {
// re-cancel); already-terminal states (COMPLETED/STOPPED/FAILED) need no
// action. An unexpected status returns an error so the caller nacks and
// redelivery settles it.
func (e *Ingestor) settleToTerminal(taskID string) error {
task, err := e.ingestionTaskSvc.GetTask(taskID)
func (e *Ingestor) settleToTerminal(ctx context.Context, taskID string) error {
task, err := e.ingestionTaskSvc.GetTask(ctx, taskID)
if err != nil {
return err
}
switch task.Status {
case common.STOPPING:
if !e.markStopped(taskID) {
if !e.markStopped(ctx, taskID) {
return fmt.Errorf("task %s: settle to STOPPED failed", taskID)
}
return nil
@@ -477,7 +477,7 @@ func (e *Ingestor) settleToTerminal(taskID string) error {
// Settlement queries the DB for the task's actual status: a terminal state
// (COMPLETED/STOPPED/FAILED) means Ack; anything else means Nack. The body's
// return value is advisory only — DB truth is authoritative (BP1).
func (e *Ingestor) settleMessage(taskCtx *taskpkg.TaskContext, body func(context.Context) bool) (terminal bool) {
func (e *Ingestor) settleMessage(ctx context.Context, taskCtx *taskpkg.TaskContext, body func(context.Context) bool) (terminal bool) {
stop := e.startHeartbeat(taskCtx)
defer func() {
stop() // stop heartbeat (and wait) before ack/nack
@@ -488,12 +488,12 @@ func (e *Ingestor) settleMessage(taskCtx *taskpkg.TaskContext, body func(context
// redelivery. The broker's redelivery limit handles deterministic
// poison messages.
common.Error(fmt.Sprintf("task %s panicked: %v", taskCtx.IngestionTask.ID, r), fmt.Errorf("%v", r))
e.markFailed(taskCtx.IngestionTask.ID)
e.markFailed(ctx, taskCtx.IngestionTask.ID)
terminal = false
}
// Settlement authority is the DB, not the in-memory bool (BP1).
// Fall back to the in-memory bool only when the DB is unavailable.
if dbTerminal, ok := e.safeGetTerminal(taskCtx.IngestionTask.ID); ok {
if dbTerminal, ok := e.safeGetTerminal(ctx, taskCtx.IngestionTask.ID); ok {
terminal = dbTerminal
}
e.ackOrNack(taskCtx, terminal)
@@ -506,9 +506,9 @@ func (e *Ingestor) settleMessage(taskCtx *taskpkg.TaskContext, body func(context
// whether it is terminal (COMPLETED/STOPPED/FAILED). A recover guards
// against nil-DB panics in test environments — in that case (false, false)
// is returned so the caller falls back to the in-memory bool.
func (e *Ingestor) safeGetTerminal(taskID string) (terminal bool, ok bool) {
func (e *Ingestor) safeGetTerminal(ctx context.Context, taskID string) (terminal bool, ok bool) {
defer func() { recover() }()
task, err := e.ingestionTaskSvc.GetTask(taskID)
task, err := e.ingestionTaskSvc.GetTask(ctx, taskID)
if err != nil {
return false, false
}
@@ -539,14 +539,14 @@ func (e *Ingestor) ackOrNack(taskCtx *taskpkg.TaskContext, terminal bool) {
// REDIS_CONN.set(f"{task_id}-cancel", "x"). Falls back to checking the
// task status in DB when Redis is unavailable — a STOPPING status
// (set by RequestStop) is treated as a cancel signal.
func (e *Ingestor) defaultCancelCheck(taskID string) bool {
func (e *Ingestor) defaultCancelCheck(ctx context.Context, taskID string) bool {
rc := redis2.Get()
if rc != nil {
if ok, _ := rc.Exist(fmt.Sprintf("%s-cancel", taskID)); ok {
return true
}
}
task, err := e.ingestionTaskSvc.GetTask(taskID)
task, err := e.ingestionTaskSvc.GetTask(ctx, taskID)
if err != nil {
return false
}
@@ -566,7 +566,7 @@ func (e *Ingestor) pollCancel(taskID string, cancel context.CancelFunc, done <-c
result := make(chan bool, 1)
go func() {
defer func() { recover() }() // goroutine may outlive pollCancel; must not crash process
result <- e.cancelCheck(taskID)
result <- e.cancelCheck(e.ctx, taskID)
}()
return result
}

View File

@@ -172,7 +172,7 @@ func TestPollCancel_ExitsWhenDoneClosed(t *testing.T) {
// Block cancelCheck until released — simulate a stuck DB call.
blocking := make(chan struct{})
released := make(chan struct{})
ingestor.cancelCheck = func(taskID string) bool {
ingestor.cancelCheck = func(ctx context.Context, taskID string) bool {
close(blocking)
<-released
return false

View File

@@ -68,20 +68,20 @@ func newProgressSink(ctx context.Context, taskSvc *servicepkg.IngestionTaskServi
func (s *progressSink) OnComponentTotal(ctx context.Context, taskID string, total int) {
s.total.Store(int64(total))
if err := s.taskSvc.UpdateComponentTotal(taskID, total); err != nil {
if err := s.taskSvc.UpdateComponentTotal(ctx, taskID, total); err != nil {
common.Error(fmt.Sprintf("progressSink: update component_total for task %s failed: %v", taskID, err), err)
}
}
func (s *progressSink) OnComponentProgress(ctx context.Context, ev pipeline.ProgressEvent) {
if err := s.taskSvc.RecordComponentProgress(ev.TaskID, ev.Component, ev.Phase, ev.Message); err != nil {
if err := s.taskSvc.RecordComponentProgress(ctx, ev.TaskID, ev.Component, ev.Phase, ev.Message); err != nil {
common.Error(fmt.Sprintf("progressSink: record component progress for task %s failed: %v", ev.TaskID, err), err)
}
if ev.DocumentID == "" {
return
}
total := s.total.Load()
agg, err := s.taskSvc.AggregateTaskProgress(ev.TaskID, int(total))
agg, err := s.taskSvc.AggregateTaskProgress(ctx, ev.TaskID, int(total))
if err != nil {
common.Error(fmt.Sprintf("progressSink: aggregate task progress for task %s failed: %v", ev.TaskID, err), err)
return

View File

@@ -170,7 +170,7 @@ func TestProgressSinkPersistsViaService(t *testing.T) {
sink.docSvc = stub
sink.OnComponentTotal(ctx, taskID, 2)
task, err := dao.NewIngestionTaskDAO().GetByID(taskID)
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -186,7 +186,7 @@ func TestProgressSinkPersistsViaService(t *testing.T) {
Message: "Parser Done",
})
logs, err := dao.NewIngestionTaskLogDAO().ListLogsByTaskID(taskID)
logs, err := dao.NewIngestionTaskLogDAO().ListLogsByTaskID(ctx, db, taskID)
if err != nil {
t.Fatalf("list logs: %v", err)
}
@@ -235,7 +235,7 @@ func TestProgressSinkEmptyDocumentIDSkipsMirror(t *testing.T) {
Message: "Chunker Done",
})
logs, err := dao.NewIngestionTaskLogDAO().ListLogsByTaskID(taskID)
logs, err := dao.NewIngestionTaskLogDAO().ListLogsByTaskID(ctx, db, taskID)
if err != nil {
t.Fatalf("list logs: %v", err)
}

View File

@@ -41,8 +41,9 @@ func TestRunTask_ContextCancelledBeforeCheckpoint(t *testing.T) {
if runDocCalled {
t.Fatal("expected runDocumentTask to be skipped on cancelled ctx")
}
testCtx := t.Context()
// Checkpoint must not have been bumped — no log row should exist.
logs, err := dao.NewIngestionTaskLogDAO().ListLogsByTaskID(taskID)
logs, err := dao.NewIngestionTaskLogDAO().ListLogsByTaskID(testCtx, db, taskID)
if err != nil {
t.Fatalf("list logs: %v", err)
}
@@ -50,7 +51,7 @@ func TestRunTask_ContextCancelledBeforeCheckpoint(t *testing.T) {
t.Fatalf("expected 0 checkpoint rows (ctx cancelled before checkpoint), got %d", len(logs))
}
// Task must be STOPPED, not left in RUNNING.
task, err := dao.NewIngestionTaskDAO().GetByID(taskID)
task, err := dao.NewIngestionTaskDAO().GetByID(testCtx, db, taskID)
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -96,7 +97,8 @@ func TestRunTask_CorruptedRunCountSkipped(t *testing.T) {
t.Fatal("expected runDocumentTask to be called (bad run_count is skipped, not fatal)")
}
task, err := dao.NewIngestionTaskDAO().GetByID(taskID)
ctx := t.Context()
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -126,7 +128,8 @@ func TestRunTask_RunDocumentTaskFailureMarksFailed(t *testing.T) {
t.Fatal("expected true (terminal: durably marked FAILED)")
}
task, err := dao.NewIngestionTaskDAO().GetByID(taskID)
ctx := t.Context()
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -163,7 +166,8 @@ func TestRunTask_PipelineCancelledMarksStopped(t *testing.T) {
t.Fatal("expected true (terminal: durably marked STOPPED)")
}
task, err := dao.NewIngestionTaskDAO().GetByID(taskID)
ctx := t.Context()
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -195,7 +199,8 @@ func TestRunTask_ComponentTimeoutMarksFailed(t *testing.T) {
t.Fatal("expected true (terminal: durably marked FAILED)")
}
task, err := dao.NewIngestionTaskDAO().GetByID(taskID)
ctx := t.Context()
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -236,7 +241,8 @@ func TestRunTask_AlreadyCompletedAcksNotRedelivers(t *testing.T) {
}
// Task must still be COMPLETED (MarkCompleted failed to transition it).
task, err := dao.NewIngestionTaskDAO().GetByID(taskID)
ctx := t.Context()
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -260,7 +266,7 @@ func TestRunTask_PipelineSucceedsConcurrentStopSettlesStopped(t *testing.T) {
ingestor := NewIngestor("test", 1, []string{"pdf"})
ingestor.runDocumentTask = func(ctx context.Context, task *entity.IngestionTask) error {
// Simulate the user pressing Stop mid-pipeline: RUNNING->STOPPING.
if _, err := ingestor.ingestionTaskSvc.RequestStop(task.ID); err != nil {
if _, err := ingestor.ingestionTaskSvc.RequestStop(ctx, task.ID); err != nil {
t.Fatalf("RequestStop: %v", err)
}
return nil // pipeline still finishes successfully
@@ -274,7 +280,8 @@ func TestRunTask_PipelineSucceedsConcurrentStopSettlesStopped(t *testing.T) {
t.Fatal("expected true (terminal: settled to STOPPED, Ack)")
}
task, err := dao.NewIngestionTaskDAO().GetByID(taskID)
ctx := t.Context()
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -304,7 +311,8 @@ func TestRunTask_SuccessfulCompletion(t *testing.T) {
t.Fatal("expected true (terminal: durably completed)")
}
task, err := dao.NewIngestionTaskDAO().GetByID(taskID)
ctx := t.Context()
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("load task: %v", err)
}

View File

@@ -202,7 +202,7 @@ func TestPipelineE2E_PipelineExecutor(t *testing.T) {
defer cleanupEngine()
// Load task context
ingestionTask, err := dao.NewIngestionTaskDAO().GetByID(taskID)
ingestionTask, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("GetByID failed: %v", err)
}
@@ -315,11 +315,11 @@ func TestPipelineE2E_PipelineExecutor(t *testing.T) {
// Verify final task status can be marked completed
ingestSvc := service.NewIngestionTaskService()
if err := ingestSvc.MarkCompleted(taskID); err != nil {
if err = ingestSvc.MarkCompleted(ctx, taskID); err != nil {
t.Fatalf("MarkCompleted failed: %v", err)
}
finalTask, err := dao.NewIngestionTaskDAO().GetByID(taskID)
finalTask, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, taskID)
if err != nil {
t.Fatalf("GetByID failed: %v", err)
}

View File

@@ -73,7 +73,7 @@ func SetupTestDB(t *testing.T, tables ...any) *gorm.DB {
}
}
if err := db.AutoMigrate(tables...); err != nil {
if err = db.AutoMigrate(tables...); err != nil {
t.Fatalf("auto-migrate: %v", err)
}
return db

View File

@@ -626,7 +626,7 @@ func (s *ChunkService) StopParsing(ctx context.Context, userID, datasetID string
return nil, common.CodeDataError, fmt.Errorf("you don't own the document %s", docID)
}
task, err := dao.NewIngestionTaskDAO().GetByDocumentID(docID)
task, err := dao.NewIngestionTaskDAO().GetByDocumentID(ctx, dao.DB, docID)
if err != nil {
return nil, common.CodeServerError, fmt.Errorf("get ingestion task for %s: %w", docID, err)
}
@@ -733,18 +733,18 @@ func (s *ChunkService) Parse(ctx context.Context, userID, datasetID string, req
}
}
if len(notFound) > 0 {
return nil, common.CodeDataError, fmt.Errorf("Documents not found: %v", notFound)
return nil, common.CodeDataError, fmt.Errorf("documents not found: %v", notFound)
}
for _, docID := range docIDs {
doc := docByID[docID]
if doc.Run != nil && *doc.Run == string(entity.TaskStatusRunning) {
return nil, common.CodeDataError, fmt.Errorf("Can't parse document that is currently being processed")
return nil, common.CodeDataError, fmt.Errorf("can't parse document that is currently being processed")
}
}
// 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 = (document.NewDocumentService().AssertIngestionTasksTerminal(docIDs)); err != nil {
if err = (document.NewDocumentService().AssertIngestionTasksTerminal(ctx, docIDs)); err != nil {
return nil, common.CodeDataError, err
}
@@ -1510,7 +1510,7 @@ func decodeChunkImageBase64(raw string) ([]byte, error) {
}
imageBinary, err := base64.StdEncoding.Strict().DecodeString(raw)
if err != nil {
return nil, fmt.Errorf("Invalid `image_base64`")
return nil, fmt.Errorf("invalid `image_base64`")
}
if len(imageBinary) == 0 {
return nil, fmt.Errorf("`image_base64` is empty")

View File

@@ -592,7 +592,7 @@ func TestAddChunkImageAndTagFeatureValidation(t *testing.T) {
Content: "chunk body",
ImageBase64: strPtr("not-base64"),
}, userID)
if err == nil || !strings.Contains(err.Error(), "Invalid `image_base64`") {
if err == nil || !strings.Contains(err.Error(), "invalid `image_base64`") {
t.Fatalf("expected invalid image error, got %v", err)
}

View File

@@ -239,13 +239,13 @@ func (s *DocumentService) deleteDocumentFull(ctx context.Context, docID string)
}
// Delete tasks from DB
ingestionTask, err := s.ingestionTaskDAO.GetByDocumentID(docID)
ingestionTask, err := s.ingestionTaskDAO.GetByDocumentID(ctx, dao.DB, 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)
taskInfo, err := s.ingestionTaskSvc.Remove(ctx, ingestionTask.ID, &ingestionTask.UserID)
if err != nil {
return err
}

View File

@@ -11,7 +11,7 @@ import (
)
func (s *DocumentService) ListIngestionTasks(ctx context.Context, userID string, datasetID *string, page, pageSize int) ([]*entity.IngestionTask, error) {
return s.ingestionTaskSvc.ListByUser(userID, datasetID, page, pageSize)
return s.ingestionTaskSvc.ListByUser(ctx, userID, datasetID, page, pageSize)
}
func (s *DocumentService) IngestDocuments(ctx context.Context, datasetID, userID string, docIDs []string) ([]*service.ParseDocumentResponse, error) {
@@ -24,11 +24,11 @@ func (s *DocumentService) IngestDocuments(ctx context.Context, datasetID, userID
}
func (s *DocumentService) StopIngestionTasks(ctx context.Context, tasks []string, userID string) ([]*entity.IngestionTask, error) {
return s.ingestionTaskSvc.RequestStopMany(tasks, &userID)
return s.ingestionTaskSvc.RequestStopMany(ctx, tasks, &userID)
}
func (s *DocumentService) RemoveIngestionTasks(ctx context.Context, tasks []string, userID string) ([]map[string]string, error) {
return s.ingestionTaskSvc.RemoveMany(tasks, &userID)
return s.ingestionTaskSvc.RemoveMany(ctx, tasks, &userID)
}
func (s *DocumentService) Ingest(ctx context.Context, userID string, req *IngestDocumentRequest) (common.ErrorCode, error) {
@@ -70,10 +70,10 @@ func (s *DocumentService) Ingest(ctx context.Context, userID string, req *Ingest
validatedIDs = append(validatedIDs, docID)
}
// Batch pre-check for re-parse with delete: use the validated doc IDs
// Batch pre-check for reparse 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 {
if err = s.AssertIngestionTasksTerminal(ctx, validatedIDs); err != nil {
return common.CodeDataError, err
}
}

View File

@@ -34,7 +34,7 @@ func (s *DocumentService) StartParseDocuments(ctx context.Context, doc *entity.D
}
if opts.RerunWithDelete {
if err := s.clearDocumentParseResults(doc, kb.TenantID); err != nil {
if err := s.clearDocumentParseResults(ctx, doc, kb.TenantID); err != nil {
return err
}
}
@@ -49,9 +49,9 @@ func (s *DocumentService) StartParseDocuments(ctx context.Context, doc *entity.D
// 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 {
func (s *DocumentService) AssertIngestionTasksTerminal(ctx context.Context, docIDs []string) error {
for _, docID := range docIDs {
task, err := s.ingestionTaskDAO.GetByDocumentID(docID)
task, err := s.ingestionTaskDAO.GetByDocumentID(ctx, dao.DB, docID)
if err != nil {
return fmt.Errorf("check ingestion task for %s: %w", docID, err)
}
@@ -65,7 +65,7 @@ func (s *DocumentService) AssertIngestionTasksTerminal(docIDs []string) error {
return nil
}
func (s *DocumentService) clearDocumentParseResults(doc *entity.Document, tenantID string) error {
func (s *DocumentService) clearDocumentParseResults(ctx context.Context, doc *entity.Document, tenantID string) error {
if doc == nil {
return fmt.Errorf("document is nil")
}
@@ -74,7 +74,7 @@ func (s *DocumentService) clearDocumentParseResults(doc *entity.Document, tenant
// (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, _ := s.ingestionTaskDAO.GetByDocumentID(ctx, dao.DB, 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)
}
@@ -84,7 +84,7 @@ func (s *DocumentService) clearDocumentParseResults(doc *entity.Document, tenant
// 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 {
if _, err := s.ingestionTaskDAO.DeleteIfTerminal(ctx, dao.DB, doc.ID); err != nil {
return err
}
@@ -296,7 +296,7 @@ func (s *DocumentService) validateDocsInDataset(ctx context.Context, docIDs []st
// CancelDocParse stops the ingestion task for the document by calling
// RequestStop (STOPPING), then marks the document run status as CANCEL.
func (s *DocumentService) CancelDocParse(ctx context.Context, doc *entity.Document) error {
task, err := s.ingestionTaskDAO.GetByDocumentID(doc.ID)
task, err := s.ingestionTaskDAO.GetByDocumentID(ctx, dao.DB, doc.ID)
if err != nil {
return fmt.Errorf("failed to get ingestion task for %s: %v", doc.ID, err)
}
@@ -304,7 +304,7 @@ func (s *DocumentService) CancelDocParse(ctx context.Context, doc *entity.Docume
return fmt.Errorf("no ingestion task found for document %s", doc.ID)
}
if _, err = s.ingestionTaskSvc.RequestStop(task.ID); err != nil {
if _, err = s.ingestionTaskSvc.RequestStop(ctx, task.ID); err != nil {
return fmt.Errorf("failed to stop ingestion task %s: %v", task.ID, err)
}

View File

@@ -1105,7 +1105,7 @@ func TestStartParseDocuments_EnqueuesIngestionTask(t *testing.T) {
t.Fatal("expected non-empty task id")
}
ingestionTask, err := svc.ingestionTaskDAO.GetByID(msg.TaskID)
ingestionTask, err := svc.ingestionTaskDAO.GetByID(ctx, db, msg.TaskID)
if err != nil {
t.Fatalf("load ingestion task: %v", err)
}
@@ -1806,7 +1806,7 @@ func TestClearDocumentParseResultsClearsCountersTasksAndChunks(t *testing.T) {
svc := testDocumentService(t)
svc.docEngine = engine
if err = svc.clearDocumentParseResults(doc, "tenant-1"); err != nil {
if err = svc.clearDocumentParseResults(ctx, doc, "tenant-1"); err != nil {
t.Fatalf("clearDocumentParseResults failed: %v", err)
}
@@ -1819,7 +1819,7 @@ func TestClearDocumentParseResultsClearsCountersTasksAndChunks(t *testing.T) {
t.Fatalf("kb counters = token:%d chunk:%d, want zero", kb.TokenNum, kb.ChunkNum)
}
// The completed ingestion task must be deleted so the new run can proceed.
remainingTask, _ := svc.ingestionTaskDAO.GetByDocumentID("doc-1")
remainingTask, _ := svc.ingestionTaskDAO.GetByDocumentID(ctx, db, "doc-1")
if remainingTask != nil {
t.Fatalf("ingestion task should be deleted, status was %q", remainingTask.Status)
}
@@ -1844,10 +1844,10 @@ func TestClearDocumentParseResultsIsIdempotentForStaleDocSnapshot(t *testing.T)
}
svc := testDocumentService(t)
if err := svc.clearDocumentParseResults(staleDoc, "tenant-1"); err != nil {
if err = svc.clearDocumentParseResults(ctx, staleDoc, "tenant-1"); err != nil {
t.Fatalf("first clearDocumentParseResults failed: %v", err)
}
if err := svc.clearDocumentParseResults(staleDoc, "tenant-1"); err != nil {
if err = svc.clearDocumentParseResults(ctx, staleDoc, "tenant-1"); err != nil {
t.Fatalf("second clearDocumentParseResults failed: %v", err)
}
@@ -1862,7 +1862,7 @@ func TestClearDocumentParseResultsIsIdempotentForStaleDocSnapshot(t *testing.T)
}
// TestClearDocumentParseResults_RejectsNonTerminalIngestionTask verifies
// that re-parsing is refused while a document's ingestion task is still
// that reparsing is refused while a document's ingestion task is still
// RUNNING or STOPPING. Deleting a non-terminal task would let the in-flight
// worker keep writing chunks and corrupt the new run's results.
func TestClearDocumentParseResults_RejectsNonTerminalIngestionTask(t *testing.T) {
@@ -1880,11 +1880,11 @@ func TestClearDocumentParseResults_RejectsNonTerminalIngestionTask(t *testing.T)
}
svc := testDocumentService(t)
if err := svc.clearDocumentParseResults(doc, "tenant-1"); err == nil {
if err = svc.clearDocumentParseResults(ctx, doc, "tenant-1"); err == nil {
t.Fatalf("expected error for %s ingestion task, got nil", status)
}
// The non-terminal task must NOT be deleted.
task, _ := svc.ingestionTaskDAO.GetByDocumentID("doc-1")
task, _ := svc.ingestionTaskDAO.GetByDocumentID(ctx, db, "doc-1")
if task == nil {
t.Fatalf("%s ingestion task must not be deleted", status)
}
@@ -1910,10 +1910,10 @@ func TestClearDocumentParseResults_DeletesTerminalIngestionTask(t *testing.T) {
}
svc := testDocumentService(t)
if err := svc.clearDocumentParseResults(doc, "tenant-1"); err != nil {
if err = svc.clearDocumentParseResults(ctx, doc, "tenant-1"); err != nil {
t.Fatalf("clearDocumentParseResults for %s task: %v", status, err)
}
task, _ := svc.ingestionTaskDAO.GetByDocumentID("doc-1")
task, _ := svc.ingestionTaskDAO.GetByDocumentID(ctx, db, "doc-1")
if task != nil {
t.Fatalf("%s ingestion task should be deleted, still present", status)
}
@@ -1930,8 +1930,9 @@ func TestAssertIngestionTasksTerminal_RejectsNonTerminal(t *testing.T) {
insertTestIngestionTaskWithStatus(t, "task-1", "user-1", "doc-1", "kb-1", common.COMPLETED)
insertTestIngestionTaskWithStatus(t, "task-2", "user-1", "doc-2", "kb-1", common.RUNNING)
ctx := t.Context()
svc := testDocumentService(t)
if err := svc.AssertIngestionTasksTerminal([]string{"doc-1", "doc-2"}); err == nil {
if err := svc.AssertIngestionTasksTerminal(ctx, []string{"doc-1", "doc-2"}); err == nil {
t.Fatal("expected error for RUNNING task, got nil")
}
}
@@ -1945,8 +1946,9 @@ func TestAssertIngestionTasksTerminal_AcceptsAllTerminal(t *testing.T) {
insertTestIngestionTaskWithStatus(t, "task-1", "user-1", "doc-1", "kb-1", common.COMPLETED)
insertTestIngestionTaskWithStatus(t, "task-2", "user-1", "doc-2", "kb-1", common.STOPPED)
ctx := t.Context()
svc := testDocumentService(t)
if err := svc.AssertIngestionTasksTerminal([]string{"doc-1", "doc-2"}); err != nil {
if err := svc.AssertIngestionTasksTerminal(ctx, []string{"doc-1", "doc-2"}); err != nil {
t.Fatalf("expected nil for all-terminal batch, got %v", err)
}
}
@@ -1976,7 +1978,7 @@ func TestIngest_RerunWithDelete_RejectsBatchWithRunningTask(t *testing.T) {
t.Fatal("expected error for batch with RUNNING task, got nil")
}
// doc-1 (terminal) task must NOT be deleted - the whole batch was rejected.
task1, _ := svc.ingestionTaskDAO.GetByDocumentID("doc-1")
task1, _ := svc.ingestionTaskDAO.GetByDocumentID(ctx, db, "doc-1")
if task1 == nil {
t.Fatal("doc-1 terminal task should not be deleted when batch is rejected")
}
@@ -2790,7 +2792,7 @@ func TestStartParseDocuments_FailsBeforeClearing(t *testing.T) {
}
// Force GetDocumentStorageAddress to fail by clearing the document location.
if err := dao.DB.Model(&entity.Document{}).Where("id = ?", "doc-1").Update("location", "").Error; err != nil {
if err = dao.DB.Model(&entity.Document{}).Where("id = ?", "doc-1").Update("location", "").Error; err != nil {
t.Fatalf("clear location: %v", err)
}
@@ -2801,7 +2803,7 @@ func TestStartParseDocuments_FailsBeforeClearing(t *testing.T) {
}
// The old ingestion task must still exist — we failed before clearing.
remaining, _ := svc.ingestionTaskDAO.GetByDocumentID("doc-1")
remaining, _ := svc.ingestionTaskDAO.GetByDocumentID(ctx, db, "doc-1")
if remaining == nil {
t.Fatal("ingestion task should NOT be deleted when storage validation fails")
}
@@ -2831,7 +2833,7 @@ func TestIngest_CancelDoesNotDeleteIngestionTask(t *testing.T) {
}
// The ingestion task must NOT be deleted — cancel only stops it.
remaining, _ := svc.ingestionTaskDAO.GetByDocumentID("doc-1")
remaining, _ := svc.ingestionTaskDAO.GetByDocumentID(ctx, db, "doc-1")
if remaining == nil {
t.Fatal("ingestion task must NOT be deleted by cancel")
}

View File

@@ -53,7 +53,7 @@ func NewFileCommitService() *FileCommitService {
// CreateCommit creates a new commit for a workspace folder
func (s *FileCommitService) CreateCommit(ctx context.Context, folderID, authorID, message string, changes []entity.FileChange) (*entity.FileCommit, error) {
// 1. Get the latest commit for this folder
latestCommit, _ := s.commitDAO.GetLatestByFolderID(folderID)
latestCommit, _ := s.commitDAO.GetLatestByFolderID(ctx, dao.DB, folderID)
// 2. Build tree state from latest commit
treeState := make(map[string]interface{})
@@ -83,7 +83,7 @@ func (s *FileCommitService) CreateCommit(ctx context.Context, folderID, authorID
// All DB operations run inside a single transaction.
var treeStr string
if err := dao.DB.Transaction(func(tx *gorm.DB) error {
if err := dao.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
// Save commit
if err := tx.Create(commit).Error; err != nil {
return fmt.Errorf("failed to create commit: %w", err)
@@ -235,26 +235,26 @@ func (s *FileCommitService) CreateCommit(ctx context.Context, folderID, authorID
// ListCommits lists commits for a workspace folder with pagination
func (s *FileCommitService) ListCommits(ctx context.Context, folderID string, page, pageSize int, orderBy string, desc bool) ([]*entity.FileCommit, int64, error) {
return s.commitDAO.ListByFolderID(folderID, page, pageSize, orderBy, desc)
return s.commitDAO.ListByFolderID(ctx, dao.DB, folderID, page, pageSize, orderBy, desc)
}
// GetCommit gets a single commit by ID
func (s *FileCommitService) GetCommit(ctx context.Context, commitID string) (*entity.FileCommit, error) {
return s.commitDAO.GetByID(commitID)
return s.commitDAO.GetByID(ctx, dao.DB, commitID)
}
// ListCommitFiles lists all file change items for a commit
func (s *FileCommitService) ListCommitFiles(ctx context.Context, commitID string) ([]*entity.FileCommitItem, error) {
return s.commitItemDAO.ListByCommitID(commitID)
return s.commitItemDAO.ListByCommitID(ctx, dao.DB, commitID)
}
// DiffCommits compares two commits and returns the diff
func (s *FileCommitService) DiffCommits(ctx context.Context, fromID, toID string) ([]entity.DiffEntry, error) {
fromItems, err := s.commitItemDAO.ListByCommitID(fromID)
fromItems, err := s.commitItemDAO.ListByCommitID(ctx, dao.DB, fromID)
if err != nil {
return nil, err
}
toItems, err := s.commitItemDAO.ListByCommitID(toID)
toItems, err := s.commitItemDAO.ListByCommitID(ctx, dao.DB, toID)
if err != nil {
return nil, err
}
@@ -269,7 +269,7 @@ func (s *FileCommitService) DiffCommits(ctx context.Context, fromID, toID string
}
// Get tree state for file names (use to commit)
toCommit, err := s.commitDAO.GetByID(toID)
toCommit, err := s.commitDAO.GetByID(ctx, dao.DB, toID)
treeState := make(map[string]interface{})
if err == nil && toCommit != nil && toCommit.TreeState != nil {
json.Unmarshal([]byte(*toCommit.TreeState), &treeState)
@@ -351,7 +351,7 @@ func (s *FileCommitService) DiffCommits(ctx context.Context, fromID, toID string
// Recursively scans all sub-folders.
func (s *FileCommitService) GetUncommittedChanges(ctx context.Context, folderID string) ([]entity.DiffEntry, error) {
// Get latest commit tree state
latest, err := s.commitDAO.GetLatestByFolderID(folderID)
latest, err := s.commitDAO.GetLatestByFolderID(ctx, dao.DB, folderID)
committedFiles := make(map[string]map[string]interface{})
if err == nil && latest != nil && latest.TreeState != nil {
var treeData map[string]interface{}
@@ -438,7 +438,7 @@ func (s *FileCommitService) collectAllFilesRecursive(ctx context.Context, folder
// GetCommitTree gets the tree state snapshot for a commit as a hierarchical tree.
func (s *FileCommitService) GetCommitTree(ctx context.Context, commitID string) (map[string]interface{}, error) {
commit, err := s.commitDAO.GetByID(commitID)
commit, err := s.commitDAO.GetByID(ctx, dao.DB, commitID)
if err != nil {
return nil, err
}
@@ -541,12 +541,12 @@ func (s *FileCommitService) buildHierarchicalTree(ctx context.Context, flat map[
// GetCommitFileContent gets file content as it existed in a given commit
func (s *FileCommitService) GetCommitFileContent(ctx context.Context, folderID, commitID, fileID string) ([]byte, error) {
_, err := s.commitDAO.GetByID(commitID)
_, err := s.commitDAO.GetByID(ctx, dao.DB, commitID)
if err != nil {
return nil, fmt.Errorf("commit not found: %w", err)
}
item, err := s.commitItemDAO.GetByCommitIDAndFileID(commitID, fileID)
item, err := s.commitItemDAO.GetByCommitIDAndFileID(ctx, dao.DB, commitID, fileID)
if err != nil {
return nil, fmt.Errorf("file not found in commit: %w", err)
}
@@ -579,7 +579,7 @@ func (s *FileCommitService) GetCommitFileContent(ctx context.Context, folderID,
// GetFileVersionHistory gets version history for a specific file
func (s *FileCommitService) GetFileVersionHistory(ctx context.Context, fileID string) ([]entity.VersionEntry, error) {
items, err := s.commitItemDAO.ListByFileID(fileID)
items, err := s.commitItemDAO.ListByFileID(ctx, dao.DB, fileID)
if err != nil {
return nil, err
}
@@ -587,7 +587,7 @@ func (s *FileCommitService) GetFileVersionHistory(ctx context.Context, fileID st
var versions []entity.VersionEntry
for _, item := range items {
var commit *entity.FileCommit
commit, err = s.commitDAO.GetByID(item.CommitID)
commit, err = s.commitDAO.GetByID(ctx, dao.DB, item.CommitID)
if err != nil {
continue
}

View File

@@ -67,11 +67,11 @@ func (s *IngestionTaskService) SetTaskPublisher(taskPublisher TaskPublisher) {
s.taskPublisher = taskPublisher
}
func (s *IngestionTaskService) ListByUser(userID string, datasetID *string, page, pageSize int) ([]*entity.IngestionTask, error) {
func (s *IngestionTaskService) ListByUser(ctx context.Context, userID string, datasetID *string, page, pageSize int) ([]*entity.IngestionTask, error) {
if datasetID == nil {
return s.ingestionTaskDAO.ListByUserID(userID, page, pageSize)
return s.ingestionTaskDAO.ListByUserID(ctx, dao.DB, userID, page, pageSize)
}
return s.ingestionTaskDAO.ListByUserIDAndDatasetID(userID, *datasetID, page, pageSize)
return s.ingestionTaskDAO.ListByUserIDAndDatasetID(ctx, dao.DB, userID, *datasetID, page, pageSize)
}
func (s *IngestionTaskService) CreateForDocuments(ctx context.Context, datasetID, userID string, docIDs []string) ([]*ParseDocumentResponse, error) {
@@ -122,11 +122,11 @@ func (s *IngestionTaskService) CreateForDocuments(ctx context.Context, datasetID
return responses, nil
}
func (s *IngestionTaskService) RequestStopMany(tasks []string, ownerUserID *string) ([]*entity.IngestionTask, error) {
func (s *IngestionTaskService) RequestStopMany(ctx context.Context, tasks []string, ownerUserID *string) ([]*entity.IngestionTask, error) {
taskResponses := make([]*entity.IngestionTask, 0, len(tasks))
for _, taskID := range tasks {
if ownerUserID != nil {
task, err := s.GetTask(taskID)
task, err := s.GetTask(ctx, taskID)
if err != nil {
return nil, err
}
@@ -134,7 +134,7 @@ func (s *IngestionTaskService) RequestStopMany(tasks []string, ownerUserID *stri
return nil, errors.New("task does not belong to the user")
}
}
task, err := s.RequestStop(taskID)
task, err := s.RequestStop(ctx, taskID)
if err != nil {
return nil, err
}
@@ -143,11 +143,11 @@ func (s *IngestionTaskService) RequestStopMany(tasks []string, ownerUserID *stri
return taskResponses, nil
}
func (s *IngestionTaskService) RemoveMany(tasks []string, ownerUserID *string) ([]map[string]string, error) {
func (s *IngestionTaskService) RemoveMany(ctx context.Context, tasks []string, ownerUserID *string) ([]map[string]string, error) {
deletedTasks := make([]map[string]string, 0, len(tasks))
for _, taskID := range tasks {
taskRecord := map[string]string{"task_id": taskID}
if _, err := s.Remove(taskID, ownerUserID); err != nil {
if _, err := s.Remove(ctx, taskID, ownerUserID); err != nil {
taskRecord["remove"] = fmt.Sprintf("fail: %s", err.Error())
} else {
taskRecord["remove"] = "success"
@@ -157,8 +157,8 @@ func (s *IngestionTaskService) RemoveMany(tasks []string, ownerUserID *string) (
return deletedTasks, nil
}
func (s *IngestionTaskService) ListAllForAdmin() ([]map[string]interface{}, error) {
ingestionTasks, err := s.ingestionTaskDAO.GetAllTasks(0, 0)
func (s *IngestionTaskService) ListAllForAdmin(ctx context.Context) ([]map[string]interface{}, error) {
ingestionTasks, err := s.ingestionTaskDAO.GetAllTasks(ctx, dao.DB, 0, 0)
if err != nil {
return nil, err
}
@@ -179,13 +179,13 @@ func (s *IngestionTaskService) ListAllForAdmin() ([]map[string]interface{}, erro
"status": task.Status,
}
if count, ok := s.lastRunCount(task.ID); ok {
if count, ok := s.lastRunCount(ctx, task.ID); ok {
showTask["run_count"] = count
}
showTask["component_total"] = task.ComponentTotal
if task.ComponentTotal > 0 {
progress, err := s.ingestionTaskLogDAO.AggregateProgress(task.ID, task.ComponentTotal)
progress, err := s.ingestionTaskLogDAO.AggregateProgress(ctx, dao.DB, task.ID, task.ComponentTotal)
if err == nil {
showTask["component_done"] = progress.Done
} else {
@@ -201,13 +201,13 @@ func (s *IngestionTaskService) ListAllForAdmin() ([]map[string]interface{}, erro
}
func (s *IngestionTaskService) StartRunning(ctx context.Context, taskID string) (*entity.IngestionTask, error) {
task, err := s.GetTask(taskID)
task, err := s.GetTask(ctx, taskID)
if err != nil {
return nil, err
}
switch task.Status {
case common.CREATED:
task, err = s.transition(taskID, common.RUNNING)
task, err = s.transition(ctx, taskID, common.RUNNING)
if err != nil {
return nil, err
}
@@ -229,7 +229,7 @@ func (s *IngestionTaskService) StartRunning(ctx context.Context, taskID string)
}
return task, nil
case common.STOPPING:
return s.transition(taskID, common.STOPPED)
return s.transition(ctx, taskID, common.STOPPED)
case common.RUNNING, common.COMPLETED, common.STOPPED, common.FAILED:
return task, nil
default:
@@ -237,16 +237,16 @@ func (s *IngestionTaskService) StartRunning(ctx context.Context, taskID string)
}
}
func (s *IngestionTaskService) RequestStop(taskID string) (*entity.IngestionTask, error) {
task, err := s.GetTask(taskID)
func (s *IngestionTaskService) RequestStop(ctx context.Context, taskID string) (*entity.IngestionTask, error) {
task, err := s.GetTask(ctx, taskID)
if err != nil {
return nil, err
}
switch task.Status {
case common.CREATED:
return s.transition(taskID, common.STOPPED)
return s.transition(ctx, taskID, common.STOPPED)
case common.RUNNING:
task, err = s.transition(taskID, common.STOPPING)
task, err = s.transition(ctx, taskID, common.STOPPING)
if err != nil {
return nil, err
}
@@ -262,51 +262,51 @@ func (s *IngestionTaskService) RequestStop(taskID string) (*entity.IngestionTask
}
}
func (s *IngestionTaskService) MarkCompleted(taskID string) error {
task, err := s.GetTask(taskID)
func (s *IngestionTaskService) MarkCompleted(ctx context.Context, taskID string) error {
task, err := s.GetTask(ctx, taskID)
if err != nil {
return err
}
if task.Status == common.COMPLETED || task.Status == common.STOPPED || task.Status == common.FAILED {
return nil // already terminal, idempotent — mirrors MarkStopped
}
_, err = s.transition(taskID, common.COMPLETED)
_, err = s.transition(ctx, taskID, common.COMPLETED)
return err
}
func (s *IngestionTaskService) MarkFailed(taskID string) error {
task, err := s.GetTask(taskID)
func (s *IngestionTaskService) MarkFailed(ctx context.Context, taskID string) error {
task, err := s.GetTask(ctx, taskID)
if err != nil {
return err
}
if task.Status == common.FAILED || task.Status == common.COMPLETED || task.Status == common.STOPPED {
return nil // already terminal, idempotent — mirrors MarkStopped
}
_, err = s.transition(taskID, common.FAILED)
_, err = s.transition(ctx, taskID, common.FAILED)
return err
}
// MarkStopped transitions the task from STOPPING to STOPPED (terminal).
// Idempotent: returns nil if the task is already in a terminal state
// (STOPPED, COMPLETED, or FAILED).
func (s *IngestionTaskService) MarkStopped(taskID string) error {
task, err := s.GetTask(taskID)
func (s *IngestionTaskService) MarkStopped(ctx context.Context, taskID string) error {
task, err := s.GetTask(ctx, taskID)
if err != nil {
return err
}
if task.Status == common.STOPPED || task.Status == common.COMPLETED || task.Status == common.FAILED {
return nil
}
_, err = s.transition(taskID, common.STOPPED)
_, err = s.transition(ctx, taskID, common.STOPPED)
return err
}
func (s *IngestionTaskService) Remove(taskID string, userID *string) (*dao.TaskInfo, error) {
return s.ingestionTaskDAO.Delete(taskID, userID)
func (s *IngestionTaskService) Remove(ctx context.Context, taskID string, userID *string) (*dao.TaskInfo, error) {
return s.ingestionTaskDAO.Delete(ctx, dao.DB, taskID, userID)
}
func (s *IngestionTaskService) GetTask(taskID string) (*entity.IngestionTask, error) {
task, err := s.ingestionTaskDAO.GetByID(taskID)
func (s *IngestionTaskService) GetTask(ctx context.Context, taskID string) (*entity.IngestionTask, error) {
task, err := s.ingestionTaskDAO.GetByID(ctx, dao.DB, taskID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, common.ErrTaskNotFound
@@ -338,8 +338,8 @@ func validateTransition(from, to string) error {
return &InvalidTaskTransitionError{From: from, To: to}
}
func (s *IngestionTaskService) newTaskStatusConflictError(taskID, expectedFrom, attemptedTo string) error {
current, err := s.GetTask(taskID)
func (s *IngestionTaskService) newTaskStatusConflictError(ctx context.Context, taskID, expectedFrom, attemptedTo string) error {
current, err := s.GetTask(ctx, taskID)
if err != nil {
return err
}
@@ -351,8 +351,8 @@ func (s *IngestionTaskService) newTaskStatusConflictError(taskID, expectedFrom,
}
}
func (s *IngestionTaskService) transition(taskID string, to string) (*entity.IngestionTask, error) {
task, err := s.GetTask(taskID)
func (s *IngestionTaskService) transition(ctx context.Context, taskID string, to string) (*entity.IngestionTask, error) {
task, err := s.GetTask(ctx, taskID)
if err != nil {
return nil, err
}
@@ -363,19 +363,19 @@ func (s *IngestionTaskService) transition(taskID string, to string) (*entity.Ing
}
return task, err
}
updated, err := s.ingestionTaskDAO.UpdateStatusIfCurrent(taskID, task.Status, to)
updated, err := s.ingestionTaskDAO.UpdateStatusIfCurrent(ctx, dao.DB, taskID, task.Status, to)
if err != nil {
return nil, err
}
if !updated {
return nil, s.newTaskStatusConflictError(taskID, task.Status, to)
return nil, s.newTaskStatusConflictError(ctx, taskID, task.Status, to)
}
task.Status = to
return task, nil
}
func (s *IngestionTaskService) CreateAndEnqueue(ctx context.Context, task *entity.IngestionTask) (*entity.IngestionTask, error) {
existing, err := s.ingestionTaskDAO.GetByDocumentID(task.DocumentID)
existing, err := s.ingestionTaskDAO.GetByDocumentID(ctx, dao.DB, task.DocumentID)
if err != nil {
return nil, err
}
@@ -383,12 +383,12 @@ func (s *IngestionTaskService) CreateAndEnqueue(ctx context.Context, task *entit
switch existing.Status {
case common.FAILED, common.STOPPED:
originalStatus := existing.Status
existing, err = s.transition(existing.ID, common.CREATED)
existing, err = s.transition(ctx, existing.ID, common.CREATED)
if err != nil {
return nil, err
}
if err = s.enqueueTask(existing.ID); err != nil {
if rollbackErr := s.rollbackRetriedTask(existing.ID, originalStatus); rollbackErr != nil {
if rollbackErr := s.rollbackRetriedTask(ctx, existing.ID, originalStatus); rollbackErr != nil {
return nil, fmt.Errorf("enqueue task %s: %w (rollback failed: %v)", existing.ID, err, rollbackErr)
}
return nil, err
@@ -398,12 +398,12 @@ func (s *IngestionTaskService) CreateAndEnqueue(ctx context.Context, task *entit
return nil, fmt.Errorf("document id %s already exists, status: %s, task id: %s", task.DocumentID, existing.Status, existing.ID)
}
}
created, err := s.ingestionTaskDAO.Create(task)
created, err := s.ingestionTaskDAO.Create(ctx, dao.DB, task)
if err != nil {
return nil, err
}
if err = s.enqueueTask(created.ID); err != nil {
if rollbackErr := s.rollbackCreatedTask(created.ID); rollbackErr != nil {
if rollbackErr := s.rollbackCreatedTask(ctx, created.ID); rollbackErr != nil {
return nil, fmt.Errorf("enqueue task %s: %w (rollback failed: %v)", created.ID, err, rollbackErr)
}
return nil, err
@@ -411,19 +411,19 @@ func (s *IngestionTaskService) CreateAndEnqueue(ctx context.Context, task *entit
return created, nil
}
func (s *IngestionTaskService) rollbackRetriedTask(taskID, status string) error {
updated, err := s.ingestionTaskDAO.UpdateStatusIfCurrent(taskID, common.CREATED, status)
func (s *IngestionTaskService) rollbackRetriedTask(ctx context.Context, taskID, status string) error {
updated, err := s.ingestionTaskDAO.UpdateStatusIfCurrent(ctx, dao.DB, taskID, common.CREATED, status)
if err != nil {
return err
}
if !updated {
return s.newTaskStatusConflictError(taskID, common.CREATED, status)
return s.newTaskStatusConflictError(ctx, taskID, common.CREATED, status)
}
return nil
}
func (s *IngestionTaskService) rollbackCreatedTask(taskID string) error {
_, err := s.ingestionTaskDAO.Delete(taskID, nil)
func (s *IngestionTaskService) rollbackCreatedTask(ctx context.Context, taskID string) error {
_, err := s.ingestionTaskDAO.Delete(ctx, dao.DB, taskID, nil)
return err
}
@@ -437,15 +437,15 @@ func (s *IngestionTaskService) enqueueTask(taskID string) error {
// UpdateComponentTotal records the number of components in the task's DSL
// graph - the authoritative denominator for progress percentage.
func (s *IngestionTaskService) UpdateComponentTotal(taskID string, total int) error {
return s.ingestionTaskDAO.UpdateComponentTotal(taskID, total)
func (s *IngestionTaskService) UpdateComponentTotal(ctx context.Context, taskID string, total int) error {
return s.ingestionTaskDAO.UpdateComponentTotal(ctx, dao.DB, taskID, total)
}
// RecordComponentProgress appends a component lifecycle row to
// ingestion_task_log (phase: 0 started / 1 done / 2 errored). The row's
// Checkpoint is empty; component progress and step checkpoints are distinct
// row models sharing the same table.
func (s *IngestionTaskService) RecordComponentProgress(taskID, component string, phase int, message string) error {
func (s *IngestionTaskService) RecordComponentProgress(ctx context.Context, taskID, component string, phase int, message string) error {
entry := &entity.IngestionTaskLog{
TaskID: taskID,
Checkpoint: entity.JSONMap{},
@@ -453,20 +453,20 @@ func (s *IngestionTaskService) RecordComponentProgress(taskID, component string,
Component: component,
Message: message,
}
return s.ingestionTaskLogDAO.Create(entry)
return s.ingestionTaskLogDAO.Create(ctx, dao.DB, entry)
}
// AggregateTaskProgress returns the SQL-aggregated component progress for a
// task (done/failed/running/percent against the given total denominator).
func (s *IngestionTaskService) AggregateTaskProgress(taskID string, total int) (*dao.TaskProgress, error) {
return s.ingestionTaskLogDAO.AggregateProgress(taskID, total)
func (s *IngestionTaskService) AggregateTaskProgress(ctx context.Context, taskID string, total int) (*dao.TaskProgress, error) {
return s.ingestionTaskLogDAO.AggregateProgress(ctx, dao.DB, taskID, total)
}
// lastRunCount scans all task logs (newest first) for a run_count entry,
// skipping component-progress rows whose Checkpoint is empty. It returns
// the counter and whether one was found.
func (s *IngestionTaskService) lastRunCount(taskID string) (int, bool) {
logs, err := s.ingestionTaskLogDAO.ListLogsByTaskID(taskID)
func (s *IngestionTaskService) lastRunCount(ctx context.Context, taskID string) (int, bool) {
logs, err := s.ingestionTaskLogDAO.ListLogsByTaskID(ctx, dao.DB, taskID)
if err != nil {
return 0, false
}
@@ -488,14 +488,14 @@ func (s *IngestionTaskService) lastRunCount(taskID string) (int, bool) {
// ignored). A failure to persist the new row is best-effort (logged) and
// does not return an error — matching the legacy semantics that the run
// proceeds even if the counter write fails.
func (s *IngestionTaskService) IncrementRunCount(taskID string) error {
prevCount, _ := s.lastRunCount(taskID)
func (s *IngestionTaskService) IncrementRunCount(ctx context.Context, taskID string) error {
prevCount, _ := s.lastRunCount(ctx, taskID)
entry := &entity.IngestionTaskLog{
TaskID: taskID,
Checkpoint: entity.JSONMap{stepKeyRunCount: prevCount + 1},
}
if err := s.ingestionTaskLogDAO.Create(entry); err != nil {
if err := s.ingestionTaskLogDAO.Create(ctx, dao.DB, entry); err != nil {
common.Error(fmt.Sprintf("Failed to persist run_count for task %s", taskID), err)
}
return nil

View File

@@ -49,7 +49,7 @@ func TestIngestionTaskServiceCreateForDocumentsPublishesTaskMessages(t *testing.
if msg.TaskType != common.TaskTypeIngestionTask {
t.Fatalf("task type = %q, want %q", msg.TaskType, common.TaskTypeIngestionTask)
}
task, err := dao.NewIngestionTaskDAO().GetByID(msg.TaskID)
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, msg.TaskID)
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -67,7 +67,8 @@ func TestIngestionTaskServiceListByUserFiltersDataset(t *testing.T) {
svc := NewIngestionTaskService()
datasetID := "kb-1"
tasks, err := svc.ListByUser("user-1", &datasetID, 0, 0)
ctx := t.Context()
tasks, err := svc.ListByUser(ctx, "user-1", &datasetID, 0, 0)
if err != nil {
t.Fatalf("ListByUser failed: %v", err)
}
@@ -84,9 +85,10 @@ func TestIngestionTaskServiceRequestStopManyStopsOwnedTasks(t *testing.T) {
pushServiceDB(t, db)
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
ctx := t.Context()
userID := "user-1"
svc := NewIngestionTaskService()
tasks, err := svc.RequestStopMany([]string{"task-1"}, &userID)
tasks, err := svc.RequestStopMany(ctx, []string{"task-1"}, &userID)
if err != nil {
t.Fatalf("RequestStopMany failed: %v", err)
}
@@ -103,12 +105,13 @@ func TestIngestionTaskServiceRequestStopManyRejectsOtherUsersTask(t *testing.T)
pushServiceDB(t, db)
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
ctx := t.Context()
userID := "user-2"
svc := NewIngestionTaskService()
if _, err := svc.RequestStopMany([]string{"task-1"}, &userID); err == nil {
if _, err := svc.RequestStopMany(ctx, []string{"task-1"}, &userID); err == nil {
t.Fatal("expected RequestStopMany to reject non-owner")
}
task, err := dao.NewIngestionTaskDAO().GetByID("task-1")
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, "task-1")
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -122,8 +125,9 @@ func TestIngestionTaskServiceRequestStopManyAllowsAdmin(t *testing.T) {
pushServiceDB(t, db)
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
ctx := t.Context()
svc := NewIngestionTaskService()
tasks, err := svc.RequestStopMany([]string{"task-1"}, nil)
tasks, err := svc.RequestStopMany(ctx, []string{"task-1"}, nil)
if err != nil {
t.Fatalf("RequestStopMany admin failed: %v", err)
}
@@ -140,16 +144,17 @@ func TestIngestionTaskServiceRemoveManyRemovesOwnedTasks(t *testing.T) {
pushServiceDB(t, db)
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
ctx := t.Context()
userID := "user-1"
svc := NewIngestionTaskService()
result, err := svc.RemoveMany([]string{"task-1"}, &userID)
result, err := svc.RemoveMany(ctx, []string{"task-1"}, &userID)
if err != nil {
t.Fatalf("RemoveMany failed: %v", err)
}
if len(result) != 1 || result[0]["remove"] != "success" {
t.Fatalf("unexpected remove result: %+v", result)
}
if _, err := dao.NewIngestionTaskDAO().GetByID("task-1"); err == nil {
if _, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, "task-1"); err == nil {
t.Fatal("task should be removed")
}
}
@@ -177,8 +182,9 @@ func TestIngestionTaskServiceListAllForAdminIncludesRunAndUserEmail(t *testing.T
t.Fatalf("insert task log: %v", err)
}
ctx := t.Context()
svc := NewIngestionTaskService()
tasks, err := svc.ListAllForAdmin()
tasks, err := svc.ListAllForAdmin(ctx)
if err != nil {
t.Fatalf("ListAllForAdmin failed: %v", err)
}
@@ -317,9 +323,10 @@ func TestIngestionTaskServiceRequestStopTransitionsCreatedTaskToStopped(t *testi
db := setupServiceTestDB(t)
pushServiceDB(t, db)
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
ctx := t.Context()
svc := NewIngestionTaskService()
task, err := svc.RequestStop("task-1")
task, err := svc.RequestStop(ctx, "task-1")
if err != nil {
t.Fatalf("RequestStop failed: %v", err)
}
@@ -332,12 +339,13 @@ func TestIngestionTaskServiceMarkCompletedRejectsNonRunningTask(t *testing.T) {
db := setupServiceTestDB(t)
pushServiceDB(t, db)
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
ctx := t.Context()
svc := NewIngestionTaskService()
if err := svc.MarkCompleted("task-1"); err == nil {
if err := svc.MarkCompleted(ctx, "task-1"); err == nil {
t.Fatal("expected MarkCompleted to reject non-running task")
}
task, err := dao.NewIngestionTaskDAO().GetByID("task-1")
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, "task-1")
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -353,12 +361,13 @@ func TestIngestionTaskServiceMarkCompletedUpdatesTaskStatus(t *testing.T) {
if err := dao.DB.Model(&entity.IngestionTask{}).Where("id = ?", "task-1").Update("status", common.RUNNING).Error; err != nil {
t.Fatalf("set running status: %v", err)
}
ctx := t.Context()
svc := NewIngestionTaskService()
if err := svc.MarkCompleted("task-1"); err != nil {
if err := svc.MarkCompleted(ctx, "task-1"); err != nil {
t.Fatalf("MarkCompleted failed: %v", err)
}
task, err := dao.NewIngestionTaskDAO().GetByID("task-1")
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, "task-1")
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -374,12 +383,13 @@ func TestIngestionTaskServiceMarkFailedUpdatesTaskStatus(t *testing.T) {
if err := dao.DB.Model(&entity.IngestionTask{}).Where("id = ?", "task-1").Update("status", common.RUNNING).Error; err != nil {
t.Fatalf("set running status: %v", err)
}
ctx := t.Context()
svc := NewIngestionTaskService()
if err := svc.MarkFailed("task-1"); err != nil {
if err := svc.MarkFailed(ctx, "task-1"); err != nil {
t.Fatalf("MarkFailed failed: %v", err)
}
task, err := dao.NewIngestionTaskDAO().GetByID("task-1")
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, "task-1")
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -395,9 +405,10 @@ func TestIngestionTaskServiceNewTaskStatusConflictErrorLoadsActualStatus(t *test
if err := dao.DB.Model(&entity.IngestionTask{}).Where("id = ?", "task-1").Update("status", common.STOPPING).Error; err != nil {
t.Fatalf("set stopping status: %v", err)
}
ctx := t.Context()
svc := NewIngestionTaskService()
err := svc.newTaskStatusConflictError("task-1", common.CREATED, common.RUNNING)
err := svc.newTaskStatusConflictError(ctx, "task-1", common.CREATED, common.RUNNING)
var conflictErr *TaskStatusConflictError
if !errors.As(err, &conflictErr) {
t.Fatalf("expected TaskStatusConflictError, got %T", err)
@@ -411,9 +422,10 @@ func TestIngestionTaskServiceMarkCompletedReturnsTaskIDInTransitionError(t *test
db := setupServiceTestDB(t)
pushServiceDB(t, db)
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
ctx := t.Context()
svc := NewIngestionTaskService()
err := svc.MarkCompleted("task-1")
err := svc.MarkCompleted(ctx, "task-1")
var transitionErr *InvalidTaskTransitionError
if !errors.As(err, &transitionErr) {
t.Fatalf("expected InvalidTaskTransitionError, got %T", err)
@@ -469,7 +481,7 @@ func TestIngestionTaskServiceCreateAndEnqueueRetriesTerminalTask(t *testing.T) {
if len(publisher.messages) != 1 || publisher.messages[0].TaskID != "task-1" {
t.Fatalf("unexpected published messages: %+v", publisher.messages)
}
reloaded, err := dao.NewIngestionTaskDAO().GetByID("task-1")
reloaded, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, "task-1")
if err != nil {
t.Fatalf("reload task: %v", err)
}
@@ -515,7 +527,7 @@ func TestIngestionTaskServiceCreateAndEnqueueRollsBackNewTaskOnPublishFailure(t
if err == nil || err.Error() != "publish failed" {
t.Fatalf("expected publish failure, got %v", err)
}
task, getErr := dao.NewIngestionTaskDAO().GetByDocumentID("doc-1")
task, getErr := dao.NewIngestionTaskDAO().GetByDocumentID(ctx, db, "doc-1")
if getErr != nil {
t.Fatalf("reload task by document id: %v", getErr)
}
@@ -546,7 +558,7 @@ func TestIngestionTaskServiceCreateAndEnqueueRollsBackRetriedTaskOnPublishFailur
if err == nil || err.Error() != "publish failed" {
t.Fatalf("expected publish failure, got %v", err)
}
reloaded, getErr := dao.NewIngestionTaskDAO().GetByID("task-1")
reloaded, getErr := dao.NewIngestionTaskDAO().GetByID(ctx, db, "task-1")
if getErr != nil {
t.Fatalf("reload task: %v", getErr)
}
@@ -559,17 +571,18 @@ func TestIngestionTaskServiceRemoveDeletesOwnedTask(t *testing.T) {
db := setupServiceTestDB(t)
pushServiceDB(t, db)
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
ctx := t.Context()
userID := "user-1"
svc := NewIngestionTaskService()
info, err := svc.Remove("task-1", &userID)
info, err := svc.Remove(ctx, "task-1", &userID)
if err != nil {
t.Fatalf("Remove failed: %v", err)
}
if info == nil || info.TaskID != "task-1" {
t.Fatalf("unexpected task info: %+v", info)
}
if _, err := dao.NewIngestionTaskDAO().GetByID("task-1"); err == nil {
if _, err = dao.NewIngestionTaskDAO().GetByID(ctx, db, "task-1"); err == nil {
t.Fatal("task should be removed")
}
}
@@ -578,12 +591,13 @@ func TestIngestionTaskServiceUpdateComponentTotalPersistsDenominator(t *testing.
db := setupServiceTestDB(t)
pushServiceDB(t, db)
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
ctx := t.Context()
svc := NewIngestionTaskService()
if err := svc.UpdateComponentTotal("task-1", 4); err != nil {
if err := svc.UpdateComponentTotal(ctx, "task-1", 4); err != nil {
t.Fatalf("UpdateComponentTotal failed: %v", err)
}
task, err := dao.NewIngestionTaskDAO().GetByID("task-1")
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, "task-1")
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -596,12 +610,13 @@ func TestIngestionTaskServiceRecordComponentProgressAppendsRow(t *testing.T) {
db := setupServiceTestDB(t)
pushServiceDB(t, db)
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
ctx := t.Context()
svc := NewIngestionTaskService()
if err := svc.RecordComponentProgress("task-1", "Parser", 1, "Parser Done"); err != nil {
if err := svc.RecordComponentProgress(ctx, "task-1", "Parser", 1, "Parser Done"); err != nil {
t.Fatalf("RecordComponentProgress failed: %v", err)
}
logs, err := dao.NewIngestionTaskLogDAO().ListLogsByTaskID("task-1")
logs, err := dao.NewIngestionTaskLogDAO().ListLogsByTaskID(ctx, db, "task-1")
if err != nil {
t.Fatalf("list logs: %v", err)
}
@@ -621,15 +636,16 @@ func TestIngestionTaskServiceAggregateTaskProgressClassifiesByPhase(t *testing.T
db := setupServiceTestDB(t)
pushServiceDB(t, db)
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
ctx := t.Context()
svc := NewIngestionTaskService()
if err := svc.RecordComponentProgress("task-1", "Parser", 1, "Parser Done"); err != nil {
if err := svc.RecordComponentProgress(ctx, "task-1", "Parser", 1, "Parser Done"); err != nil {
t.Fatalf("record Parser: %v", err)
}
if err := svc.RecordComponentProgress("task-1", "Chunker", 0, "Chunker Started"); err != nil {
if err := svc.RecordComponentProgress(ctx, "task-1", "Chunker", 0, "Chunker Started"); err != nil {
t.Fatalf("record Chunker: %v", err)
}
agg, err := svc.AggregateTaskProgress("task-1", 2)
agg, err := svc.AggregateTaskProgress(ctx, "task-1", 2)
if err != nil {
t.Fatalf("AggregateTaskProgress failed: %v", err)
}
@@ -645,21 +661,22 @@ func TestIngestionTaskServiceIncrementRunCountInitializesAndBumps(t *testing.T)
db := setupServiceTestDB(t)
pushServiceDB(t, db)
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
ctx := t.Context()
svc := NewIngestionTaskService()
if err := svc.IncrementRunCount("task-1"); err != nil {
if err := svc.IncrementRunCount(ctx, "task-1"); err != nil {
t.Fatalf("IncrementRunCount (first call) failed: %v", err)
}
run, ok := svc.lastRunCount("task-1")
run, ok := svc.lastRunCount(ctx, "task-1")
if !ok || run != 1 {
t.Fatalf("run_count = %v (ok=%v), want 1", run, ok)
}
// Second call bumps the existing counter to 2.
if err := svc.IncrementRunCount("task-1"); err != nil {
if err := svc.IncrementRunCount(ctx, "task-1"); err != nil {
t.Fatalf("IncrementRunCount (second call) failed: %v", err)
}
run, _ = svc.lastRunCount("task-1")
run, _ = svc.lastRunCount(ctx, "task-1")
if run != 2 {
t.Fatalf("run_count after second bump = %v, want 2", run)
}
@@ -675,13 +692,14 @@ func TestIngestionTaskServiceIncrementRunCountSkippedCorruptedRunCount(t *testin
}).Error; err != nil {
t.Fatalf("insert bad task log: %v", err)
}
ctx := t.Context()
svc := NewIngestionTaskService()
// Corrupted value is skipped; a fresh run_count=1 row is created.
if err := svc.IncrementRunCount("task-1"); err != nil {
if err := svc.IncrementRunCount(ctx, "task-1"); err != nil {
t.Fatalf("IncrementRunCount should skip corrupted value, got: %v", err)
}
run, ok := svc.lastRunCount("task-1")
run, ok := svc.lastRunCount(ctx, "task-1")
if !ok || run != 1 {
t.Fatalf("run_count = %v (ok=%v), want 1", run, ok)
}
@@ -691,18 +709,19 @@ func TestIngestionTaskServiceIncrementRunCountRecoversFromComponentProgressLog(t
db := setupServiceTestDB(t)
pushServiceDB(t, db)
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
ctx := t.Context()
// Simulate a previous run that created some component-progress logs
// but died before recording a run_count row. The latest log has no run_count.
svc := NewIngestionTaskService()
if err := svc.RecordComponentProgress("task-1", "Parser", 1, "Parser Done"); err != nil {
if err := svc.RecordComponentProgress(ctx, "task-1", "Parser", 1, "Parser Done"); err != nil {
t.Fatalf("record Parser: %v", err)
}
if err := svc.RecordComponentProgress("task-1", "Chunker", 1, "Chunker Done"); err != nil {
if err := svc.RecordComponentProgress(ctx, "task-1", "Chunker", 1, "Chunker Done"); err != nil {
t.Fatalf("record Chunker: %v", err)
}
// Verify latest log has empty checkpoint (no run_count).
latest, err := dao.NewIngestionTaskLogDAO().LatestLogByTaskID("task-1")
latest, err := dao.NewIngestionTaskLogDAO().LatestLogByTaskID(ctx, db, "task-1")
if err != nil {
t.Fatalf("load latest: %v", err)
}
@@ -711,17 +730,17 @@ func TestIngestionTaskServiceIncrementRunCountRecoversFromComponentProgressLog(t
}
// IncrementRunCount should create a new row with run_count=1.
if err := svc.IncrementRunCount("task-1"); err != nil {
if err = svc.IncrementRunCount(ctx, "task-1"); err != nil {
t.Fatalf("IncrementRunCount failed: %v", err)
}
run, ok := svc.lastRunCount("task-1")
run, ok := svc.lastRunCount(ctx, "task-1")
if !ok || run != 1 {
t.Fatalf("run_count = %v (ok=%v), want 1", run, ok)
}
// AggregateProgress should still work (run_count row with component=""
// has phase=0, which doesn't affect counts).
agg, err := svc.AggregateTaskProgress("task-1", 2)
agg, err := svc.AggregateTaskProgress(ctx, "task-1", 2)
if err != nil {
t.Fatalf("AggregateTaskProgress: %v", err)
}
@@ -736,31 +755,32 @@ func TestIngestionTaskServiceIncrementRunCountAccumulatesAcrossRetries(t *testin
insertTestIngestionTask(t, "task-1", "user-1", "doc-1", "kb-1")
svc := NewIngestionTaskService()
ctx := t.Context()
// First attempt: IncrementRunCount creates run_count=1.
if err := svc.IncrementRunCount("task-1"); err != nil {
if err := svc.IncrementRunCount(ctx, "task-1"); err != nil {
t.Fatalf("first IncrementRunCount: %v", err)
}
// Simulate first run: some components progress, then failure.
if err := svc.RecordComponentProgress("task-1", "Parser", 1, "Parser Done"); err != nil {
if err := svc.RecordComponentProgress(ctx, "task-1", "Parser", 1, "Parser Done"); err != nil {
t.Fatalf("record Parser: %v", err)
}
// Second attempt (retry): should find previous run_count=1 and create row with run_count=2.
if err := svc.IncrementRunCount("task-1"); err != nil {
if err := svc.IncrementRunCount(ctx, "task-1"); err != nil {
t.Fatalf("second IncrementRunCount: %v", err)
}
// More progress, then failure.
if err := svc.RecordComponentProgress("task-1", "Chunker", 1, "Chunker Done"); err != nil {
if err := svc.RecordComponentProgress(ctx, "task-1", "Chunker", 1, "Chunker Done"); err != nil {
t.Fatalf("record Chunker: %v", err)
}
// Third attempt (retry): should find previous run_count=2 and create row with run_count=3.
if err := svc.IncrementRunCount("task-1"); err != nil {
if err := svc.IncrementRunCount(ctx, "task-1"); err != nil {
t.Fatalf("third IncrementRunCount: %v", err)
}
run, ok := svc.lastRunCount("task-1")
run, ok := svc.lastRunCount(ctx, "task-1")
if !ok || run != 3 {
t.Fatalf("run_count = %v (ok=%v), want 3", run, ok)
}
@@ -778,7 +798,7 @@ func TestIngestionTaskServiceIncrementRunCountAccumulatesAcrossRetries(t *testin
}).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
adminTasks, err := svc.ListAllForAdmin()
adminTasks, err := svc.ListAllForAdmin(ctx)
if err != nil {
t.Fatalf("ListAllForAdmin: %v", err)
}
@@ -799,13 +819,14 @@ func TestIngestionTaskServiceMarkStoppedTransitionsStoppingTask(t *testing.T) {
Update("status", common.STOPPING).Error; err != nil {
t.Fatalf("set task STOPPING: %v", err)
}
ctx := t.Context()
svc := NewIngestionTaskService()
if err := svc.MarkStopped("task-1"); err != nil {
if err := svc.MarkStopped(ctx, "task-1"); err != nil {
t.Fatalf("MarkStopped failed: %v", err)
}
task, err := dao.NewIngestionTaskDAO().GetByID("task-1")
task, err := dao.NewIngestionTaskDAO().GetByID(ctx, db, "task-1")
if err != nil {
t.Fatalf("load task: %v", err)
}
@@ -824,7 +845,8 @@ func TestIngestionTaskServiceMarkStoppedIdempotentOnAlreadyStopped(t *testing.T)
}
svc := NewIngestionTaskService()
if err := svc.MarkStopped("task-1"); err != nil {
ctx := t.Context()
if err := svc.MarkStopped(ctx, "task-1"); err != nil {
t.Fatalf("MarkStopped on already STOPPED task should be idempotent, got: %v", err)
}
}
@@ -839,7 +861,8 @@ func TestIngestionTaskServiceMarkFailedIdempotentOnAlreadyTerminal(t *testing.T)
}
svc := NewIngestionTaskService()
if err := svc.MarkFailed("task-1"); err != nil {
ctx := t.Context()
if err := svc.MarkFailed(ctx, "task-1"); err != nil {
t.Fatalf("MarkFailed on already COMPLETED task should be idempotent, got: %v", err)
}
}
@@ -853,8 +876,9 @@ func TestIngestionTaskServiceMarkCompletedIdempotentOnAlreadyTerminal(t *testing
t.Fatalf("set task FAILED: %v", err)
}
ctx := t.Context()
svc := NewIngestionTaskService()
if err := svc.MarkCompleted("task-1"); err != nil {
if err := svc.MarkCompleted(ctx, "task-1"); err != nil {
t.Fatalf("MarkCompleted on already FAILED task should be idempotent, got: %v", err)
}
}