151 lines
3.6 KiB
Go
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)
|
|
}
|