feat: add MySQL lease queue and CAS completion (#9)
This commit is contained in:
@@ -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
|
||||
)
|
||||
|
||||
|
||||
@@ -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=
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user