Something went wrong. Try again.
This repository has no description
Something went wrong. Try again.
10 kB · 358 lines
Go
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359package statedb
import ( "context" "encoding/json" "errors" "fmt" "strings" "time"
"gorm.io/gorm")
// TaskStatus represents the status of a task in the queuetype TaskStatus string
const ( TaskStatusPending TaskStatus = "PENDING" TaskStatusProcessing TaskStatus = "PROCESSING" TaskStatusCompleted TaskStatus = "COMPLETED" TaskStatusFailed TaskStatus = "FAILED" TaskStatusRetrying TaskStatus = "RETRYING")
// AppTask represents a task in the queuetype AppTask struct { ID uint `gorm:"column:id;primarykey"` Type string `gorm:"column:type;not null;index"` TaskKey *string `gorm:"column:task_key;index:idx_task_dedup,unique"` Status TaskStatus `gorm:"column:status;not null;index;default:'PENDING'"` Payload json.RawMessage `gorm:"column:payload;type:jsonb"` Priority int `gorm:"column:priority;default:0;index"` TryCount int `gorm:"column:try_count;default:0"` MaxTries int `gorm:"column:max_tries;default:3"` LockExpires *time.Time `gorm:"column:lock_expires"` WorkerID *string `gorm:"column:worker_id"` Error *string `gorm:"column:error"` CreatedAt time.Time `gorm:"column:created_at"` UpdatedAt time.Time `gorm:"column:updated_at"` ScheduledAt *time.Time `gorm:"column:scheduled_at"` // for delayed tasks}
// EnqueueTask adds a new task to the queuefunc (state *StatefulDB) EnqueueTask(ctx context.Context, taskType string, payload any, options ...TaskOption) (*AppTask, error) { payloadBytes, err := json.Marshal(payload) if err != nil { return nil, fmt.Errorf("failed to marshal payload: %w", err) }
task := &AppTask{ Type: taskType, Status: TaskStatusPending, Payload: payloadBytes, Priority: 0, MaxTries: 3, }
// Apply options for _, opt := range options { opt(task) }
// If task has a key, check for deduplication if task.TaskKey != nil { existingTask, err := state.GetTaskByKey(ctx, *task.TaskKey) if err != nil { return nil, fmt.Errorf("failed to check for existing task: %w", err) } if existingTask != nil { // Task already exists, return the existing one return existingTask, nil } }
if err := state.DB.WithContext(ctx).Create(task).Error; err != nil { // Handle unique constraint violation gracefully if strings.Contains(err.Error(), "duplicate") || strings.Contains(err.Error(), "UNIQUE constraint") { // Another node beat us to it, try to fetch the existing task if task.TaskKey != nil { existingTask, fetchErr := state.GetTaskByKey(ctx, *task.TaskKey) if fetchErr == nil && existingTask != nil { return existingTask, nil } } } return nil, fmt.Errorf("failed to enqueue task: %w", err) }
go func() { select { case state.pokeQueue <- struct{}{}: // wake up the queue processor default: // queue is already awake, do nothing } }()
return task, nil}
// DequeueTask retrieves the next available task from the queue and locks itfunc (state *StatefulDB) DequeueTask(ctx context.Context, workerID string, taskTypes ...string) (*AppTask, error) { var task AppTask
err := state.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error { query := tx.Where("status = ?", TaskStatusPending). Where("try_count < max_tries"). Where("(lock_expires IS NULL OR lock_expires < ?)", time.Now().UTC()). Where("(scheduled_at IS NULL OR scheduled_at <= ?)", time.Now().UTC())
if len(taskTypes) > 0 { query = query.Where("type IN ?", taskTypes) }
// Use raw SQL for PostgreSQL-specific locking if state.Type == DBTypePostgres { baseQuery := "SELECT * FROM app_tasks WHERE status = ? AND try_count < max_tries AND (lock_expires IS NULL OR lock_expires < ?) AND (scheduled_at IS NULL OR scheduled_at <= ?)" if len(taskTypes) > 0 { baseQuery += " AND type IN ?" params := []interface{}{TaskStatusPending, time.Now(), time.Now(), taskTypes} err := tx.Raw(baseQuery+" ORDER BY priority DESC, created_at ASC LIMIT 1 FOR UPDATE SKIP LOCKED", params...). Scan(&task).Error if err != nil { return err } } else { err := tx.Raw(baseQuery+" ORDER BY priority DESC, created_at ASC LIMIT 1 FOR UPDATE SKIP LOCKED", TaskStatusPending, time.Now(), time.Now()). Scan(&task).Error if err != nil { return err } } } else { // Fallback for SQLite (no SKIP LOCKED support) err := query.Order("priority DESC, created_at ASC").First(&task).Error if err != nil { return err } }
if task.ID == 0 { return gorm.ErrRecordNotFound }
// Lock the task lockExpires := time.Now().Add(30 * time.Minute) // 30-minute lock updates := map[string]interface{}{ "status": TaskStatusProcessing, "worker_id": workerID, "lock_expires": lockExpires, "try_count": task.TryCount + 1, }
return tx.Model(&task).Updates(updates).Error })
if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, nil // No tasks available } return nil, fmt.Errorf("failed to dequeue task: %w", err) }
// Reload the task to get updated fields if err := state.DB.WithContext(ctx).First(&task, task.ID).Error; err != nil { return nil, fmt.Errorf("failed to reload task: %w", err) }
return &task, nil}
// CompleteTask marks a task as completedfunc (state *StatefulDB) CompleteTask(ctx context.Context, taskID uint) error { result := state.DB.WithContext(ctx).Model(&AppTask{}). Where("id = ?", taskID). Updates(map[string]interface{}{ "status": TaskStatusCompleted, "lock_expires": nil, "worker_id": nil, })
if result.Error != nil { return fmt.Errorf("failed to complete task: %w", result.Error) }
if result.RowsAffected == 0 { return errors.New("task not found") }
return nil}
// FailTask marks a task as failed and optionally retries itfunc (state *StatefulDB) FailTask(ctx context.Context, taskID uint, errorMsg string) error { var task AppTask err := state.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error { if err := tx.First(&task, taskID).Error; err != nil { return err }
updates := map[string]interface{}{ "error": errorMsg, "lock_expires": nil, "worker_id": nil, }
if task.TryCount >= task.MaxTries { updates["status"] = TaskStatusFailed } else { updates["status"] = TaskStatusPending }
return tx.Model(&task).Updates(updates).Error })
if err != nil { return fmt.Errorf("failed to mark task as failed: %w", err) }
return nil}
// ReleaseTask releases a locked task back to the queue (e.g., worker shutdown)func (state *StatefulDB) ReleaseTask(ctx context.Context, taskID uint) error { result := state.DB.WithContext(ctx).Model(&AppTask{}). Where("id = ?", taskID). Updates(map[string]interface{}{ "status": TaskStatusPending, "lock_expires": nil, "worker_id": nil, })
if result.Error != nil { return fmt.Errorf("failed to release task: %w", result.Error) }
if result.RowsAffected == 0 { return errors.New("task not found") }
return nil}
// GetTask retrieves a task by IDfunc (state *StatefulDB) GetTask(ctx context.Context, taskID uint) (*AppTask, error) { var task AppTask if err := state.DB.WithContext(ctx).First(&task, taskID).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, nil } return nil, fmt.Errorf("failed to get task: %w", err) } return &task, nil}
// GetTaskByKey retrieves a task by its unique task keyfunc (state *StatefulDB) GetTaskByKey(ctx context.Context, taskKey string) (*AppTask, error) { var task AppTask if err := state.DB.WithContext(ctx).Where("task_key = ?", taskKey).First(&task).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, nil } return nil, fmt.Errorf("failed to get task by key: %w", err) } return &task, nil}
// ListTasks retrieves tasks with optional filtersfunc (state *StatefulDB) ListTasks(ctx context.Context, filters TaskFilters) ([]AppTask, error) { var tasks []AppTask query := state.DB.WithContext(ctx).Model(&AppTask{})
if filters.Status != "" { query = query.Where("status = ?", filters.Status) } if filters.Type != "" { query = query.Where("type = ?", filters.Type) } if filters.TaskKey != "" { query = query.Where("task_key = ?", filters.TaskKey) } if filters.WorkerID != "" { query = query.Where("worker_id = ?", filters.WorkerID) } if filters.Limit > 0 { query = query.Limit(filters.Limit) } if filters.Offset > 0 { query = query.Offset(filters.Offset) }
query = query.Order("created_at DESC")
if err := query.Find(&tasks).Error; err != nil { return nil, fmt.Errorf("failed to list tasks: %w", err) }
return tasks, nil}
// CleanupExpiredLocks releases tasks with expired locksfunc (state *StatefulDB) CleanupExpiredLocks(ctx context.Context) (int64, error) { result := state.DB.WithContext(ctx).Model(&AppTask{}). Where("status = ? AND lock_expires < ?", TaskStatusProcessing, time.Now()). Updates(map[string]interface{}{ "status": TaskStatusPending, "lock_expires": nil, "worker_id": nil, })
if result.Error != nil { return 0, fmt.Errorf("failed to cleanup expired locks: %w", result.Error) }
return result.RowsAffected, nil}
// TaskOption is a function that configures a tasktype TaskOption func(*AppTask)
// WithPriority sets the task priority (higher numbers = higher priority)func WithPriority(priority int) TaskOption { return func(t *AppTask) { t.Priority = priority }}
// WithMaxTries sets the maximum number of retry attemptsfunc WithMaxTries(maxTries int) TaskOption { return func(t *AppTask) { t.MaxTries = maxTries }}
// WithScheduledAt sets when the task should be processed (for delayed tasks)func WithScheduledAt(scheduledAt time.Time) TaskOption { return func(t *AppTask) { t.ScheduledAt = &scheduledAt }}
// WithTaskKey sets a unique key for task deduplicationfunc WithTaskKey(taskKey string) TaskOption { return func(t *AppTask) { t.TaskKey = &taskKey }}
// TaskFilters holds filters for listing taskstype TaskFilters struct { Status TaskStatus Type string TaskKey string WorkerID string Limit int Offset int}