feat: add MySQL lease queue and CAS completion (#9)

This commit is contained in:
ila
2026-08-20 23:48:59 +08:00
parent 94468c5c44
commit 7dfb528166
10 changed files with 781 additions and 9 deletions
+1
View File
@@ -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
)
+6
View File
@@ -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=
+6 -6
View File
@@ -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"`
}
+35
View File
@@ -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)
}
+45
View File
@@ -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)
}
}
@@ -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)
}
}
+312
View File
@@ -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)
+5 -3
View File
@@ -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)
}
+54
View File
@@ -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
}
+26
View File
@@ -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)
}
}