feat: add portal authentication and generation API (#11)
This commit is contained in:
@@ -3,6 +3,7 @@ package queue
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
@@ -141,6 +142,38 @@ func TestMySQLIdempotencyIsScopedToUser(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestMySQLPreparedInputsRollbackAndReplay(t *testing.T) {
|
||||
repository, db := integrationRepository(t)
|
||||
user := createIntegrationUser(t, db, "prepared")
|
||||
generation := pendingGeneration(user.ID, "prepared-error")
|
||||
callbackError := fmt.Errorf("synthetic prepare failure")
|
||||
created, err := repository.CreateIdempotentPrepared(context.Background(), &generation, func(uint64) ([]model.GenerationInput, error) {
|
||||
return nil, callbackError
|
||||
})
|
||||
if created || !errors.Is(err, callbackError) {
|
||||
t.Fatalf("prepared failure = created %v error %v", created, err)
|
||||
}
|
||||
var count int64
|
||||
if err := db.Model(&model.Generation{}).Where("user_id=? AND idempotency_key=?", user.ID, "prepared-error").Count(&count).Error; err != nil || count != 0 {
|
||||
t.Fatalf("failed prepared create persisted generation: count=%d error=%v", count, err)
|
||||
}
|
||||
|
||||
generation = pendingGeneration(user.ID, "prepared-replay")
|
||||
callbackCalls := 0
|
||||
created, err = repository.CreateIdempotentPrepared(context.Background(), &generation, func(id uint64) ([]model.GenerationInput, error) {
|
||||
callbackCalls++
|
||||
return []model.GenerationInput{{GenerationID: id, Position: 0, Role: model.RolePrimary, OriginalName: "input.png", MIMEType: "image/png", StorageKey: "synthetic/input", SizeBytes: 10}}, nil
|
||||
})
|
||||
if err != nil || !created || callbackCalls != 1 {
|
||||
t.Fatalf("prepared create=%v calls=%d err=%v", created, callbackCalls, err)
|
||||
}
|
||||
replay := pendingGeneration(user.ID, "prepared-replay")
|
||||
created, err = repository.CreateIdempotentPrepared(context.Background(), &replay, func(uint64) ([]model.GenerationInput, error) { callbackCalls++; return nil, nil })
|
||||
if err != nil || created || replay.ID != generation.ID || callbackCalls != 1 {
|
||||
t.Fatalf("prepared replay=%v id=%d calls=%d err=%v", created, replay.ID, callbackCalls, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMySQLConcurrentClaimAndStaleTokenCAS(t *testing.T) {
|
||||
repository, db := integrationRepository(t)
|
||||
user := createIntegrationUser(t, db, "claim")
|
||||
|
||||
@@ -25,6 +25,8 @@ type MySQLRepository struct {
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
type PrepareInputs func(generationID uint64) ([]model.GenerationInput, error)
|
||||
|
||||
func NewMySQLRepository(db *gorm.DB) (*MySQLRepository, error) {
|
||||
if db == nil {
|
||||
return nil, fmt.Errorf("queue database is required")
|
||||
@@ -33,9 +35,18 @@ func NewMySQLRepository(db *gorm.DB) (*MySQLRepository, error) {
|
||||
}
|
||||
|
||||
func (r *MySQLRepository) CreateIdempotent(ctx context.Context, generation *model.Generation, inputs []model.GenerationInput) (bool, error) {
|
||||
return r.CreateIdempotentPrepared(ctx, generation, func(uint64) ([]model.GenerationInput, error) {
|
||||
return inputs, nil
|
||||
})
|
||||
}
|
||||
|
||||
func (r *MySQLRepository) CreateIdempotentPrepared(ctx context.Context, generation *model.Generation, prepare PrepareInputs) (bool, error) {
|
||||
if generation == nil || generation.UserID == 0 || strings.TrimSpace(generation.IdempotencyKey) == "" || !generation.Kind.Valid() || strings.TrimSpace(generation.RenderedPrompt) == "" {
|
||||
return false, ErrInvalidGeneration
|
||||
}
|
||||
if prepare == nil {
|
||||
return false, ErrInvalidGeneration
|
||||
}
|
||||
var created bool
|
||||
err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var user model.User
|
||||
@@ -62,6 +73,10 @@ func (r *MySQLRepository) CreateIdempotent(ctx context.Context, generation *mode
|
||||
if err := tx.Create(generation).Error; err != nil {
|
||||
return fmt.Errorf("create generation: %w", err)
|
||||
}
|
||||
inputs, err := prepare(generation.ID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("prepare generation inputs: %w", err)
|
||||
}
|
||||
for index := range inputs {
|
||||
inputs[index].GenerationID = generation.ID
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user