package tasks import ( "bytes" "context" "database/sql" "encoding/json" "fmt" "log/slog" "net/http" "runtime/debug" "strings" "time" "atlas9.dev/c/core" "atlas9.dev/c/core/dbi" "atlas9.dev/c/core/iam" "atlas9.dev/c/demo/lib" "atlas9.dev/c/demo/lib/access" ) type Worker struct { DB *sql.DB Consumer func(dbi.DBI) Consumer Handler http.Handler Route string Access access.Access Interval time.Duration LeaseDuration time.Duration MaxAttempts int BatchSize int } func (w *Worker) Run(ctx context.Context) { for { if ctx.Err() != nil { return } processed := w.poll(ctx) if !processed { select { case <-time.After(w.interval()): case <-ctx.Done(): return } } } } func (w *Worker) interval() time.Duration { if w.Interval > 0 { return w.Interval } return 5 * time.Second } func (w *Worker) batchSize() int { if w.BatchSize > 0 { return w.BatchSize } return 1 } func (w *Worker) leaseDuration() time.Duration { if w.LeaseDuration > 0 { return w.LeaseDuration } return 30 * time.Second } func (w *Worker) poll(ctx context.Context) bool { var claimed []Task err := dbi.ReadWrite(ctx, w.DB, func(tx dbi.DBI) error { var err error claimed, err = w.Consumer(tx).Claim(ctx, w.leaseDuration(), w.batchSize()) return err }) if err != nil { slog.Error("claiming tasks", "err", err) return false } if len(claimed) == 0 { return false } for _, task := range claimed { res := w.execute(ctx, task) if res.Outcome == OutcomeRetry && w.MaxAttempts > 0 && task.Attempts >= w.MaxAttempts { res = Result{Outcome: OutcomeFailed, Error: fmt.Sprintf("max attempts reached: %s", res.Error)} } switch res.Outcome { case OutcomeCompleted: err := dbi.ReadWrite(ctx, w.DB, func(tx dbi.DBI) error { return w.Consumer(tx).Complete(ctx, task.TaskID) }) if err != nil { slog.Error("completing task", "task_id", task.TaskID, "err", err) } case OutcomeFailed: slog.Error("task failed permanently", "task_id", task.TaskID, "attempts", task.Attempts, "err", res.Error) err := dbi.ReadWrite(ctx, w.DB, func(tx dbi.DBI) error { return w.Consumer(tx).Fail(ctx, task.TaskID) }) if err != nil { slog.Error("failing task", "task_id", task.TaskID, "err", err) } default: // OutcomeRetry slog.Error("task will retry", "task_id", task.TaskID, "attempts", task.Attempts, "err", res.Error) retryAfter := time.Now().Add(backoff(task.Attempts)) err := dbi.ReadWrite(ctx, w.DB, func(tx dbi.DBI) error { return w.Consumer(tx).Retry(ctx, task.TaskID, retryAfter) }) if err != nil { slog.Error("retrying task", "task_id", task.TaskID, "err", err) } } } return true } func (w *Worker) execute(ctx context.Context, task Task) (res Result) { method, path, _ := strings.Cut(w.Route, " ") // Each execution is its own trace, linked to the request that enqueued // the task via the parent attribute. Task requests bypass the HTTP // middleware, so the worker emits the request/response envelope events // itself. ctx = lib.PutRequestID(ctx, core.NewID("req").String()) ctx = lib.PutTaskID(ctx, task.TaskID) start := time.Now() slog.InfoContext(ctx, "http request", "method", method, "path", path, "task_id", task.TaskID, "attempt", task.Attempts, "parent", task.Origin, ) req, err := http.NewRequestWithContext(ctx, method, path, bytes.NewReader(task.Payload)) if err != nil { return Result{Outcome: OutcomeFailed, Error: fmt.Sprintf("building request: %v", err)} } req.Header.Set("Content-Type", "application/json") // Grant the configured capabilities to the request, and identify it as // the worker so it passes RequireAuth on the way to the task endpoint. ctx = access.Put(req.Context(), w.Access) ctx = iam.PutPrincipal(ctx, iam.Principal{Subject: "task-worker"}) req = req.WithContext(ctx) rec := &responseWriter{status: http.StatusOK} defer func() { if r := recover(); r != nil { res = Result{Outcome: OutcomeRetry, Error: fmt.Sprintf("task handler panicked: %v\n%s", r, debug.Stack())} } }() w.Handler.ServeHTTP(rec, req) res = outcome(rec.status, rec.body) logFn := slog.InfoContext if res.Outcome != OutcomeCompleted { logFn = slog.WarnContext } if res.Outcome != OutcomeCompleted { logFn(ctx, "http response", "status", rec.status, "written", rec.written, "duration", time.Since(start), "outcome", res.Outcome, "err", res.Error) } else { logFn(ctx, "http response", "status", rec.status, "written", rec.written, "duration", time.Since(start), "outcome", res.Outcome) } return res } // outcome maps a task endpoint's response to a Result. Responses carrying a // Task envelope (see Result) decide explicitly; otherwise HTTP status codes // decide: 2xx completed, 429/5xx retry, other 4xx failed (a deterministic // rejection that retrying cannot fix). func outcome(status int, body []byte) Result { var env struct{ Task Result } if json.Unmarshal(body, &env) == nil && env.Task.Outcome != OutcomeUnknown { return env.Task } switch { case status < 400: return Result{Outcome: OutcomeCompleted} case status == http.StatusTooManyRequests || status >= 500: return Result{Outcome: OutcomeRetry, Error: fmt.Sprintf("task handler returned %d: %s", status, bodyExcerpt(body))} default: return Result{Outcome: OutcomeFailed, Error: fmt.Sprintf("task handler returned %d: %s", status, bodyExcerpt(body))} } } // responseWriter is a minimal http.ResponseWriter that captures the status // code and body of a task endpoint's response. type responseWriter struct { status int written int64 body []byte } // maxBodyCapture bounds how much of a task response body is retained for // envelope parsing. const maxBodyCapture = 1 << 20 func (w *responseWriter) Header() http.Header { return http.Header{} } func (w *responseWriter) Write(b []byte) (int, error) { w.written += int64(len(b)) if len(w.body) < maxBodyCapture { w.body = append(w.body, b[:min(len(b), maxBodyCapture-len(w.body))]...) } return len(b), nil } func (w *responseWriter) WriteHeader(statusCode int) { w.status = statusCode } func backoff(attempts int) time.Duration { d := time.Duration(1< max { b = b[:max] } return strings.TrimSpace(string(b)) }