feat: add river-backed async task queue
This commit is contained in:
@@ -2,13 +2,49 @@ package runner
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/clients"
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||
)
|
||||
|
||||
type localRateLimitError struct {
|
||||
clientErr *clients.ClientError
|
||||
cause error
|
||||
retryAfter time.Duration
|
||||
}
|
||||
|
||||
func (e *localRateLimitError) Error() string {
|
||||
if e == nil || e.clientErr == nil {
|
||||
return store.ErrRateLimited.Error()
|
||||
}
|
||||
return e.clientErr.Error()
|
||||
}
|
||||
|
||||
func (e *localRateLimitError) Unwrap() []error {
|
||||
if e == nil || e.clientErr == nil {
|
||||
if e != nil && e.cause != nil {
|
||||
return []error{e.cause}
|
||||
}
|
||||
return []error{store.ErrRateLimited}
|
||||
}
|
||||
if e.cause != nil {
|
||||
return []error{e.clientErr, e.cause}
|
||||
}
|
||||
return []error{e.clientErr, store.ErrRateLimited}
|
||||
}
|
||||
|
||||
func localRateLimitRetryAfter(err error) time.Duration {
|
||||
var limitErr *localRateLimitError
|
||||
if errors.As(err, &limitErr) && limitErr.retryAfter > 0 {
|
||||
return limitErr.retryAfter
|
||||
}
|
||||
return store.RateLimitRetryAfter(err)
|
||||
}
|
||||
|
||||
func (s *Service) rateLimitReservations(ctx context.Context, user *auth.User, candidate store.RuntimeModelCandidate, body map[string]any) []store.RateLimitReservation {
|
||||
out := make([]store.RateLimitReservation, 0)
|
||||
out = append(out, reservationsFromPolicy("platform_model", candidate.PlatformModelID, effectiveRateLimitPolicy(candidate), body)...)
|
||||
|
||||
@@ -0,0 +1,216 @@
|
||||
package runner
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||
"github.com/riverqueue/river"
|
||||
"github.com/riverqueue/river/riverdriver/riverpgxv5"
|
||||
"github.com/riverqueue/river/rivermigrate"
|
||||
"github.com/riverqueue/river/rivertype"
|
||||
)
|
||||
|
||||
const asyncTaskQueueName = "gateway_tasks"
|
||||
|
||||
type asyncTaskArgs struct {
|
||||
TaskID string `json:"task_id" river:"unique"`
|
||||
}
|
||||
|
||||
func (asyncTaskArgs) Kind() string { return "gateway_task_run" }
|
||||
|
||||
type asyncTaskWorker struct {
|
||||
river.WorkerDefaults[asyncTaskArgs]
|
||||
|
||||
service *Service
|
||||
}
|
||||
|
||||
func (w *asyncTaskWorker) Work(ctx context.Context, job *river.Job[asyncTaskArgs]) error {
|
||||
task, err := w.service.store.GetTask(ctx, job.Args.TaskID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if task.Status == "succeeded" || task.Status == "failed" || task.Status == "cancelled" {
|
||||
return nil
|
||||
}
|
||||
result, runErr := w.service.Execute(ctx, task, authUserFromTask(task))
|
||||
if runErr == nil {
|
||||
w.service.logger.Debug("river async task completed", "taskID", task.ID, "status", result.Task.Status, "riverJobID", job.ID)
|
||||
return nil
|
||||
}
|
||||
var queuedErr *TaskQueuedError
|
||||
if errors.As(runErr, &queuedErr) {
|
||||
return river.JobSnooze(queuedErr.Delay)
|
||||
}
|
||||
if ctx.Err() != nil {
|
||||
queued, queueErr := w.service.requeueInterruptedAsyncTask(context.WithoutCancel(ctx), task)
|
||||
if queueErr != nil {
|
||||
return queueErr
|
||||
}
|
||||
w.service.logger.Debug("river async task interrupted and requeued", "taskID", task.ID, "status", queued.Status, "riverJobID", job.ID)
|
||||
return river.JobSnooze(0)
|
||||
}
|
||||
w.service.logger.Warn("river async task completed with failure", "taskID", task.ID, "error", runErr, "riverJobID", job.ID)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) StartAsyncQueueWorker(ctx context.Context) {
|
||||
if err := s.startRiverQueue(ctx); err != nil {
|
||||
s.logger.Error("start river async queue failed", "error", err)
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) startRiverQueue(ctx context.Context) error {
|
||||
driver := riverpgxv5.New(s.store.Pool())
|
||||
migrator, err := rivermigrate.New(driver, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := migrator.Migrate(ctx, rivermigrate.DirectionUp, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
workers := river.NewWorkers()
|
||||
if err := river.AddWorkerSafely(workers, &asyncTaskWorker{service: s}); err != nil {
|
||||
return err
|
||||
}
|
||||
riverClient, err := river.NewClient(driver, &river.Config{
|
||||
ID: asyncWorkerID(),
|
||||
JobTimeout: -1,
|
||||
Logger: s.logger,
|
||||
CompletedJobRetentionPeriod: 24 * time.Hour,
|
||||
Queues: map[string]river.QueueConfig{
|
||||
asyncTaskQueueName: {MaxWorkers: 32},
|
||||
},
|
||||
RescueStuckJobsAfter: 30 * time.Second,
|
||||
TestOnly: s.cfg.AppEnv == "test",
|
||||
Workers: workers,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.riverClient = riverClient
|
||||
if err := riverClient.Start(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.recoverAsyncRiverJobs(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
go func() {
|
||||
<-ctx.Done()
|
||||
stopCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
if err := riverClient.StopAndCancel(stopCtx); err != nil {
|
||||
s.logger.Warn("stop river async queue failed", "error", err)
|
||||
}
|
||||
}()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) EnqueueAsyncTask(ctx context.Context, task store.GatewayTask) error {
|
||||
if s.riverClient == nil {
|
||||
return errors.New("river async queue is not started")
|
||||
}
|
||||
result, err := s.riverClient.Insert(ctx, asyncTaskArgs{TaskID: task.ID}, asyncTaskInsertOpts(task))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if result.Job != nil {
|
||||
return s.store.SetTaskRiverJobID(ctx, task.ID, result.Job.ID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) WakeAsyncQueueAfter(ctx context.Context, delay time.Duration) {
|
||||
}
|
||||
|
||||
func (s *Service) RunAsyncTask(ctx context.Context, task store.GatewayTask, user *auth.User) {
|
||||
if err := s.EnqueueAsyncTask(ctx, task); err != nil {
|
||||
s.logger.Warn("enqueue river async task failed", "taskID", task.ID, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) recoverAsyncRiverJobs(ctx context.Context) error {
|
||||
items, err := s.store.ListRecoverableAsyncTasks(ctx, 1000)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, item := range items {
|
||||
task := store.GatewayTask{ID: item.ID}
|
||||
result, err := s.riverClient.Insert(ctx, asyncTaskArgs{TaskID: item.ID}, asyncTaskInsertOpts(task))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if result.Job != nil {
|
||||
if err := s.store.SetTaskRiverJobID(ctx, item.ID, result.Job.ID); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(items) > 0 {
|
||||
s.logger.Info("river async queue recovered persisted tasks", "count", len(items))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func asyncTaskInsertOpts(task store.GatewayTask) *river.InsertOpts {
|
||||
priority := 2
|
||||
if task.ID == "" {
|
||||
priority = 3
|
||||
}
|
||||
return &river.InsertOpts{
|
||||
MaxAttempts: 1000,
|
||||
Priority: priority,
|
||||
Queue: asyncTaskQueueName,
|
||||
Tags: []string{"gateway-task"},
|
||||
UniqueOpts: river.UniqueOpts{
|
||||
ByArgs: true,
|
||||
ByQueue: true,
|
||||
ByState: []rivertype.JobState{
|
||||
rivertype.JobStateAvailable,
|
||||
rivertype.JobStatePending,
|
||||
rivertype.JobStateRetryable,
|
||||
rivertype.JobStateRunning,
|
||||
rivertype.JobStateScheduled,
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func authUserFromTask(task store.GatewayTask) *auth.User {
|
||||
roles := []string{"user"}
|
||||
if strings.TrimSpace(task.UserID) == "" {
|
||||
roles = nil
|
||||
}
|
||||
return &auth.User{
|
||||
ID: firstNonEmptyString(task.GatewayUserID, task.UserID),
|
||||
Roles: roles,
|
||||
TenantID: task.TenantID,
|
||||
GatewayTenantID: task.GatewayTenantID,
|
||||
TenantKey: task.TenantKey,
|
||||
Source: firstNonEmptyString(task.UserSource, "gateway"),
|
||||
GatewayUserID: task.GatewayUserID,
|
||||
UserGroupID: task.UserGroupID,
|
||||
UserGroupKey: task.UserGroupKey,
|
||||
APIKeyID: task.APIKeyID,
|
||||
APIKeyName: task.APIKeyName,
|
||||
APIKeyPrefix: task.APIKeyPrefix,
|
||||
}
|
||||
}
|
||||
|
||||
func asyncWorkerID() string {
|
||||
host, _ := os.Hostname()
|
||||
host = strings.TrimSpace(host)
|
||||
if host == "" {
|
||||
host = "localhost"
|
||||
}
|
||||
return fmt.Sprintf("%s:%d:%d", host, os.Getpid(), time.Now().UnixNano())
|
||||
}
|
||||
|
||||
var _ river.Worker[asyncTaskArgs] = (*asyncTaskWorker)(nil)
|
||||
@@ -1,6 +1,7 @@
|
||||
package runner
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
@@ -54,6 +55,9 @@ func shouldRetrySameClient(candidate store.RuntimeModelCandidate, err error) boo
|
||||
func retryDecisionForCandidate(candidate store.RuntimeModelCandidate, err error) retryDecision {
|
||||
policy := effectiveRetryPolicy(candidate)
|
||||
info := failureInfoFromError(err)
|
||||
if errors.Is(err, store.ErrRateLimited) {
|
||||
return retryDecision{Retry: false, Reason: "local_rate_limit_wait_queue", Match: policyRuleMatch{Source: "gateway_rate_limits", Policy: "rateLimitPolicy", Rule: "localCapacity", Value: "exceeded"}, Info: info}
|
||||
}
|
||||
if !boolFromPolicy(policy, "enabled", true) {
|
||||
return retryDecision{Retry: false, Reason: "retry_disabled", Match: policyRuleMatch{Source: "model_runtime_policy_sets.retry_policy", Policy: "retryPolicy", Rule: "enabled", Value: "false"}, Info: info}
|
||||
}
|
||||
@@ -94,6 +98,9 @@ func failoverDecisionForCandidate(runnerPolicy store.RunnerPolicy, candidate sto
|
||||
if cooldownSeconds <= 0 {
|
||||
cooldownSeconds = 300
|
||||
}
|
||||
if errors.Is(err, store.ErrRateLimited) && store.RateLimitRetryable(err) {
|
||||
return failoverDecision{Retry: true, Action: "next", Reason: "local_rate_limit_try_next_candidate", CooldownSeconds: cooldownSeconds, Match: policyRuleMatch{Source: "gateway_rate_limits", Policy: "rateLimitPolicy", Rule: "localCapacity", Value: "exceeded"}, Info: info}
|
||||
}
|
||||
if match, ok := failoverAllowMatchWithSources(runnerPolicy.FailoverPolicy, overridePolicy, info); ok {
|
||||
return failoverDecision{Retry: true, Action: action, Reason: "failover_allow_policy", CooldownSeconds: cooldownSeconds, Match: match, Info: info}
|
||||
}
|
||||
|
||||
@@ -2,6 +2,7 @@ package runner
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/clients"
|
||||
@@ -65,6 +66,9 @@ func (s *Service) applyFailoverAction(ctx context.Context, taskID string, candid
|
||||
}
|
||||
|
||||
func (s *Service) applyPriorityDemotePolicy(ctx context.Context, taskID string, attemptNo int, runnerPolicy store.RunnerPolicy, candidate store.RuntimeModelCandidate, cause error, simulated bool) {
|
||||
if errors.Is(cause, store.ErrRateLimited) {
|
||||
return
|
||||
}
|
||||
decision := priorityDemoteDecisionForCandidate(runnerPolicy, cause)
|
||||
if !decision.Demote {
|
||||
return
|
||||
|
||||
@@ -3,6 +3,7 @@ package runner
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -12,6 +13,8 @@ import (
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/clients"
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/config"
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/riverqueue/river"
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
@@ -20,6 +23,7 @@ type Service struct {
|
||||
logger *slog.Logger
|
||||
clients map[string]clients.Client
|
||||
httpClients *httpClientCache
|
||||
riverClient *river.Client[pgx.Tx]
|
||||
}
|
||||
|
||||
type Result struct {
|
||||
@@ -27,6 +31,20 @@ type Result struct {
|
||||
Output map[string]any
|
||||
}
|
||||
|
||||
var ErrTaskQueued = errors.New("task queued")
|
||||
|
||||
type TaskQueuedError struct {
|
||||
Delay time.Duration
|
||||
}
|
||||
|
||||
func (e *TaskQueuedError) Error() string {
|
||||
return ErrTaskQueued.Error()
|
||||
}
|
||||
|
||||
func (e *TaskQueuedError) Is(target error) bool {
|
||||
return target == ErrTaskQueued
|
||||
}
|
||||
|
||||
func New(cfg config.Config, db *store.Store, logger *slog.Logger) *Service {
|
||||
httpClients := newHTTPClientCache()
|
||||
return &Service{
|
||||
@@ -55,6 +73,14 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
|
||||
executeStartedAt := time.Now()
|
||||
body := normalizeRequest(task.Kind, task.Request)
|
||||
modelType := modelTypeFromKind(task.Kind, body)
|
||||
if err := s.store.MarkTaskRunning(ctx, task.ID, modelType, body); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
if task.Status != "running" {
|
||||
if err := s.emit(ctx, task.ID, "task.running", "running", "starting", 0.12, "task pulled from queue and started", map[string]any{"modelType": modelType}, task.RunMode == "simulation"); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
}
|
||||
if err := validateRequest(task.Kind, body); err != nil {
|
||||
failed, finishErr := s.failTask(ctx, task.ID, "bad_request", err.Error(), task.RunMode == "simulation", err)
|
||||
if finishErr != nil {
|
||||
@@ -83,9 +109,6 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
|
||||
return Result{}, err
|
||||
}
|
||||
}
|
||||
if err := s.store.MarkTaskRunning(ctx, task.ID, modelType, body); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
if err := s.emit(ctx, task.ID, "task.progress", "running", "normalizing", 0.15, "request normalized", map[string]any{"modelType": modelType}, task.RunMode == "simulation"); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
@@ -96,7 +119,7 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
|
||||
}
|
||||
maxPlatforms := maxPlatformsForCandidates(candidates, runnerPolicy)
|
||||
maxFailoverDuration := maxFailoverDurationForCandidates(candidates, runnerPolicy)
|
||||
attemptNo := 0
|
||||
attemptNo := task.AttemptCount
|
||||
var lastErr error
|
||||
for index, candidate := range candidates {
|
||||
if index >= maxPlatforms {
|
||||
@@ -251,6 +274,20 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
|
||||
if lastErr != nil {
|
||||
message = lastErr.Error()
|
||||
}
|
||||
if task.AsyncMode && ctx.Err() != nil {
|
||||
queued, queueErr := s.requeueInterruptedAsyncTask(context.WithoutCancel(ctx), task)
|
||||
if queueErr != nil {
|
||||
return Result{}, queueErr
|
||||
}
|
||||
return Result{Task: queued, Output: queued.Result}, &TaskQueuedError{Delay: 0}
|
||||
}
|
||||
if task.AsyncMode && errors.Is(lastErr, store.ErrRateLimited) && store.RateLimitRetryable(lastErr) {
|
||||
queued, delay, queueErr := s.requeueRateLimitedTask(ctx, task, lastErr)
|
||||
if queueErr != nil {
|
||||
return Result{}, queueErr
|
||||
}
|
||||
return Result{Task: queued, Output: queued.Result}, &TaskQueuedError{Delay: delay}
|
||||
}
|
||||
failed, err := s.failTask(ctx, task.ID, code, message, task.RunMode == "simulation", lastErr)
|
||||
if err != nil {
|
||||
return Result{}, err
|
||||
@@ -261,7 +298,7 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
|
||||
func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user *auth.User, body map[string]any, candidate store.RuntimeModelCandidate, attemptNo int, onDelta clients.StreamDelta) (clients.Response, error) {
|
||||
simulated := isSimulation(task, candidate)
|
||||
if err := s.emit(ctx, task.ID, "task.attempt.started", "running", "submitting", 0.25, "client attempt started", map[string]any{"attempt": attemptNo, "clientId": candidate.ClientID}, simulated); err != nil {
|
||||
return clients.Response{}, err
|
||||
return clients.Response{}, fmt.Errorf("emit attempt started: %w", err)
|
||||
}
|
||||
attemptID, err := s.store.CreateTaskAttempt(ctx, store.CreateTaskAttemptInput{
|
||||
TaskID: task.ID,
|
||||
@@ -276,21 +313,22 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
|
||||
Metrics: attemptMetrics(candidate, attemptNo, simulated),
|
||||
})
|
||||
if err != nil {
|
||||
return clients.Response{}, err
|
||||
return clients.Response{}, fmt.Errorf("create task attempt: %w", err)
|
||||
}
|
||||
reservations := s.rateLimitReservations(ctx, user, candidate, body)
|
||||
limitResult, err := s.store.ReserveRateLimits(ctx, task.ID, attemptID, reservations)
|
||||
if err != nil {
|
||||
clientErr := &clients.ClientError{Code: "rate_limit", Message: err.Error(), Retryable: false}
|
||||
retryable := store.RateLimitRetryable(err)
|
||||
clientErr := &clients.ClientError{Code: "rate_limit", Message: err.Error(), Retryable: retryable}
|
||||
_ = s.store.FinishTaskAttempt(ctx, store.FinishTaskAttemptInput{
|
||||
AttemptID: attemptID,
|
||||
Status: "failed",
|
||||
Retryable: false,
|
||||
Metrics: mergeMetrics(attemptMetrics(candidate, attemptNo, simulated), map[string]any{"error": err.Error(), "retryable": false, "trace": []any{failureTraceEntry(clientErr, false)}}),
|
||||
Retryable: retryable,
|
||||
Metrics: mergeMetrics(attemptMetrics(candidate, attemptNo, simulated), map[string]any{"error": err.Error(), "retryable": retryable, "retryAfterMs": localRateLimitRetryAfter(err).Milliseconds(), "trace": []any{failureTraceEntry(clientErr, retryable)}}),
|
||||
ErrorCode: "rate_limit",
|
||||
ErrorMessage: err.Error(),
|
||||
})
|
||||
return clients.Response{}, clientErr
|
||||
return clients.Response{}, &localRateLimitError{clientErr: clientErr, cause: err, retryAfter: localRateLimitRetryAfter(err)}
|
||||
}
|
||||
rateReservationsFinalized := false
|
||||
defer func() {
|
||||
@@ -301,7 +339,7 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
|
||||
defer s.store.ReleaseConcurrencyLeases(context.WithoutCancel(ctx), limitResult.LeaseIDs)
|
||||
|
||||
if err := s.store.RecordClientAssignment(ctx, candidate); err != nil {
|
||||
return clients.Response{}, err
|
||||
return clients.Response{}, fmt.Errorf("record client assignment: %w", err)
|
||||
}
|
||||
defer s.store.RecordClientRelease(context.WithoutCancel(ctx), candidate.ClientID, "")
|
||||
|
||||
@@ -315,17 +353,25 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
|
||||
ErrorCode: clients.ErrorCode(err),
|
||||
ErrorMessage: err.Error(),
|
||||
})
|
||||
return clients.Response{}, err
|
||||
return clients.Response{}, fmt.Errorf("prepare http client: %w", err)
|
||||
}
|
||||
client := s.clientFor(candidate, simulated)
|
||||
callStartedAt := time.Now()
|
||||
response, err := client.Run(ctx, clients.Request{
|
||||
Kind: task.Kind,
|
||||
ModelType: candidate.ModelType,
|
||||
Model: task.Model,
|
||||
Body: body,
|
||||
Candidate: candidate,
|
||||
HTTPClient: requestHTTPClient,
|
||||
Kind: task.Kind,
|
||||
ModelType: candidate.ModelType,
|
||||
Model: task.Model,
|
||||
Body: body,
|
||||
Candidate: candidate,
|
||||
HTTPClient: requestHTTPClient,
|
||||
RemoteTaskID: task.RemoteTaskID,
|
||||
RemoteTaskPayload: task.RemoteTaskPayload,
|
||||
OnRemoteTaskSubmitted: func(remoteTaskID string, payload map[string]any) error {
|
||||
if strings.TrimSpace(remoteTaskID) == "" {
|
||||
return nil
|
||||
}
|
||||
return s.store.SetTaskRemoteTask(context.WithoutCancel(ctx), task.ID, attemptID, remoteTaskID, payload)
|
||||
},
|
||||
Stream: boolFromMap(body, "stream"),
|
||||
StreamDelta: onDelta,
|
||||
})
|
||||
@@ -400,11 +446,11 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
|
||||
response.Result = uploadedResult
|
||||
for _, progress := range response.Progress {
|
||||
if err := s.emit(ctx, task.ID, "task.progress", "running", progress.Phase, progress.Progress, progress.Message, progress.Payload, simulated); err != nil {
|
||||
return clients.Response{}, err
|
||||
return clients.Response{}, fmt.Errorf("emit task progress: %w", err)
|
||||
}
|
||||
}
|
||||
if err := s.store.CommitRateLimitReservations(ctx, limitResult.Reservations, tokenUsageAmounts(response.Usage)); err != nil {
|
||||
return clients.Response{}, err
|
||||
return clients.Response{}, fmt.Errorf("commit rate limit reservations: %w", err)
|
||||
}
|
||||
rateReservationsFinalized = true
|
||||
if err := s.store.FinishTaskAttempt(ctx, store.FinishTaskAttemptInput{
|
||||
@@ -418,7 +464,7 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
|
||||
ResponseFinishedAt: response.ResponseFinishedAt,
|
||||
ResponseDurationMS: response.ResponseDurationMS,
|
||||
}); err != nil {
|
||||
return clients.Response{}, err
|
||||
return clients.Response{}, fmt.Errorf("finish task attempt: %w", err)
|
||||
}
|
||||
return response, nil
|
||||
}
|
||||
@@ -459,6 +505,41 @@ func (s *Service) failTask(ctx context.Context, taskID string, code string, mess
|
||||
return failed, nil
|
||||
}
|
||||
|
||||
func (s *Service) requeueRateLimitedTask(ctx context.Context, task store.GatewayTask, cause error) (store.GatewayTask, time.Duration, error) {
|
||||
delay := localRateLimitRetryAfter(cause)
|
||||
if delay <= 0 {
|
||||
delay = 5 * time.Second
|
||||
}
|
||||
queued, err := s.store.RequeueTask(ctx, task.ID, delay)
|
||||
if err != nil {
|
||||
return store.GatewayTask{}, 0, err
|
||||
}
|
||||
payload := map[string]any{
|
||||
"code": "rate_limit",
|
||||
"message": cause.Error(),
|
||||
"retryAfterMs": delay.Milliseconds(),
|
||||
}
|
||||
if eventErr := s.emit(ctx, task.ID, "task.queued", "queued", "rate_limited", 0.2, "task queued by local rate limit", payload, task.RunMode == "simulation"); eventErr != nil {
|
||||
return store.GatewayTask{}, 0, eventErr
|
||||
}
|
||||
return queued, delay, nil
|
||||
}
|
||||
|
||||
func (s *Service) requeueInterruptedAsyncTask(ctx context.Context, task store.GatewayTask) (store.GatewayTask, error) {
|
||||
queued, err := s.store.RequeueTask(ctx, task.ID, 0)
|
||||
if err != nil {
|
||||
return store.GatewayTask{}, err
|
||||
}
|
||||
payload := map[string]any{"code": "worker_interrupted"}
|
||||
if task.RemoteTaskID != "" {
|
||||
payload["remoteTaskId"] = task.RemoteTaskID
|
||||
}
|
||||
if eventErr := s.emit(ctx, task.ID, "task.queued", "queued", "worker_interrupted", 0.2, "async task queued after worker interruption", payload, task.RunMode == "simulation"); eventErr != nil {
|
||||
return store.GatewayTask{}, eventErr
|
||||
}
|
||||
return queued, nil
|
||||
}
|
||||
|
||||
func (s *Service) withAttemptHistory(ctx context.Context, taskID string, metrics map[string]any) map[string]any {
|
||||
attempts, err := s.store.ListTaskAttempts(ctx, taskID)
|
||||
if err != nil {
|
||||
|
||||
Reference in New Issue
Block a user