diff --git a/conf/service_conf.yaml b/conf/service_conf.yaml index 59f7de50e2..db12d36bcb 100644 --- a/conf/service_conf.yaml +++ b/conf/service_conf.yaml @@ -71,7 +71,7 @@ ingestor: max_concurrent_workers: 1 compiler_pool_size: 0 file_syncer: - max_concurrent_syncs: 1 + max_concurrent_syncs: 5 sync_interval: 3 user_default_llm: default_models: diff --git a/docker/service_conf.yaml.template b/docker/service_conf.yaml.template index 531f29a725..745b0def57 100644 --- a/docker/service_conf.yaml.template +++ b/docker/service_conf.yaml.template @@ -89,7 +89,7 @@ ingestor: max_concurrent_workers: 1 compiler_pool_size: 0 file_syncer: - max_concurrent_syncs: 1 + max_concurrent_syncs: 5 sync_interval: 3 user_default_llm: default_models: diff --git a/internal/common/task.go b/internal/common/task.go index c50bfba9f3..d6b1237e9e 100644 --- a/internal/common/task.go +++ b/internal/common/task.go @@ -27,6 +27,8 @@ const ( TaskTypeIngestionTask = "ingestion_task" TaskTypeIngestionTest = "ingestion_test" + // TaskTypeSyncer is the NATS wake-up message type for datasource sync_logs tasks. + TaskTypeSyncer = "syncer" // TaskTypeMemory is the async memory-extraction task type. Memory tasks // share the tasks.RAGFLOW subject and the Ingestor's consumer + worker // pool with ingestion tasks; processMessage dispatches them by TaskType. diff --git a/internal/dao/connector.go b/internal/dao/connector.go index 18b33d8ed3..e7dfd5b5fe 100644 --- a/internal/dao/connector.go +++ b/internal/dao/connector.go @@ -161,12 +161,12 @@ func (dao *ConnectorDAO) LinkDatasetConnectorsTx(ctx context.Context, tx *gorm.D return err } - if err := scheduleConnectorTask(ctx, tx, connector.ID, kbID, connectorTaskTypeSync, true); err != nil { + if _, err := scheduleConnectorTask(ctx, tx, connector.ID, kbID, connectorTaskTypeSync, true); err != nil { return err } if connectorConfigBool(fullConnector.Config, "sync_deleted_files") { - if err := scheduleConnectorTask(ctx, tx, connector.ID, kbID, connectorTaskTypePrune, false); err != nil { + if _, err := scheduleConnectorTask(ctx, tx, connector.ID, kbID, connectorTaskTypePrune, false); err != nil { return err } } @@ -223,8 +223,9 @@ func (dao *ConnectorDAO) CancelRunningOrScheduledLogs(ctx context.Context, db *g } // ScheduleConnectorTasks schedules sync and optional prune tasks for a connector. -func (dao *ConnectorDAO) ScheduleConnectorTasks(ctx context.Context, db *gorm.DB, connectorID string) error { - return db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { +func (dao *ConnectorDAO) ScheduleConnectorTasks(ctx context.Context, db *gorm.DB, connectorID string) ([]string, error) { + taskIDs := []string{} + err := db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { var connector entity.Connector if err := tx.WithContext(ctx).Where("id = ?", connectorID).First(&connector).Error; err != nil { return err @@ -236,13 +237,21 @@ func (dao *ConnectorDAO) ScheduleConnectorTasks(ctx context.Context, db *gorm.DB } for _, mapping := range mappings { - if err := scheduleConnectorTask(ctx, tx, connectorID, mapping.KbID, connectorTaskTypeSync, false); err != nil { + taskID, err := scheduleConnectorTask(ctx, tx, connectorID, mapping.KbID, connectorTaskTypeSync, false) + if err != nil { return err } + if taskID != "" { + taskIDs = append(taskIDs, taskID) + } if connectorConfigBool(connector.Config, "sync_deleted_files") { - if err := scheduleConnectorTask(ctx, tx, connectorID, mapping.KbID, connectorTaskTypePrune, false); err != nil { + taskID, err = scheduleConnectorTask(ctx, tx, connectorID, mapping.KbID, connectorTaskTypePrune, false) + if err != nil { return err } + if taskID != "" { + taskIDs = append(taskIDs, taskID) + } } } @@ -250,6 +259,7 @@ func (dao *ConnectorDAO) ScheduleConnectorTasks(ctx context.Context, db *gorm.DB Where("id = ?", connectorID). Update("status", string(entity.TaskStatusSchedule)).Error }) + return taskIDs, err } // ListDocumentsByKBAndSourceType lists connector documents in a dataset. @@ -260,8 +270,9 @@ func (dao *ConnectorDAO) ListDocumentsByKBAndSourceType(ctx context.Context, db } // RebuildConnector replaces old connector documents with scheduled sync tasks. -func (dao *ConnectorDAO) RebuildConnector(ctx context.Context, db *gorm.DB, connector *entity.Connector, kbID string, documents []*entity.Document) error { - return db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { +func (dao *ConnectorDAO) RebuildConnector(ctx context.Context, db *gorm.DB, connector *entity.Connector, kbID string, documents []*entity.Document) ([]string, error) { + taskIDs := []string{} + err := db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { if err := tx.WithContext(ctx).Where("connector_id = ? AND kb_id = ?", connector.ID, kbID).Delete(&entity.SyncLogs{}).Error; err != nil { return err } @@ -326,16 +337,21 @@ func (dao *ConnectorDAO) RebuildConnector(ctx context.Context, db *gorm.DB, conn return err } - if err := createRebuildSyncLog(ctx, tx, connector.ID, kbID, connectorTaskTypeSync, true); err != nil { + taskID, err := createRebuildSyncLog(ctx, tx, connector.ID, kbID, connectorTaskTypeSync, true) + if err != nil { return err } + taskIDs = append(taskIDs, taskID) if syncDeletedFiles, _ := connector.Config["sync_deleted_files"].(bool); syncDeletedFiles { - if err := createRebuildSyncLog(ctx, tx, connector.ID, kbID, connectorTaskTypePrune, false); err != nil { + taskID, err = createRebuildSyncLog(ctx, tx, connector.ID, kbID, connectorTaskTypePrune, false) + if err != nil { return err } + taskIDs = append(taskIDs, taskID) } return nil }) + return taskIDs, err } const ( @@ -343,14 +359,15 @@ const ( connectorTaskTypePrune = "prune" ) -func createRebuildSyncLog(ctx context.Context, tx *gorm.DB, connectorID, kbID, taskType string, reindex bool) error { +func createRebuildSyncLog(ctx context.Context, tx *gorm.DB, connectorID, kbID, taskType string, reindex bool) (string, error) { fromBeginning := "0" if reindex { fromBeginning = "1" } now := time.Now().Local() - return tx.WithContext(ctx).Create(&entity.SyncLogs{ - ID: utility.GenerateToken(), + taskID := utility.GenerateToken() + return taskID, tx.WithContext(ctx).Create(&entity.SyncLogs{ + ID: taskID, ConnectorID: connectorID, KbID: kbID, TaskType: taskType, @@ -362,15 +379,27 @@ func createRebuildSyncLog(ctx context.Context, tx *gorm.DB, connectorID, kbID, t }).Error } -func scheduleConnectorTask(ctx context.Context, tx *gorm.DB, connectorID, kbID, taskType string, reindex bool) error { - var existing int64 - if err := tx.WithContext(ctx).Model(&entity.SyncLogs{}). +func scheduleConnectorTask(ctx context.Context, tx *gorm.DB, connectorID, kbID, taskType string, reindex bool) (string, error) { + var scheduled entity.SyncLogs + err := tx.WithContext(ctx). Where("connector_id = ? AND kb_id = ? AND task_type = ? AND status = ?", connectorID, kbID, taskType, string(entity.TaskStatusSchedule)). - Count(&existing).Error; err != nil { - return err + Order("update_time DESC"). + First(&scheduled).Error + if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { + return "", err } - if existing > 0 { - return nil + if err == nil { + return scheduled.ID, nil + } + + var running int64 + if err := tx.WithContext(ctx).Model(&entity.SyncLogs{}). + Where("connector_id = ? AND kb_id = ? AND task_type = ? AND status = ?", connectorID, kbID, taskType, string(entity.TaskStatusRunning)). + Count(&running).Error; err != nil { + return "", err + } + if running > 0 { + return "", nil } var pollRangeStart *time.Time @@ -381,7 +410,7 @@ func scheduleConnectorTask(ctx context.Context, tx *gorm.DB, connectorID, kbID, Order("update_time DESC"). First(&latest).Error if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { - return err + return "", err } if err == nil { pollRangeStart = latest.PollRangeEnd @@ -394,8 +423,9 @@ func scheduleConnectorTask(ctx context.Context, tx *gorm.DB, connectorID, kbID, fromBeginning = "1" } now := time.Now().Local() - return tx.WithContext(ctx).Create(&entity.SyncLogs{ - ID: utility.GenerateToken(), + taskID := utility.GenerateToken() + return taskID, tx.WithContext(ctx).Create(&entity.SyncLogs{ + ID: taskID, ConnectorID: connectorID, KbID: kbID, TaskType: taskType, diff --git a/internal/dao/sync_task.go b/internal/dao/sync_task.go index 748bc09d01..00825859b1 100644 --- a/internal/dao/sync_task.go +++ b/internal/dao/sync_task.go @@ -121,15 +121,72 @@ func (d *SyncTaskDAO) ListDueTasks(ctx context.Context, now time.Time, limit int return tasks, nil } +// ListStartupTasks returns scheduled tasks that should be published during NATS startup reconciliation. +func (d *SyncTaskDAO) ListStartupTasks(ctx context.Context, limit int) ([]entity.SyncLogs, error) { + var tasks []entity.SyncLogs + err := d.db.WithContext(ctx). + Model(&entity.SyncLogs{}). + Select("sync_logs.*"). + Joins("JOIN connector ON sync_logs.connector_id = connector.id"). + Joins("JOIN connector2kb ON sync_logs.connector_id = connector2kb.connector_id AND sync_logs.kb_id = connector2kb.kb_id"). + Joins("JOIN knowledgebase ON sync_logs.kb_id = knowledgebase.id"). + Where("sync_logs.status = ? AND connector.status = ? AND sync_logs.task_type IN ?", SyncStatusSchedule, SyncStatusSchedule, []string{TaskTypeSync, TaskTypePrune}). + Order("sync_logs.update_time DESC"). + Limit(limit). + Find(&tasks).Error + return tasks, err +} + // ClaimTask conditionally marks a scheduled task as running. func (d *SyncTaskDAO) ClaimTask(ctx context.Context, taskID string, now time.Time) (bool, error) { - result := d.db.WithContext(ctx).Model(&entity.SyncLogs{}). - Where("id = ? AND status = ?", taskID, SyncStatusSchedule). - Updates(map[string]any{"status": SyncStatusRunning, "time_started": now}) - if result.Error != nil { - return false, result.Error - } - return result.RowsAffected == 1, nil + var claimed bool + err := d.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var task entity.SyncLogs + query := tx.WithContext(ctx) + if tx.Dialector.Name() != "sqlite" { + query = query.Clauses(clause.Locking{Strength: "UPDATE"}) + } + if err := query.Where("id = ?", taskID).First(&task).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil + } + return err + } + if task.Status != SyncStatusSchedule { + return nil + } + + var mapping entity.Connector2Kb + lockQuery := tx.WithContext(ctx) + if tx.Dialector.Name() != "sqlite" { + lockQuery = lockQuery.Clauses(clause.Locking{Strength: "UPDATE"}) + } + if err := lockQuery. + Where("connector_id = ? AND kb_id = ?", task.ConnectorID, task.KbID). + First(&mapping).Error; err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { + return err + } + + var running int64 + if err := tx.WithContext(ctx).Model(&entity.SyncLogs{}). + Where("id <> ? AND connector_id = ? AND kb_id = ? AND status = ? AND task_type IN ?", taskID, task.ConnectorID, task.KbID, SyncStatusRunning, []string{TaskTypeSync, TaskTypePrune}). + Count(&running).Error; err != nil { + return err + } + if running > 0 { + return nil + } + + result := tx.Model(&entity.SyncLogs{}). + Where("id = ? AND status = ?", taskID, SyncStatusSchedule). + Updates(map[string]any{"status": SyncStatusRunning, "time_started": now}) + if result.Error != nil { + return result.Error + } + claimed = result.RowsAffected == 1 + return nil + }) + return claimed, err } // GetTaskContext loads a task with connector, mapping, and knowledgebase rows. @@ -157,6 +214,15 @@ func (d *SyncTaskDAO) GetTaskContext(ctx context.Context, taskID string) (SyncTa return SyncTaskContext{Task: task, Connector: connector, Connector2Kb: connector2Kb, Knowledgebase: kb}, nil } +// IsTaskCanceled reports whether a sync_logs task has been canceled. +func (d *SyncTaskDAO) IsTaskCanceled(ctx context.Context, taskID string) (bool, error) { + var task entity.SyncLogs + if err := d.db.WithContext(ctx).Select("status").Where("id = ?", taskID).First(&task).Error; err != nil { + return false, err + } + return task.Status == SyncStatusCancel, nil +} + // MarkConnectorRunning marks a connector running. func (d *SyncTaskDAO) MarkConnectorRunning(ctx context.Context, connectorID string) error { return d.db.WithContext(ctx).Model(&entity.Connector{}).Where("id = ?", connectorID).Update("status", SyncStatusRunning).Error @@ -184,12 +250,16 @@ func (d *SyncTaskDAO) RescheduleClaimed(ctx context.Context, taskID string) erro // FailTask marks a task failed without advancing its poll waterline. func (d *SyncTaskDAO) FailTask(ctx context.Context, taskID, connectorID, message string, errorCount int64) error { return d.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Model(&entity.SyncLogs{}).Where("id = ?", taskID).Updates(map[string]any{ + result := tx.Model(&entity.SyncLogs{}).Where("id = ? AND status <> ?", taskID, SyncStatusCancel).Updates(map[string]any{ "status": SyncStatusFail, "error_msg": message, "error_count": errorCount, - }).Error; err != nil { - return err + }) + if result.Error != nil { + return result.Error + } + if result.RowsAffected == 0 { + return nil } if connectorID == "" { return nil @@ -201,15 +271,19 @@ func (d *SyncTaskDAO) FailTask(ctx context.Context, taskID, connectorID, message // CompleteSyncTask marks SYNC done and creates the next schedule task. func (d *SyncTaskDAO) CompleteSyncTask(ctx context.Context, taskContext SyncTaskContext, pollRangeEnd time.Time, newDocs, totalDocs, errorCount int64, errorMsg string) error { return d.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Model(&entity.SyncLogs{}).Where("id = ?", taskContext.Task.ID).Updates(map[string]any{ + result := tx.Model(&entity.SyncLogs{}).Where("id = ? AND status = ?", taskContext.Task.ID, SyncStatusRunning).Updates(map[string]any{ "status": SyncStatusDone, "poll_range_end": pollRangeEnd, "new_docs_indexed": newDocs, "total_docs_indexed": gorm.Expr("total_docs_indexed + ?", totalDocs), "error_msg": errorMsg, "error_count": errorCount, - }).Error; err != nil { - return err + }) + if result.Error != nil { + return result.Error + } + if result.RowsAffected == 0 { + return nil } if err := tx.Model(&entity.Connector{}).Where("id = ?", taskContext.Connector.ID).Update("status", SyncStatusDone).Error; err != nil { return err @@ -222,11 +296,15 @@ func (d *SyncTaskDAO) CompleteSyncTask(ctx context.Context, taskContext SyncTask // CompletePruneTask marks PRUNE done and creates the next schedule task. func (d *SyncTaskDAO) CompletePruneTask(ctx context.Context, taskContext SyncTaskContext, removed int64) error { return d.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Model(&entity.SyncLogs{}).Where("id = ?", taskContext.Task.ID).Updates(map[string]any{ + result := tx.Model(&entity.SyncLogs{}).Where("id = ? AND status = ?", taskContext.Task.ID, SyncStatusRunning).Updates(map[string]any{ "status": SyncStatusDone, "docs_removed_from_index": gorm.Expr("docs_removed_from_index + ?", removed), - }).Error; err != nil { - return err + }) + if result.Error != nil { + return result.Error + } + if result.RowsAffected == 0 { + return nil } if err := tx.Model(&entity.Connector{}).Where("id = ?", taskContext.Connector.ID).Update("status", SyncStatusDone).Error; err != nil { return err @@ -288,6 +366,44 @@ func (d *SyncTaskDAO) RecoverStaleRunning(ctx context.Context, now time.Time) (i return recovered, d.db.WithContext(ctx).Model(&entity.Connector{}).Where("id IN ? AND status = ?", ids, SyncStatusRunning).Update("status", SyncStatusSchedule).Error } +// RecoverRunning restores running sync tasks during syncer startup. +func (d *SyncTaskDAO) RecoverRunning(ctx context.Context) (int64, error) { + type runningTaskRow struct { + ID string `gorm:"column:id"` + ConnectorID string `gorm:"column:connector_id"` + } + + var rows []runningTaskRow + if err := d.db.WithContext(ctx). + Model(&entity.SyncLogs{}). + Select("id, connector_id"). + Where("status = ? AND task_type IN ?", SyncStatusRunning, []string{TaskTypeSync, TaskTypePrune}). + Scan(&rows).Error; err != nil { + return 0, err + } + if len(rows) == 0 { + return 0, nil + } + + connectorIDs := map[string]struct{}{} + return int64(len(rows)), d.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + for _, row := range rows { + if err := tx.Model(&entity.SyncLogs{}). + Where("id = ? AND status = ?", row.ID, SyncStatusRunning). + Update("status", SyncStatusSchedule).Error; err != nil { + return err + } + connectorIDs[row.ConnectorID] = struct{}{} + } + + ids := make([]string, 0, len(connectorIDs)) + for connectorID := range connectorIDs { + ids = append(ids, connectorID) + } + return tx.Model(&entity.Connector{}).Where("id IN ? AND status = ?", ids, SyncStatusRunning).Update("status", SyncStatusSchedule).Error + }) +} + // createScheduledTask creates the next Python-compatible scheduled task. func createScheduledTask(ctx context.Context, tx *gorm.DB, connectorID, kbID, taskType string, fromBeginning bool, pollRangeStart *time.Time, totalDocsIndexed int64) error { var lockRow entity.Connector2Kb @@ -317,6 +433,11 @@ func createScheduledTask(ctx context.Context, tx *gorm.DB, connectorID, kbID, ta } now := time.Now().Local() + if err := tx.WithContext(ctx).Model(&entity.Connector{}). + Where("id = ?", connectorID). + Update("status", SyncStatusSchedule).Error; err != nil { + return err + } return tx.WithContext(ctx).Create(&entity.SyncLogs{ ID: utility.GenerateToken(), ConnectorID: connectorID, diff --git a/internal/engine/nats/nats.go b/internal/engine/nats/nats.go index a0269b1cc3..d809699898 100644 --- a/internal/engine/nats/nats.go +++ b/internal/engine/nats/nats.go @@ -24,6 +24,7 @@ import ( "ragflow/internal/common" "strconv" "strings" + "sync" "time" "github.com/nats-io/nats.go" @@ -42,6 +43,10 @@ type NatsEngine struct { knowledgeCompileStream jetstream.Stream knowledgeCompileConsumer jetstream.Consumer kv jetstream.KeyValue + + syncerStream jetstream.Stream + syncerConsumer jetstream.Consumer + syncerMu sync.Mutex } func NewNatsEngine(host string, port int) *NatsEngine { diff --git a/internal/engine/nats/syncer.go b/internal/engine/nats/syncer.go new file mode 100644 index 0000000000..5b109d5171 --- /dev/null +++ b/internal/engine/nats/syncer.go @@ -0,0 +1,157 @@ +// +// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// + +package nats + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "strings" + "time" + + "ragflow/internal/common" + + "github.com/nats-io/nats.go" + "github.com/nats-io/nats.go/jetstream" +) + +const ( + // SyncerTaskSubject is the JetStream subject carrying sync_logs task IDs. + SyncerTaskSubject = "sync.tasks.RAGFLOW" + + syncerStreamName = "RAGFLOW_SYNC_TASKS" + syncerConsumerName = "RAGFLOW_SYNCER_CONSUMER" + syncerSubjectPattern = "sync.tasks.>" +) + +// InitSyncerStream creates the datasource syncer task stream. +func (n *NatsEngine) InitSyncerStream() error { + n.syncerMu.Lock() + defer n.syncerMu.Unlock() + return n.initSyncerStreamLocked() +} + +func (n *NatsEngine) initSyncerStreamLocked() error { + if n.jetStream == nil { + return fmt.Errorf("syncer: jetStream not initialized") + } + if n.syncerStream != nil { + return nil + } + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + // create jetStream + stream, err := n.jetStream.CreateStream(ctx, jetstream.StreamConfig{ + Name: syncerStreamName, + Subjects: []string{syncerSubjectPattern}, + Retention: jetstream.WorkQueuePolicy, + Storage: jetstream.FileStorage, + MaxMsgs: 1024 * 128, + MaxBytes: 1024 * 1024 * 64, + Duplicates: 10 * time.Minute, + }) + if err != nil { + if !strings.Contains(err.Error(), "already exists") { + return fmt.Errorf("syncer: create stream: %w", err) + } + stream, err = n.jetStream.Stream(ctx, syncerStreamName) + if err != nil { + return fmt.Errorf("syncer: get existing stream: %w", err) + } + } + n.syncerStream = stream + return nil +} + +// InitSyncerConsumer creates the durable pull consumer for syncer tasks. +func (n *NatsEngine) InitSyncerConsumer() error { + n.syncerMu.Lock() + defer n.syncerMu.Unlock() + if err := n.initSyncerStreamLocked(); err != nil { + return err + } + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + consumer, err := n.syncerStream.CreateOrUpdateConsumer(ctx, jetstream.ConsumerConfig{ + Name: syncerConsumerName, + Durable: syncerConsumerName, + AckPolicy: jetstream.AckExplicitPolicy, + MaxDeliver: 16, + MaxAckPending: 1024 * 128, + FilterSubject: syncerSubjectPattern, + }) + if err != nil { + if strings.Contains(err.Error(), "max waiting can not be updated") { + consumer, err = n.syncerStream.Consumer(ctx, syncerConsumerName) + if err != nil { + return fmt.Errorf("syncer: get existing consumer: %w", err) + } + } else { + return fmt.Errorf("syncer: create consumer: %w", err) + } + } + n.syncerConsumer = consumer + return nil +} + +// PublishSyncerTask publishes one sync_logs task wake-up. +func (n *NatsEngine) PublishSyncerTask(taskID string) error { + if err := n.InitSyncerStream(); err != nil { + return err + } + + payload, err := json.Marshal(common.TaskMessage{TaskID: taskID, TaskType: common.TaskTypeSyncer}) + if err != nil { + return err + } + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + // publish to nats + _, err = n.jetStream.Publish(ctx, SyncerTaskSubject, payload, jetstream.WithMsgID(taskID), jetstream.WithExpectStream(syncerStreamName)) + return err +} + +// FetchSyncerTasks pulls syncer task messages from JetStream. +func (n *NatsEngine) FetchSyncerTasks(batchSize int) ([]common.TaskHandle, error) { + n.syncerMu.Lock() + consumer := n.syncerConsumer + n.syncerMu.Unlock() + if consumer == nil { + return nil, fmt.Errorf("syncer: consumer not initialized") + } + // fetch task from nats(jetStream) + messages, err := consumer.Fetch(batchSize, jetstream.FetchMaxWait(1*time.Second)) + if err != nil { + if errors.Is(err, nats.ErrTimeout) { + return nil, nil + } + return nil, err + } + + handles := make([]common.TaskHandle, 0, batchSize) + for msg := range messages.Messages() { + handles = append(handles, NewNatsMessageHandle(msg)) + } + if err = messages.Error(); err != nil { + return handles, err + } + return handles, nil +} diff --git a/internal/engine/nats/syncer_test.go b/internal/engine/nats/syncer_test.go new file mode 100644 index 0000000000..03ae16adbc --- /dev/null +++ b/internal/engine/nats/syncer_test.go @@ -0,0 +1,71 @@ +package nats + +import ( + "net" + "testing" + "time" + + "ragflow/internal/common" + + "github.com/nats-io/nats-server/v2/server" +) + +// TestSyncerTaskStreamPublishesAndFetches verifies the dedicated syncer stream path. +func TestSyncerTaskStreamPublishesAndFetches(t *testing.T) { + engine := setupSyncerNATSEngine(t) + if err := engine.InitSyncerStream(); err != nil { + t.Fatalf("InitSyncerStream: %v", err) + } + if err := engine.InitSyncerConsumer(); err != nil { + t.Fatalf("InitSyncerConsumer: %v", err) + } + if err := engine.PublishSyncerTask("task-1"); err != nil { + t.Fatalf("PublishSyncerTask: %v", err) + } + + handles, err := engine.FetchSyncerTasks(1) + if err != nil { + t.Fatalf("FetchSyncerTasks: %v", err) + } + if len(handles) != 1 { + t.Fatalf("handles len = %d, want 1", len(handles)) + } + message := handles[0].GetMessage() + if message.TaskID != "task-1" || message.TaskType != common.TaskTypeSyncer { + t.Fatalf("message = %+v", message) + } + if err := handles[0].Ack(); err != nil { + t.Fatalf("Ack: %v", err) + } +} + +func setupSyncerNATSEngine(t *testing.T) *NatsEngine { + t.Helper() + opts := &server.Options{ + Port: -1, + JetStream: true, + StoreDir: t.TempDir(), + NoLog: true, + NoSigs: true, + } + ns, err := server.NewServer(opts) + if err != nil { + t.Fatalf("create embedded NATS server: %v", err) + } + ns.Start() + if !ns.ReadyForConnections(10 * time.Second) { + ns.Shutdown() + t.Fatal("embedded NATS server did not become ready") + } + t.Cleanup(func() { + ns.Shutdown() + ns.WaitForShutdown() + }) + + addr := ns.Addr().(*net.TCPAddr) + engine := NewNatsEngine("127.0.0.1", addr.Port) + if err := engine.Init(); err != nil { + t.Fatalf("NatsEngine.Init: %v", err) + } + return engine +} diff --git a/internal/server/config/syncer_config.go b/internal/server/config/syncer_config.go index 164b0ed73b..16f1911de4 100644 --- a/internal/server/config/syncer_config.go +++ b/internal/server/config/syncer_config.go @@ -25,7 +25,7 @@ type SyncerConfig struct { func (c *Config) ParseSyncerConfig(v *viper.Viper) error { // Default Syncer config - c.syncer.MaxConcurrentSyncs = 1 + c.syncer.MaxConcurrentSyncs = 5 c.syncer.SyncInterval = 3 if !v.IsSet("file_syncer") { diff --git a/internal/service/connector.go b/internal/service/connector.go index 210848938f..17fabc2149 100644 --- a/internal/service/connector.go +++ b/internal/service/connector.go @@ -87,6 +87,15 @@ type ConnectorService struct { userTenantDAO *dao.UserTenantDAO } +type syncTaskPublisher interface { + PublishSyncerTask(taskID string) error +} + +var getSyncerTaskPublisher = func() (syncTaskPublisher, bool) { + publisher, ok := engine.GetMessageQueueEngine().(syncTaskPublisher) + return publisher, ok +} + // NewConnectorService create connector service func NewConnectorService() *ConnectorService { return &ConnectorService{ @@ -921,7 +930,11 @@ func (s *ConnectorService) UpdateConnector(ctx context.Context, connectorID, use if err = s.cancelConnectorTasks(ctx, connectorID); err != nil { return nil, common.CodeServerError, err } - if err = s.connectorDAO.ScheduleConnectorTasks(ctx, dao.DB, connectorID); err != nil { + taskIDs, err := s.connectorDAO.ScheduleConnectorTasks(ctx, dao.DB, connectorID) + if err != nil { + return nil, common.CodeServerError, err + } + if err = publishSyncerTasks(taskIDs); err != nil { return nil, common.CodeServerError, err } } else if isConnectorCancelStatus(req.Status) { @@ -929,7 +942,11 @@ func (s *ConnectorService) UpdateConnector(ctx context.Context, connectorID, use return nil, common.CodeServerError, err } } else if isConnectorScheduleStatus(req.Status) { - if err = s.connectorDAO.ScheduleConnectorTasks(ctx, dao.DB, connectorID); err != nil { + taskIDs, err := s.connectorDAO.ScheduleConnectorTasks(ctx, dao.DB, connectorID) + if err != nil { + return nil, common.CodeServerError, err + } + if err = publishSyncerTasks(taskIDs); err != nil { return nil, common.CodeServerError, err } } @@ -989,12 +1006,35 @@ func (s *ConnectorService) RebuildConnector(ctx context.Context, connectorID, us s.deleteConnectorDocumentChunks(ctx, connector.TenantID, kbID, documents) - if err = s.connectorDAO.RebuildConnector(ctx, dao.DB, connector, kbID, documents); err != nil { + taskIDs, err := s.connectorDAO.RebuildConnector(ctx, dao.DB, connector, kbID, documents) + if err != nil { + return false, common.CodeServerError, err + } + if err = publishSyncerTasks(taskIDs); err != nil { return false, common.CodeServerError, err } return true, common.CodeSuccess, nil } +func publishSyncerTasks(taskIDs []string) error { + if len(taskIDs) == 0 { + return nil + } + publisher, ok := getSyncerTaskPublisher() + if !ok { + return fmt.Errorf("syncer task publisher is not configured") + } + for _, taskID := range taskIDs { + if taskID == "" { + continue + } + if err := publisher.PublishSyncerTask(taskID); err != nil { + return err + } + } + return nil +} + func (s *ConnectorService) deleteConnectorDocumentChunks(ctx context.Context, tenantID, kbID string, documents []*entity.Document) { docEngine := engine.Get() if docEngine == nil { diff --git a/internal/service/connector_syncer_test.go b/internal/service/connector_syncer_test.go new file mode 100644 index 0000000000..a29824830f --- /dev/null +++ b/internal/service/connector_syncer_test.go @@ -0,0 +1,167 @@ +package service + +import ( + "context" + "testing" + + "ragflow/internal/common" + "ragflow/internal/dao" + "ragflow/internal/entity" +) + +type fakeSyncerTaskPublisher struct { + taskIDs []string +} + +func (p *fakeSyncerTaskPublisher) PublishSyncerTask(taskID string) error { + p.taskIDs = append(p.taskIDs, taskID) + return nil +} + +func TestUpdateConnectorSchedulePublishesSyncerTask(t *testing.T) { + db := setupServiceTestDB(t) + pushServiceDB(t, db) + if err := db.AutoMigrate(&entity.Connector{}, &entity.Connector2Kb{}, &entity.Knowledgebase{}, &entity.SyncLogs{}); err != nil { + t.Fatalf("migrate connector tables: %v", err) + } + + if err := db.Create(&entity.Connector{ + ID: "conn-1", + TenantID: "user-1", + Name: "conn-1", + Source: "rss", + InputType: "poll", + Config: entity.JSONMap{}, + Status: string(entity.TaskStatusCancel), + RefreshFreq: 0, + PruneFreq: 0, + TimeoutSecs: 60, + }).Error; err != nil { + t.Fatalf("insert connector: %v", err) + } + if err := db.Create(&entity.Knowledgebase{ + ID: "kb-1", + TenantID: "user-1", + Name: "kb-1", + CreatedBy: "user-1", + EmbdID: "embd", + }).Error; err != nil { + t.Fatalf("insert kb: %v", err) + } + if err := db.Create(&entity.Connector2Kb{ + ID: "conn-1-kb-1", + ConnectorID: "conn-1", + KbID: "kb-1", + AutoParse: "1", + }).Error; err != nil { + t.Fatalf("insert connector2kb: %v", err) + } + + publisher := &fakeSyncerTaskPublisher{} + previousPublisher := getSyncerTaskPublisher + getSyncerTaskPublisher = func() (syncTaskPublisher, bool) { + return publisher, true + } + t.Cleanup(func() { getSyncerTaskPublisher = previousPublisher }) + + _, code, err := NewConnectorService().UpdateConnector(context.Background(), "conn-1", "user-1", &UpdateConnectorRequest{ + Status: string(entity.TaskStatusSchedule), + }) + if err != nil { + t.Fatalf("UpdateConnector error: %v", err) + } + if code != common.CodeSuccess { + t.Fatalf("code = %v, want success", code) + } + if len(publisher.taskIDs) != 1 { + t.Fatalf("published task IDs = %v, want one", publisher.taskIDs) + } + + var task entity.SyncLogs + if err := db.First(&task, "id = ?", publisher.taskIDs[0]).Error; err != nil { + t.Fatalf("load published task: %v", err) + } + if task.Status != string(entity.TaskStatusSchedule) || task.TaskType != dao.TaskTypeSync { + t.Fatalf("task status/type = %s/%s, want schedule/sync", task.Status, task.TaskType) + } +} + +func TestUpdateConnectorScheduleDoesNotDuplicateRunningTask(t *testing.T) { + db := setupServiceTestDB(t) + pushServiceDB(t, db) + if err := db.AutoMigrate(&entity.Connector{}, &entity.Connector2Kb{}, &entity.SyncLogs{}); err != nil { + t.Fatalf("migrate connector tables: %v", err) + } + + if err := db.Create(&entity.Connector{ + ID: "conn-1", + TenantID: "user-1", + Name: "conn-1", + Source: "rss", + InputType: "poll", + Config: entity.JSONMap{}, + Status: string(entity.TaskStatusRunning), + RefreshFreq: 0, + PruneFreq: 0, + TimeoutSecs: 60, + }).Error; err != nil { + t.Fatalf("insert connector: %v", err) + } + if err := db.Create(&entity.Knowledgebase{ + ID: "kb-1", + TenantID: "user-1", + Name: "kb-1", + CreatedBy: "user-1", + EmbdID: "embd", + }).Error; err != nil { + t.Fatalf("insert kb: %v", err) + } + if err := db.Create(&entity.Connector2Kb{ + ID: "conn-1-kb-1", + ConnectorID: "conn-1", + KbID: "kb-1", + AutoParse: "1", + }).Error; err != nil { + t.Fatalf("insert connector2kb: %v", err) + } + if err := db.Create(&entity.SyncLogs{ + ID: "running-task", + ConnectorID: "conn-1", + KbID: "kb-1", + TaskType: dao.TaskTypeSync, + Status: string(entity.TaskStatusRunning), + ErrorMsg: "", + }).Error; err != nil { + t.Fatalf("insert running task: %v", err) + } + + publisher := &fakeSyncerTaskPublisher{} + previousPublisher := getSyncerTaskPublisher + getSyncerTaskPublisher = func() (syncTaskPublisher, bool) { + return publisher, true + } + t.Cleanup(func() { getSyncerTaskPublisher = previousPublisher }) + + _, code, err := NewConnectorService().UpdateConnector(context.Background(), "conn-1", "user-1", &UpdateConnectorRequest{ + Status: string(entity.TaskStatusSchedule), + }) + if err != nil { + t.Fatalf("UpdateConnector error: %v", err) + } + if code != common.CodeSuccess { + t.Fatalf("code = %v, want success", code) + } + if len(publisher.taskIDs) != 0 { + t.Fatalf("published task IDs = %v, want none", publisher.taskIDs) + } + + var activeCount int64 + if err := db.Model(&entity.SyncLogs{}). + Where("connector_id = ? AND kb_id = ? AND task_type = ? AND status IN ?", "conn-1", "kb-1", dao.TaskTypeSync, []string{string(entity.TaskStatusSchedule), string(entity.TaskStatusRunning)}). + Count(&activeCount).Error; err != nil { + t.Fatalf("count active tasks: %v", err) + } + if activeCount != 1 { + t.Fatalf("active tasks = %d, want 1", activeCount) + } +} diff --git a/internal/service/sync_task_service.go b/internal/service/sync_task_service.go index c9ee375cc6..b3cf3e8064 100644 --- a/internal/service/sync_task_service.go +++ b/internal/service/sync_task_service.go @@ -92,6 +92,19 @@ func (s *SyncTaskService) ListDueTasks(ctx context.Context, now time.Time) ([]da return out, nil } +// ListStartupTasks returns scheduled tasks for NATS startup reconciliation. +func (s *SyncTaskService) ListStartupTasks(ctx context.Context) ([]dao.SyncTask, error) { + tasks, err := s.taskDAO.ListStartupTasks(ctx, 128) + if err != nil { + return nil, err + } + out := make([]dao.SyncTask, 0, len(tasks)) + for _, task := range tasks { + out = append(out, dao.SyncTask{SyncLogs: task}) + } + return out, nil +} + // Claim marks a scheduled task running if no other scanner claimed it first. func (s *SyncTaskService) Claim(ctx context.Context, taskID string) (bool, error) { claimed, err := s.taskDAO.ClaimTask(ctx, taskID, time.Now().Local()) @@ -107,13 +120,16 @@ func (s *SyncTaskService) Claim(ctx context.Context, taskID string) (bool, error return true, s.taskDAO.MarkConnectorRunning(ctx, taskContext.Connector.ID) } -// TODO: refactor some needless func - // GetContext loads a task execution context. func (s *SyncTaskService) GetContext(ctx context.Context, taskID string) (SyncTaskContext, error) { return s.taskDAO.GetTaskContext(ctx, taskID) } +// IsCanceled reports whether a task was canceled while a worker is running it. +func (s *SyncTaskService) IsCanceled(ctx context.Context, taskID string) (bool, error) { + return s.taskDAO.IsTaskCanceled(ctx, taskID) +} + // RescheduleClaimed puts a claimed task back into schedule state. func (s *SyncTaskService) RescheduleClaimed(ctx context.Context, taskID string) error { return s.taskDAO.RescheduleClaimed(ctx, taskID) @@ -145,6 +161,12 @@ func (s *SyncTaskService) RecoverStaleRunning(ctx context.Context, now time.Time return err } +// RecoverRunning restores running tasks during syncer startup. +func (s *SyncTaskService) RecoverRunning(ctx context.Context) error { + _, err := s.taskDAO.RecoverRunning(ctx) + return err +} + // IsFromBeginning reports whether a task is a full sync. func IsFromBeginning(value *string) bool { if value == nil { diff --git a/internal/syncer/config.go b/internal/syncer/config.go index 4abd51b219..2b9604b0e6 100644 --- a/internal/syncer/config.go +++ b/internal/syncer/config.go @@ -22,25 +22,25 @@ import ( // Config contains runtime limits for the datasource syncer. type Config struct { - PollInterval time.Duration - TaskConcurrency int - TaskQueueSize int - PerTaskItemConcurrency int - GlobalItemConcurrency int - ItemRetryCount int - ItemRetryBaseDelay time.Duration + PollInterval time.Duration + TaskWorkerCount int + TaskQueueSize int + JobWorkerCount int + JobQueueSize int + ItemRetryCount int + ItemRetryBaseDelay time.Duration } // DefaultConfig returns the first-version syncer defaults. func DefaultConfig() Config { return Config{ - PollInterval: 3 * time.Second, - TaskConcurrency: 3, - TaskQueueSize: 32, - PerTaskItemConcurrency: 4, - GlobalItemConcurrency: 12, - ItemRetryCount: 3, - ItemRetryBaseDelay: time.Second, + PollInterval: 3 * time.Second, + TaskWorkerCount: 5, + TaskQueueSize: 10, + JobWorkerCount: 400, + JobQueueSize: 400, + ItemRetryCount: 3, + ItemRetryBaseDelay: time.Second, } } @@ -51,20 +51,20 @@ func (c Config) Normalize() Config { c.PollInterval = def.PollInterval } - if c.TaskConcurrency <= 0 { - c.TaskConcurrency = def.TaskConcurrency + if c.TaskWorkerCount <= 0 { + c.TaskWorkerCount = def.TaskWorkerCount } if c.TaskQueueSize <= 0 { c.TaskQueueSize = def.TaskQueueSize } - if c.PerTaskItemConcurrency <= 0 { - c.PerTaskItemConcurrency = def.PerTaskItemConcurrency + if c.JobWorkerCount <= 0 { + c.JobWorkerCount = def.JobWorkerCount } - if c.GlobalItemConcurrency <= 0 { - c.GlobalItemConcurrency = def.GlobalItemConcurrency + if c.JobQueueSize <= 0 { + c.JobQueueSize = def.JobQueueSize } if c.ItemRetryCount <= 0 { diff --git a/internal/syncer/recovery.go b/internal/syncer/connector/fingerprint.go similarity index 68% rename from internal/syncer/recovery.go rename to internal/syncer/connector/fingerprint.go index 1fea849a1f..b463cca42b 100644 --- a/internal/syncer/recovery.go +++ b/internal/syncer/connector/fingerprint.go @@ -14,15 +14,21 @@ // limitations under the License. // -package syncer +package connector import ( - "context" - "ragflow/internal/service" - "time" + "encoding/hex" + "encoding/json" + "fmt" + + "github.com/zeebo/xxh3" ) -// RecoverStaleRunning restores timed-out running sync tasks to schedule. -func RecoverStaleRunning(ctx context.Context, taskService *service.SyncTaskService, now time.Time) error { - return taskService.RecoverStaleRunning(ctx, now) +func stableFingerprint(value any) string { + data, err := json.Marshal(value) + if err != nil { + data = []byte(fmt.Sprint(value)) + } + sum := xxh3.Hash128(data).Bytes() + return hex.EncodeToString(sum[:]) } diff --git a/internal/syncer/connector/github.go b/internal/syncer/connector/github.go index 0c0ab8bfe8..ae48aebd24 100644 --- a/internal/syncer/connector/github.go +++ b/internal/syncer/connector/github.go @@ -24,6 +24,7 @@ import ( "net/http" "net/url" "os" + "sort" "strconv" "strings" "time" @@ -308,6 +309,7 @@ func (s *githubSyncSession) NextBatch(ctx context.Context) (SyncBatch, error) { documents = append(documents, s.buffer[:n]...) s.buffer = s.buffer[n:] } + for len(documents) < s.batchSize { if s.repoIndex >= len(s.repos) { if len(documents) == 0 { @@ -516,6 +518,9 @@ type githubPullRequest struct { // toSourceDocument converts a pull request into the syncer model. func (p githubPullRequest) toSourceDocument(repo string) SourceDocument { body := []byte(p.Body) + labels := githubLabelNames(p.Labels) + user := p.User.metadata() + assignees := githubUsersMetadata(p.Assignees) return SourceDocument{ SourceID: p.HTMLURL, SemanticIdentifier: fmt.Sprintf("%d:%s", p.Number, sanitizeGitHubName(p.Title, "md")), @@ -529,10 +534,22 @@ func (p githubPullRequest) toSourceDocument(repo string) SourceDocument { "state": p.State, "repo": repo, "merged": strconv.FormatBool(p.MergedAt != nil), - "labels": githubLabelNames(p.Labels), - "user": p.User.metadata(), - "assignees": githubUsersMetadata(p.Assignees), + "labels": labels, + "user": user, + "assignees": assignees, }, + Fingerprint: stableFingerprint(map[string]any{ + "type": "PullRequest", + "url": p.HTMLURL, + "title": p.Title, + "body": p.Body, + "state": p.State, + "updated_at": p.UpdatedAt.UTC(), + "merged_at": p.MergedAt, + "labels": labels, + "user": user, + "assignees": assignees, + }), } } @@ -554,6 +571,9 @@ type githubIssue struct { // toSourceDocument converts an issue into the syncer model. func (i githubIssue) toSourceDocument(repo string) SourceDocument { body := []byte(i.Body) + labels := githubLabelNames(i.Labels) + user := i.User.metadata() + assignees := githubUsersMetadata(i.Assignees) return SourceDocument{ SourceID: i.HTMLURL, SemanticIdentifier: fmt.Sprintf("%d:%s", i.Number, sanitizeGitHubName(i.Title, "md")), @@ -566,10 +586,22 @@ func (i githubIssue) toSourceDocument(repo string) SourceDocument { "id": strconv.Itoa(i.Number), "state": i.State, "repo": repo, - "labels": githubLabelNames(i.Labels), - "user": i.User.metadata(), - "assignees": githubUsersMetadata(i.Assignees), + "labels": labels, + "user": user, + "assignees": assignees, }, + Fingerprint: stableFingerprint(map[string]any{ + "type": "Issue", + "url": i.HTMLURL, + "title": i.Title, + "body": i.Body, + "state": i.State, + "updated_at": i.UpdatedAt.UTC(), + "closed_at": i.ClosedAt, + "labels": labels, + "user": user, + "assignees": assignees, + }), } } @@ -609,11 +641,16 @@ func githubLabelNames(labels []githubLabel) []string { out = append(out, label.Name) } } + sort.Strings(out) return out } // githubUsersMetadata returns metadata for users. func githubUsersMetadata(users []githubUser) []map[string]string { + users = append([]githubUser(nil), users...) + sort.Slice(users, func(i, j int) bool { + return githubUserSortKey(users[i]) < githubUserSortKey(users[j]) + }) out := make([]map[string]string, 0, len(users)) for i := range users { out = append(out, (&users[i]).metadata()) @@ -621,6 +658,16 @@ func githubUsersMetadata(users []githubUser) []map[string]string { return out } +func githubUserSortKey(user githubUser) string { + if user.Login != "" { + return user.Login + } + if user.Email != "" { + return user.Email + } + return user.Name +} + // hasNextPage reports whether a GitHub Link header has rel next. func hasNextPage(headers http.Header) bool { return strings.Contains(headers.Get("Link"), `rel="next"`) diff --git a/internal/syncer/connector/github_test.go b/internal/syncer/connector/github_test.go index 7ceae7b19f..18aa5915cf 100644 --- a/internal/syncer/connector/github_test.go +++ b/internal/syncer/connector/github_test.go @@ -3,11 +3,95 @@ package connector import ( "context" "encoding/json" + "errors" + "io" "net/http" "net/url" "testing" + "time" ) +// TestGitHubConnectorOpenSyncUsesWindowAndFingerprint verifies incremental sync emits only updated docs with fingerprints. +func TestGitHubConnectorOpenSyncUsesWindowAndFingerprint(t *testing.T) { + connector, err := NewGitHubConnector(map[string]any{ + "repository_owner": "openai", + "repository_name": "ragflow", + "include_pull_requests": true, + "include_issues": true, + "batch_size": 10, + "credentials": map[string]any{"github_access_token": "token"}, + }) + if err != nil { + t.Fatalf("NewGitHubConnector failed: %v", err) + } + connector.baseURL = "https://api.github.test" + connector.doJSON = githubFixtureDoJSON(t) + + start := mustTime(t, "2026-01-02T12:00:00Z") + end := mustTime(t, "2026-01-04T00:00:00Z") + session, err := connector.OpenSync(context.Background(), SyncRequest{WindowStart: &start, WindowEnd: end}) + if err != nil { + t.Fatalf("OpenSync failed: %v", err) + } + batch, err := session.NextBatch(context.Background()) + if err != nil { + t.Fatalf("NextBatch failed: %v", err) + } + if len(batch.Documents) != 1 { + t.Fatalf("documents len = %d, want 1", len(batch.Documents)) + } + doc := batch.Documents[0] + if doc.SourceID != "https://github.com/openai/ragflow/pull/7" { + t.Fatalf("source id = %q", doc.SourceID) + } + if doc.Fingerprint == "" { + t.Fatalf("fingerprint is empty") + } + if _, err = session.NextBatch(context.Background()); !errors.Is(err, io.EOF) { + t.Fatalf("NextBatch EOF = %v", err) + } +} + +// TestGitHubFingerprintStable verifies GitHub fingerprints are stable and content-sensitive. +func TestGitHubFingerprintStable(t *testing.T) { + updatedAt := time.Date(2026, 1, 3, 0, 0, 0, 0, time.UTC) + pr := githubPullRequest{ + HTMLURL: "https://github.com/openai/ragflow/pull/7", + Number: 7, + Title: "Add syncer", + Body: "PR body", + State: "open", + UpdatedAt: updatedAt, + User: &githubUser{Login: "alice"}, + Assignees: []githubUser{ + {Login: "zoe"}, + {Login: "bob"}, + }, + Labels: []githubLabel{ + {Name: "sync"}, + {Name: "bug"}, + }, + } + fp1 := pr.toSourceDocument("openai/ragflow").Fingerprint + fp2 := pr.toSourceDocument("openai/ragflow").Fingerprint + if fp1 == "" || fp1 != fp2 { + t.Fatalf("fingerprint unstable: %q %q", fp1, fp2) + } + + reordered := pr + reordered.Labels = []githubLabel{{Name: "bug"}, {Name: "sync"}} + reordered.Assignees = []githubUser{{Login: "bob"}, {Login: "zoe"}} + if got := reordered.toSourceDocument("openai/ragflow").Fingerprint; got != fp1 { + t.Fatalf("fingerprint changed after order-only change: %q != %q", got, fp1) + } + + changed := pr + changed.Title = "Add syncer v2" + if got := changed.toSourceDocument("openai/ragflow").Fingerprint; got == fp1 { + t.Fatalf("fingerprint did not change after title update") + } +} + // TestGitHubConnectorOpenPrune verifies PRUNE returns Python-compatible html_url IDs. func TestGitHubConnectorOpenPrune(t *testing.T) { connector, err := NewGitHubConnector(map[string]any{ diff --git a/internal/syncer/connector/gmail.go b/internal/syncer/connector/gmail.go index 408112a19c..9f18f6f5c0 100644 --- a/internal/syncer/connector/gmail.go +++ b/internal/syncer/connector/gmail.go @@ -26,6 +26,7 @@ import ( "net/mail" "net/url" "os" + "sort" "strings" "sync" "time" @@ -47,10 +48,6 @@ var gmailScopes = []string{ "https://www.googleapis.com/auth/admin.directory.group.readonly", } -// FIXME: IDK why everytime, gmail do sync, It will update all file's Metadata. -// FIXME: I think this need to be checked or fixed after all data syncer is done -// FIXME: Some file sync from gmail have no content, this need to be checked too - // GmailConnector reads Gmail threads from a Workspace domain or one Gmail account. type GmailConnector struct { primaryAdminEmail string @@ -526,6 +523,12 @@ func (t gmailThread) toSourceDocument(userEmail string) (SourceDocument, bool) { UpdatedAt: updatedAt, SizeBytes: int64(len(blob)), Metadata: metadata, + Fingerprint: stableFingerprint(map[string]any{ + "thread_id": t.ID, + "updated_at": updatedAt, + "blob": string(blob), + "metadata": metadata, + }), }, true } @@ -659,7 +662,13 @@ func parseGmailAddress(value string) (string, string) { // gmailOwnersMetadata converts email owners to compact metadata. func gmailOwnersMetadata(owners map[string]string) []map[string]string { out := make([]map[string]string, 0, len(owners)) - for email, name := range owners { + emails := make([]string, 0, len(owners)) + for email := range owners { + emails = append(emails, email) + } + sort.Strings(emails) + for _, email := range emails { + name := owners[email] item := map[string]string{"email": email} if name != "" { parts := strings.Fields(name) @@ -695,12 +704,65 @@ func isGmailDisabled(err error) bool { // isGoogleForbiddenOrNotFound reports item-level Google visibility failures. func isGoogleForbiddenOrNotFound(err error) bool { + return isGooglePermissionDeniedOrNotFound(err) +} + +func isGooglePermissionDeniedOrNotFound(err error) bool { if httpErr, ok := err.(googleHTTPError); ok { - return httpErr.status == http.StatusForbidden || httpErr.status == http.StatusNotFound + if httpErr.status == http.StatusNotFound { + return true + } + return httpErr.status == http.StatusForbidden && !isGoogleRateLimited(err) } return false } +func isGoogleRateLimited(err error) bool { + httpErr, ok := err.(googleHTTPError) + if !ok { + return false + } + if httpErr.status == http.StatusTooManyRequests { + return true + } + if httpErr.status != http.StatusForbidden { + return false + } + if strings.Contains(httpErr.body, "rateLimitExceeded") || strings.Contains(httpErr.body, "userRateLimitExceeded") || strings.Contains(httpErr.body, "quotaExceeded") { + return true + } + for _, reason := range googleErrorReasons(httpErr.body) { + switch reason { + case "rateLimitExceeded", "userRateLimitExceeded", "quotaExceeded", "dailyLimitExceeded", "RESOURCE_EXHAUSTED": + } + } + return false +} + +func googleErrorReasons(body string) []string { + var response struct { + Error struct { + Errors []struct { + Reason string `json:"reason"` + } `json:"errors"` + Status string `json:"status"` + } `json:"error"` + } + if err := json.Unmarshal([]byte(body), &response); err != nil { + return nil + } + reasons := make([]string, 0, len(response.Error.Errors)+1) + for _, item := range response.Error.Errors { + if item.Reason != "" { + reasons = append(reasons, item.Reason) + } + } + if response.Error.Status != "" { + reasons = append(reasons, response.Error.Status) + } + return reasons +} + type googleHTTPError struct { status int body string diff --git a/internal/syncer/connector/gmail_test.go b/internal/syncer/connector/gmail_test.go index f266dbb87d..6bfa76f96f 100644 --- a/internal/syncer/connector/gmail_test.go +++ b/internal/syncer/connector/gmail_test.go @@ -51,11 +51,74 @@ func TestGmailConnectorOpenSync(t *testing.T) { if doc.Metadata["external_user_emails"].([]string)[0] != "admin@example.com" { t.Fatalf("external user metadata = %v", doc.Metadata["external_user_emails"]) } + if doc.Fingerprint == "" { + t.Fatalf("fingerprint is empty") + } if _, err = session.NextBatch(context.Background()); !errors.Is(err, io.EOF) { t.Fatalf("NextBatch EOF = %v", err) } } +// TestGmailFingerprintStable verifies Gmail fingerprints are stable and content-sensitive. +func TestGmailFingerprintStable(t *testing.T) { + thread := gmailThread{ + ID: "thread-1", + Messages: []gmailMessage{{ + ID: "msg-1", + Payload: gmailPayload{ + Headers: []gmailHeader{ + {Name: "From", Value: "Alice Example "}, + {Name: "To", Value: "Bob "}, + {Name: "Cc", Value: "Carol "}, + {Name: "Subject", Value: "Hello"}, + {Name: "Date", Value: "Fri, 02 Jan 2026 03:04:05 +0000"}, + }, + Parts: []gmailPart{{ + MimeType: "text/plain", + Body: gmailBody{Data: base64.RawURLEncoding.EncodeToString([]byte("Body text"))}, + }}, + }, + }}, + } + doc1, ok := thread.toSourceDocument("admin@example.com") + if !ok { + t.Fatalf("thread did not produce document") + } + doc2, ok := thread.toSourceDocument("admin@example.com") + if !ok { + t.Fatalf("thread did not produce second document") + } + if doc1.Fingerprint == "" || doc1.Fingerprint != doc2.Fingerprint { + t.Fatalf("fingerprint unstable: %q %q", doc1.Fingerprint, doc2.Fingerprint) + } + + changed := thread + changed.Messages = append([]gmailMessage(nil), thread.Messages...) + changed.Messages[0].Payload.Headers = append([]gmailHeader(nil), thread.Messages[0].Payload.Headers...) + changed.Messages[0].Payload.Headers[3].Value = "Hello v2" + doc3, ok := changed.toSourceDocument("admin@example.com") + if !ok { + t.Fatalf("changed thread did not produce document") + } + if doc3.Fingerprint == doc1.Fingerprint { + t.Fatalf("fingerprint did not change after subject update") + } +} + +// TestGmailOwnersMetadataSorted verifies owner metadata is deterministic. +func TestGmailOwnersMetadataSorted(t *testing.T) { + owners := gmailOwnersMetadata(map[string]string{ + "carol@example.com": "Carol Example", + "alice@example.com": "Alice Example", + }) + if len(owners) != 2 { + t.Fatalf("owners len = %d, want 2", len(owners)) + } + if owners[0]["email"] != "alice@example.com" || owners[1]["email"] != "carol@example.com" { + t.Fatalf("owners not sorted: %+v", owners) + } +} + // TestGmailConnectorOpenPrune verifies Gmail prune emits thread IDs only. func TestGmailConnectorOpenPrune(t *testing.T) { connector := newFixtureGmailConnector() diff --git a/internal/syncer/connector/google_drive.go b/internal/syncer/connector/google_drive.go new file mode 100644 index 0000000000..96da7d9113 --- /dev/null +++ b/internal/syncer/connector/google_drive.go @@ -0,0 +1,1078 @@ +// +// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// + +package connector + +import ( + "context" + "encoding/json" + "fmt" + "io" + "mime" + "net/http" + "net/url" + "os" + "path" + "sort" + "strconv" + "strings" + "sync" + "time" + + "golang.org/x/oauth2" + "golang.org/x/oauth2/google" +) + +const ( + defaultGoogleDriveBatchSize = 32 + defaultGoogleDriveSizeThreshold = 10 * 1024 * 1024 + googleDriveItemsPerPage = 100 + googleDriveRequestTimeout = 60 * time.Second + googleDriveOAuthTokenURL = "https://oauth2.googleapis.com/token" + googleDriveListRetryCount = 4 + googleDriveListRetryBaseDelay = 200 * time.Millisecond + + googleDriveFolderMimeType = "application/vnd.google-apps.folder" + googleDriveShortcutMimeType = "application/vnd.google-apps.shortcut" +) + +var googleDriveScopes = []string{ + "https://www.googleapis.com/auth/drive.readonly", + "https://www.googleapis.com/auth/drive.metadata.readonly", + "https://www.googleapis.com/auth/admin.directory.group.readonly", + "https://www.googleapis.com/auth/admin.directory.user.readonly", +} + +var googleDriveNativeExports = map[string]googleDriveExport{ + "application/vnd.google-apps.document": {mimeType: "application/vnd.openxmlformats-officedocument.wordprocessingml.document", extension: ".docx"}, + "application/vnd.google-apps.spreadsheet": {mimeType: "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet", extension: ".xlsx"}, + "application/vnd.google-apps.presentation": {mimeType: "application/vnd.openxmlformats-officedocument.presentationml.presentation", extension: ".pptx"}, +} + +// GoogleDriveConnector reads files from Google Drive. +type GoogleDriveConnector struct { + includeSharedDrives bool + includeMyDrives bool + includeFilesSharedWithMe bool + allowImages bool + sharedDriveIDs []string + myDriveEmails []string + sharedFolderIDs []string + specificUserEmails []string + specificRequests bool + primaryAdminEmail string + credentials map[string]any + batchSize int + sizeThreshold int64 + clientsMu sync.Mutex + clients map[string]*http.Client + httpClientForUser func(ctx context.Context, userEmail string) (*http.Client, error) + listUsers func(ctx context.Context) ([]string, error) + listDrives func(ctx context.Context, userEmail string) ([]googleDriveDrive, error) + listFiles func(ctx context.Context, userEmail string, request googleDriveListRequest) (googleDriveFilePage, error) + listFolders func(ctx context.Context, userEmail, parentID string) ([]string, error) + downloadFile func(ctx context.Context, userEmail string, file googleDriveFile) ([]byte, string, error) +} + +// NewGoogleDriveConnector creates a Google Drive connector from Python-compatible config. +func NewGoogleDriveConnector(config map[string]any) (*GoogleDriveConnector, error) { + credentials, _ := config["credentials"].(map[string]any) + specificRequests := stringConfig(config["shared_drive_urls"]) != "" || stringConfig(config["my_drive_emails"]) != "" || stringConfig(config["shared_folder_urls"]) != "" + includeSharedWithMe := configBoolDefault(config["include_files_shared_with_me"], false) + if specificRequests { + includeSharedWithMe = false + } + sizeThreshold := int64(configInt(config["size_threshold"], defaultGoogleDriveSizeThreshold)) + if sizeThreshold <= 0 { + sizeThreshold = defaultGoogleDriveSizeThreshold + } + return &GoogleDriveConnector{ + includeSharedDrives: !specificRequests && configBoolDefault(config["include_shared_drives"], false), + includeMyDrives: !specificRequests && configBoolDefault(config["include_my_drives"], false), + includeFilesSharedWithMe: includeSharedWithMe, + allowImages: configBoolDefault(config["allow_images"], false), + sharedDriveIDs: googleDriveIDsFromURLs(stringConfig(config["shared_drive_urls"])), + myDriveEmails: splitCommaList(stringConfig(config["my_drive_emails"])), + sharedFolderIDs: googleDriveIDsFromURLs(stringConfig(config["shared_folder_urls"])), + specificUserEmails: splitCommaList(stringConfig(config["specific_user_emails"])), + specificRequests: specificRequests, + primaryAdminEmail: strings.TrimSpace(stringConfig(credentials["google_primary_admin"])), + credentials: credentials, + batchSize: configInt(firstNonEmpty(stringConfig(config["sync_batch_size"]), stringConfig(config["batch_size"])), defaultGoogleDriveBatchSize), + sizeThreshold: sizeThreshold, + clients: map[string]*http.Client{}, + }, nil +} + +// Validate validates Google Drive connector settings and credentials. +func (c *GoogleDriveConnector) Validate(ctx context.Context) error { + if c == nil { + return fmt.Errorf("google drive connector is nil") + } + if c.primaryAdminEmail == "" { + return fmt.Errorf("Google Drive connector is missing google_primary_admin") + } + if len(c.credentials) == 0 { + return fmt.Errorf("Google Drive connector is missing credentials") + } + if c.batchSize <= 0 { + return fmt.Errorf("batch_size must be a positive integer") + } + if !c.hasRetrievalScope() { + return fmt.Errorf("Nothing to index. Please specify include_shared_drives, include_my_drives, include_files_shared_with_me, shared_drive_urls, shared_folder_urls, or my_drive_emails") + } + _, err := c.clientForUser(ctx, c.primaryAdminEmail) + return err +} + +// OpenSync opens one Google Drive sync session. +func (c *GoogleDriveConnector) OpenSync(ctx context.Context, request SyncRequest) (SyncSession, error) { + scopes, err := c.buildScopes(ctx) + if err != nil { + return nil, err + } + return &googleDriveSyncSession{ + connector: c, + scopes: scopes, + scopeIndex: 0, + batchSize: c.batchSize, + windowStart: request.WindowStart, + windowEnd: request.WindowEnd, + }, nil +} + +// OpenPrune opens one complete Google Drive prune snapshot session. +func (c *GoogleDriveConnector) OpenPrune(ctx context.Context, request PruneRequest) (PruneSession, error) { + scopes, err := c.buildScopes(ctx) + if err != nil { + return nil, err + } + return &googleDrivePruneSession{connector: c, scopes: scopes, batchSize: c.batchSize}, nil +} + +// Fetch downloads a Google Drive file after the sync runner determines it changed. +func (c *GoogleDriveConnector) Fetch(ctx context.Context, ref FetchReference) ([]byte, error) { + var fetch googleDriveFetchReference + if err := json.Unmarshal([]byte(ref.Key), &fetch); err != nil { + return nil, err + } + file := googleDriveFile{ + ID: fetch.FileID, + Name: fetch.Name, + MimeType: fetch.MimeType, + Size: fetch.Size, + } + blob, _, err := c.fetchBlob(ctx, fetch.UserEmail, file) + return blob, err +} + +func (c *GoogleDriveConnector) hasRetrievalScope() bool { + return c.includeSharedDrives || c.includeMyDrives || c.includeFilesSharedWithMe || len(c.sharedDriveIDs) > 0 || len(c.sharedFolderIDs) > 0 || len(c.myDriveEmails) > 0 +} + +func (c *GoogleDriveConnector) buildScopes(ctx context.Context) ([]googleDriveScope, error) { + users, err := c.userEmails(ctx) + if err != nil { + return nil, err + } + scopes := make([]googleDriveScope, 0) + includeSharedWithMe := c.effectiveIncludeSharedWithMe() + for _, email := range users { + if c.includeMyDrives || containsString(c.myDriveEmails, email) { + scopes = append(scopes, googleDriveScope{userEmail: email, corpora: "user", includeSharedWithMe: includeSharedWithMe}) + } + } + if includeSharedWithMe && !c.includeMyDrives { + for _, email := range users { + scopes = append(scopes, googleDriveScope{userEmail: email, corpora: "user", sharedWithMeOnly: true}) + } + } + + driveIDs := append([]string(nil), c.sharedDriveIDs...) + if c.includeSharedDrives { + allDrives, err := c.allDriveIDs(ctx, c.primaryAdminEmail) + if err != nil { + return nil, err + } + driveIDs = append(driveIDs, allDrives...) + } + driveIDs = uniqueSorted(driveIDs) + for _, driveID := range driveIDs { + scopes = append(scopes, googleDriveScope{userEmail: c.primaryAdminEmail, corpora: "drive", driveID: driveID}) + } + for _, folderID := range c.sharedFolderIDs { + scopes = append(scopes, googleDriveScope{userEmail: c.primaryAdminEmail, corpora: "folder", folderID: folderID}) + } + if len(scopes) == 0 { + return nil, fmt.Errorf("Google Drive connector has no effective retrieval scope") + } + return scopes, nil +} + +func (c *GoogleDriveConnector) userEmails(ctx context.Context) ([]string, error) { + if len(c.specificUserEmails) > 0 { + return uniqueSorted(c.specificUserEmails), nil + } + if len(c.myDriveEmails) > 0 { + return uniqueSorted(c.myDriveEmails), nil + } + if c.listUsers != nil { + return c.listUsers(ctx) + } + if !c.isServiceAccount() { + return []string{c.primaryAdminEmail}, nil + } + client, err := c.clientForUser(ctx, c.primaryAdminEmail) + if err != nil { + return nil, err + } + domain := c.primaryAdminEmail + if _, after, ok := strings.Cut(c.primaryAdminEmail, "@"); ok { + domain = after + } + users := []string{c.primaryAdminEmail} + for _, queryText := range []string{"isAdmin=true", "isAdmin=false"} { + pageToken := "" + for { + query := url.Values{"domain": {domain}, "fields": {"nextPageToken,users(primaryEmail)"}, "maxResults": {"500"}, "query": {queryText}} + if pageToken != "" { + query.Set("pageToken", pageToken) + } + var page googleDriveUsersPage + if err = c.getJSON(ctx, client, "https://admin.googleapis.com/admin/directory/v1/users?"+query.Encode(), &page); err != nil { + return nil, err + } + for _, user := range page.Users { + if user.PrimaryEmail != "" && !containsString(users, user.PrimaryEmail) { + users = append(users, user.PrimaryEmail) + } + } + if page.NextPageToken == "" { + break + } + pageToken = page.NextPageToken + } + } + return users, nil +} + +func (c *GoogleDriveConnector) allDriveIDs(ctx context.Context, userEmail string) ([]string, error) { + if c.listDrives != nil { + drives, err := c.listDrives(ctx, userEmail) + if err != nil { + return nil, err + } + return googleDriveIDs(drives), nil + } + client, err := c.clientForUser(ctx, userEmail) + if err != nil { + return nil, err + } + ids := []string{} + pageToken := "" + for { + query := url.Values{"fields": {"nextPageToken,drives(id)"}, "pageSize": {"100"}} + if c.isServiceAccount() { + query.Set("useDomainAdminAccess", "true") + } + if pageToken != "" { + query.Set("pageToken", pageToken) + } + var page googleDriveDrivePage + if err = c.getJSON(ctx, client, "https://www.googleapis.com/drive/v3/drives?"+query.Encode(), &page); err != nil { + return nil, err + } + ids = append(ids, googleDriveIDs(page.Drives)...) + if page.NextPageToken == "" { + break + } + pageToken = page.NextPageToken + } + return uniqueSorted(ids), nil +} + +func (c *GoogleDriveConnector) listFilePage(ctx context.Context, scope googleDriveScope, pageToken string, windowStart *time.Time, windowEnd time.Time) (googleDriveFilePage, error) { + if c.listFiles != nil { + return c.listFiles(ctx, scope.userEmail, googleDriveListRequest{Scope: scope, PageToken: pageToken, WindowStart: windowStart, WindowEnd: windowEnd, PageSize: googleDriveItemsPerPage}) + } + client, err := c.clientForUser(ctx, scope.userEmail) + if err != nil { + return googleDriveFilePage{}, err + } + query := url.Values{ + "fields": {"nextPageToken,files(id,name,mimeType,modifiedTime,createdTime,webViewLink,shortcutDetails,owners(emailAddress),size,md5Checksum,fileExtension)"}, + "pageSize": {strconv.Itoa(googleDriveItemsPerPage)}, + "orderBy": {"modifiedTime"}, + "q": {googleDriveFileQuery(scope, windowStart, windowEnd)}, + } + if pageToken != "" { + query.Set("pageToken", pageToken) + } + if scope.corpora == "drive" { + query.Set("corpora", "drive") + query.Set("driveId", scope.driveID) + query.Set("supportsAllDrives", "true") + query.Set("includeItemsFromAllDrives", "true") + } else if scope.corpora == "folder" { + query.Set("corpora", "allDrives") + query.Set("supportsAllDrives", "true") + query.Set("includeItemsFromAllDrives", "true") + } else if scope.includeSharedWithMe && !scope.sharedWithMeOnly { + query.Set("corpora", "allDrives") + query.Set("supportsAllDrives", "true") + query.Set("includeItemsFromAllDrives", "true") + } else { + query.Set("corpora", "user") + } + var page googleDriveFilePage + err = c.getJSON(ctx, client, "https://www.googleapis.com/drive/v3/files?"+query.Encode(), &page) + return page, err +} + +func (c *GoogleDriveConnector) listFilePageWithRetry(ctx context.Context, scope googleDriveScope, pageToken string, windowStart *time.Time, windowEnd time.Time) (googleDriveFilePage, error) { + var lastErr error + for attempt := 1; attempt <= googleDriveListRetryCount; attempt++ { + page, err := c.listFilePage(ctx, scope, pageToken, windowStart, windowEnd) + if err == nil { + return page, nil + } + lastErr = err + if !isGoogleRateLimited(err) || attempt == googleDriveListRetryCount { + break + } + delay := googleDriveListRetryBaseDelay * time.Duration(1<<(attempt-1)) + select { + case <-ctx.Done(): + return googleDriveFilePage{}, ctx.Err() + case <-time.After(delay): + } + } + return googleDriveFilePage{}, lastErr +} + +func (c *GoogleDriveConnector) listFolderIDs(ctx context.Context, userEmail, parentID string) ([]string, error) { + if c.listFolders != nil { + return c.listFolders(ctx, userEmail, parentID) + } + client, err := c.clientForUser(ctx, userEmail) + if err != nil { + return nil, err + } + ids := []string{} + pageToken := "" + for { + query := url.Values{ + "corpora": {"allDrives"}, + "fields": {"nextPageToken,files(id)"}, + "includeItemsFromAllDrives": {"true"}, + "supportsAllDrives": {"true"}, + "pageSize": {strconv.Itoa(googleDriveItemsPerPage)}, + "q": {fmt.Sprintf("mimeType = '%s' and trashed = false and '%s' in parents", googleDriveFolderMimeType, parentID)}, + } + if pageToken != "" { + query.Set("pageToken", pageToken) + } + var page googleDriveFilePage + if err = c.getJSON(ctx, client, "https://www.googleapis.com/drive/v3/files?"+query.Encode(), &page); err != nil { + return nil, err + } + for _, file := range page.Files { + if file.ID != "" { + ids = append(ids, file.ID) + } + } + if page.NextPageToken == "" { + break + } + pageToken = page.NextPageToken + } + return uniqueSorted(ids), nil +} + +func (c *GoogleDriveConnector) fetchBlob(ctx context.Context, userEmail string, file googleDriveFile) ([]byte, string, error) { + if c.downloadFile != nil { + return c.downloadFile(ctx, userEmail, file) + } + if file.isFolderLike() { + return nil, "", fmt.Errorf("Google Drive file %s is not downloadable", file.ID) + } + if file.isImage() && !c.allowImages { + return nil, "", fmt.Errorf("Google Drive image %s is disabled", file.ID) + } + if file.sizeInt() > c.sizeThreshold && file.sizeInt() > 0 { + return nil, "", fmt.Errorf("Google Drive file %s exceeds size threshold", file.ID) + } + client, err := c.clientForUser(ctx, userEmail) + if err != nil { + return nil, "", err + } + export, isNative := googleDriveNativeExports[file.MimeType] + extension := file.extension() + apiURL := "https://www.googleapis.com/drive/v3/files/" + url.PathEscape(file.ID) + "?alt=media" + if isNative { + extension = export.extension + apiURL = "https://www.googleapis.com/drive/v3/files/" + url.PathEscape(file.ID) + "/export?mimeType=" + url.QueryEscape(export.mimeType) + } else if strings.HasPrefix(file.MimeType, "application/vnd.google-apps") { + extension = ".pdf" + apiURL = "https://www.googleapis.com/drive/v3/files/" + url.PathEscape(file.ID) + "/export?mimeType=" + url.QueryEscape("application/pdf") + } + blob, err := c.getBytes(ctx, client, apiURL, c.sizeThreshold) + return blob, extension, err +} + +func (c *GoogleDriveConnector) clientForUser(ctx context.Context, userEmail string) (*http.Client, error) { + c.clientsMu.Lock() + defer c.clientsMu.Unlock() + if c.clients == nil { + c.clients = map[string]*http.Client{} + } + if client := c.clients[userEmail]; client != nil { + return client, nil + } + var client *http.Client + var err error + if c.httpClientForUser != nil { + client, err = c.httpClientForUser(ctx, userEmail) + } else { + var tokenSource oauth2.TokenSource + tokenSource, err = c.tokenSource(ctx, userEmail) + if err == nil { + client = oauth2.NewClient(ctx, tokenSource) + } + } + if err != nil { + return nil, err + } + c.clients[userEmail] = client + return client, nil +} + +func (c *GoogleDriveConnector) tokenSource(ctx context.Context, userEmail string) (oauth2.TokenSource, error) { + if value := stringConfig(c.credentials["google_service_account_key"]); value != "" { + config, err := google.JWTConfigFromJSON([]byte(value), googleDriveScopes...) + if err != nil { + return nil, err + } + config.Subject = userEmail + return config.TokenSource(ctx), nil + } + tokenJSON := stringConfig(c.credentials["google_tokens"]) + if tokenJSON == "" { + return nil, fmt.Errorf("Google Drive connector credentials must include google_tokens or google_service_account_key") + } + var tokenData gmailOAuthToken + if err := json.Unmarshal([]byte(tokenJSON), &tokenData); err != nil { + return nil, err + } + if tokenData.ClientID == "" { + tokenData.ClientID = os.Getenv("OAUTH_GOOGLE_DRIVE_CLIENT_ID") + } + if tokenData.ClientSecret == "" { + tokenData.ClientSecret = os.Getenv("OAUTH_GOOGLE_DRIVE_CLIENT_SECRET") + } + if tokenData.ClientID == "" || tokenData.ClientSecret == "" || tokenData.RefreshToken == "" { + return nil, fmt.Errorf("Google Drive OAuth credentials are incomplete") + } + return (&oauth2.Config{ + ClientID: tokenData.ClientID, + ClientSecret: tokenData.ClientSecret, + Endpoint: oauth2.Endpoint{TokenURL: googleDriveOAuthTokenURL}, + Scopes: googleDriveScopes, + }).TokenSource(ctx, tokenData.token()), nil +} + +func (c *GoogleDriveConnector) getJSON(ctx context.Context, client *http.Client, apiURL string, out any) error { + ctx, cancel := context.WithTimeout(ctx, googleDriveRequestTimeout) + defer cancel() + req, err := http.NewRequestWithContext(ctx, http.MethodGet, apiURL, nil) + if err != nil { + return err + } + resp, err := client.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + if resp.StatusCode >= 400 { + body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096)) + return googleHTTPError{status: resp.StatusCode, body: strings.TrimSpace(string(body))} + } + return json.NewDecoder(resp.Body).Decode(out) +} + +func (c *GoogleDriveConnector) getBytes(ctx context.Context, client *http.Client, apiURL string, sizeThreshold int64) ([]byte, error) { + ctx, cancel := context.WithTimeout(ctx, googleDriveRequestTimeout) + defer cancel() + req, err := http.NewRequestWithContext(ctx, http.MethodGet, apiURL, nil) + if err != nil { + return nil, err + } + resp, err := client.Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + if resp.StatusCode >= 400 { + body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096)) + return nil, googleHTTPError{status: resp.StatusCode, body: strings.TrimSpace(string(body))} + } + blob, err := io.ReadAll(io.LimitReader(resp.Body, sizeThreshold+1)) + if err != nil { + return nil, err + } + if int64(len(blob)) > sizeThreshold { + return nil, fmt.Errorf("Google Drive file exceeds size threshold") + } + return blob, nil +} + +type googleDriveSyncSession struct { + connector *GoogleDriveConnector + scopes []googleDriveScope + scopeIndex int + pageToken string + batchSize int + windowStart *time.Time + windowEnd time.Time + buffer []SourceDocument + seen map[string]struct{} + folderSeen map[string]struct{} +} + +// NextBatch returns the next Google Drive document batch. +func (s *googleDriveSyncSession) NextBatch(ctx context.Context) (SyncBatch, error) { + if s.seen == nil { + s.seen = map[string]struct{}{} + } + if s.folderSeen == nil { + s.folderSeen = map[string]struct{}{} + } + documents := make([]SourceDocument, 0, s.batchSize) + if len(s.buffer) > 0 { + n := min(s.batchSize, len(s.buffer)) + documents = append(documents, s.buffer[:n]...) + s.buffer = s.buffer[n:] + } + for len(documents) < s.batchSize { + if s.scopeIndex >= len(s.scopes) { + if len(documents) == 0 { + return SyncBatch{}, io.EOF + } + break + } + batch, err := s.nextDocumentPage(ctx) + if err != nil { + return SyncBatch{}, err + } + remaining := s.batchSize - len(documents) + if len(batch) > remaining { + documents = append(documents, batch[:remaining]...) + s.buffer = append(s.buffer, batch[remaining:]...) + break + } + documents = append(documents, batch...) + } + return SyncBatch{Documents: documents}, nil +} + +func (s *googleDriveSyncSession) nextDocumentPage(ctx context.Context) ([]SourceDocument, error) { + scope := s.scopes[s.scopeIndex] + page, err := s.connector.listFilePageWithRetry(ctx, scope, s.pageToken, s.windowStart, s.windowEnd) + if err != nil { + if isGooglePermissionDeniedOrNotFound(err) { + s.advanceScope() + return nil, nil + } + return nil, err + } + docs := make([]SourceDocument, 0, len(page.Files)) + for _, file := range page.Files { + doc, ok := file.toSourceDocument(scope.userEmail, s.connector.allowImages) + if !ok { + continue + } + if _, ok = s.seen[doc.SourceID]; ok { + continue + } + s.seen[doc.SourceID] = struct{}{} + docs = append(docs, doc) + } + if page.NextPageToken == "" { + if err = s.finishScope(ctx, scope); err != nil { + return nil, err + } + } else { + s.pageToken = page.NextPageToken + } + return docs, nil +} + +func (s *googleDriveSyncSession) finishScope(ctx context.Context, scope googleDriveScope) error { + if scope.corpora == "folder" { + childIDs, err := s.connector.listFolderIDs(ctx, scope.userEmail, scope.folderID) + if err != nil { + if !isGooglePermissionDeniedOrNotFound(err) { + return err + } + } + for _, folderID := range childIDs { + if _, ok := s.folderSeen[folderID]; ok { + continue + } + s.folderSeen[folderID] = struct{}{} + s.scopes = append(s.scopes, googleDriveScope{userEmail: scope.userEmail, corpora: "folder", folderID: folderID}) + } + } + s.advanceScope() + return nil +} + +func (s *googleDriveSyncSession) advanceScope() { + s.scopeIndex++ + s.pageToken = "" +} + +// Close closes the Google Drive sync session. +func (s *googleDriveSyncSession) Close() error { + return nil +} + +// Fetch downloads a Google Drive file for this sync session. +func (s *googleDriveSyncSession) Fetch(ctx context.Context, ref FetchReference) ([]byte, error) { + return s.connector.Fetch(ctx, ref) +} + +type googleDrivePruneSession struct { + connector *GoogleDriveConnector + scopes []googleDriveScope + scopeIndex int + pageToken string + batchSize int + buffer []SlimDocument + seen map[string]struct{} + folderSeen map[string]struct{} +} + +// NextBatch returns the next Google Drive prune snapshot batch. +func (s *googleDrivePruneSession) NextBatch(ctx context.Context) (PruneBatch, error) { + if s.seen == nil { + s.seen = map[string]struct{}{} + } + if s.folderSeen == nil { + s.folderSeen = map[string]struct{}{} + } + documents := make([]SlimDocument, 0, s.batchSize) + if len(s.buffer) > 0 { + n := min(s.batchSize, len(s.buffer)) + documents = append(documents, s.buffer[:n]...) + s.buffer = s.buffer[n:] + } + for len(documents) < s.batchSize { + if s.scopeIndex >= len(s.scopes) { + if len(documents) == 0 { + return PruneBatch{}, io.EOF + } + break + } + batch, err := s.nextSlimPage(ctx) + if err != nil { + return PruneBatch{}, err + } + remaining := s.batchSize - len(documents) + if len(batch) > remaining { + documents = append(documents, batch[:remaining]...) + s.buffer = append(s.buffer, batch[remaining:]...) + break + } + documents = append(documents, batch...) + } + return PruneBatch{Documents: documents}, nil +} + +func (s *googleDrivePruneSession) nextSlimPage(ctx context.Context) ([]SlimDocument, error) { + scope := s.scopes[s.scopeIndex] + page, err := s.connector.listFilePageWithRetry(ctx, scope, s.pageToken, nil, time.Time{}) + if err != nil { + if isGooglePermissionDeniedOrNotFound(err) { + s.advanceScope() + return nil, nil + } + return nil, err + } + docs := make([]SlimDocument, 0, len(page.Files)) + for _, file := range page.Files { + sourceID, ok := file.sourceID() + if !ok { + continue + } + if _, ok = s.seen[sourceID]; ok { + continue + } + s.seen[sourceID] = struct{}{} + docs = append(docs, SlimDocument{SourceID: sourceID}) + } + if page.NextPageToken == "" { + if err = s.finishScope(ctx, scope); err != nil { + return nil, err + } + } else { + s.pageToken = page.NextPageToken + } + return docs, nil +} + +func (s *googleDrivePruneSession) finishScope(ctx context.Context, scope googleDriveScope) error { + if scope.corpora == "folder" { + childIDs, err := s.connector.listFolderIDs(ctx, scope.userEmail, scope.folderID) + if err != nil { + if !isGooglePermissionDeniedOrNotFound(err) { + return err + } + } + for _, folderID := range childIDs { + if _, ok := s.folderSeen[folderID]; ok { + continue + } + s.folderSeen[folderID] = struct{}{} + s.scopes = append(s.scopes, googleDriveScope{userEmail: scope.userEmail, corpora: "folder", folderID: folderID}) + } + } + s.advanceScope() + return nil +} + +func (s *googleDrivePruneSession) advanceScope() { + s.scopeIndex++ + s.pageToken = "" +} + +// Close closes the Google Drive prune session. +func (s *googleDrivePruneSession) Close() error { + return nil +} + +type googleDriveScope struct { + userEmail string + corpora string + driveID string + folderID string + includeSharedWithMe bool + sharedWithMeOnly bool +} + +type googleDriveListRequest struct { + Scope googleDriveScope + PageToken string + WindowStart *time.Time + WindowEnd time.Time + PageSize int +} + +type googleDriveFilePage struct { + NextPageToken string `json:"nextPageToken"` + Files []googleDriveFile `json:"files"` +} + +type googleDriveDrivePage struct { + NextPageToken string `json:"nextPageToken"` + Drives []googleDriveDrive `json:"drives"` +} + +type googleDriveDrive struct { + ID string `json:"id"` +} + +type googleDriveUsersPage struct { + NextPageToken string `json:"nextPageToken"` + Users []struct { + PrimaryEmail string `json:"primaryEmail"` + } `json:"users"` +} + +type googleDriveFile struct { + ID string `json:"id"` + Name string `json:"name"` + MimeType string `json:"mimeType"` + ModifiedTime string `json:"modifiedTime"` + CreatedTime string `json:"createdTime"` + WebViewLink string `json:"webViewLink"` + Size string `json:"size"` + MD5Checksum string `json:"md5Checksum"` + FileExtension string `json:"fileExtension"` + ShortcutDetails struct { + TargetID string `json:"targetId"` + TargetMimeType string `json:"targetMimeType"` + } `json:"shortcutDetails"` + Owners []struct { + EmailAddress string `json:"emailAddress"` + } `json:"owners"` +} + +func (f googleDriveFile) toSourceDocument(userEmail string, allowImages bool) (SourceDocument, bool) { + if f.isFolderLike() || (f.isImage() && !allowImages) { + return SourceDocument{}, false + } + sourceID, ok := f.sourceID() + if !ok { + return SourceDocument{}, false + } + updatedAt := f.updatedAt() + fetch := googleDriveFetchReference{FileID: f.ID, UserEmail: userEmail, MimeType: f.MimeType, Name: f.Name, Size: f.Size} + fetchKey, _ := json.Marshal(fetch) + metadata := map[string]any{ + "file_id": f.ID, + "mime_type": f.MimeType, + "web_view_link": f.WebViewLink, + "owners": f.ownerEmails(), + } + return SourceDocument{ + SourceID: sourceID, + SemanticIdentifier: f.Name, + Extension: f.extension(), + FetchRef: &FetchReference{Key: string(fetchKey), SizeHint: f.sizeInt()}, + UpdatedAt: updatedAt, + SizeBytes: f.sizeInt(), + Metadata: metadata, + Fingerprint: f.fingerprint(), + }, true +} + +func (f googleDriveFile) sourceID() (string, bool) { + link := f.WebViewLink + if link == "" && f.ID != "" { + if template := googleDriveFallbackLinkTemplate(f.MimeType); template != "" { + link = fmt.Sprintf(template, f.ID) + } else { + link = "https://drive.google.com/file/d/" + url.PathEscape(f.ID) + "/view" + } + } + if link == "" { + return "", false + } + parsed, err := url.Parse(link) + if err != nil { + return strings.TrimSpace(link), true + } + parsed.RawQuery = "" + parsed.Fragment = "" + parts := strings.Split(strings.TrimRight(parsed.Path, "/"), "/") + if len(parts) > 0 { + last := parts[len(parts)-1] + if last == "edit" || last == "view" || last == "preview" { + parsed.Path = strings.TrimRight(strings.TrimSuffix(parsed.Path, "/"+last), "/") + } + } + return parsed.String(), true +} + +func (f googleDriveFile) fingerprint() string { + return stableFingerprint(map[string]any{ + "id": f.ID, + "name": f.Name, + "mime_type": f.MimeType, + "modified_time": f.ModifiedTime, + "created_time": f.CreatedTime, + "size": f.Size, + "md5": f.MD5Checksum, + "owners": f.ownerEmails(), + }) +} + +func (f googleDriveFile) updatedAt() time.Time { + for _, value := range []string{f.ModifiedTime, f.CreatedTime} { + if value == "" { + continue + } + if parsed, err := time.Parse(time.RFC3339Nano, value); err == nil { + return parsed.UTC() + } + } + return time.Now().UTC() +} + +func (f googleDriveFile) extension() string { + if export, ok := googleDriveNativeExports[f.MimeType]; ok { + return export.extension + } + if strings.HasPrefix(f.MimeType, "application/vnd.google-apps") { + return ".pdf" + } + if f.FileExtension != "" { + return "." + strings.TrimPrefix(f.FileExtension, ".") + } + if ext := path.Ext(f.Name); ext != "" { + return ext + } + if exts, err := mime.ExtensionsByType(f.MimeType); err == nil && len(exts) > 0 { + return exts[0] + } + return ".bin" +} + +func (f googleDriveFile) isFolderLike() bool { + return f.MimeType == googleDriveFolderMimeType || f.MimeType == googleDriveShortcutMimeType +} + +func (f googleDriveFile) isImage() bool { + return strings.HasPrefix(f.MimeType, "image/") && f.MimeType != "image/bmp" && f.MimeType != "image/tiff" && f.MimeType != "image/gif" && f.MimeType != "image/svg+xml" && f.MimeType != "image/avif" +} + +func (f googleDriveFile) sizeInt() int64 { + if f.Size == "" { + return 0 + } + size, _ := strconv.ParseInt(f.Size, 10, 64) + return size +} + +func (f googleDriveFile) ownerEmails() []string { + owners := make([]string, 0, len(f.Owners)) + for _, owner := range f.Owners { + if owner.EmailAddress != "" { + owners = append(owners, owner.EmailAddress) + } + } + sort.Strings(owners) + return owners +} + +type googleDriveFetchReference struct { + FileID string `json:"file_id"` + UserEmail string `json:"user_email"` + MimeType string `json:"mime_type"` + Name string `json:"name"` + Size string `json:"size"` +} + +type googleDriveExport struct { + mimeType string + extension string +} + +func googleDriveQuoteLiteral(value string) string { + value = strings.ReplaceAll(value, `\`, `\\`) + return strings.ReplaceAll(value, `'`, `\'`) +} + +func googleDriveFileQuery(scope googleDriveScope, windowStart *time.Time, windowEnd time.Time) string { + parts := []string{fmt.Sprintf("mimeType != '%s'", googleDriveFolderMimeType), "trashed = false"} + if scope.corpora == "folder" { + parts = append(parts, fmt.Sprintf("'%s' in parents", googleDriveQuoteLiteral(scope.folderID))) + } + if scope.corpora == "user" { + if scope.sharedWithMeOnly { + parts = append(parts, "not 'me' in owners") + } else if !scope.includeSharedWithMe { + parts = append(parts, "'me' in owners") + } + } + if windowStart != nil { + start := windowStart.UTC().Format(time.RFC3339) + parts = append(parts, fmt.Sprintf("(modifiedTime > '%s' or createdTime >= '%s')", start, start)) + } + if !windowEnd.IsZero() { + parts = append(parts, fmt.Sprintf("modifiedTime <= '%s'", windowEnd.UTC().Format(time.RFC3339))) + } + return strings.Join(parts, " and ") +} + +func googleDriveIDsFromURLs(value string) []string { + ids := []string{} + for _, item := range splitCommaList(value) { + parsed, err := url.Parse(item) + if err != nil || parsed.Path == "" { + ids = append(ids, strings.Trim(item, "/")) + continue + } + parts := strings.Split(strings.Trim(parsed.Path, "/"), "/") + if len(parts) > 0 && parts[len(parts)-1] != "" { + ids = append(ids, parts[len(parts)-1]) + } + } + return uniqueSorted(ids) +} + +func googleDriveIDs(drives []googleDriveDrive) []string { + ids := make([]string, 0, len(drives)) + for _, drive := range drives { + if drive.ID != "" { + ids = append(ids, drive.ID) + } + } + return ids +} + +func googleDriveFallbackLinkTemplate(mimeType string) string { + switch mimeType { + case "application/vnd.google-apps.document": + return "https://docs.google.com/document/d/%s/view" + case "application/vnd.google-apps.spreadsheet": + return "https://docs.google.com/spreadsheets/d/%s/view" + case "application/vnd.google-apps.presentation": + return "https://docs.google.com/presentation/d/%s/view" + default: + return "" + } +} + +func (c *GoogleDriveConnector) isServiceAccount() bool { + return stringConfig(c.credentials["google_service_account_key"]) != "" +} + +func (c *GoogleDriveConnector) effectiveIncludeSharedWithMe() bool { + if c.isServiceAccount() && !c.specificRequests { + return true + } + return c.includeFilesSharedWithMe +} + +func splitCommaList(value string) []string { + values := []string{} + for _, part := range strings.Split(value, ",") { + part = strings.TrimSpace(part) + if part != "" { + values = append(values, part) + } + } + return values +} + +func uniqueSorted(values []string) []string { + seen := map[string]struct{}{} + out := make([]string, 0, len(values)) + for _, value := range values { + value = strings.TrimSpace(value) + if value == "" { + continue + } + if _, ok := seen[value]; ok { + continue + } + seen[value] = struct{}{} + out = append(out, value) + } + sort.Strings(out) + return out +} + +func containsString(values []string, value string) bool { + for _, item := range values { + if item == value { + return true + } + } + return false +} diff --git a/internal/syncer/connector/google_drive_test.go b/internal/syncer/connector/google_drive_test.go new file mode 100644 index 0000000000..045188938f --- /dev/null +++ b/internal/syncer/connector/google_drive_test.go @@ -0,0 +1,235 @@ +package connector + +import ( + "context" + "errors" + "io" + "net/http" + "strings" + "testing" + "time" +) + +// TestGoogleDriveConnectorOpenSyncUsesWindowFingerprintAndFetch verifies incremental listing and lazy download. +func TestGoogleDriveConnectorOpenSyncUsesWindowFingerprintAndFetch(t *testing.T) { + connector, err := NewGoogleDriveConnector(map[string]any{ + "my_drive_emails": "admin@example.com", + "batch_size": 2, + "credentials": map[string]any{ + "google_primary_admin": "admin@example.com", + "google_tokens": `{"client_id":"client","client_secret":"secret","refresh_token":"refresh"}`, + }, + }) + if err != nil { + t.Fatalf("NewGoogleDriveConnector failed: %v", err) + } + var gotRequest googleDriveListRequest + connector.listFiles = func(ctx context.Context, userEmail string, request googleDriveListRequest) (googleDriveFilePage, error) { + gotRequest = request + return googleDriveFilePage{Files: []googleDriveFile{{ + ID: "file-1", + Name: "Plan.txt", + MimeType: "text/plain", + ModifiedTime: "2026-01-03T00:00:00Z", + CreatedTime: "2026-01-01T00:00:00Z", + WebViewLink: "https://drive.google.com/file/d/file-1/view?usp=sharing", + Size: "9", + MD5Checksum: "md5-1", + Owners: []struct { + EmailAddress string `json:"emailAddress"` + }{{EmailAddress: "owner@example.com"}}, + }}}, nil + } + connector.downloadFile = func(ctx context.Context, userEmail string, file googleDriveFile) ([]byte, string, error) { + if userEmail != "admin@example.com" || file.ID != "file-1" { + t.Fatalf("unexpected fetch user/file: %s %s", userEmail, file.ID) + } + return []byte("plan body"), ".txt", nil + } + + start := mustTime(t, "2026-01-02T00:00:00Z") + end := mustTime(t, "2026-01-04T00:00:00Z") + session, err := connector.OpenSync(context.Background(), SyncRequest{WindowStart: &start, WindowEnd: end}) + if err != nil { + t.Fatalf("OpenSync failed: %v", err) + } + batch, err := session.NextBatch(context.Background()) + if err != nil { + t.Fatalf("NextBatch failed: %v", err) + } + if gotRequest.WindowStart == nil || !gotRequest.WindowStart.Equal(start) || !gotRequest.WindowEnd.Equal(end) { + t.Fatalf("window = %v %v", gotRequest.WindowStart, gotRequest.WindowEnd) + } + if gotRequest.Scope.userEmail != "admin@example.com" || gotRequest.Scope.corpora != "user" { + t.Fatalf("scope = %+v", gotRequest.Scope) + } + if len(batch.Documents) != 1 { + t.Fatalf("documents len = %d, want 1", len(batch.Documents)) + } + doc := batch.Documents[0] + if doc.SourceID != "https://drive.google.com/file/d/file-1" { + t.Fatalf("source id = %q", doc.SourceID) + } + if doc.Fingerprint == "" { + t.Fatalf("fingerprint is empty") + } + if doc.FetchRef == nil { + t.Fatalf("fetch ref is nil") + } + fetcher, ok := session.(Fetcher) + if !ok { + t.Fatalf("session does not implement Fetcher") + } + blob, err := fetcher.Fetch(context.Background(), *doc.FetchRef) + if err != nil { + t.Fatalf("Fetch failed: %v", err) + } + if string(blob) != "plan body" { + t.Fatalf("blob = %q", string(blob)) + } + if _, err = session.NextBatch(context.Background()); !errors.Is(err, io.EOF) { + t.Fatalf("NextBatch EOF = %v", err) + } +} + +// TestGoogleDriveSharedFolderScopesRecurse verifies shared folders walk child folders. +func TestGoogleDriveSharedFolderScopesRecurse(t *testing.T) { + connector, err := NewGoogleDriveConnector(map[string]any{ + "shared_folder_urls": "https://drive.google.com/drive/folders/root-folder", + "batch_size": 10, + "credentials": map[string]any{ + "google_primary_admin": "admin@example.com", + "google_tokens": `{"client_id":"client","client_secret":"secret","refresh_token":"refresh"}`, + }, + }) + if err != nil { + t.Fatalf("NewGoogleDriveConnector failed: %v", err) + } + connector.listFiles = func(ctx context.Context, userEmail string, request googleDriveListRequest) (googleDriveFilePage, error) { + switch request.Scope.folderID { + case "root-folder": + return googleDriveFilePage{}, nil + case "child-folder": + return googleDriveFilePage{Files: []googleDriveFile{{ + ID: "child-file", + Name: "Child.txt", + MimeType: "text/plain", + ModifiedTime: "2026-01-03T00:00:00Z", + WebViewLink: "https://drive.google.com/file/d/child-file/view", + }}}, nil + default: + t.Fatalf("unexpected folder scope %q", request.Scope.folderID) + return googleDriveFilePage{}, nil + } + } + connector.listFolders = func(ctx context.Context, userEmail, parentID string) ([]string, error) { + if parentID == "root-folder" { + return []string{"child-folder"}, nil + } + return nil, nil + } + + session, err := connector.OpenSync(context.Background(), SyncRequest{FromBeginning: true}) + if err != nil { + t.Fatalf("OpenSync failed: %v", err) + } + batch, err := session.NextBatch(context.Background()) + if err != nil { + t.Fatalf("NextBatch failed: %v", err) + } + if len(batch.Documents) != 1 || batch.Documents[0].SourceID != "https://drive.google.com/file/d/child-file" { + t.Fatalf("unexpected recursive documents: %+v", batch.Documents) + } +} + +// TestGoogleDriveRateLimitRetries verifies rate limits do not truncate a scope. +func TestGoogleDriveRateLimitRetries(t *testing.T) { + connector, err := NewGoogleDriveConnector(map[string]any{ + "my_drive_emails": "admin@example.com", + "batch_size": 10, + "credentials": map[string]any{ + "google_primary_admin": "admin@example.com", + "google_tokens": `{"client_id":"client","client_secret":"secret","refresh_token":"refresh"}`, + }, + }) + if err != nil { + t.Fatalf("NewGoogleDriveConnector failed: %v", err) + } + calls := 0 + connector.listFiles = func(ctx context.Context, userEmail string, request googleDriveListRequest) (googleDriveFilePage, error) { + calls++ + if calls == 1 { + return googleDriveFilePage{}, googleHTTPError{ + status: http.StatusForbidden, + body: `{"error":{"errors":[{"reason":"rateLimitExceeded"}],"status":"RESOURCE_EXHAUSTED"}}`, + } + } + return googleDriveFilePage{Files: []googleDriveFile{{ + ID: "file-1", + Name: "Plan.txt", + MimeType: "text/plain", + ModifiedTime: "2026-01-03T00:00:00Z", + WebViewLink: "https://drive.google.com/file/d/file-1/view", + }}}, nil + } + + session, err := connector.OpenSync(context.Background(), SyncRequest{FromBeginning: true}) + if err != nil { + t.Fatalf("OpenSync failed: %v", err) + } + batch, err := session.NextBatch(context.Background()) + if err != nil { + t.Fatalf("NextBatch failed: %v", err) + } + if calls != 2 { + t.Fatalf("list calls = %d, want retry", calls) + } + if len(batch.Documents) != 1 { + t.Fatalf("documents len = %d, want 1", len(batch.Documents)) + } +} + +// TestGoogleDriveFingerprintStable verifies fingerprints are stable and metadata-sensitive. +func TestGoogleDriveFingerprintStable(t *testing.T) { + file := googleDriveFile{ + ID: "file-1", + Name: "Plan.txt", + MimeType: "text/plain", + ModifiedTime: "2026-01-03T00:00:00Z", + CreatedTime: "2026-01-01T00:00:00Z", + MD5Checksum: "md5-1", + Owners: []struct { + EmailAddress string `json:"emailAddress"` + }{{EmailAddress: "zoe@example.com"}, {EmailAddress: "alice@example.com"}}, + } + fp1 := file.fingerprint() + fp2 := file.fingerprint() + if fp1 == "" || fp1 != fp2 { + t.Fatalf("fingerprint unstable: %q %q", fp1, fp2) + } + reordered := file + reordered.Owners = []struct { + EmailAddress string `json:"emailAddress"` + }{{EmailAddress: "alice@example.com"}, {EmailAddress: "zoe@example.com"}} + if got := reordered.fingerprint(); got != fp1 { + t.Fatalf("fingerprint changed after owner order-only change: %q != %q", got, fp1) + } + changed := file + changed.MD5Checksum = "md5-2" + if got := changed.fingerprint(); got == fp1 { + t.Fatalf("fingerprint did not change after checksum update") + } +} + +// TestGoogleDriveFileQueryUsesIncrementalWindow verifies Python-compatible Drive time filters. +func TestGoogleDriveFileQueryUsesIncrementalWindow(t *testing.T) { + start := time.Date(2026, 1, 2, 0, 0, 0, 0, time.UTC) + end := time.Date(2026, 1, 4, 0, 0, 0, 0, time.UTC) + query := googleDriveFileQuery(googleDriveScope{corpora: "user", includeSharedWithMe: false}, &start, end) + if !strings.Contains(query, "modifiedTime > '2026-01-02T00:00:00Z'") || + !strings.Contains(query, "createdTime >= '2026-01-02T00:00:00Z'") || + !strings.Contains(query, "modifiedTime <= '2026-01-04T00:00:00Z'") || + !strings.Contains(query, "'me' in owners") { + t.Fatalf("query = %q", query) + } +} diff --git a/internal/syncer/connector/rss_test.go b/internal/syncer/connector/rss_test.go index bfae3c2adc..ad644a3f30 100644 --- a/internal/syncer/connector/rss_test.go +++ b/internal/syncer/connector/rss_test.go @@ -65,6 +65,9 @@ func TestRSSConnectorOpenSyncFullAndIncremental(t *testing.T) { if doc.Metadata["link"] != "https://example.com/new" { t.Fatalf("link metadata = %v", doc.Metadata["link"]) } + if doc.Fingerprint == "" { + t.Fatalf("fingerprint is empty") + } } // TestRSSConnectorOpenPrune verifies complete slim snapshot generation. diff --git a/internal/syncer/connector_lock.go b/internal/syncer/connector_lock.go index 4bacaffa78..b89a3fad43 100644 --- a/internal/syncer/connector_lock.go +++ b/internal/syncer/connector_lock.go @@ -16,39 +16,91 @@ package syncer -import "sync" +import ( + "context" + "fmt" + "ragflow/internal/engine/redis" + "ragflow/internal/utility" + "sync" + "time" +) -// ConnectorLocker serializes work for a connector. +// ConnectorLocker serializes work for one connector and knowledge base. type ConnectorLocker interface { - TryLock(connectorID string) bool - Unlock(connectorID string) + TryLock(connectorID, kbID string) (ConnectorLockLease, bool) + Unlock(connectorID, kbID string) } -// ConnectorLock is a process-local connector mutex registry. +const connectorLockTTL = 24 * time.Hour + +// ConnectorLockLease describes the bounded lifetime of a connector/KB lock. +type ConnectorLockLease struct { + ExpiresAt time.Time +} + +// ConnectorLock serializes connector/KB work through Redis when available. type ConnectorLock struct { + holder string mu sync.Mutex - locked map[string]struct{} + local map[string]struct{} + redis map[string]*redis.DistributedLock } -// NewConnectorLock creates an empty process-local connector lock. +// NewConnectorLock creates an empty connector/KB lock. func NewConnectorLock() *ConnectorLock { - return &ConnectorLock{locked: map[string]struct{}{}} + return &ConnectorLock{holder: utility.GenerateUUID(), local: map[string]struct{}{}, redis: map[string]*redis.DistributedLock{}} } -// TryLock attempts to acquire the connector lock without blocking. -func (l *ConnectorLock) TryLock(connectorID string) bool { - l.mu.Lock() - defer l.mu.Unlock() - if _, ok := l.locked[connectorID]; ok { - return false +// TryLock attempts to acquire the connector/KB lock without blocking. +func (l *ConnectorLock) TryLock(connectorID, kbID string) (ConnectorLockLease, bool) { + if l == nil { + return ConnectorLockLease{}, false } - l.locked[connectorID] = struct{}{} - return true + key := connectorLockKey(connectorID, kbID) + l.mu.Lock() + if _, ok := l.local[key]; ok { + l.mu.Unlock() + return ConnectorLockLease{}, false + } + l.local[key] = struct{}{} + l.mu.Unlock() + + if client := redis.Get(); client != nil { + lock := redis.NewDistributedLock(key, l.holder, connectorLockTTL, 0) + if lock == nil || !lock.Acquire(context.Background()) { + l.mu.Lock() + delete(l.local, key) + l.mu.Unlock() + return ConnectorLockLease{}, false + } + l.mu.Lock() + l.redis[key] = lock + l.mu.Unlock() + return newConnectorLockLease(), true + } + return newConnectorLockLease(), true } -// Unlock releases the connector lock. -func (l *ConnectorLock) Unlock(connectorID string) { +// Unlock releases the connector/KB lock. +func (l *ConnectorLock) Unlock(connectorID, kbID string) { + if l == nil { + return + } + key := connectorLockKey(connectorID, kbID) l.mu.Lock() - defer l.mu.Unlock() - delete(l.locked, connectorID) + lock := l.redis[key] + delete(l.redis, key) + delete(l.local, key) + l.mu.Unlock() + if lock != nil { + lock.Release(context.Background()) + } +} + +func connectorLockKey(connectorID, kbID string) string { + return fmt.Sprintf("syncer:connector-lock:%s:%s", connectorID, kbID) +} + +func newConnectorLockLease() ConnectorLockLease { + return ConnectorLockLease{ExpiresAt: time.Now().Add(connectorLockTTL)} } diff --git a/internal/syncer/job_executor.go b/internal/syncer/job_executor.go new file mode 100644 index 0000000000..e6f7c8f6fb --- /dev/null +++ b/internal/syncer/job_executor.go @@ -0,0 +1,340 @@ +// +// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// + +package syncer + +import ( + "context" + "errors" + "fmt" + "ragflow/internal/service" + "sync" + "time" +) + +var errSyncJobExecutorClosed = errors.New("sync job executor is closed") + +// SyncJobExecutorConfig controls the shared batch job executor. +type SyncJobExecutorConfig struct { + WorkerCount int // num of the workers + JobQueueSize int // jobs channel size + PerTaskQueueSize int // sync_task's queue size +} + +// normalize prevent the channel from malfunctioning when set to 0 +func (c SyncJobExecutorConfig) normalize() SyncJobExecutorConfig { + if c.WorkerCount <= 0 { + c.WorkerCount = 1 + } + if c.JobQueueSize <= 0 { + c.JobQueueSize = c.WorkerCount + } + if c.PerTaskQueueSize <= 0 { + c.PerTaskQueueSize = c.JobQueueSize + } + return c +} + +// syncJobFunc the func that the batch job to execute +type syncJobFunc func(context.Context) (service.SyncStats, error) + +// syncJobResult the stats +type syncJobResult struct { + stats service.SyncStats + err error +} + +type syncJob struct { + ctx context.Context + fn syncJobFunc + done chan syncJobResult +} + +// SyncJobQueue is one Coordinator-owned queue feeding the fair dispatcher. +type SyncJobQueue struct { + taskID string + jobs chan *syncJob + close func(string) + once sync.Once +} + +// Submit adds one BatchJob to this task's dispatcher queue. +func (q *SyncJobQueue) Submit(ctx context.Context, fn syncJobFunc) (<-chan syncJobResult, error) { + if q == nil { + return nil, fmt.Errorf("sync job queue is nil") + } + + done := make(chan syncJobResult, 1) // done channel + job := &syncJob{ctx: ctx, fn: fn, done: done} + select { + case <-ctx.Done(): + return nil, ctx.Err() + case q.jobs <- job: + return done, nil + } +} + +// Close unregisters this task from the fair dispatcher. +func (q *SyncJobQueue) Close() { + if q == nil { + return + } + q.once.Do(func() { + if q.close != nil { + q.close(q.taskID) + } + }) +} + +type executorCommandKind int + +const ( + executorRegister executorCommandKind = iota + executorUnregister +) + +type executorCommand struct { + kind executorCommandKind + queue *SyncJobQueue + task string + err chan error + done chan struct{} +} + +type syncTaskState struct { + jobs <-chan *syncJob +} + +// SyncJobExecutor fairly dispatches per-task BatchJobs into one shared worker channel. +type SyncJobExecutor struct { + perTaskQueueSize int + jobs chan *syncJob + commands chan executorCommand + stop chan struct{} + done chan struct{} + stopOnce sync.Once + workerGroup sync.WaitGroup +} + +// NewSyncJobExecutor creates a global BatchJob executor. +func NewSyncJobExecutor(config SyncJobExecutorConfig) *SyncJobExecutor { + config = config.normalize() + executor := &SyncJobExecutor{ + perTaskQueueSize: config.PerTaskQueueSize, + jobs: make(chan *syncJob, config.JobQueueSize), + commands: make(chan executorCommand), + stop: make(chan struct{}), + done: make(chan struct{}), + } + + go executor.dispatch() + for i := 0; i < config.WorkerCount; i++ { + executor.workerGroup.Add(1) + go executor.work() + } + return executor +} + +// RegisterTask creates the bounded Coordinator queue for one running task. +func (e *SyncJobExecutor) RegisterTask(ctx context.Context, taskID string) (*SyncJobQueue, error) { + if e == nil { + return nil, fmt.Errorf("sync job executor is nil") + } + queue := &SyncJobQueue{taskID: taskID, jobs: make(chan *syncJob, e.perTaskQueueSize), close: e.unregisterTask} + reply := make(chan error, 1) + command := executorCommand{kind: executorRegister, queue: queue, err: reply} + + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-e.stop: + return nil, errSyncJobExecutorClosed + case e.commands <- command: + } + + select { + case <-ctx.Done(): + return nil, ctx.Err() + case err := <-reply: + if err != nil { + return nil, err + } + return queue, nil + } +} + +// Close stops the dispatcher and waits for fixed workers to exit. +func (e *SyncJobExecutor) Close() { + if e == nil { + return + } + e.stopOnce.Do(func() { + close(e.stop) + <-e.done + e.workerGroup.Wait() + }) +} + +func (e *SyncJobExecutor) unregisterTask(taskID string) { + done := make(chan struct{}) + command := executorCommand{kind: executorUnregister, task: taskID, done: done} + select { + case <-e.stop: + return + case e.commands <- command: + <-done + } +} + +func (e *SyncJobExecutor) dispatch() { + defer close(e.done) + tasks := map[string]*syncTaskState{} + order := []string{} + cursor := 0 + var pending *syncJob + + for { + if pending == nil { + pending = popReadyJob(tasks, order, &cursor) + } + if pending == nil { + select { + case <-e.stop: + settleQueuedJobs(nil, tasks) + close(e.jobs) + return + case command := <-e.commands: + order = applyExecutorCommand(command, tasks, order, &cursor) + case <-time.After(time.Millisecond): + } + continue + } + + select { + case <-e.stop: + settleQueuedJobs(pending, tasks) + close(e.jobs) + return + case command := <-e.commands: + order = applyExecutorCommand(command, tasks, order, &cursor) + case e.jobs <- pending: + pending = nil + } + } +} + +func settleQueuedJobs(pending *syncJob, tasks map[string]*syncTaskState) { + if pending != nil { + pending.done <- syncJobResult{err: errSyncJobExecutorClosed} + } + for _, task := range tasks { + if task == nil { + continue + } + settleTaskJobs(task) + } +} + +func settleTaskJobs(task *syncTaskState) { + for { + select { + case job := <-task.jobs: + job.done <- syncJobResult{err: errSyncJobExecutorClosed} + default: + return + } + } +} + +func applyExecutorCommand(command executorCommand, tasks map[string]*syncTaskState, order []string, cursor *int) []string { + switch command.kind { + case executorRegister: + err := registerExecutorTask(command.queue, tasks, &order) + command.err <- err + case executorUnregister: + order = unregisterExecutorTask(command.task, tasks, order, cursor) + close(command.done) + } + return order +} + +func registerExecutorTask(queue *SyncJobQueue, tasks map[string]*syncTaskState, order *[]string) error { + if queue == nil || queue.taskID == "" { + return fmt.Errorf("sync job task id is required") + } + if tasks[queue.taskID] != nil { + return fmt.Errorf("sync job task %q is already registered", queue.taskID) + } + tasks[queue.taskID] = &syncTaskState{jobs: queue.jobs} + *order = append(*order, queue.taskID) + return nil +} + +func unregisterExecutorTask(taskID string, tasks map[string]*syncTaskState, order []string, cursor *int) []string { + if tasks[taskID] == nil { + return order + } + delete(tasks, taskID) + for i, existing := range order { + if existing != taskID { + continue + } + order = append(order[:i], order[i+1:]...) + if len(order) == 0 { + *cursor = 0 + } else if *cursor >= len(order) { + *cursor = 0 + } + return order + } + return order +} + +// popReadyJob pop the job fairly +func popReadyJob(tasks map[string]*syncTaskState, order []string, cursor *int) *syncJob { + if len(order) == 0 { + return nil + } + + for i := 0; i < len(order); i++ { + index := (*cursor + i) % len(order) + task := tasks[order[index]] + if task == nil { + continue + } + select { + case job := <-task.jobs: + *cursor = (index + 1) % len(order) + return job + default: + } + } + return nil +} + +// work execute the job +func (e *SyncJobExecutor) work() { + defer e.workerGroup.Done() + + for job := range e.jobs { // + if err := job.ctx.Err(); err != nil { + job.done <- syncJobResult{err: err} + continue + } + stats, err := job.fn(job.ctx) // run the job + job.done <- syncJobResult{stats: stats, err: err} + } +} diff --git a/internal/syncer/job_executor_test.go b/internal/syncer/job_executor_test.go new file mode 100644 index 0000000000..8f8c9a4146 --- /dev/null +++ b/internal/syncer/job_executor_test.go @@ -0,0 +1,184 @@ +package syncer + +import ( + "context" + "errors" + "ragflow/internal/service" + "testing" + "time" +) + +// TestSyncJobExecutorUsesIdleWorkers verifies one task can use the full shared pool. +func TestSyncJobExecutorUsesIdleWorkers(t *testing.T) { + executor := NewSyncJobExecutor(SyncJobExecutorConfig{WorkerCount: 3}) + defer executor.Close() + + queue, err := executor.RegisterTask(t.Context(), "task-1") + if err != nil { + t.Fatalf("register task: %v", err) + } + defer queue.Close() + + started := make(chan struct{}, 3) + release := make(chan struct{}) + results := make([]<-chan syncJobResult, 0, 3) + for i := 0; i < 3; i++ { + result, err := queue.Submit(t.Context(), func(ctx context.Context) (service.SyncStats, error) { + started <- struct{}{} + <-release + return service.SyncStats{Added: 1}, nil + }) + if err != nil { + t.Fatalf("submit: %v", err) + } + results = append(results, result) + } + + for i := 0; i < 3; i++ { + select { + case <-started: + case <-time.After(time.Second): + t.Fatalf("started workers = %d, want 3", i) + } + } + close(release) + for _, result := range results { + if item := <-result; item.err != nil || item.stats.Added != 1 { + t.Fatalf("job result = %+v", item) + } + } +} + +// TestSyncJobExecutorDispatchesRoundRobin verifies a waiting task is not hidden behind one large task. +func TestSyncJobExecutorDispatchesRoundRobin(t *testing.T) { + executor := NewSyncJobExecutor(SyncJobExecutorConfig{WorkerCount: 1}) + defer executor.Close() + + task1, err := executor.RegisterTask(t.Context(), "task-1") + if err != nil { + t.Fatalf("register task-1: %v", err) + } + defer task1.Close() + task2, err := executor.RegisterTask(t.Context(), "task-2") + if err != nil { + t.Fatalf("register task-2: %v", err) + } + defer task2.Close() + + started := make(chan string, 3) + releaseFirst := make(chan struct{}) + firstResult, err := task1.Submit(t.Context(), func(ctx context.Context) (service.SyncStats, error) { + started <- "task-1-a" + <-releaseFirst + return service.SyncStats{Added: 1}, nil + }) + if err != nil { + t.Fatalf("submit first task-1: %v", err) + } + if got := waitStarted(t, started); got != "task-1-a" { + t.Fatalf("first started = %s", got) + } + + task2Result, err := task2.Submit(t.Context(), func(ctx context.Context) (service.SyncStats, error) { + started <- "task-2" + return service.SyncStats{Updated: 1}, nil + }) + if err != nil { + t.Fatalf("submit task-2: %v", err) + } + task1Result, err := task1.Submit(t.Context(), func(ctx context.Context) (service.SyncStats, error) { + started <- "task-1-b" + return service.SyncStats{Skipped: 1}, nil + }) + if err != nil { + t.Fatalf("submit second task-1: %v", err) + } + + close(releaseFirst) + if item := <-firstResult; item.err != nil { + t.Fatalf("first task-1 result: %v", item.err) + } + if got := waitStarted(t, started); got != "task-2" { + t.Fatalf("second started = %s, want task-2", got) + } + if item := <-task2Result; item.err != nil || item.stats.Updated != 1 { + t.Fatalf("task-2 result = %+v", item) + } + if got := waitStarted(t, started); got != "task-1-b" { + t.Fatalf("third started = %s, want task-1-b", got) + } + if item := <-task1Result; item.err != nil || item.stats.Skipped != 1 { + t.Fatalf("second task-1 result = %+v", item) + } +} + +// TestSyncJobExecutorCloseSettlesQueuedJobs verifies shutdown replies to jobs still in task queues. +func TestSyncJobExecutorCloseSettlesQueuedJobs(t *testing.T) { + executor := NewSyncJobExecutor(SyncJobExecutorConfig{WorkerCount: 1, JobQueueSize: 1, PerTaskQueueSize: 4}) + queue, err := executor.RegisterTask(t.Context(), "task-1") + if err != nil { + t.Fatalf("register task: %v", err) + } + + started := make(chan struct{}) + release := make(chan struct{}) + results := make([]<-chan syncJobResult, 0, 4) + first, err := queue.Submit(t.Context(), func(ctx context.Context) (service.SyncStats, error) { + close(started) + <-release + return service.SyncStats{Added: 1}, nil + }) + if err != nil { + t.Fatalf("submit first: %v", err) + } + results = append(results, first) + <-started + + for i := 0; i < 3; i++ { + result, err := queue.Submit(t.Context(), func(ctx context.Context) (service.SyncStats, error) { + return service.SyncStats{Updated: 1}, nil + }) + if err != nil { + t.Fatalf("submit queued %d: %v", i, err) + } + results = append(results, result) + } + + closed := make(chan struct{}) + go func() { + executor.Close() + close(closed) + }() + close(release) + + closedErrs := 0 + for _, result := range results { + select { + case item := <-result: + if errors.Is(item.err, errSyncJobExecutorClosed) { + closedErrs++ + } + case <-time.After(time.Second): + t.Fatalf("timed out waiting for shutdown result") + } + } + if closedErrs == 0 { + t.Fatalf("queued jobs were not settled with closed error") + } + select { + case <-closed: + case <-time.After(time.Second): + t.Fatalf("executor did not close") + } +} + +func waitStarted(t *testing.T, started <-chan string) string { + t.Helper() + select { + case task := <-started: + return task + case <-time.After(time.Second): + t.Fatalf("timed out waiting for job start") + return "" + } +} diff --git a/internal/syncer/prune_runner.go b/internal/syncer/prune_runner.go index e99daa19fc..219d42e057 100644 --- a/internal/syncer/prune_runner.go +++ b/internal/syncer/prune_runner.go @@ -53,6 +53,9 @@ func (r *PruneRunner) Run(ctx context.Context, taskContext service.SyncTaskConte retain := map[string]struct{}{} for { + if err := r.checkCanceled(ctx, taskContext.Task.ID); err != nil { + return err + } batch, nextErr := session.NextBatch(ctx) if errors.Is(nextErr, io.EOF) { break @@ -65,9 +68,26 @@ func (r *PruneRunner) Run(ctx context.Context, taskContext service.SyncTaskConte } } + if err := r.checkCanceled(ctx, taskContext.Task.ID); err != nil { + return err + } removed, err := r.pruneService.DeleteStale(ctx, taskContext, retain) if err != nil { return err } return r.taskService.CompletePrune(ctx, taskContext, removed) } + +func (r *PruneRunner) checkCanceled(ctx context.Context, taskID string) error { + if err := ctx.Err(); err != nil { + return err + } + canceled, err := r.taskService.IsCanceled(ctx, taskID) + if err != nil { + return err + } + if canceled { + return errSyncTaskCanceled + } + return nil +} diff --git a/internal/syncer/scheduler.go b/internal/syncer/scheduler.go index 0cc775b646..f51da67ec5 100644 --- a/internal/syncer/scheduler.go +++ b/internal/syncer/scheduler.go @@ -26,14 +26,25 @@ import ( // TaskEnvelope is the only payload sent through the task queue. type TaskEnvelope struct { - TaskID string + TaskID string + Handle common.TaskHandle + stopHeartbeat func() } -// Scheduler scans due work and enqueues claimed task IDs. +// SyncTaskBroker publishes and pulls syncer task wake-up messages. +type SyncTaskBroker interface { + InitSyncerStream() error + InitSyncerConsumer() error + PublishSyncerTask(taskID string) error + FetchSyncerTasks(batchSize int) ([]common.TaskHandle, error) +} + +// Scheduler discovers due work and enqueues task IDs for workers. type Scheduler struct { pollInterval time.Duration queue chan<- TaskEnvelope taskService *service.SyncTaskService + broker SyncTaskBroker } // NewScheduler creates a global scheduler for datasource sync tasks. @@ -41,50 +52,152 @@ func NewScheduler(pollInterval time.Duration, queue chan<- TaskEnvelope, taskSer return &Scheduler{pollInterval: pollInterval, queue: queue, taskService: taskService} } -// Run starts recovery and periodic task discovery. -func (s *Scheduler) Run(ctx context.Context) error { - if err := s.scan(ctx); err != nil && ctx.Err() == nil { - common.Error("syncer scheduler scan failed", err) +// NewNATSScheduler creates a JetStream-driven scheduler with DB reconciliation. +func NewNATSScheduler(pollInterval time.Duration, queue chan<- TaskEnvelope, taskService *service.SyncTaskService, broker SyncTaskBroker) *Scheduler { + return &Scheduler{ + pollInterval: pollInterval, + queue: queue, + taskService: taskService, + broker: broker, } +} + +// Run starts the NATS listener. +func (s *Scheduler) Run(ctx context.Context) error { + if s.broker != nil { + return s.runNATS(ctx) + } + return errors.New("syncer scheduler requires a NATS broker") +} + +func (s *Scheduler) runNATS(ctx context.Context) error { + if err := s.broker.InitSyncerStream(); err != nil { + return err + } + if err := s.broker.InitSyncerConsumer(); err != nil { + return err + } + + // scan DB for first time run + if err := s.publishStartupTasks(ctx); err != nil && ctx.Err() == nil { + common.Error("syncer scheduler startup publish failed", err) + } + ticker := time.NewTicker(s.pollInterval) defer ticker.Stop() for { + if err := ctx.Err(); err != nil { + return err + } + select { case <-ctx.Done(): return ctx.Err() - case <-ticker.C: - if err := s.scan(ctx); err != nil && ctx.Err() == nil { - common.Error("syncer scheduler scan failed", err) + case <-ticker.C: // Run `publishDueTasks` periodically + if err := s.publishDueTasks(ctx); err != nil && ctx.Err() == nil { + common.Error("syncer scheduler due publish failed", err) } + continue + default: + } + + // check `scheduler`'s task queue's slot + available := s.queueAvailable() + if available <= 0 { + if err := waitNATSFetchCapacity(ctx); err != nil { + return err + } + continue + } + + // pull `available` tasks from nats + handles, err := s.broker.FetchSyncerTasks(available) + if err != nil { + common.Error("syncer scheduler fetch failed", err) + if waitErr := waitNATSFetchCapacity(ctx); waitErr != nil { + return waitErr + } + } else if len(handles) == 0 { + if waitErr := waitNATSFetchCapacity(ctx); waitErr != nil { + return waitErr + } + } else if err = s.enqueueHandles(ctx, handles); err != nil { // put tasks to task queue + return err } } } -// scan claims due tasks and places their IDs on the bounded queue. -func (s *Scheduler) scan(ctx context.Context) error { - now := time.Now() - if err := s.taskService.RecoverStaleRunning(ctx, now); err != nil { +func waitNATSFetchCapacity(ctx context.Context) error { + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(100 * time.Millisecond): + return nil + } +} + +// enqueueHandles put task to `scheduler`'s task queue +func (s *Scheduler) enqueueHandles(ctx context.Context, handles []common.TaskHandle) error { + for index, handle := range handles { + message := handle.GetMessage() + if message.TaskID == "" { + _ = handle.Ack() + continue + } + + stopHeartbeat := startHandleHeartbeat(ctx, handle) + select { + case <-ctx.Done(): + stopHeartbeat() + for _, pending := range handles[index:] { + _ = pending.Nack() + } + return ctx.Err() + case s.queue <- TaskEnvelope{TaskID: message.TaskID, Handle: handle, stopHeartbeat: stopHeartbeat}: + } + } + return nil +} + +func (s *Scheduler) queueAvailable() int { + return cap(s.queue) - len(s.queue) +} + +// publishStartupTasks scan DB for first time run +func (s *Scheduler) publishStartupTasks(ctx context.Context) error { + // TODO restores running tasks during syncer startup. + if err := s.taskService.RecoverRunning(ctx); err != nil { return err } + + tasks, err := s.taskService.ListStartupTasks(ctx) + if err != nil { + return err + } + + // publish task to nats + var publishErr error + for _, task := range tasks { + if err = s.broker.PublishSyncerTask(task.ID); err != nil { + publishErr = errors.Join(publishErr, err) + } + } + return publishErr +} + +// publishDueTasks publishes due scheduled DB tasks to nats +func (s *Scheduler) publishDueTasks(ctx context.Context) error { + now := time.Now() tasks, err := s.taskService.ListDueTasks(ctx, now) if err != nil { return err } - var claimErr error + + var publishErr error for _, task := range tasks { - claimed, err := s.taskService.Claim(ctx, task.ID) - if err != nil { - claimErr = errors.Join(claimErr, err) - continue - } - if !claimed { - continue - } - select { - case <-ctx.Done(): - return ctx.Err() - case s.queue <- TaskEnvelope{TaskID: task.ID}: + if err = s.broker.PublishSyncerTask(task.ID); err != nil { + publishErr = errors.Join(publishErr, err) } } - return claimErr + return publishErr } diff --git a/internal/syncer/sync_runner.go b/internal/syncer/sync_runner.go index 41edd8a0a3..e2b7696765 100644 --- a/internal/syncer/sync_runner.go +++ b/internal/syncer/sync_runner.go @@ -24,32 +24,36 @@ import ( "math/rand" "ragflow/internal/service" syncerconnector "ragflow/internal/syncer/connector" - "sync" "time" ) -// SyncRunner executes one SYNC task with serial batches and parallel items. +var errSyncTaskCanceled = errors.New("sync task canceled") + +const syncCancelCheckInterval = time.Second + +// SyncRunner executes one SYNC task by submitting source batches as BatchJobs. type SyncRunner struct { config TaskCoordinatorConfig taskService *service.SyncTaskService sink service.DocumentSink idResolver *service.DocumentIDResolver - globalItems chan struct{} + queue *SyncJobQueue } // NewSyncRunner creates a SYNC runner. -func NewSyncRunner(config TaskCoordinatorConfig, taskService *service.SyncTaskService, sink service.DocumentSink, idResolver *service.DocumentIDResolver, globalItems chan struct{}) *SyncRunner { - return &SyncRunner{config: config, taskService: taskService, sink: sink, idResolver: idResolver, globalItems: globalItems} +func NewSyncRunner(config TaskCoordinatorConfig, taskService *service.SyncTaskService, sink service.DocumentSink, idResolver *service.DocumentIDResolver, queue *SyncJobQueue) *SyncRunner { + return &SyncRunner{config: config, taskService: taskService, sink: sink, idResolver: idResolver, queue: queue} } // Run executes all sync batches and commits the final waterline. func (r *SyncRunner) Run(ctx context.Context, taskContext service.SyncTaskContext, connector syncerconnector.Connector) error { + // sink is nil means this syncer task cannot write it to document, it will fail anyway if r.sink == nil { return errors.New("document sink is not configured") } windowEnd := time.Now().UTC() - var windowStart *time.Time + var windowStart *time.Time // = nil if it is `Full synchronisation` if !service.IsFromBeginning(taskContext.Task.FromBeginning) { windowStart = taskContext.Task.PollRangeStart } @@ -67,79 +71,114 @@ func (r *SyncRunner) Run(ctx context.Context, taskContext service.SyncTaskContex } defer session.Close() + // prepare sourceType, waterline, stats, resultChan sourceType := service.SourceType(taskContext.Connector.Source, taskContext.Connector.ID) - candidateEnd := windowStart - stats := service.SyncStats{} + candidateEnd := windowStart // the waterLine that will write to DB + stats := service.SyncStats{} // count `add`, `updated`, `skipped` + resultChans := make([]<-chan syncJobResult, 0) + for { + // check if task has been canceled + if err := r.checkCanceled(ctx, taskContext.Task.ID); err != nil { + return err + } + + // get a batch of files batch, nextErr := session.NextBatch(ctx) - if errors.Is(nextErr, io.EOF) { + if errors.Is(nextErr, io.EOF) { // end of file break } if nextErr != nil { return nextErr } + for _, doc := range batch.Documents { if candidateEnd == nil || doc.UpdatedAt.After(*candidateEnd) { updatedAt := doc.UpdatedAt - candidateEnd = &updatedAt + candidateEnd = &updatedAt // use `max_update` to push `waterLine` } } - batchStats, err := r.processBatch(ctx, taskContext, sourceType, session, batch) + // a batch, a syncJob + resultChan, err := r.submitBatch(ctx, taskContext, sourceType, session, batch) if err != nil { return err } - stats.Add(batchStats) + resultChans = append(resultChans, resultChan) } + // run sync Job + var firstErr error + for _, resultChan := range resultChans { + var jobResult syncJobResult + select { + case <-ctx.Done(): + return ctx.Err() + case jobResult = <-resultChan: + } + stats.Add(jobResult.stats) + if jobResult.err != nil && firstErr == nil { + firstErr = jobResult.err + } + } + if firstErr != nil { + return firstErr + } + + if err := r.checkCanceled(ctx, taskContext.Task.ID); err != nil { + return err + } if candidateEnd == nil { candidateEnd = &windowEnd } return r.taskService.CompleteSync(ctx, taskContext, *candidateEnd, stats) } -// processBatch processes one batch with bounded document concurrency. -func (r *SyncRunner) processBatch(ctx context.Context, taskContext service.SyncTaskContext, sourceType string, session syncerconnector.SyncSession, batch syncerconnector.SyncBatch) (service.SyncStats, error) { - sem := make(chan struct{}, r.config.PerTaskItemConcurrency) - results := make(chan service.DocumentUpsertResult, len(batch.Documents)) - errs := make(chan error, len(batch.Documents)) - - var wg sync.WaitGroup - for _, sourceDocument := range batch.Documents { - sourceDocument := sourceDocument - - sem <- struct{}{} - wg.Add(1) - go func() { - defer wg.Done() - defer func() { <-sem }() - result, err := r.processDocumentWithRetry(ctx, taskContext, sourceType, session, sourceDocument) - if err != nil { - errs <- err - return - } - results <- result - }() +// submitBatch submits one source batch as one BatchJob. +func (r *SyncRunner) submitBatch(ctx context.Context, taskContext service.SyncTaskContext, sourceType string, session syncerconnector.SyncSession, batch syncerconnector.SyncBatch) (<-chan syncJobResult, error) { + resultChan, err := r.queue.Submit(ctx, func(jobCtx context.Context) (service.SyncStats, error) { + return r.processDocuments(jobCtx, taskContext, sourceType, session, batch.Documents) + }) + if err != nil { + return nil, err } + return resultChan, nil +} - wg.Wait() - close(results) - close(errs) - - if len(errs) > 0 { - return service.SyncStats{}, <-errs - } +// processDocuments +func (r *SyncRunner) processDocuments(ctx context.Context, taskContext service.SyncTaskContext, sourceType string, session syncerconnector.SyncSession, documents []syncerconnector.SourceDocument) (service.SyncStats, error) { stats := service.SyncStats{} - for result := range results { + var firstErr error + lastCancelCheck := time.Time{} + for _, sourceDocument := range documents { + if err := ctx.Err(); err != nil { + return stats, err + } + if lastCancelCheck.IsZero() || time.Since(lastCancelCheck) >= syncCancelCheckInterval { + if err := r.checkCanceled(ctx, taskContext.Task.ID); err != nil { + return stats, err + } + lastCancelCheck = time.Now() + } + result, err := r.processDocumentWithRetry(ctx, taskContext, sourceType, session, sourceDocument) + if err != nil { + if firstErr == nil { + firstErr = err + } + continue + } stats.AddResult(result) } - return stats, nil + return stats, firstErr } // processDocumentWithRetry retries transient item failures. func (r *SyncRunner) processDocumentWithRetry(ctx context.Context, taskContext service.SyncTaskContext, sourceType string, session syncerconnector.SyncSession, sourceDocument syncerconnector.SourceDocument) (service.DocumentUpsertResult, error) { var lastErr error for attempt := 1; attempt <= r.config.ItemRetryCount; attempt++ { + if err := ctx.Err(); err != nil { + return service.DocumentUpsertResult{}, err + } result, err := r.processDocument(ctx, taskContext, sourceType, session, sourceDocument) if err == nil { return result, nil @@ -171,17 +210,25 @@ func (r *SyncRunner) processDocumentWithRetry(ctx context.Context, taskContext s return service.DocumentUpsertResult{}, lastErr } -// processDocument resolves IDs, skips unchanged fingerprints, fetches blobs, and upserts. -func (r *SyncRunner) processDocument(ctx context.Context, taskContext service.SyncTaskContext, sourceType string, session syncerconnector.SyncSession, sourceDocument syncerconnector.SourceDocument) (service.DocumentUpsertResult, error) { - if r.globalItems != nil { - select { - case <-ctx.Done(): - return service.DocumentUpsertResult{}, ctx.Err() - case r.globalItems <- struct{}{}: - defer func() { <-r.globalItems }() - } +// checkCanceled check if the task has been canceled +func (r *SyncRunner) checkCanceled(ctx context.Context, taskID string) error { + if err := ctx.Err(); err != nil { + return err } + canceled, err := r.taskService.IsCanceled(ctx, taskID) + if err != nil { + return err + } + + if canceled { + return errSyncTaskCanceled + } + return nil +} + +// processDocument resolves IDs, skips unchanged fingerprints, fetches blobs, and upserts. +func (r *SyncRunner) processDocument(ctx context.Context, taskContext service.SyncTaskContext, sourceType string, session syncerconnector.SyncSession, sourceDocument syncerconnector.SourceDocument) (service.DocumentUpsertResult, error) { resolved, err := r.idResolver.Resolve(ctx, taskContext.Knowledgebase.ID, taskContext.Connector.ID, sourceType, sourceDocument.SourceID) if err != nil { return service.DocumentUpsertResult{}, err diff --git a/internal/syncer/syncer.go b/internal/syncer/syncer.go index 53afa04a20..d41c7735dd 100644 --- a/internal/syncer/syncer.go +++ b/internal/syncer/syncer.go @@ -19,11 +19,13 @@ package syncer import ( "context" "errors" + "fmt" "sync" "time" "ragflow/internal/common" "ragflow/internal/dao" + "ragflow/internal/engine" "ragflow/internal/service" documentservice "ragflow/internal/service/document" syncerconnector "ragflow/internal/syncer/connector" @@ -32,13 +34,14 @@ import ( "go.uber.org/zap" ) -// Syncer owns the scheduler, task queue, and bounded worker pool. +// Syncer owns NATS/DB scheduling, task workers, and the shared batch job executor. type Syncer struct { id string config Config queue chan TaskEnvelope scheduler *Scheduler worker *TaskWorker + executor *SyncJobExecutor cancel context.CancelFunc workerGroup sync.WaitGroup stopOnce sync.Once @@ -46,10 +49,10 @@ type Syncer struct { } // NewSyncer creates a server-compatible syncer with default dependencies. -func NewSyncer(maxConcurrency int, pollInterval time.Duration) *Syncer { +func NewSyncer(taskWorkerCount int, pollInterval time.Duration) *Syncer { // init the config config := DefaultConfig() - config.TaskConcurrency = maxConcurrency + config.TaskWorkerCount = taskWorkerCount config.PollInterval = pollInterval taskDAO := dao.NewSyncTaskDAO(nil) @@ -69,22 +72,30 @@ func New(config Config, taskDAO *dao.SyncTaskDAO, registry ConnectorRegistry, si queue := make(chan TaskEnvelope, config.TaskQueueSize) locker := NewConnectorLock() - globalItems := make(chan struct{}, config.GlobalItemConcurrency) + executor := NewSyncJobExecutor(SyncJobExecutorConfig{ + WorkerCount: config.JobWorkerCount, + JobQueueSize: config.JobQueueSize, + }) taskService := service.NewSyncTaskService(taskDAO) idResolver := service.NewDocumentIDResolver(service.NewGormDocumentStore()) coordinator := NewTaskCoordinator(TaskCoordinatorConfig{ - PerTaskItemConcurrency: config.PerTaskItemConcurrency, - ItemRetryCount: config.ItemRetryCount, - ItemRetryBaseDelay: config.ItemRetryBaseDelay, - }, taskService, registry, sink, pruneService, idResolver, globalItems) + ItemRetryCount: config.ItemRetryCount, + ItemRetryBaseDelay: config.ItemRetryBaseDelay, + }, taskService, registry, sink, pruneService, idResolver, executor) + + scheduler := NewScheduler(config.PollInterval, queue, taskService) + if broker, ok := engine.GetMessageQueueEngine().(SyncTaskBroker); ok { + scheduler = NewNATSScheduler(config.PollInterval, queue, taskService, broker) + } return &Syncer{ id: utility.GenerateUUID(), config: config, queue: queue, - scheduler: NewScheduler(config.PollInterval, queue, taskService), + scheduler: scheduler, worker: NewTaskWorker(queue, taskService, coordinator, locker), + executor: executor, ShutdownCh: make(chan struct{}), } } @@ -112,13 +123,18 @@ func (s *Syncer) StartContext(ctx context.Context) error { s.cancel = cancel s.workerGroup.Add(2) + // run scheduler go func() { defer s.workerGroup.Done() - _ = s.scheduler.Run(runCtx) + if err := s.scheduler.Run(runCtx); err != nil && !errors.Is(err, context.Canceled) { + common.Error("syncer scheduler stopped", err) + } }() + + // run worker poll go func() { defer s.workerGroup.Done() - s.worker.Run(runCtx, s.config.TaskConcurrency) + s.worker.Run(runCtx, s.config.TaskWorkerCount) }() return nil } @@ -133,18 +149,18 @@ func (s *Syncer) Stop() { s.cancel() } s.workerGroup.Wait() + s.executor.Close() close(s.ShutdownCh) }) } -// logTemporarySyncTaskDuration records sync task wall time while concurrency tuning is in progress. -func logTemporarySyncTaskDuration(taskContext service.SyncTaskContext, startedAt time.Time) { +// logSyncTaskDuration test run time, delete it soon +func logSyncTaskDuration(taskContext service.SyncTaskContext, startedAt time.Time) { if taskContext.Task.TaskType != service.TaskTypeSync { return } common.Info( "sync task duration", - zap.String("temporary_code", "remove_after_sync_concurrency_optimization"), zap.String("task_id", taskContext.Task.ID), zap.String("connector_id", taskContext.Connector.ID), zap.String("kb_id", taskContext.Knowledgebase.ID), @@ -155,25 +171,19 @@ func logTemporarySyncTaskDuration(taskContext service.SyncTaskContext, startedAt // registerBuiltInConnectors registers datasource connectors available in the server binary. func registerBuiltInConnectors(registry *syncerconnector.Registry) { - registry.Register("rss", func(ctx context.Context, taskContext any) (syncerconnector.Connector, error) { + registerDAOConnector(registry, "rss", syncerconnector.NewRSSConnector) + registerDAOConnector(registry, "github", syncerconnector.NewGitHubConnector) + registerDAOConnector(registry, "gmail", syncerconnector.NewGmailConnector) + registerDAOConnector(registry, "google-drive", syncerconnector.NewGoogleDriveConnector) + registerDAOConnector(registry, "google_drive", syncerconnector.NewGoogleDriveConnector) +} + +func registerDAOConnector[T syncerconnector.Connector](registry *syncerconnector.Registry, source string, factory func(map[string]any) (T, error)) { + registry.Register(source, func(ctx context.Context, taskContext any) (syncerconnector.Connector, error) { row, ok := taskContext.(dao.SyncTaskContext) if !ok { - return nil, errors.New("rss connector received an invalid task context") + return nil, fmt.Errorf("%s connector received an invalid task context", source) } - return syncerconnector.NewRSSConnector(map[string]any(row.Connector.Config)) - }) - registry.Register("github", func(ctx context.Context, taskContext any) (syncerconnector.Connector, error) { - row, ok := taskContext.(dao.SyncTaskContext) - if !ok { - return nil, errors.New("github connector received an invalid task context") - } - return syncerconnector.NewGitHubConnector(map[string]any(row.Connector.Config)) - }) - registry.Register("gmail", func(ctx context.Context, taskContext any) (syncerconnector.Connector, error) { - row, ok := taskContext.(dao.SyncTaskContext) - if !ok { - return nil, errors.New("gmail connector received an invalid task context") - } - return syncerconnector.NewGmailConnector(map[string]any(row.Connector.Config)) + return factory(map[string]any(row.Connector.Config)) }) } diff --git a/internal/syncer/syncer_test.go b/internal/syncer/syncer_test.go index 8e4077274e..a6bec5a3d5 100644 --- a/internal/syncer/syncer_test.go +++ b/internal/syncer/syncer_test.go @@ -20,11 +20,13 @@ import ( "context" "errors" "io" + "ragflow/internal/common" "ragflow/internal/dao" "ragflow/internal/entity" "ragflow/internal/service" syncerconnector "ragflow/internal/syncer/connector" connectormock "ragflow/internal/syncer/connector/mock" + "strings" "sync" "testing" "time" @@ -38,6 +40,7 @@ import ( type fakeSink struct { mu sync.Mutex delay time.Duration + onUpsert func(input service.DocumentUpsertInput) current int maxConcurrent int calls []service.DocumentUpsertInput @@ -58,6 +61,9 @@ func (s *fakeSink) Upsert(ctx context.Context, input service.DocumentUpsertInput } s.autoParseByDoc[input.SourceDocument.SourceID] = input.AutoParse s.mu.Unlock() + if s.onUpsert != nil { + s.onUpsert(input) + } if s.delay > 0 { select { case <-ctx.Done(): @@ -115,6 +121,63 @@ type fakeDeleter struct { deleted []string } +type fakeTaskHandle struct { + mu sync.Mutex + msg common.TaskMessage + acks int + nacks int + inProgress int +} + +func (h *fakeTaskHandle) GetMessage() common.TaskMessage { return h.msg } +func (h *fakeTaskHandle) Ack() error { + h.mu.Lock() + defer h.mu.Unlock() + h.acks++ + return nil +} +func (h *fakeTaskHandle) Nack() error { + h.mu.Lock() + defer h.mu.Unlock() + h.nacks++ + return nil +} +func (h *fakeTaskHandle) InProgress() error { + h.mu.Lock() + defer h.mu.Unlock() + h.inProgress++ + return nil +} + +func (h *fakeTaskHandle) counts() (int, int) { + h.mu.Lock() + defer h.mu.Unlock() + return h.acks, h.nacks +} + +type fakeSyncTaskBroker struct { + published []string + fetched []common.TaskHandle +} + +func (b *fakeSyncTaskBroker) InitSyncerStream() error { return nil } +func (b *fakeSyncTaskBroker) InitSyncerConsumer() error { return nil } +func (b *fakeSyncTaskBroker) PublishSyncerTask(taskID string) error { + b.published = append(b.published, taskID) + return nil +} +func (b *fakeSyncTaskBroker) FetchSyncerTasks(batchSize int) ([]common.TaskHandle, error) { + if len(b.fetched) == 0 { + return nil, nil + } + if batchSize > len(b.fetched) { + batchSize = len(b.fetched) + } + out := b.fetched[:batchSize] + b.fetched = b.fetched[batchSize:] + return out, nil +} + // DeleteDocument records one delete. func (d *fakeDeleter) DeleteDocument(ctx context.Context, docID string) error { d.mu.Lock() @@ -159,7 +222,7 @@ func insertTaskContext(t *testing.T, db *gorm.DB, connectorID, kbID, taskID, tas Config: entity.JSONMap{"sync_deleted_files": true}, RefreshFreq: 0, PruneFreq: 0, - TimeoutSecs: 1, + TimeoutSecs: 60, Status: dao.SyncStatusSchedule, BaseModel: entity.BaseModel{UpdateDate: &now, UpdateTime: &ts}, }).Error; err != nil { @@ -195,6 +258,44 @@ func insertTaskContext(t *testing.T, db *gorm.DB, connectorID, kbID, taskID, tas } } +func insertKnowledgebaseMapping(t *testing.T, db *gorm.DB, connectorID, kbID string) { + t.Helper() + if err := db.Create(&entity.Knowledgebase{ + ID: kbID, + TenantID: "tenant-1", + Name: kbID, + EmbdID: "embd", + CreatedBy: "tenant-1", + ParserID: "naive", + ParserConfig: entity.JSONMap{}, + }).Error; err != nil { + t.Fatalf("insert kb: %v", err) + } + if err := db.Create(&entity.Connector2Kb{ID: connectorID + kbID, ConnectorID: connectorID, KbID: kbID, AutoParse: "1"}).Error; err != nil { + t.Fatalf("insert mapping: %v", err) + } +} + +func insertSyncLog(t *testing.T, db *gorm.DB, connectorID, kbID, taskID, taskType string) { + t.Helper() + now := time.Now().Add(-time.Hour).Truncate(time.Second) + ts := now.UnixMilli() + fromBeginning := "0" + if err := db.Create(&entity.SyncLogs{ + ID: taskID, + ConnectorID: connectorID, + KbID: kbID, + TaskType: taskType, + Status: dao.SyncStatusSchedule, + FromBeginning: &fromBeginning, + TimeStarted: &now, + ErrorMsg: "", + BaseModel: entity.BaseModel{UpdateDate: &now, UpdateTime: &ts}, + }).Error; err != nil { + t.Fatalf("insert task: %v", err) + } +} + // newTestRegistry creates a mock connector registry. func newTestRegistry(connectors map[string]*connectormock.Connector) *syncerconnector.Registry { registry := syncerconnector.NewRegistry() @@ -206,33 +307,213 @@ func newTestRegistry(connectors map[string]*connectormock.Connector) *syncerconn } // newCoordinator creates a test coordinator. -func newCoordinator(taskService *service.SyncTaskService, registry *syncerconnector.Registry, sink service.DocumentSink, pruneService *service.SyncPruneService, store service.DocumentStore, perTask int) *TaskCoordinator { - return NewTaskCoordinator(TaskCoordinatorConfig{PerTaskItemConcurrency: perTask, ItemRetryCount: 1, ItemRetryBaseDelay: time.Millisecond}, taskService, registry, sink, pruneService, service.NewDocumentIDResolver(store), make(chan struct{}, 16)) +func newCoordinator(taskService *service.SyncTaskService, registry *syncerconnector.Registry, sink service.DocumentSink, pruneService *service.SyncPruneService, store service.DocumentStore) *TaskCoordinator { + executor := NewSyncJobExecutor(SyncJobExecutorConfig{WorkerCount: 16}) + return NewTaskCoordinator(TaskCoordinatorConfig{ItemRetryCount: 1, ItemRetryBaseDelay: time.Millisecond}, taskService, registry, sink, pruneService, service.NewDocumentIDResolver(store), executor) } -// TestSchedulerClaimsDueTasks verifies conditional task claiming. -func TestSchedulerClaimsDueTasks(t *testing.T) { +// TestClaimBlocksSameConnectorKBRunningTasks verifies DB-backed task mutual exclusion. +func TestClaimBlocksSameConnectorKBRunningTasks(t *testing.T) { + db := setupSyncerDB(t) + insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) + insertSyncLog(t, db, "conn-1", "kb-1", "task-2", dao.TaskTypePrune) + taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(db)) + claimed, err := taskService.Claim(t.Context(), "task-1") + if err != nil { + t.Fatalf("claim task-1: %v", err) + } + if !claimed { + t.Fatalf("task-1 not claimed") + } + claimed, err = taskService.Claim(t.Context(), "task-2") + if err != nil { + t.Fatalf("claim task-2: %v", err) + } + if claimed { + t.Fatalf("task-2 claimed while task-1 is running") + } +} + +// TestClaimAllowsSameConnectorDifferentKB verifies connector work can run for different KBs. +func TestClaimAllowsSameConnectorDifferentKB(t *testing.T) { + db := setupSyncerDB(t) + insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) + insertKnowledgebaseMapping(t, db, "conn-1", "kb-2") + insertSyncLog(t, db, "conn-1", "kb-2", "task-2", dao.TaskTypeSync) + taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(db)) + for _, taskID := range []string{"task-1", "task-2"} { + claimed, err := taskService.Claim(t.Context(), taskID) + if err != nil { + t.Fatalf("claim %s: %v", taskID, err) + } + if !claimed { + t.Fatalf("%s not claimed", taskID) + } + } +} + +// TestSchedulerRequiresBroker verifies JetStream is mandatory for the scheduler. +func TestSchedulerRequiresBroker(t *testing.T) { + scheduler := NewScheduler(time.Hour, make(chan TaskEnvelope, 1), service.NewSyncTaskService(dao.NewSyncTaskDAO(nil))) + err := scheduler.Run(t.Context()) + if err == nil || !strings.Contains(err.Error(), "NATS broker") { + t.Fatalf("Run error = %v, want missing broker", err) + } +} + +// TestNATSSchedulerPublishesDueTasksWithoutClaiming verifies NATS mode leaves MySQL claim to workers. +func TestNATSSchedulerPublishesDueTasksWithoutClaiming(t *testing.T) { db := setupSyncerDB(t) insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(db)) - queue := make(chan TaskEnvelope, 1) - scheduler := NewScheduler(time.Hour, queue, taskService) - if err := scheduler.scan(t.Context()); err != nil { - t.Fatalf("scan: %v", err) + broker := &fakeSyncTaskBroker{} + scheduler := NewNATSScheduler(time.Hour, make(chan TaskEnvelope, 1), taskService, broker) + if err := scheduler.publishStartupTasks(t.Context()); err != nil { + t.Fatalf("publish startup tasks: %v", err) } - envelope := <-queue - if envelope.TaskID != "task-1" { - t.Fatalf("TaskID = %s, want task-1", envelope.TaskID) + if len(broker.published) != 1 || broker.published[0] != "task-1" { + t.Fatalf("published = %v", broker.published) } var task entity.SyncLogs if err := db.First(&task, "id = ?", "task-1").Error; err != nil { t.Fatalf("load task: %v", err) } - if task.Status != dao.SyncStatusRunning { - t.Fatalf("status = %s, want running", task.Status) + if task.Status != dao.SyncStatusSchedule { + t.Fatalf("status = %s, want schedule", task.Status) } } +// TestNATSSchedulerStartupPublishesScheduledTaskWithRefreshFreq verifies startup does not wait for the next refresh window. +func TestNATSSchedulerStartupPublishesScheduledTaskWithRefreshFreq(t *testing.T) { + db := setupSyncerDB(t) + insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) + now := time.Now() + updateTime := now.UnixMilli() + if err := db.Model(&entity.Connector{}).Where("id = ?", "conn-1").Update("refresh_freq", 5).Error; err != nil { + t.Fatalf("set refresh freq: %v", err) + } + if err := db.Model(&entity.SyncLogs{}).Where("id = ?", "task-1").Updates(map[string]any{ + "update_date": now, + "update_time": updateTime, + }).Error; err != nil { + t.Fatalf("set task update time: %v", err) + } + + taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(db)) + broker := &fakeSyncTaskBroker{} + scheduler := NewNATSScheduler(time.Hour, make(chan TaskEnvelope, 1), taskService, broker) + if err := scheduler.publishStartupTasks(t.Context()); err != nil { + t.Fatalf("publish startup tasks: %v", err) + } + if len(broker.published) != 1 || broker.published[0] != "task-1" { + t.Fatalf("published = %v", broker.published) + } +} + +// TestNATSSchedulerRecoversRunningTasksOnStartup verifies startup reconciliation publishes interrupted tasks. +func TestNATSSchedulerRecoversRunningTasksOnStartup(t *testing.T) { + db := setupSyncerDB(t) + insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) + if err := db.Model(&entity.SyncLogs{}).Where("id = ?", "task-1").Update("status", dao.SyncStatusRunning).Error; err != nil { + t.Fatalf("mark task running: %v", err) + } + if err := db.Model(&entity.Connector{}).Where("id = ?", "conn-1").Update("status", dao.SyncStatusRunning).Error; err != nil { + t.Fatalf("mark connector running: %v", err) + } + taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(db)) + broker := &fakeSyncTaskBroker{} + scheduler := NewNATSScheduler(time.Hour, make(chan TaskEnvelope, 1), taskService, broker) + + if err := scheduler.publishStartupTasks(t.Context()); err != nil { + t.Fatalf("publish startup tasks: %v", err) + } + if len(broker.published) != 1 || broker.published[0] != "task-1" { + t.Fatalf("published = %v", broker.published) + } + var task entity.SyncLogs + if err := db.First(&task, "id = ?", "task-1").Error; err != nil { + t.Fatalf("load task: %v", err) + } + if task.Status != dao.SyncStatusSchedule { + t.Fatalf("task status = %s, want schedule", task.Status) + } + var connector entity.Connector + if err := db.First(&connector, "id = ?", "conn-1").Error; err != nil { + t.Fatalf("load connector: %v", err) + } + if connector.Status != dao.SyncStatusSchedule { + t.Fatalf("connector status = %s, want schedule", connector.Status) + } +} + +// TestNATSSchedulerPublishesOnlyDueTasks verifies periodic NATS publishing respects refresh windows. +func TestNATSSchedulerPublishesOnlyDueTasks(t *testing.T) { + db := setupSyncerDB(t) + insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) + now := time.Now() + updateTime := now.UnixMilli() + if err := db.Model(&entity.Connector{}).Where("id = ?", "conn-1").Update("refresh_freq", 5).Error; err != nil { + t.Fatalf("set refresh freq: %v", err) + } + if err := db.Model(&entity.SyncLogs{}).Where("id = ?", "task-1").Updates(map[string]any{ + "update_date": now, + "update_time": updateTime, + }).Error; err != nil { + t.Fatalf("set fresh task update time: %v", err) + } + taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(db)) + broker := &fakeSyncTaskBroker{} + scheduler := NewNATSScheduler(10*time.Millisecond, make(chan TaskEnvelope, 1), taskService, broker) + + if err := scheduler.publishDueTasks(t.Context()); err != nil { + t.Fatalf("publish fresh due tasks: %v", err) + } + if len(broker.published) != 0 { + t.Fatalf("fresh task published = %v, want none", broker.published) + } + + dueAt := now.Add(-10 * time.Minute) + if err := db.Model(&entity.SyncLogs{}).Where("id = ?", "task-1").Updates(map[string]any{ + "update_date": dueAt, + "update_time": dueAt.UnixMilli(), + }).Error; err != nil { + t.Fatalf("set due task update time: %v", err) + } + + if err := scheduler.publishDueTasks(t.Context()); err != nil { + t.Fatalf("publish due tasks: %v", err) + } + if len(broker.published) != 1 || broker.published[0] != "task-1" { + t.Fatalf("published = %v, want due task", broker.published) + } +} + +// TestNATSSchedulerBuffersFetchedTasks verifies NATS mode stores excess tasks in the local queue. +func TestNATSSchedulerBuffersFetchedTasks(t *testing.T) { + taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(nil)) + queue := make(chan TaskEnvelope, 2) + scheduler := NewNATSScheduler(time.Hour, queue, taskService, &fakeSyncTaskBroker{}) + handles := []common.TaskHandle{ + &fakeTaskHandle{msg: common.TaskMessage{TaskID: "task-1", TaskType: common.TaskTypeSyncer}}, + &fakeTaskHandle{msg: common.TaskMessage{TaskID: "task-2", TaskType: common.TaskTypeSyncer}}, + } + + if err := scheduler.enqueueHandles(t.Context(), handles); err != nil { + t.Fatalf("enqueue handles: %v", err) + } + if got := scheduler.queueAvailable(); got != 0 { + t.Fatalf("queue capacity while buffered = %d, want 0", got) + } + + first := <-queue + stopEnvelopeHeartbeat(first) + if got := scheduler.queueAvailable(); got != 1 { + t.Fatalf("queue capacity after dequeue = %d, want 1", got) + } + second := <-queue + stopEnvelopeHeartbeat(second) +} + // TestWorkersRunDifferentConnectorsInParallel verifies task-level parallelism. func TestWorkersRunDifferentConnectorsInParallel(t *testing.T) { db := setupSyncerDB(t) @@ -253,7 +534,7 @@ func TestWorkersRunDifferentConnectorsInParallel(t *testing.T) { "conn-2": {SyncBatches: []syncerconnector.SyncBatch{{Documents: []syncerconnector.SourceDocument{{SourceID: "b", UpdatedAt: now}}}}}, } queue := make(chan TaskEnvelope, 2) - worker := NewTaskWorker(queue, taskService, newCoordinator(taskService, newTestRegistry(connectors), sink, nil, fakeStore{}, 1), NewConnectorLock()) + worker := NewTaskWorker(queue, taskService, newCoordinator(taskService, newTestRegistry(connectors), sink, nil, fakeStore{}), NewConnectorLock()) ctx, cancel := context.WithCancel(t.Context()) defer cancel() go worker.Run(ctx, 2) @@ -268,8 +549,52 @@ func TestWorkersRunDifferentConnectorsInParallel(t *testing.T) { } } -// TestConnectorLockSerializesSameConnector verifies connector-level mutual exclusion. -func TestConnectorLockSerializesSameConnector(t *testing.T) { +// TestNATSTaskWorkerClaimsAndAcksOnSuccess verifies NATS messages are ACKed after durable completion. +func TestNATSTaskWorkerClaimsAndAcksOnSuccess(t *testing.T) { + db := setupSyncerDB(t) + insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) + taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(db)) + now := time.Now() + connector := &connectormock.Connector{SyncBatches: []syncerconnector.SyncBatch{{Documents: []syncerconnector.SourceDocument{{SourceID: "a", UpdatedAt: now}}}}} + handle := &fakeTaskHandle{msg: common.TaskMessage{TaskID: "task-1", TaskType: common.TaskTypeSyncer}} + worker := NewTaskWorker( + make(chan TaskEnvelope, 1), + taskService, + newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), &fakeSink{}, nil, fakeStore{}), + NewConnectorLock(), + ) + worker.handle(t.Context(), TaskEnvelope{TaskID: "task-1", Handle: handle}) + + if handle.acks != 1 || handle.nacks != 0 { + t.Fatalf("settlement acks=%d nacks=%d", handle.acks, handle.nacks) + } + var task entity.SyncLogs + if err := db.First(&task, "id = ?", "task-1").Error; err != nil { + t.Fatalf("load task: %v", err) + } + if task.Status != dao.SyncStatusDone { + t.Fatalf("status = %s, want done", task.Status) + } +} + +// TestNATSTaskWorkerAcksUnclaimableMessage verifies duplicate/stale messages do not redeliver forever. +func TestNATSTaskWorkerAcksUnclaimableMessage(t *testing.T) { + db := setupSyncerDB(t) + insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) + if err := db.Model(&entity.SyncLogs{}).Where("id = ?", "task-1").Update("status", dao.SyncStatusRunning).Error; err != nil { + t.Fatalf("mark running: %v", err) + } + taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(db)) + handle := &fakeTaskHandle{msg: common.TaskMessage{TaskID: "task-1", TaskType: common.TaskTypeSyncer}} + worker := NewTaskWorker(make(chan TaskEnvelope, 1), taskService, newCoordinator(taskService, newTestRegistry(nil), &fakeSink{}, nil, fakeStore{}), NewConnectorLock()) + worker.handle(t.Context(), TaskEnvelope{TaskID: "task-1", Handle: handle}) + if handle.acks != 1 || handle.nacks != 0 { + t.Fatalf("settlement acks=%d nacks=%d", handle.acks, handle.nacks) + } +} + +// TestSameConnectorDifferentKBsRunInParallel verifies one datasource can sync into different KBs concurrently. +func TestSameConnectorDifferentKBsRunInParallel(t *testing.T) { db := setupSyncerDB(t) insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) insertTaskContext(t, db, "conn-1", "kb-2", "task-2", dao.TaskTypeSync) @@ -283,30 +608,94 @@ func TestConnectorLockSerializesSameConnector(t *testing.T) { sink := &fakeSink{delay: 100 * time.Millisecond} connector := &connectormock.Connector{SyncBatches: []syncerconnector.SyncBatch{{Documents: []syncerconnector.SourceDocument{{SourceID: "a", UpdatedAt: time.Now()}}}}} queue := make(chan TaskEnvelope, 2) - worker := NewTaskWorker(queue, taskService, newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, nil, fakeStore{}, 1), NewConnectorLock()) + worker := NewTaskWorker(queue, taskService, newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, nil, fakeStore{}), NewConnectorLock()) ctx, cancel := context.WithCancel(t.Context()) defer cancel() go worker.Run(ctx, 2) queue <- TaskEnvelope{TaskID: "task-1"} queue <- TaskEnvelope{TaskID: "task-2"} - time.Sleep(180 * time.Millisecond) - var scheduled int64 - if err := db.Model(&entity.SyncLogs{}).Where("connector_id = ? AND status = ?", "conn-1", dao.SyncStatusSchedule).Count(&scheduled).Error; err != nil { - t.Fatalf("count scheduled: %v", err) - } - if scheduled == 0 { - t.Fatalf("expected one same-connector task to be rescheduled") - } + time.Sleep(40 * time.Millisecond) sink.mu.Lock() maxConcurrent := sink.maxConcurrent sink.mu.Unlock() - if maxConcurrent > 1 { - t.Fatalf("same connector ran concurrently: %d", maxConcurrent) + if maxConcurrent < 2 { + t.Fatalf("max concurrent same-connector different-kb sink calls = %d, want >= 2", maxConcurrent) } } -// TestSyncRunnerReadsBatchesSerially verifies the next batch waits for current items. -func TestSyncRunnerReadsBatchesSerially(t *testing.T) { +// TestConnectorKBLockSerializesSyncAndPrune verifies sync and prune for the same connector/KB do not overlap. +func TestConnectorKBLockSerializesSyncAndPrune(t *testing.T) { + db := setupSyncerDB(t) + insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) + now := time.Now().Add(-time.Hour).Truncate(time.Second) + ts := now.UnixMilli() + fromBeginning := "0" + if err := db.Create(&entity.SyncLogs{ + ID: "task-2", + ConnectorID: "conn-1", + KbID: "kb-1", + TaskType: dao.TaskTypePrune, + Status: dao.SyncStatusSchedule, + FromBeginning: &fromBeginning, + TimeStarted: &now, + ErrorMsg: "", + BaseModel: entity.BaseModel{UpdateDate: &now, UpdateTime: &ts}, + }).Error; err != nil { + t.Fatalf("insert prune task: %v", err) + } + taskDAO := dao.NewSyncTaskDAO(db) + taskService := service.NewSyncTaskService(taskDAO) + for _, id := range []string{"task-1", "task-2"} { + if _, err := taskDAO.ClaimTask(t.Context(), id, time.Now()); err != nil { + t.Fatalf("claim %s: %v", id, err) + } + } + sink := &fakeSink{delay: 120 * time.Millisecond} + connector := &connectormock.Connector{ + SyncBatches: []syncerconnector.SyncBatch{{Documents: []syncerconnector.SourceDocument{{SourceID: "a", UpdatedAt: time.Now()}}}}, + PruneBatches: []syncerconnector.PruneBatch{{Documents: []syncerconnector.SlimDocument{{SourceID: "a"}}}}, + } + pruneService := service.NewSyncPruneService(&fakeDeleter{}, fakeStore{ids: map[string]struct{}{}}) + worker := NewTaskWorker(make(chan TaskEnvelope, 2), taskService, newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, pruneService, fakeStore{}), NewConnectorLock()) + done := make(chan struct{}) + go func() { + defer close(done) + worker.handle(t.Context(), TaskEnvelope{TaskID: "task-1"}) + }() + waitForSinkConcurrency(t, sink, 1) + worker.handle(t.Context(), TaskEnvelope{TaskID: "task-2"}) + var task entity.SyncLogs + if err := db.First(&task, "id = ?", "task-2").Error; err != nil { + t.Fatalf("load prune task: %v", err) + } + if task.Status != dao.SyncStatusSchedule { + t.Fatalf("same connector/kb prune status = %s, want schedule", task.Status) + } + <-done +} + +func waitForSinkConcurrency(t *testing.T, sink *fakeSink, want int) { + t.Helper() + deadline := time.After(time.Second) + ticker := time.NewTicker(time.Millisecond) + defer ticker.Stop() + for { + select { + case <-deadline: + t.Fatalf("timed out waiting for sink concurrency >= %d", want) + case <-ticker.C: + sink.mu.Lock() + current := sink.current + sink.mu.Unlock() + if current >= want { + return + } + } + } +} + +// TestSyncRunnerSubmitsBatchesBeforeWaiting verifies source reads are not blocked by prior batch jobs. +func TestSyncRunnerSubmitsBatchesBeforeWaiting(t *testing.T) { db := setupSyncerDB(t) insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) taskDAO := dao.NewSyncTaskDAO(db) @@ -326,32 +715,39 @@ func TestSyncRunnerReadsBatchesSerially(t *testing.T) { }, } sink := &fakeSink{delay: 80 * time.Millisecond} - coordinator := newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, nil, fakeStore{}, 2) + coordinator := newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, nil, fakeStore{}) taskContext, _ := taskService.GetContext(t.Context(), "task-1") - if err := coordinator.Execute(t.Context(), taskContext); err != nil { + if err := coordinator.Execute(t.Context(), taskContext, testLockLease()); err != nil { t.Fatalf("execute: %v", err) } - if secondBatchAt.Sub(start) < 70*time.Millisecond { - t.Fatalf("second batch was read before first batch completed") + if secondBatchAt.Sub(start) >= 70*time.Millisecond { + t.Fatalf("second batch was blocked by first batch for %s", secondBatchAt.Sub(start)) } } -// TestSyncRunnerProcessesBatchItemsInParallel verifies per-batch item concurrency. -func TestSyncRunnerProcessesBatchItemsInParallel(t *testing.T) { +// TestSyncRunnerProcessesBatchJobsInParallel verifies source batches run as parallel BatchJobs. +func TestSyncRunnerProcessesBatchJobsInParallel(t *testing.T) { db := setupSyncerDB(t) insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) _ = db.Model(&entity.SyncLogs{}).Where("id = ?", "task-1").Update("status", dao.SyncStatusRunning).Error taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(db)) now := time.Now() - connector := &connectormock.Connector{SyncBatches: []syncerconnector.SyncBatch{{Documents: []syncerconnector.SourceDocument{{SourceID: "a", UpdatedAt: now}, {SourceID: "b", UpdatedAt: now}, {SourceID: "c", UpdatedAt: now}}}}} + connector := &connectormock.Connector{SyncBatches: []syncerconnector.SyncBatch{ + {Documents: []syncerconnector.SourceDocument{{SourceID: "a", UpdatedAt: now}}}, + {Documents: []syncerconnector.SourceDocument{{SourceID: "b", UpdatedAt: now}}}, + {Documents: []syncerconnector.SourceDocument{{SourceID: "c", UpdatedAt: now}}}, + }} sink := &fakeSink{delay: 80 * time.Millisecond} - coordinator := newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, nil, fakeStore{}, 3) + coordinator := newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, nil, fakeStore{}) taskContext, _ := taskService.GetContext(t.Context(), "task-1") - if err := coordinator.Execute(t.Context(), taskContext); err != nil { + if err := coordinator.Execute(t.Context(), taskContext, testLockLease()); err != nil { t.Fatalf("execute: %v", err) } - if sink.maxConcurrent < 2 { - t.Fatalf("max concurrent items = %d, want >= 2", sink.maxConcurrent) + sink.mu.Lock() + maxConcurrent := sink.maxConcurrent + sink.mu.Unlock() + if maxConcurrent < 2 { + t.Fatalf("max concurrent batches = %d, want >= 2", maxConcurrent) } } @@ -380,9 +776,9 @@ func TestFingerprintSkipsUnchangedDocument(t *testing.T) { store := fakeStore{ids: map[string]struct{}{legacyID: {}}, fingerprints: map[string]string{legacyID: "fp-1"}} connector := &connectormock.Connector{SyncBatches: []syncerconnector.SyncBatch{{Documents: []syncerconnector.SourceDocument{{SourceID: "source-1", Fingerprint: "fp-1", FetchRef: &syncerconnector.FetchReference{Key: "lazy"}, UpdatedAt: time.Now()}}}}} sink := &fakeSink{} - coordinator := newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, nil, store, 1) + coordinator := newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, nil, store) taskContext, _ := taskService.GetContext(t.Context(), "task-1") - if err := coordinator.Execute(t.Context(), taskContext); err != nil { + if err := coordinator.Execute(t.Context(), taskContext, testLockLease()); err != nil { t.Fatalf("execute: %v", err) } if sink.callCount() != 0 { @@ -401,9 +797,9 @@ func TestAutoParseFlagFlowsToSink(t *testing.T) { taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(db)) connector := &connectormock.Connector{SyncBatches: []syncerconnector.SyncBatch{{Documents: []syncerconnector.SourceDocument{{SourceID: "source-1", Blob: []byte("x"), UpdatedAt: time.Now()}}}}} sink := &fakeSink{} - coordinator := newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, nil, fakeStore{}, 1) + coordinator := newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, nil, fakeStore{}) taskContext, _ := taskService.GetContext(t.Context(), "task-1") - if err := coordinator.Execute(t.Context(), taskContext); err != nil { + if err := coordinator.Execute(t.Context(), taskContext, testLockLease()); err != nil { t.Fatalf("execute: %v", err) } if sink.autoParseByDoc["source-1"] { @@ -411,6 +807,133 @@ func TestAutoParseFlagFlowsToSink(t *testing.T) { } } +// TestCompleteSyncSchedulesNextRun verifies completing one sync keeps the connector schedulable. +func TestCompleteSyncSchedulesNextRun(t *testing.T) { + db := setupSyncerDB(t) + insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) + _ = db.Model(&entity.SyncLogs{}).Where("id = ?", "task-1").Update("status", dao.SyncStatusRunning).Error + _ = db.Model(&entity.Connector{}).Where("id = ?", "conn-1").Update("status", dao.SyncStatusRunning).Error + taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(db)) + now := time.Now() + connector := &connectormock.Connector{SyncBatches: []syncerconnector.SyncBatch{{Documents: []syncerconnector.SourceDocument{{SourceID: "source-1", Blob: []byte("x"), UpdatedAt: now}}}}} + coordinator := newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), &fakeSink{}, nil, fakeStore{}) + taskContext, _ := taskService.GetContext(t.Context(), "task-1") + if err := coordinator.Execute(t.Context(), taskContext, testLockLease()); err != nil { + t.Fatalf("execute: %v", err) + } + + var connectorRow entity.Connector + if err := db.First(&connectorRow, "id = ?", "conn-1").Error; err != nil { + t.Fatalf("load connector: %v", err) + } + if connectorRow.Status != dao.SyncStatusSchedule { + t.Fatalf("connector status = %s, want schedule", connectorRow.Status) + } + var scheduled int64 + if err := db.Model(&entity.SyncLogs{}). + Where("connector_id = ? AND kb_id = ? AND task_type = ? AND status = ?", "conn-1", "kb-1", dao.TaskTypeSync, dao.SyncStatusSchedule). + Count(&scheduled).Error; err != nil { + t.Fatalf("count scheduled sync logs: %v", err) + } + if scheduled != 1 { + t.Fatalf("scheduled sync logs = %d, want 1", scheduled) + } +} + +// TestCancelStopsRunningSync verifies a stop request prevents further work and completion. +func TestCancelStopsRunningSync(t *testing.T) { + db := setupSyncerDB(t) + insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) + _ = db.Model(&entity.SyncLogs{}).Where("id = ?", "task-1").Update("status", dao.SyncStatusRunning).Error + _ = db.Model(&entity.Connector{}).Where("id = ?", "conn-1").Update("status", dao.SyncStatusRunning).Error + taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(db)) + now := time.Now() + connector := &connectormock.Connector{SyncBatches: []syncerconnector.SyncBatch{{Documents: []syncerconnector.SourceDocument{ + {SourceID: "source-1", Blob: []byte("one"), UpdatedAt: now}, + {SourceID: "source-2", Blob: []byte("two"), UpdatedAt: now}, + }}}} + var cancelOnce sync.Once + sink := &fakeSink{ + onUpsert: func(input service.DocumentUpsertInput) { + cancelOnce.Do(func() { + if err := db.Model(&entity.SyncLogs{}).Where("id = ?", input.TaskContext.Task.ID).Update("status", dao.SyncStatusCancel).Error; err != nil { + t.Errorf("cancel sync log: %v", err) + } + if err := db.Model(&entity.Connector{}).Where("id = ?", input.TaskContext.Connector.ID).Update("status", dao.SyncStatusCancel).Error; err != nil { + t.Errorf("cancel connector: %v", err) + } + time.Sleep(syncCancelCheckInterval) + }) + }, + } + worker := NewTaskWorker(make(chan TaskEnvelope, 1), taskService, newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, nil, fakeStore{}), NewConnectorLock()) + worker.handle(t.Context(), TaskEnvelope{TaskID: "task-1"}) + + if calls := sink.callCount(); calls != 1 { + t.Fatalf("sink calls = %d, want 1", calls) + } + var task entity.SyncLogs + if err := db.First(&task, "id = ?", "task-1").Error; err != nil { + t.Fatalf("load task: %v", err) + } + if task.Status != dao.SyncStatusCancel { + t.Fatalf("task status = %s, want cancel", task.Status) + } + var scheduled int64 + if err := db.Model(&entity.SyncLogs{}). + Where("connector_id = ? AND kb_id = ? AND task_type = ? AND status = ?", "conn-1", "kb-1", dao.TaskTypeSync, dao.SyncStatusSchedule). + Count(&scheduled).Error; err != nil { + t.Fatalf("count scheduled sync logs: %v", err) + } + if scheduled != 0 { + t.Fatalf("scheduled sync logs = %d, want 0", scheduled) + } + var connectorRow entity.Connector + if err := db.First(&connectorRow, "id = ?", "conn-1").Error; err != nil { + t.Fatalf("load connector: %v", err) + } + if connectorRow.Status != dao.SyncStatusCancel { + t.Fatalf("connector status = %s, want cancel", connectorRow.Status) + } +} + +// TestSyncRunnerResultWaitHonorsCancel verifies Run does not block forever waiting for job results after cancellation. +func TestSyncRunnerResultWaitHonorsCancel(t *testing.T) { + db := setupSyncerDB(t) + insertTaskContext(t, db, "conn-1", "kb-1", "task-1", dao.TaskTypeSync) + _ = db.Model(&entity.SyncLogs{}).Where("id = ?", "task-1").Update("status", dao.SyncStatusRunning).Error + taskService := service.NewSyncTaskService(dao.NewSyncTaskDAO(db)) + taskContext, err := taskService.GetContext(t.Context(), "task-1") + if err != nil { + t.Fatalf("get context: %v", err) + } + connector := &connectormock.Connector{SyncBatches: []syncerconnector.SyncBatch{{Documents: []syncerconnector.SourceDocument{{SourceID: "source-1", Blob: []byte("x"), UpdatedAt: time.Now()}}}}} + queue := &SyncJobQueue{taskID: "task-1", jobs: make(chan *syncJob, 1)} + runner := NewSyncRunner(TaskCoordinatorConfig{ItemRetryCount: 1, ItemRetryBaseDelay: time.Millisecond}, taskService, &fakeSink{}, service.NewDocumentIDResolver(fakeStore{}), queue) + + ctx, cancel := context.WithCancel(context.Background()) + result := make(chan error, 1) + go func() { + result <- runner.Run(ctx, taskContext, connector) + }() + + select { + case <-queue.jobs: + case <-time.After(time.Second): + t.Fatalf("timed out waiting for submitted job") + } + cancel() + + select { + case err = <-result: + if !errors.Is(err, context.Canceled) { + t.Fatalf("Run error = %v, want context canceled", err) + } + case <-time.After(time.Second): + t.Fatalf("Run did not return after cancellation") + } +} + // TestBatchFailureDoesNotAdvanceWaterline verifies failed tasks keep poll_range_end. func TestBatchFailureDoesNotAdvanceWaterline(t *testing.T) { db := setupSyncerDB(t) @@ -420,7 +943,7 @@ func TestBatchFailureDoesNotAdvanceWaterline(t *testing.T) { connector := &connectormock.Connector{SyncBatches: []syncerconnector.SyncBatch{{Documents: []syncerconnector.SourceDocument{{SourceID: "bad", Blob: []byte("x"), UpdatedAt: time.Now()}}}}} sink := &fakeSink{errBySourceID: map[string]error{"bad": errors.New("boom")}} queue := make(chan TaskEnvelope, 1) - worker := NewTaskWorker(queue, taskService, newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, nil, fakeStore{}, 1), NewConnectorLock()) + worker := NewTaskWorker(queue, taskService, newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), sink, nil, fakeStore{}), NewConnectorLock()) worker.handle(t.Context(), TaskEnvelope{TaskID: "task-1"}) var task entity.SyncLogs if err := db.First(&task, "id = ?", "task-1").Error; err != nil { @@ -465,7 +988,7 @@ func TestPruneSourceFailureDoesNotDelete(t *testing.T) { pruneService := service.NewSyncPruneService(deleter, fakeStore{ids: map[string]struct{}{"stale": {}}}) connector := &connectormock.Connector{PruneErrAt: 1, PruneBatches: []syncerconnector.PruneBatch{{Documents: []syncerconnector.SlimDocument{{SourceID: "keep"}}}}} queue := make(chan TaskEnvelope, 1) - worker := NewTaskWorker(queue, taskService, newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), &fakeSink{}, pruneService, fakeStore{}, 1), NewConnectorLock()) + worker := NewTaskWorker(queue, taskService, newCoordinator(taskService, newTestRegistry(map[string]*connectormock.Connector{"conn-1": connector}), &fakeSink{}, pruneService, fakeStore{}), NewConnectorLock()) worker.handle(t.Context(), TaskEnvelope{TaskID: "task-1"}) if len(deleter.deleted) != 0 { t.Fatalf("deleted %v despite incomplete snapshot", deleter.deleted) @@ -493,3 +1016,26 @@ func TestMockSessionEOF(t *testing.T) { t.Fatalf("err = %v, want io.EOF", err) } } + +// TestTaskExecutionDeadlineCapsConnectorTimeout verifies tasks cannot outlive the connector lock lease. +func TestTaskExecutionDeadlineCapsConnectorTimeout(t *testing.T) { + now := time.Date(2026, 8, 10, 10, 0, 0, 0, time.UTC) + taskContext := service.SyncTaskContext{} + taskContext.Connector.TimeoutSecs = int64(connectorLockTTL.Seconds()) + 60 + lease := ConnectorLockLease{ExpiresAt: now.Add(10 * time.Minute)} + if got, want := taskExecutionDeadline(now, taskContext, lease), lease.ExpiresAt.Add(-connectorLockSafetyMargin); !got.Equal(want) { + t.Fatalf("deadline = %s, want %s", got, want) + } + taskContext.Connector.TimeoutSecs = 30 + if got, want := taskExecutionDeadline(now, taskContext, lease), now.Add(30*time.Second); !got.Equal(want) { + t.Fatalf("deadline = %s, want %s", got, want) + } + lease.ExpiresAt = now.Add(time.Second) + if got := taskExecutionDeadline(now, taskContext, lease); !got.Equal(now) { + t.Fatalf("deadline = %s, want %s", got, now) + } +} + +func testLockLease() ConnectorLockLease { + return ConnectorLockLease{ExpiresAt: time.Now().Add(connectorLockTTL)} +} diff --git a/internal/syncer/task_coordinator.go b/internal/syncer/task_coordinator.go index f129ea5da2..8b7ffd1419 100644 --- a/internal/syncer/task_coordinator.go +++ b/internal/syncer/task_coordinator.go @@ -24,6 +24,8 @@ import ( "time" ) +const connectorLockSafetyMargin = 5 * time.Second + // ConnectorRegistry opens registered connectors by source. type ConnectorRegistry interface { // Open creates a connector for a task context. @@ -32,9 +34,8 @@ type ConnectorRegistry interface { // TaskCoordinatorConfig controls per-task document processing. type TaskCoordinatorConfig struct { - PerTaskItemConcurrency int - ItemRetryCount int - ItemRetryBaseDelay time.Duration + ItemRetryCount int + ItemRetryBaseDelay time.Duration } // TaskCoordinator owns one task execution window. @@ -45,15 +46,11 @@ type TaskCoordinator struct { sink service.DocumentSink pruneService *service.SyncPruneService idResolver *service.DocumentIDResolver - globalItems chan struct{} + executor *SyncJobExecutor } // NewTaskCoordinator creates a coordinator for one claimed task at a time. -func NewTaskCoordinator(config TaskCoordinatorConfig, taskService *service.SyncTaskService, registry ConnectorRegistry, sink service.DocumentSink, pruneService *service.SyncPruneService, idResolver *service.DocumentIDResolver, globalItems chan struct{}) *TaskCoordinator { - if config.PerTaskItemConcurrency <= 0 { - config.PerTaskItemConcurrency = 1 - } - +func NewTaskCoordinator(config TaskCoordinatorConfig, taskService *service.SyncTaskService, registry ConnectorRegistry, sink service.DocumentSink, pruneService *service.SyncPruneService, idResolver *service.DocumentIDResolver, executor *SyncJobExecutor) *TaskCoordinator { if config.ItemRetryCount <= 0 { config.ItemRetryCount = 1 } @@ -61,12 +58,19 @@ func NewTaskCoordinator(config TaskCoordinatorConfig, taskService *service.SyncT if config.ItemRetryBaseDelay <= 0 { config.ItemRetryBaseDelay = time.Second } + if executor == nil { + panic("task coordinator executor must not be nil") + } - return &TaskCoordinator{config: config, taskService: taskService, registry: registry, sink: sink, pruneService: pruneService, idResolver: idResolver, globalItems: globalItems} + return &TaskCoordinator{config: config, taskService: taskService, registry: registry, sink: sink, pruneService: pruneService, idResolver: idResolver, executor: executor} } // Execute dispatches a sync_logs task by task type. -func (c *TaskCoordinator) Execute(ctx context.Context, taskContext service.SyncTaskContext) error { +func (c *TaskCoordinator) Execute(ctx context.Context, taskContext service.SyncTaskContext, lease ConnectorLockLease) error { + runCtx, cancel := context.WithDeadline(ctx, taskExecutionDeadline(time.Now(), taskContext, lease)) + defer cancel() + ctx = runCtx + connector, err := c.registry.Open(ctx, taskContext) if err != nil { return err @@ -74,9 +78,16 @@ func (c *TaskCoordinator) Execute(ctx context.Context, taskContext service.SyncT if err = connector.Validate(ctx); err != nil { return err } + switch taskContext.Task.TaskType { case service.TaskTypeSync: - runner := NewSyncRunner(c.config, c.taskService, c.sink, c.idResolver, c.globalItems) + queue, err := c.executor.RegisterTask(ctx, taskContext.Task.ID) + if err != nil { + return err + } + defer queue.Close() + + runner := NewSyncRunner(c.config, c.taskService, c.sink, c.idResolver, queue) return runner.Run(ctx, taskContext, connector) case service.TaskTypePrune: runner := NewPruneRunner(c.taskService, c.pruneService) @@ -85,3 +96,27 @@ func (c *TaskCoordinator) Execute(ctx context.Context, taskContext service.SyncT return fmt.Errorf("unsupported sync task type %q", taskContext.Task.TaskType) } } + +func taskExecutionDeadline(now time.Time, taskContext service.SyncTaskContext, lease ConnectorLockLease) time.Time { + timeout := connectorLockTTL + if seconds := taskContext.Connector.TimeoutSecs; seconds > 0 && seconds <= int64(connectorLockTTL/time.Second) { + timeout = time.Duration(seconds) * time.Second + } + + deadline := now.Add(timeout) + if lease.ExpiresAt.IsZero() { + if timeout > connectorLockTTL { + return now.Add(connectorLockTTL) + } + return deadline + } + + lockDeadline := lease.ExpiresAt.Add(-connectorLockSafetyMargin) + if lockDeadline.Before(now) { + return now + } + if lockDeadline.Before(deadline) { + return lockDeadline + } + return deadline +} diff --git a/internal/syncer/task_worker.go b/internal/syncer/task_worker.go index 34b3931e77..044a5c17bc 100644 --- a/internal/syncer/task_worker.go +++ b/internal/syncer/task_worker.go @@ -18,7 +18,9 @@ package syncer import ( "context" + "errors" "fmt" + "ragflow/internal/common" "ragflow/internal/service" "sync" "time" @@ -65,30 +67,120 @@ func (w *TaskWorker) loop(ctx context.Context) { // handle loads a claimed task and executes it under the connector lock. func (w *TaskWorker) handle(ctx context.Context, envelope TaskEnvelope) { + // start heartbeat, communication with nats + if envelope.Handle != nil && envelope.stopHeartbeat == nil { + envelope.stopHeartbeat = startHandleHeartbeat(ctx, envelope.Handle) + } + defer stopEnvelopeHeartbeat(envelope) + + if envelope.Handle != nil { + // claim a task(sync/ prune) + claimed, err := w.taskService.Claim(ctx, envelope.TaskID) + if err != nil { + _ = envelope.Handle.Nack() + return + } + if !claimed { + _ = envelope.Handle.Ack() // this task has been claimed by other worker + return + } + } + + // get the whole context by task_id from nats taskContext, err := w.taskService.GetContext(ctx, envelope.TaskID) if err != nil { - if ctx.Err() != nil { + if ctx.Err() != nil { // exiting _ = w.taskService.RescheduleClaimed(context.WithoutCancel(ctx), envelope.TaskID) + nackEnvelope(envelope) return } - _ = w.taskService.Fail(ctx, envelope.TaskID, "", err) + if failErr := w.taskService.Fail(ctx, envelope.TaskID, "", err); failErr != nil { // getContext failed + _ = w.taskService.RescheduleClaimed(context.WithoutCancel(ctx), envelope.TaskID) + nackEnvelope(envelope) + return + } + ackEnvelope(envelope) return } - if !w.locker.TryLock(taskContext.Connector.ID) { + + // lock the connector and the KB + lease, locked := w.locker.TryLock(taskContext.Connector.ID, taskContext.Knowledgebase.ID) + if !locked { _ = w.taskService.RescheduleClaimed(ctx, taskContext.Task.ID) + ackEnvelope(envelope) return } - defer w.locker.Unlock(taskContext.Connector.ID) + defer w.locker.Unlock(taskContext.Connector.ID, taskContext.Knowledgebase.ID) startedAt := time.Now() - if err = w.coordinator.Execute(ctx, taskContext); err != nil { - logTemporarySyncTaskDuration(taskContext, startedAt) - if ctx.Err() != nil { - _ = w.taskService.RescheduleClaimed(context.WithoutCancel(ctx), taskContext.Task.ID) + if err = w.coordinator.Execute(ctx, taskContext, lease); err != nil { // execute the task(sync/ prune) + logSyncTaskDuration(taskContext, startedAt) + if errors.Is(err, errSyncTaskCanceled) { // task is canceled by the user + ackEnvelope(envelope) return } - _ = w.taskService.Fail(ctx, taskContext.Task.ID, taskContext.Connector.ID, fmt.Errorf("sync task failed: %w", err)) + if ctx.Err() != nil { // the task is canceled by system, this need to rerun + _ = w.taskService.RescheduleClaimed(context.WithoutCancel(ctx), taskContext.Task.ID) + nackEnvelope(envelope) + return + } + if failErr := w.taskService.Fail(ctx, taskContext.Task.ID, taskContext.Connector.ID, fmt.Errorf("sync task failed: %w", err)); failErr != nil { + _ = w.taskService.RescheduleClaimed(context.WithoutCancel(ctx), taskContext.Task.ID) + nackEnvelope(envelope) + return + } + ackEnvelope(envelope) return } - logTemporarySyncTaskDuration(taskContext, startedAt) + logSyncTaskDuration(taskContext, startedAt) // Todo delete soon + ackEnvelope(envelope) +} + +func ackEnvelope(envelope TaskEnvelope) { + if envelope.Handle != nil { + _ = envelope.Handle.Ack() + } +} + +func nackEnvelope(envelope TaskEnvelope) { + if envelope.Handle != nil { + _ = envelope.Handle.Nack() + } +} + +func stopEnvelopeHeartbeat(envelope TaskEnvelope) { + if envelope.stopHeartbeat != nil { + envelope.stopHeartbeat() + } +} + +// startHandleHeartbeat start handle heartbeat +func startHandleHeartbeat(ctx context.Context, handle common.TaskHandle) func() { + if handle == nil { + return func() {} + } + + done := make(chan struct{}) + stopped := make(chan struct{}) + + go func() { + defer close(stopped) + ticker := time.NewTicker(10 * time.Second) + defer ticker.Stop() + + for { + select { + case <-done: + return + case <-ctx.Done(): + return + case <-ticker.C: + _ = handle.InProgress() + } + } + }() + return func() { + close(done) + <-stopped + } }