diff --git a/apps/api/internal/store/postgres.go b/apps/api/internal/store/postgres.go index 7f31b48..e219f78 100644 --- a/apps/api/internal/store/postgres.go +++ b/apps/api/internal/store/postgres.go @@ -21,7 +21,8 @@ import ( ) type Store struct { - pool *pgxpool.Pool + pool *pgxpool.Pool + walletReservationLocks walletReservationLockSet } const ( diff --git a/apps/api/internal/store/wallet.go b/apps/api/internal/store/wallet.go index 439ceff..24d2e6f 100644 --- a/apps/api/internal/store/wallet.go +++ b/apps/api/internal/store/wallet.go @@ -5,14 +5,30 @@ import ( "encoding/json" "errors" "fmt" + "hash/fnv" "strconv" "strings" + "sync" "time" "github.com/easyai/easyai-ai-gateway/apps/api/internal/auth" "github.com/jackc/pgx/v5" ) +const walletReservationLockStripes = 256 + +type walletReservationLockSet struct { + stripes [walletReservationLockStripes]sync.Mutex +} + +func (locks *walletReservationLockSet) lock(gatewayUserID string) func() { + hasher := fnv.New32a() + _, _ = hasher.Write([]byte(strings.TrimSpace(gatewayUserID))) + stripe := &locks.stripes[hasher.Sum32()%walletReservationLockStripes] + stripe.Lock() + return stripe.Unlock +} + type GatewayWalletAccount struct { ID string `json:"id"` GatewayTenantID string `json:"gatewayTenantId,omitempty"` @@ -153,12 +169,16 @@ func (s *Store) ReserveTaskBilling(ctx context.Context, task GatewayTask, user * pricingSnapshot = emptyObjectIfNil(pricingSnapshots[0]) } if exactAmount := walletString(pricingSnapshot["reservationAmount"]); exactAmount != "" { + unlock := s.walletReservationLocks.lock(gatewayUserID) + defer unlock() return s.reserveTaskBillingExact(ctx, task, gatewayUserID, exactAmount, pricingSnapshot) } amounts := walletBillingAmounts(billings) if len(amounts) == 0 { return nil, nil } + unlock := s.walletReservationLocks.lock(gatewayUserID) + defer unlock() reservations := make([]WalletBillingReservation, 0, len(amounts)) pricingSnapshotJSON, _ := json.Marshal(sanitizeJSONForStorage(pricingSnapshot)) diff --git a/apps/api/internal/store/wallet_reservation_test.go b/apps/api/internal/store/wallet_reservation_test.go index b9ec265..5aadd46 100644 --- a/apps/api/internal/store/wallet_reservation_test.go +++ b/apps/api/internal/store/wallet_reservation_test.go @@ -113,6 +113,31 @@ func TestReserveTaskBillingSerializesConcurrentWalletReservations(t *testing.T) } } +func TestWalletReservationLockSetSerializesSameWallet(t *testing.T) { + var locks walletReservationLockSet + unlock := locks.lock("wallet-user") + acquired := make(chan struct{}) + released := make(chan struct{}) + go func() { + defer close(released) + unlockSecond := locks.lock("wallet-user") + close(acquired) + unlockSecond() + }() + select { + case <-acquired: + t.Fatal("same-wallet reservation lock was acquired concurrently") + case <-time.After(20 * time.Millisecond): + } + unlock() + select { + case <-acquired: + case <-time.After(time.Second): + t.Fatal("same-wallet reservation lock did not unblock") + } + <-released +} + func seedWalletReservationUser(t *testing.T, ctx context.Context, db *Store) (string, string) { t.Helper() suffix := strconv.FormatInt(time.Now().UnixNano(), 10)