diff --git a/go.mod b/go.mod index 9f94c43..33e1445 100644 --- a/go.mod +++ b/go.mod @@ -6,6 +6,7 @@ require ( github.com/disintegration/imaging v1.6.2 github.com/go-sql-driver/mysql v1.10.0 golang.org/x/crypto v0.55.0 + gorm.io/driver/mysql v1.6.0 gorm.io/gorm v1.31.2 ) diff --git a/go.sum b/go.sum index f67c1ef..15a246d 100644 --- a/go.sum +++ b/go.sum @@ -8,6 +8,8 @@ github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkryuEj+Srlc= github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ= github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8= +github.com/mattn/go-sqlite3 v1.14.22 h1:2gZY6PC6kBnID23Tichd1K+Z0oS6nE/XwU+Vz/5o4kU= +github.com/mattn/go-sqlite3 v1.14.22/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y= golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M= golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis= golang.org/x/image v0.0.0-20191009234506-e7c1f5e7dbb8/go.mod h1:FeLwcggjj3mMvU+oOTbSwawSJRM1uh48EjtB4UJZlP0= @@ -16,5 +18,9 @@ golang.org/x/image v0.41.0/go.mod h1:uIc348UZMSvS5Z65CVZ7iDPaNobNFEPeJ4kbqTOszmA golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= +gorm.io/driver/mysql v1.6.0 h1:eNbLmNTpPpTOVZi8MMxCi2aaIm0ZpInbORNXDwyLGvg= +gorm.io/driver/mysql v1.6.0/go.mod h1:D/oCC2GWK3M/dqoLxnOlaNKmXz8WNTfcS9y5ovaSqKo= +gorm.io/driver/sqlite v1.6.0 h1:WHRRrIiulaPiPFmDcod6prc4l2VGVWHz80KspNsxSfQ= +gorm.io/driver/sqlite v1.6.0/go.mod h1:AO9V1qIQddBESngQUKWL9yoH93HIeA1X6V633rBwyT8= gorm.io/gorm v1.31.2 h1:3o8FXNo9v9S858gil+3LlZA1LkCOzgb4g5BL64FgaCo= gorm.io/gorm v1.31.2/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs= diff --git a/internal/core/model/models.go b/internal/core/model/models.go index 89f20bb..a157e56 100644 --- a/internal/core/model/models.go +++ b/internal/core/model/models.go @@ -120,10 +120,10 @@ type GenerationOutput struct { func (GenerationOutput) TableName() string { return "generation_outputs" } type Attempt struct { - ProviderModelID uint64 `json:"provider_model_id"` - ErrorCode string `json:"error_code,omitempty"` - ErrorMessage string `json:"error_message,omitempty"` - LatencyMS int64 `json:"latency_ms"` - StartedAt time.Time `json:"started_at"` - FinishedAt time.Time `json:"finished_at"` + ProviderModelID uint64 `json:"provider_model_id,omitempty"` + ErrorCode string `json:"error_code,omitempty"` + ErrorMessage string `json:"error_message,omitempty"` + LatencyMS int64 `json:"latency_ms,omitempty"` + StartedAt *time.Time `json:"started_at,omitempty"` + FinishedAt *time.Time `json:"finished_at,omitempty"` } diff --git a/internal/core/queue/controller.go b/internal/core/queue/controller.go new file mode 100644 index 0000000..4c39c14 --- /dev/null +++ b/internal/core/queue/controller.go @@ -0,0 +1,35 @@ +package queue + +import ( + "context" + "errors" + "sync" + "time" +) + +var ErrClaimsStopped = errors.New("queue claims are stopped") + +type Controller struct { + repository Repository + mutex sync.RWMutex + stopped bool +} + +func NewController(repository Repository) *Controller { + return &Controller{repository: repository} +} + +func (c *Controller) StopClaims() { + c.mutex.Lock() + defer c.mutex.Unlock() + c.stopped = true +} + +func (c *Controller) ClaimNext(ctx context.Context, owner string, leaseDuration time.Duration) (*Claim, error) { + c.mutex.RLock() + defer c.mutex.RUnlock() + if c.stopped { + return nil, ErrClaimsStopped + } + return c.repository.ClaimNext(ctx, owner, leaseDuration) +} diff --git a/internal/core/queue/controller_test.go b/internal/core/queue/controller_test.go new file mode 100644 index 0000000..1fdeafa --- /dev/null +++ b/internal/core/queue/controller_test.go @@ -0,0 +1,45 @@ +package queue + +import ( + "context" + "testing" + "time" + + "git.ilapage.cn/OPC/chorus/internal/core/model" +) + +type fakeRepository struct { + claims int +} + +func (f *fakeRepository) CreateIdempotent(context.Context, *model.Generation, []model.GenerationInput) (bool, error) { + return false, nil +} +func (f *fakeRepository) ClaimNext(context.Context, string, time.Duration) (*Claim, error) { + f.claims++ + return nil, nil +} +func (f *fakeRepository) AssignProvider(context.Context, uint64, string, uint64) (bool, error) { + return false, nil +} +func (f *fakeRepository) Succeed(context.Context, uint64, string, []model.GenerationOutput, model.Attempt) (bool, error) { + return false, nil +} +func (f *fakeRepository) Fail(context.Context, uint64, string, string, string, model.Attempt) (bool, error) { + return false, nil +} + +func TestControllerStopsNewClaims(t *testing.T) { + repository := &fakeRepository{} + controller := NewController(repository) + if _, err := controller.ClaimNext(context.Background(), "worker", time.Second); err != nil { + t.Fatal(err) + } + controller.StopClaims() + if _, err := controller.ClaimNext(context.Background(), "worker", time.Second); err != ErrClaimsStopped { + t.Fatalf("ClaimNext() error = %v", err) + } + if repository.claims != 1 { + t.Fatalf("repository claims = %d", repository.claims) + } +} diff --git a/internal/core/queue/mysql_integration_test.go b/internal/core/queue/mysql_integration_test.go new file mode 100644 index 0000000..a057e01 --- /dev/null +++ b/internal/core/queue/mysql_integration_test.go @@ -0,0 +1,291 @@ +package queue + +import ( + "context" + "encoding/json" + "fmt" + "os" + "strings" + "sync" + "testing" + "time" + + "git.ilapage.cn/OPC/chorus/internal/core/model" + "gorm.io/driver/mysql" + "gorm.io/gorm" + "gorm.io/gorm/logger" +) + +func integrationRepository(t *testing.T) (*MySQLRepository, *gorm.DB) { + t.Helper() + dsn := os.Getenv("CHORUS_TEST_DSN") + if dsn == "" { + t.Skip("CHORUS_TEST_DSN is not set") + } + db, err := gorm.Open(mysql.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) + if err != nil { + t.Fatalf("open MySQL integration database: %v", err) + } + repository, err := NewMySQLRepository(db) + if err != nil { + t.Fatal(err) + } + return repository, db +} + +func createIntegrationUser(t *testing.T, db *gorm.DB, suffix string) model.User { + t.Helper() + user := model.User{ + Email: fmt.Sprintf("queue-%s-%d@chorus.invalid", suffix, time.Now().UnixNano()), + PasswordHash: "synthetic-hash", DisplayName: "Queue Test", Status: "active", + } + if err := db.Create(&user).Error; err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + _ = db.Where("user_id = ?", user.ID).Delete(&model.Generation{}).Error + _ = db.Delete(&model.User{}, user.ID).Error + }) + return user +} + +func createIntegrationProviderModel(t *testing.T, db *gorm.DB, suffix string) model.ProviderModel { + t.Helper() + provider := model.Provider{ + Slug: fmt.Sprintf("queue-%s-%d", suffix, time.Now().UnixNano()), Name: "Queue Test Provider", + BaseURL: "https://mock.invalid/v1", AuthType: "none", Enabled: true, + } + if err := db.Create(&provider).Error; err != nil { + t.Fatal(err) + } + providerModel := model.ProviderModel{ + ProviderID: provider.ID, Name: "Queue Test Model", ModelID: "queue-test", + APIType: model.APIChat, Kind: model.KindText, ExtraBody: json.RawMessage("{}"), + TimeoutMS: 1000, Weight: 100, Enabled: true, + } + if err := db.Create(&providerModel).Error; err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + _ = db.Delete(&model.Provider{}, provider.ID).Error + }) + return providerModel +} + +func pendingGeneration(userID uint64, key string) model.Generation { + return model.Generation{ + UserID: userID, Kind: model.KindText, IdempotencyKey: key, + UserPrompt: "synthetic prompt", RenderedPrompt: "synthetic prompt", + } +} + +func createPending(t *testing.T, repository *MySQLRepository, userID uint64, key string) model.Generation { + t.Helper() + generation := pendingGeneration(userID, key) + created, err := repository.CreateIdempotent(context.Background(), &generation, nil) + if err != nil || !created { + t.Fatalf("CreateIdempotent() = %v, %v", created, err) + } + return generation +} + +func textOutput(text string) model.GenerationOutput { + return model.GenerationOutput{Kind: model.KindText, TextContent: &text} +} + +func TestMySQLIdempotencyIsScopedToUser(t *testing.T) { + repository, db := integrationRepository(t) + firstUser := createIntegrationUser(t, db, "idempotent-a") + secondUser := createIntegrationUser(t, db, "idempotent-b") + + first := pendingGeneration(firstUser.ID, "same-key") + created, err := repository.CreateIdempotent(context.Background(), &first, nil) + if err != nil || !created { + t.Fatalf("first CreateIdempotent() = %v, %v", created, err) + } + replay := pendingGeneration(firstUser.ID, "same-key") + created, err = repository.CreateIdempotent(context.Background(), &replay, []model.GenerationInput{{OriginalName: "must-not-write"}}) + if err != nil || created || replay.ID != first.ID { + t.Fatalf("replay CreateIdempotent() = created %v id %d error %v", created, replay.ID, err) + } + var inputCount int64 + if err := db.Model(&model.GenerationInput{}).Where("generation_id = ?", first.ID).Count(&inputCount).Error; err != nil || inputCount != 0 { + t.Fatalf("replay wrote inputs: count=%d error=%v", inputCount, err) + } + other := pendingGeneration(secondUser.ID, "same-key") + created, err = repository.CreateIdempotent(context.Background(), &other, nil) + if err != nil || !created || other.ID == first.ID { + t.Fatalf("other user CreateIdempotent() = created %v id %d error %v", created, other.ID, err) + } + + start := make(chan struct{}) + type result struct { + id uint64 + created bool + err error + } + results := make(chan result, 2) + for range 2 { + go func() { + generation := pendingGeneration(firstUser.ID, "concurrent-key") + <-start + wasCreated, createErr := repository.CreateIdempotent(context.Background(), &generation, nil) + results <- result{id: generation.ID, created: wasCreated, err: createErr} + }() + } + close(start) + firstResult := <-results + secondResult := <-results + if firstResult.err != nil || secondResult.err != nil || firstResult.id == 0 || firstResult.id != secondResult.id || firstResult.created == secondResult.created { + t.Fatalf("concurrent idempotency results = %#v, %#v", firstResult, secondResult) + } +} + +func TestMySQLConcurrentClaimAndStaleTokenCAS(t *testing.T) { + repository, db := integrationRepository(t) + user := createIntegrationUser(t, db, "claim") + providerModel := createIntegrationProviderModel(t, db, "claim") + generation := createPending(t, repository, user.ID, "claim-once") + + start := make(chan struct{}) + claims := make(chan *Claim, 2) + errorsCh := make(chan error, 2) + var wait sync.WaitGroup + for _, owner := range []string{"worker-a", "worker-b"} { + wait.Add(1) + go func(owner string) { + defer wait.Done() + <-start + claim, err := repository.ClaimNext(context.Background(), owner, 5*time.Second) + claims <- claim + errorsCh <- err + }(owner) + } + close(start) + wait.Wait() + close(claims) + close(errorsCh) + for err := range errorsCh { + if err != nil { + t.Fatal(err) + } + } + var firstClaim *Claim + claimed := 0 + for claim := range claims { + if claim != nil { + claimed++ + firstClaim = claim + } + } + if claimed != 1 || firstClaim.Generation.ID != generation.ID { + t.Fatalf("claimed=%d claim=%#v", claimed, firstClaim) + } + owned, err := repository.AssignProvider(context.Background(), generation.ID, firstClaim.LeaseToken, providerModel.ID) + if err != nil || !owned { + t.Fatalf("first AssignProvider() = %v, %v", owned, err) + } + + if err := db.Model(&model.Generation{}).Where("id = ?", generation.ID).Update("lease_until", time.Now().Add(-time.Second)).Error; err != nil { + t.Fatal(err) + } + secondClaim, err := repository.ClaimNext(context.Background(), "worker-c", 5*time.Second) + if err != nil || secondClaim == nil || secondClaim.LeaseToken == firstClaim.LeaseToken { + t.Fatalf("reclaim = %#v, %v", secondClaim, err) + } + if secondClaim.Generation.Status != model.StatusRunning || secondClaim.Generation.AttemptCount != 2 { + t.Fatalf("reclaimed generation = %#v", secondClaim.Generation) + } + owned, err = repository.AssignProvider(context.Background(), generation.ID, firstClaim.LeaseToken, providerModel.ID) + if err != nil || owned { + t.Fatalf("stale AssignProvider() = %v, %v", owned, err) + } + owned, err = repository.AssignProvider(context.Background(), generation.ID, secondClaim.LeaseToken, providerModel.ID) + if err != nil || !owned { + t.Fatalf("second AssignProvider() = %v, %v", owned, err) + } + owned, err = repository.Succeed(context.Background(), generation.ID, firstClaim.LeaseToken, []model.GenerationOutput{textOutput("stale")}, model.Attempt{}) + if err != nil || owned { + t.Fatalf("stale Succeed() = %v, %v", owned, err) + } + owned, err = repository.Fail(context.Background(), generation.ID, firstClaim.LeaseToken, "stale", "stale", model.Attempt{}) + if err != nil || owned { + t.Fatalf("stale Fail() = %v, %v", owned, err) + } + owned, err = repository.Succeed(context.Background(), generation.ID, secondClaim.LeaseToken, []model.GenerationOutput{textOutput("current")}, model.Attempt{LatencyMS: 12}) + if err != nil || !owned { + t.Fatalf("current Succeed() = %v, %v", owned, err) + } + var stored model.Generation + if err := db.First(&stored, generation.ID).Error; err != nil { + t.Fatal(err) + } + if stored.Status != model.StatusSucceeded || stored.AttemptCount != 2 || stored.LeaseToken != nil { + t.Fatalf("stored generation = %#v", stored) + } + var attempts []model.Attempt + if err := json.Unmarshal(stored.Attempts, &attempts); err != nil || len(attempts) != 2 { + t.Fatalf("stored attempts = %#v, error=%v", attempts, err) + } + if attempts[0].ErrorCode != "lease_expired" || attempts[0].FinishedAt == nil || attempts[1].FinishedAt == nil { + t.Fatalf("reclaim attempt chain is incomplete: %#v", attempts) + } + var outputCount int64 + if err := db.Model(&model.GenerationOutput{}).Where("generation_id = ?", generation.ID).Count(&outputCount).Error; err != nil || outputCount != 1 { + t.Fatalf("outputs=%d error=%v", outputCount, err) + } +} + +func TestMySQLFailureSanitizationAndCompletionRollback(t *testing.T) { + repository, db := integrationRepository(t) + user := createIntegrationUser(t, db, "failure") + providerModel := createIntegrationProviderModel(t, db, "failure") + failing := createPending(t, repository, user.ID, "failure") + claim, err := repository.ClaimNext(context.Background(), "worker-fail", time.Second) + if err != nil || claim == nil { + t.Fatalf("ClaimNext() = %#v, %v", claim, err) + } + if owned, err := repository.AssignProvider(context.Background(), failing.ID, claim.LeaseToken, providerModel.ID); err != nil || !owned { + t.Fatalf("AssignProvider() = %v, %v", owned, err) + } + secret := "do-not-persist" + owned, err := repository.Fail(context.Background(), failing.ID, claim.LeaseToken, "upstream_timeout", "token="+secret+"\nrequest failed", model.Attempt{ErrorMessage: "Bearer " + secret, LatencyMS: 25}) + if err != nil || !owned { + t.Fatalf("Fail() = %v, %v", owned, err) + } + var failed model.Generation + if err := db.First(&failed, failing.ID).Error; err != nil { + t.Fatal(err) + } + if failed.Status != model.StatusFailed || failed.ErrorCode == nil || failed.ErrorMessage == nil || strings.Contains(*failed.ErrorMessage, secret) || strings.Contains(string(failed.Attempts), secret) { + t.Fatalf("failed generation leaked or mismatched: %#v attempts=%s", failed, failed.Attempts) + } + var attempts []model.Attempt + if err := json.Unmarshal(failed.Attempts, &attempts); err != nil || len(attempts) != int(failed.AttemptCount) || attempts[0].FinishedAt == nil { + t.Fatalf("attempts=%#v count=%d error=%v", attempts, failed.AttemptCount, err) + } + + rollbackGeneration := createPending(t, repository, user.ID, "rollback") + rollbackClaim, err := repository.ClaimNext(context.Background(), "worker-rollback", time.Second) + if err != nil || rollbackClaim == nil || rollbackClaim.Generation.ID != rollbackGeneration.ID { + t.Fatalf("rollback claim = %#v, %v", rollbackClaim, err) + } + if owned, err := repository.AssignProvider(context.Background(), rollbackGeneration.ID, rollbackClaim.LeaseToken, providerModel.ID); err != nil || !owned { + t.Fatalf("rollback AssignProvider() = %v, %v", owned, err) + } + owned, err = repository.Succeed(context.Background(), rollbackGeneration.ID, rollbackClaim.LeaseToken, []model.GenerationOutput{{Kind: model.KindImage}}, model.Attempt{}) + if err == nil || owned { + t.Fatalf("invalid output Succeed() = %v, %v", owned, err) + } + var afterRollback model.Generation + if err := db.First(&afterRollback, rollbackGeneration.ID).Error; err != nil { + t.Fatal(err) + } + if afterRollback.Status != model.StatusRunning || afterRollback.LeaseToken == nil || *afterRollback.LeaseToken != rollbackClaim.LeaseToken { + t.Fatalf("completion transaction did not roll back: %#v", afterRollback) + } + owned, err = repository.Succeed(context.Background(), rollbackGeneration.ID, rollbackClaim.LeaseToken, []model.GenerationOutput{textOutput("valid")}, model.Attempt{}) + if err != nil || !owned { + t.Fatalf("valid Succeed() = %v, %v", owned, err) + } +} diff --git a/internal/core/queue/mysql_repository.go b/internal/core/queue/mysql_repository.go new file mode 100644 index 0000000..dcc4c25 --- /dev/null +++ b/internal/core/queue/mysql_repository.go @@ -0,0 +1,312 @@ +package queue + +import ( + "context" + "crypto/rand" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "strings" + "time" + + "git.ilapage.cn/OPC/chorus/internal/core/model" + "gorm.io/gorm" + "gorm.io/gorm/clause" +) + +var ( + ErrInvalidLease = errors.New("queue lease configuration is invalid") + ErrInvalidAttempt = errors.New("queue attempt is invalid") + ErrInvalidGeneration = errors.New("generation is invalid") +) + +type MySQLRepository struct { + db *gorm.DB +} + +func NewMySQLRepository(db *gorm.DB) (*MySQLRepository, error) { + if db == nil { + return nil, fmt.Errorf("queue database is required") + } + return &MySQLRepository{db: db}, nil +} + +func (r *MySQLRepository) CreateIdempotent(ctx context.Context, generation *model.Generation, inputs []model.GenerationInput) (bool, error) { + if generation == nil || generation.UserID == 0 || strings.TrimSpace(generation.IdempotencyKey) == "" || !generation.Kind.Valid() || strings.TrimSpace(generation.RenderedPrompt) == "" { + return false, ErrInvalidGeneration + } + var created bool + err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var user model.User + if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").First(&user, generation.UserID).Error; err != nil { + return fmt.Errorf("lock generation user: %w", err) + } + var existing model.Generation + err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("user_id = ? AND idempotency_key = ?", generation.UserID, generation.IdempotencyKey).First(&existing).Error + if err == nil { + *generation = existing + created = false + return nil + } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return fmt.Errorf("find idempotent generation: %w", err) + } + + generation.Status = model.StatusPending + generation.Attempts = json.RawMessage("[]") + generation.AttemptCount = 0 + generation.LeaseOwner = nil + generation.LeaseToken = nil + generation.LeaseUntil = nil + if err := tx.Create(generation).Error; err != nil { + return fmt.Errorf("create generation: %w", err) + } + for index := range inputs { + inputs[index].GenerationID = generation.ID + } + if len(inputs) > 0 { + if err := tx.Create(&inputs).Error; err != nil { + return fmt.Errorf("create generation inputs: %w", err) + } + } + created = true + return nil + }) + return created, err +} + +func (r *MySQLRepository) ClaimNext(ctx context.Context, owner string, leaseDuration time.Duration) (*Claim, error) { + owner = strings.TrimSpace(owner) + if owner == "" || len(owner) > 128 || leaseDuration <= 0 { + return nil, ErrInvalidLease + } + microseconds := leaseDuration.Microseconds() + if microseconds <= 0 { + return nil, ErrInvalidLease + } + var claim *Claim + err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var generation model.Generation + selection := tx.Raw(` + SELECT * FROM generations + WHERE status = 'pending' + OR (status = 'running' AND lease_until < NOW(6)) + ORDER BY created_at, id + LIMIT 1 + FOR UPDATE SKIP LOCKED`).Scan(&generation) + if selection.Error != nil { + return fmt.Errorf("select queue generation: %w", selection.Error) + } + if selection.RowsAffected == 0 { + claim = nil + return nil + } + token, err := newLeaseToken() + if err != nil { + return err + } + started := time.Now().UTC() + attempt := model.Attempt{StartedAt: &started} + if generation.ProviderModelID != nil { + attempt.ProviderModelID = *generation.ProviderModelID + } + attemptJSON, err := json.Marshal(attempt) + if err != nil { + return fmt.Errorf("encode queue attempt: %w", err) + } + attemptsExpression := "JSON_ARRAY_APPEND(attempts, '$', CAST(? AS JSON))" + arguments := []any{owner, token, microseconds} + if generation.Status == model.StatusRunning { + expiredAt := started + expiredAttempt, marshalErr := json.Marshal(model.Attempt{ + ErrorCode: "lease_expired", ErrorMessage: "worker lease expired", FinishedAt: &expiredAt, + }) + if marshalErr != nil { + return fmt.Errorf("encode expired queue attempt: %w", marshalErr) + } + attemptsExpression = `JSON_ARRAY_APPEND( + JSON_SET( + attempts, + CONCAT('$[', attempt_count - 1, ']'), + JSON_MERGE_PATCH( + JSON_EXTRACT(attempts, CONCAT('$[', attempt_count - 1, ']')), + CAST(? AS JSON) + ) + ), + '$', CAST(? AS JSON) + )` + arguments = append(arguments, string(expiredAttempt), string(attemptJSON), generation.ID) + } else { + arguments = append(arguments, string(attemptJSON), generation.ID) + } + updateSQL := fmt.Sprintf(` + UPDATE generations + SET status = 'running', + lease_owner = ?, lease_token = ?, + lease_until = DATE_ADD(NOW(6), INTERVAL ? MICROSECOND), + started_at = COALESCE(started_at, NOW(6)), + attempts = %s, + attempt_count = attempt_count + 1, + error_code = NULL, error_message = NULL + WHERE id = ?`, attemptsExpression) + update := tx.Exec(updateSQL, arguments...) + if update.Error != nil { + return fmt.Errorf("claim queue generation: %w", update.Error) + } + if update.RowsAffected != 1 { + return fmt.Errorf("claim queue generation affected %d rows", update.RowsAffected) + } + if err := tx.First(&generation, generation.ID).Error; err != nil { + return fmt.Errorf("reload claimed generation: %w", err) + } + claim = &Claim{Generation: generation, LeaseToken: token} + return nil + }) + return claim, err +} + +func (r *MySQLRepository) AssignProvider(ctx context.Context, generationID uint64, leaseToken string, providerModelID uint64) (bool, error) { + if generationID == 0 || leaseToken == "" || providerModelID == 0 { + return false, ErrInvalidAttempt + } + result := r.db.WithContext(ctx).Exec(` + UPDATE generations + SET provider_model_id = ?, + attempts = JSON_SET( + attempts, + CONCAT('$[', attempt_count - 1, '].provider_model_id'), + CAST(? AS UNSIGNED) + ) + WHERE id = ? AND status = 'running' AND lease_token = ? AND attempt_count > 0`, + providerModelID, providerModelID, generationID, leaseToken) + if result.Error != nil { + return false, fmt.Errorf("assign queue provider model: %w", result.Error) + } + if result.RowsAffected == 1 { + return true, nil + } + var matches int64 + count := r.db.WithContext(ctx).Raw(` + SELECT COUNT(*) + FROM generations + WHERE id = ? AND status = 'running' AND lease_token = ? + AND provider_model_id = ? + AND CAST(JSON_UNQUOTE(JSON_EXTRACT(attempts, CONCAT('$[', attempt_count - 1, '].provider_model_id'))) AS UNSIGNED) = ?`, + generationID, leaseToken, providerModelID, providerModelID).Scan(&matches) + if count.Error != nil { + return false, fmt.Errorf("verify queue provider model assignment: %w", count.Error) + } + return matches == 1, nil +} + +func (r *MySQLRepository) Succeed(ctx context.Context, generationID uint64, leaseToken string, outputs []model.GenerationOutput, attempt model.Attempt) (bool, error) { + if generationID == 0 || leaseToken == "" || len(outputs) == 0 { + return false, ErrInvalidGeneration + } + finished := time.Now().UTC() + attempt.ProviderModelID = 0 + attempt.StartedAt = nil + if attempt.LatencyMS < 0 { + attempt.LatencyMS = 0 + } + attempt.FinishedAt = &finished + attempt.ErrorCode = "" + attempt.ErrorMessage = "" + attemptJSON, err := json.Marshal(attempt) + if err != nil { + return false, fmt.Errorf("encode successful attempt: %w", err) + } + owned := false + err = r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + result := tx.Exec(` + UPDATE generations + SET status = 'succeeded', completed_at = NOW(6), + lease_owner = NULL, lease_token = NULL, lease_until = NULL, + error_code = NULL, error_message = NULL, + attempts = JSON_SET( + attempts, + CONCAT('$[', attempt_count - 1, ']'), + JSON_MERGE_PATCH( + JSON_EXTRACT(attempts, CONCAT('$[', attempt_count - 1, ']')), + CAST(? AS JSON) + ) + ) + WHERE id = ? AND status = 'running' AND lease_token = ? + AND JSON_EXTRACT(attempts, CONCAT('$[', attempt_count - 1, '].provider_model_id')) IS NOT NULL`, string(attemptJSON), generationID, leaseToken) + if result.Error != nil { + return fmt.Errorf("complete queue generation: %w", result.Error) + } + if result.RowsAffected == 0 { + return nil + } + owned = true + for index := range outputs { + outputs[index].ID = 0 + outputs[index].GenerationID = generationID + } + if err := tx.Create(&outputs).Error; err != nil { + return fmt.Errorf("create generation outputs: %w", err) + } + return nil + }) + if err != nil { + return false, err + } + return owned, nil +} + +func (r *MySQLRepository) Fail(ctx context.Context, generationID uint64, leaseToken string, code, message string, attempt model.Attempt) (bool, error) { + if generationID == 0 || leaseToken == "" { + return false, ErrInvalidGeneration + } + code = sanitizeCode(code) + message = SanitizeErrorMessage(message, 1024) + finished := time.Now().UTC() + attempt.ProviderModelID = 0 + attempt.StartedAt = nil + if attempt.LatencyMS < 0 { + attempt.LatencyMS = 0 + } + attempt.FinishedAt = &finished + attempt.ErrorCode = code + attempt.ErrorMessage = SanitizeErrorMessage(attempt.ErrorMessage, 512) + if attempt.ErrorMessage == "upstream request failed" { + attempt.ErrorMessage = message + } + attemptJSON, err := json.Marshal(attempt) + if err != nil { + return false, fmt.Errorf("encode failed attempt: %w", err) + } + result := r.db.WithContext(ctx).Exec(` + UPDATE generations + SET status = 'failed', completed_at = NOW(6), + lease_owner = NULL, lease_token = NULL, lease_until = NULL, + error_code = ?, error_message = ?, + attempts = JSON_SET( + attempts, + CONCAT('$[', attempt_count - 1, ']'), + JSON_MERGE_PATCH( + JSON_EXTRACT(attempts, CONCAT('$[', attempt_count - 1, ']')), + CAST(? AS JSON) + ) + ) + WHERE id = ? AND status = 'running' AND lease_token = ? + AND JSON_EXTRACT(attempts, CONCAT('$[', attempt_count - 1, '].provider_model_id')) IS NOT NULL`, + code, message, string(attemptJSON), generationID, leaseToken) + if result.Error != nil { + return false, fmt.Errorf("fail queue generation: %w", result.Error) + } + return result.RowsAffected == 1, nil +} + +func newLeaseToken() (string, error) { + bytes := make([]byte, 16) + if _, err := rand.Read(bytes); err != nil { + return "", fmt.Errorf("generate queue lease token: %w", err) + } + hexValue := hex.EncodeToString(bytes) + return hexValue[0:8] + "-" + hexValue[8:12] + "-" + hexValue[12:16] + "-" + hexValue[16:20] + "-" + hexValue[20:32], nil +} + +var _ Repository = (*MySQLRepository)(nil) diff --git a/internal/core/queue/queue.go b/internal/core/queue/queue.go index 98d7227..3bbc6f2 100644 --- a/internal/core/queue/queue.go +++ b/internal/core/queue/queue.go @@ -13,7 +13,9 @@ type Claim struct { } type Repository interface { - ClaimNext(ctx context.Context, owner string, now time.Time, leaseDuration time.Duration) (*Claim, error) - Succeed(ctx context.Context, generationID uint64, leaseToken string, outputs []model.GenerationOutput, attempts []model.Attempt) (bool, error) - Fail(ctx context.Context, generationID uint64, leaseToken string, code, message string, attempts []model.Attempt) (bool, error) + CreateIdempotent(ctx context.Context, generation *model.Generation, inputs []model.GenerationInput) (created bool, err error) + ClaimNext(ctx context.Context, owner string, leaseDuration time.Duration) (*Claim, error) + AssignProvider(ctx context.Context, generationID uint64, leaseToken string, providerModelID uint64) (bool, error) + Succeed(ctx context.Context, generationID uint64, leaseToken string, outputs []model.GenerationOutput, attempt model.Attempt) (bool, error) + Fail(ctx context.Context, generationID uint64, leaseToken string, code, message string, attempt model.Attempt) (bool, error) } diff --git a/internal/core/queue/sanitize.go b/internal/core/queue/sanitize.go new file mode 100644 index 0000000..431501b --- /dev/null +++ b/internal/core/queue/sanitize.go @@ -0,0 +1,54 @@ +package queue + +import ( + "regexp" + "strings" + "unicode" + "unicode/utf8" +) + +var ( + credentialPattern = regexp.MustCompile(`(?i)\b(api[_-]?key|authorization|password|token)\s*[:=]\s*[^\s,;]+`) + bearerPattern = regexp.MustCompile(`(?i)\bbearer\s+[^\s,;]+`) + urlPattern = regexp.MustCompile(`(?i)https?://[^\s]+`) + codePattern = regexp.MustCompile(`^[a-z0-9_]+$`) +) + +func SanitizeErrorMessage(message string, maxBytes int) string { + message = urlPattern.ReplaceAllString(message, "[REDACTED-URL]") + message = bearerPattern.ReplaceAllString(message, "Bearer [REDACTED]") + message = credentialPattern.ReplaceAllString(message, "$1=[REDACTED]") + message = strings.Map(func(value rune) rune { + if unicode.IsControl(value) { + return ' ' + } + return value + }, message) + message = strings.Join(strings.Fields(message), " ") + if message == "" { + message = "upstream request failed" + } + return truncateUTF8(message, maxBytes) +} + +func sanitizeCode(code string) string { + code = strings.ToLower(strings.TrimSpace(code)) + if len(code) == 0 || len(code) > 64 || !codePattern.MatchString(code) { + return "upstream_error" + } + return code +} + +func truncateUTF8(value string, maxBytes int) string { + if maxBytes <= 0 { + return "" + } + if len(value) <= maxBytes { + return value + } + value = value[:maxBytes] + for !utf8.ValidString(value) { + value = value[:len(value)-1] + } + return value +} diff --git a/internal/core/queue/sanitize_test.go b/internal/core/queue/sanitize_test.go new file mode 100644 index 0000000..093d109 --- /dev/null +++ b/internal/core/queue/sanitize_test.go @@ -0,0 +1,26 @@ +package queue + +import ( + "strings" + "testing" +) + +func TestSanitizeErrorMessage(t *testing.T) { + message := "Authorization: Bearer-secret\r\n api_key=topsecret token=abc password=hunter2 https://example.test/path?X-Amz-Signature=signed 请求失败" + got := SanitizeErrorMessage(message, 90) + for _, secret := range []string{"Bearer-secret", "topsecret", "abc", "hunter2", "X-Amz-Signature", "signed", "\r", "\n"} { + if strings.Contains(got, secret) { + t.Fatalf("SanitizeErrorMessage() leaked %q in %q", secret, got) + } + } + if len(got) > 90 { + t.Fatalf("SanitizeErrorMessage() length = %d", len(got)) + } +} + +func TestSanitizeErrorMessageKeepsValidUTF8WhenTruncated(t *testing.T) { + got := SanitizeErrorMessage(strings.Repeat("中", 20), 10) + if !strings.HasPrefix(strings.Repeat("中", 20), got) || len(got) > 10 { + t.Fatalf("truncated message = %q", got) + } +}