diff --git a/backend/cmd/api/main.go b/backend/cmd/api/main.go index 28bbb26..4c149ff 100644 --- a/backend/cmd/api/main.go +++ b/backend/cmd/api/main.go @@ -1,9 +1,11 @@ package main import ( + "context" "log" "senlinai-agent/backend/internal/config" + "senlinai-agent/backend/internal/crontab" "senlinai-agent/backend/internal/httpx" "senlinai-agent/backend/internal/initdb" "senlinai-agent/backend/internal/logic/ai" @@ -29,7 +31,9 @@ func main() { } authService := auth.NewService(cfg.AuthSecret) - datasetHandler := dataset.NewHandler(dataset.NewService(models.DBService)) + datasetService := dataset.NewService(models.DBService) + datasetHandler := dataset.NewHandler(datasetService) + datasetScheduler := crontab.NewDatasetScheduler(datasetService, log.Default()) projectService := projects.NewService() fileService := files.NewService(cfg.StorageDir, models.DBService) taskService := tasks.NewService(models.DBService) @@ -57,6 +61,10 @@ func main() { searchHandler, aiHandler, ) + schedulerContext, stopScheduler := context.WithCancel(context.Background()) + defer stopScheduler() + go datasetScheduler.Run(schedulerContext) + if err := appRouter.Run(":" + cfg.Port); err != nil { log.Fatal(err) } diff --git a/backend/internal/crontab/dataset.go b/backend/internal/crontab/dataset.go new file mode 100644 index 0000000..4e3444a --- /dev/null +++ b/backend/internal/crontab/dataset.go @@ -0,0 +1,89 @@ +package crontab + +import ( + "context" + "log" + "sync" + "sync/atomic" + "time" + + "senlinai-agent/backend/internal/models" +) + +const datasetInterval = 5 * time.Minute + +type DatasetCollector interface { + SyncAllSources(context.Context) ([]models.SaDatasetCron, error) +} + +type DatasetScheduler struct { + collector DatasetCollector + interval time.Duration + logger *log.Logger +} + +func NewDatasetScheduler(collector DatasetCollector, logger *log.Logger) *DatasetScheduler { + if logger == nil { + logger = log.Default() + } + return &DatasetScheduler{ + collector: collector, + interval: datasetInterval, + logger: logger, + } +} + +// Run collects once at startup, then on a fixed five-minute interval. +// A tick is skipped when the previous collection is still running. +func (s *DatasetScheduler) Run(ctx context.Context) { + if ctx.Err() != nil { + return + } + ticker := time.NewTicker(s.interval) + defer ticker.Stop() + + var running atomic.Bool + var workers sync.WaitGroup + start := func() { + if !running.CompareAndSwap(false, true) { + s.logger.Print("dataset collection skipped: previous run is still active") + return + } + workers.Add(1) + go func() { + defer workers.Done() + defer running.Store(false) + s.collect(ctx) + }() + } + start() + + for { + select { + case <-ctx.Done(): + workers.Wait() + return + case <-ticker.C: + start() + } + } +} + +func (s *DatasetScheduler) collect(ctx context.Context) { + crons, err := s.collector.SyncAllSources(ctx) + completed := 0 + failed := 0 + for _, cron := range crons { + switch cron.Status { + case "completed": + completed++ + case "failed": + failed++ + } + } + if err != nil { + s.logger.Printf("dataset collection failed: tasks=%d completed=%d failed=%d error=%v", len(crons), completed, failed, err) + return + } + s.logger.Printf("dataset collection finished: tasks=%d completed=%d failed=%d", len(crons), completed, failed) +} diff --git a/backend/internal/crontab/dataset_test.go b/backend/internal/crontab/dataset_test.go new file mode 100644 index 0000000..243b164 --- /dev/null +++ b/backend/internal/crontab/dataset_test.go @@ -0,0 +1,150 @@ +package crontab + +import ( + "context" + "io" + "log" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + "senlinai-agent/backend/internal/models" +) + +func TestDatasetSchedulerCollectsImmediatelyAndRepeats(t *testing.T) { + collector := &recordingDatasetCollector{calls: make(chan time.Time, 3)} + scheduler := &DatasetScheduler{ + collector: collector, + interval: 10 * time.Millisecond, + logger: log.New(io.Discard, "", 0), + } + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + defer close(done) + scheduler.Run(ctx) + }() + + first := receiveCollectionCall(t, collector.calls) + second := receiveCollectionCall(t, collector.calls) + require.GreaterOrEqual(t, second.Sub(first), 8*time.Millisecond) + + cancel() + require.Eventually(t, func() bool { + select { + case <-done: + return true + default: + return false + } + }, time.Second, 5*time.Millisecond) +} + +func TestDatasetSchedulerStopsWhenContextIsAlreadyCanceled(t *testing.T) { + collector := &recordingDatasetCollector{calls: make(chan time.Time, 1)} + scheduler := &DatasetScheduler{ + collector: collector, + interval: 10 * time.Millisecond, + logger: log.New(io.Discard, "", 0), + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + scheduler.Run(ctx) + + require.Zero(t, collector.callCount()) +} + +func TestDatasetSchedulerSkipsTicksWhileCollectionIsRunning(t *testing.T) { + collector := &blockingDatasetCollector{ + started: make(chan struct{}, 2), + release: make(chan struct{}), + } + scheduler := &DatasetScheduler{ + collector: collector, + interval: 5 * time.Millisecond, + logger: log.New(io.Discard, "", 0), + } + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + defer close(done) + scheduler.Run(ctx) + }() + + select { + case <-collector.started: + case <-time.After(time.Second): + t.Fatal("timed out waiting for first collection") + } + time.Sleep(20 * time.Millisecond) + require.Equal(t, 1, collector.callCount()) + + close(collector.release) + require.Eventually(t, func() bool { + return collector.callCount() >= 2 + }, time.Second, 5*time.Millisecond) + cancel() + <-done +} + +type recordingDatasetCollector struct { + mu sync.Mutex + count int + calls chan time.Time +} + +func (c *recordingDatasetCollector) SyncAllSources(context.Context) ([]models.SaDatasetCron, error) { + now := time.Now() + c.mu.Lock() + c.count++ + c.mu.Unlock() + c.calls <- now + return []models.SaDatasetCron{{Status: "completed"}}, nil +} + +func (c *recordingDatasetCollector) callCount() int { + c.mu.Lock() + defer c.mu.Unlock() + return c.count +} + +func receiveCollectionCall(t *testing.T, calls <-chan time.Time) time.Time { + t.Helper() + select { + case calledAt := <-calls: + return calledAt + case <-time.After(time.Second): + t.Fatal("timed out waiting for dataset collection") + return time.Time{} + } +} + +type blockingDatasetCollector struct { + mu sync.Mutex + count int + started chan struct{} + release chan struct{} +} + +func (c *blockingDatasetCollector) SyncAllSources(ctx context.Context) ([]models.SaDatasetCron, error) { + c.mu.Lock() + c.count++ + c.mu.Unlock() + select { + case c.started <- struct{}{}: + default: + } + select { + case <-c.release: + case <-ctx.Done(): + } + return nil, nil +} + +func (c *blockingDatasetCollector) callCount() int { + c.mu.Lock() + defer c.mu.Unlock() + return c.count +} diff --git a/backend/internal/logic/dataset/feed_test.go b/backend/internal/logic/dataset/feed_test.go index 7e72ec2..ccbf134 100644 --- a/backend/internal/logic/dataset/feed_test.go +++ b/backend/internal/logic/dataset/feed_test.go @@ -119,6 +119,27 @@ func TestSyncSourcesRecordsFeedFailure(t *testing.T) { require.Contains(t, crons[0].LastResult, "feed unavailable") } +func TestSyncAllSourcesCollectsEnabledRSSSourcesForEveryUser(t *testing.T) { + database := newDatasetTestDatabase(t) + firstUser := createDatasetTestUser(t, database, "first-feed-owner@example.com") + secondUser := createDatasetTestUser(t, database, "second-feed-owner@example.com") + for _, source := range []models.SaDatasetSource{ + {CreatedBy: firstUser.ID, Name: "First RSS", Kind: "rss", URL: "https://example.com/first", Enabled: true}, + {CreatedBy: secondUser.ID, Name: "Second RSS", Kind: "rss", URL: "https://example.com/second", Enabled: true}, + {CreatedBy: firstUser.ID, Name: "Manual", Kind: "manual", Enabled: true}, + } { + require.NoError(t, database.Create(&source).Error) + } + + crons, err := newServiceWithFetcher(database, staticFeedFetcher{}).SyncAllSources(context.Background()) + + require.NoError(t, err) + require.Len(t, crons, 2) + var itemCount int64 + require.NoError(t, database.Model(&models.SaDatasetItem{}).Count(&itemCount).Error) + require.Equal(t, int64(2), itemCount) +} + type failingFeedFetcher struct{} func (failingFeedFetcher) Fetch(context.Context, string) (ParsedFeed, error) { diff --git a/backend/internal/logic/dataset/service.go b/backend/internal/logic/dataset/service.go index 2a84415..52ba3c8 100644 --- a/backend/internal/logic/dataset/service.go +++ b/backend/internal/logic/dataset/service.go @@ -6,6 +6,7 @@ import ( "fmt" "net/url" "strings" + "sync" "time" "gorm.io/gorm" @@ -23,6 +24,7 @@ var ( type Service struct { db *gorm.DB fetcher FeedFetcher + syncMu sync.Mutex } type SourceInput struct { @@ -216,6 +218,38 @@ func (s *Service) ListCrons(userID uint) ([]models.SaDatasetCron, error) { } func (s *Service) SyncSources(ctx context.Context, userID uint) ([]models.SaDatasetCron, error) { + s.syncMu.Lock() + defer s.syncMu.Unlock() + + return s.syncSources(ctx, userID) +} + +func (s *Service) SyncAllSources(ctx context.Context) ([]models.SaDatasetCron, error) { + var userIDs []uint + if err := s.db.Model(&models.SaDatasetSource{}). + Distinct("created_by"). + Where("enabled = ? AND kind = ?", true, "rss"). + Order("created_by asc"). + Pluck("created_by", &userIDs).Error; err != nil { + return nil, err + } + + result := make([]models.SaDatasetCron, 0) + var syncErrors []error + for _, userID := range userIDs { + crons, err := s.SyncSources(ctx, userID) + result = append(result, crons...) + if err != nil { + syncErrors = append(syncErrors, fmt.Errorf("sync dataset sources for user %d: %w", userID, err)) + } + if ctx.Err() != nil { + break + } + } + return result, errors.Join(syncErrors...) +} + +func (s *Service) syncSources(ctx context.Context, userID uint) ([]models.SaDatasetCron, error) { result := make([]models.SaDatasetCron, 0) var sources []models.SaDatasetSource if err := s.db.Where("created_by = ? AND enabled = ? AND kind = ?", userID, true, "rss"). @@ -223,6 +257,9 @@ func (s *Service) SyncSources(ctx context.Context, userID uint) ([]models.SaData return nil, err } for _, source := range sources { + if err := ctx.Err(); err != nil { + return result, err + } now := time.Now().UTC() cron := models.SaDatasetCron{ SourceID: source.ID, CreatedBy: userID, Schedule: "@once", @@ -244,6 +281,9 @@ func (s *Service) SyncSources(ctx context.Context, userID uint) ([]models.SaData return nil, err } result = append(result, cron) + if err := ctx.Err(); err != nil { + return result, err + } continue }