Files
TaskPool/internal/services/cron_service.go
T
2025-12-20 09:30:16 +08:00

151 lines
3.6 KiB
Go

package services
import (
"sync"
"time"
"baihu/internal/database"
"baihu/internal/logger"
"baihu/internal/models"
"github.com/robfig/cron/v3"
)
// CronService manages scheduled tasks using robfig/cron
type CronService struct {
cron *cron.Cron
taskService *TaskService
executorService *ExecutorService
entryMap map[uint]cron.EntryID // task ID -> cron entry ID
mu sync.RWMutex
}
// NewCronService creates a new cron service
func NewCronService(taskService *TaskService, executorService *ExecutorService) *CronService {
// 使用秒级精度的 cron parser,支持 5 位和 6 位表达式
c := cron.New(cron.WithParser(cron.NewParser(
cron.Minute | cron.Hour | cron.Dom | cron.Month | cron.Dow | cron.Descriptor,
)))
return &CronService{
cron: c,
taskService: taskService,
executorService: executorService,
entryMap: make(map[uint]cron.EntryID),
}
}
// Start starts the cron service and loads all enabled tasks
func (cs *CronService) Start() {
cs.loadTasks()
cs.cron.Start()
logger.Info("Cron service started")
}
// Stop stops the cron service
func (cs *CronService) Stop() {
ctx := cs.cron.Stop()
<-ctx.Done()
logger.Info("Cron service stopped")
}
// loadTasks loads all enabled tasks from database
func (cs *CronService) loadTasks() {
tasks := cs.taskService.GetTasks()
for _, task := range tasks {
if task.Enabled {
err := cs.AddTask(&task)
if err != nil {
return
}
}
}
}
// AddTask adds a task to the cron scheduler
func (cs *CronService) AddTask(task *models.Task) error {
cs.mu.Lock()
// 如果已存在,先移除
if entryID, exists := cs.entryMap[task.ID]; exists {
cs.cron.Remove(entryID)
delete(cs.entryMap, task.ID)
}
taskID := task.ID
entryID, err := cs.cron.AddFunc(task.Schedule, func() {
cs.runTask(taskID)
})
if err != nil {
cs.mu.Unlock()
logger.Errorf("Failed to add task %d: %v", task.ID, err)
return err
}
cs.entryMap[task.ID] = entryID
cs.mu.Unlock()
logger.Infof("Task %d (%s) scheduled with cron: %s", task.ID, task.Name, task.Schedule)
// 更新下次运行时间
cs.updateNextRun(task.ID)
return nil
}
// RemoveTask removes a task from the cron scheduler
func (cs *CronService) RemoveTask(taskID uint) {
cs.mu.Lock()
defer cs.mu.Unlock()
if entryID, exists := cs.entryMap[taskID]; exists {
cs.cron.Remove(entryID)
delete(cs.entryMap, taskID)
logger.Infof("Task %d removed from scheduler", taskID)
}
}
// runTask executes a task and updates its status
func (cs *CronService) runTask(taskID uint) {
logger.Infof("Running task %d", taskID)
// 更新 last_run
now := time.Now()
database.DB.Model(&models.Task{}).Where("id = ?", taskID).Update("last_run", now)
// 执行任务
cs.executorService.ExecuteTask(int(taskID))
// 更新 next_run
cs.updateNextRun(taskID)
}
// updateNextRun updates the next run time for a task
func (cs *CronService) updateNextRun(taskID uint) {
cs.mu.RLock()
entryID, exists := cs.entryMap[taskID]
cs.mu.RUnlock()
if !exists {
return
}
entry := cs.cron.Entry(entryID)
if !entry.Next.IsZero() {
database.DB.Model(&models.Task{}).Where("id = ?", taskID).Update("next_run", entry.Next)
}
}
// ValidateCron validates a cron expression
func (cs *CronService) ValidateCron(expression string) error {
parser := cron.NewParser(cron.Minute | cron.Hour | cron.Dom | cron.Month | cron.Dow | cron.Descriptor)
_, err := parser.Parse(expression)
return err
}
// GetScheduledCount returns the number of scheduled tasks
func (cs *CronService) GetScheduledCount() int {
cs.mu.RLock()
defer cs.mu.RUnlock()
return len(cs.entryMap)
}