fix(admission): 避免高并发准入锁占满连接池
将同进程相同准入键的请求先在内存中串行化,跨节点使用 pg_try_advisory_xact_lock 非阻塞竞争并在事务外退避,避免等待 advisory lock 时长期占用 PostgreSQL 连接。\n\n新增 64 路竞争回归测试,验证被占用锁下连接池峰值仅保留持锁者和单个尝试者;256 路 Gemini Base64 端到端压力通过,advisory wait 峰值为 0。
This commit is contained in:
@@ -0,0 +1,170 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sort"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
)
|
||||
|
||||
const (
|
||||
admissionLockRetryMin = 10 * time.Millisecond
|
||||
admissionLockRetryMax = 250 * time.Millisecond
|
||||
)
|
||||
|
||||
var (
|
||||
errAdmissionLockBusy = errors.New("task admission lock is busy")
|
||||
processAdmissionLocks = newAdmissionLocalLockSet()
|
||||
)
|
||||
|
||||
type admissionLocalLock struct {
|
||||
token chan struct{}
|
||||
refs int
|
||||
}
|
||||
|
||||
type admissionLocalLockSet struct {
|
||||
mu sync.Mutex
|
||||
locks map[string]*admissionLocalLock
|
||||
}
|
||||
|
||||
func newAdmissionLocalLockSet() *admissionLocalLockSet {
|
||||
return &admissionLocalLockSet{locks: make(map[string]*admissionLocalLock)}
|
||||
}
|
||||
|
||||
// acquire serializes contenders inside one API or Worker process before they
|
||||
// consume a PostgreSQL connection. Sorted acquisition keeps operations that
|
||||
// span task, platform-model, and user-group keys deadlock-free.
|
||||
func (s *admissionLocalLockSet) acquire(ctx context.Context, keys []string) (func(), error) {
|
||||
keys = normalizedAdmissionLockKeys(keys)
|
||||
acquired := make([]struct {
|
||||
key string
|
||||
lock *admissionLocalLock
|
||||
}, 0, len(keys))
|
||||
for _, key := range keys {
|
||||
lock := s.reference(key)
|
||||
select {
|
||||
case lock.token <- struct{}{}:
|
||||
acquired = append(acquired, struct {
|
||||
key string
|
||||
lock *admissionLocalLock
|
||||
}{key: key, lock: lock})
|
||||
case <-ctx.Done():
|
||||
s.unreference(key, lock)
|
||||
for index := len(acquired) - 1; index >= 0; index-- {
|
||||
<-acquired[index].lock.token
|
||||
s.unreference(acquired[index].key, acquired[index].lock)
|
||||
}
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
}
|
||||
return func() {
|
||||
for index := len(acquired) - 1; index >= 0; index-- {
|
||||
<-acquired[index].lock.token
|
||||
s.unreference(acquired[index].key, acquired[index].lock)
|
||||
}
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *admissionLocalLockSet) reference(key string) *admissionLocalLock {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
lock := s.locks[key]
|
||||
if lock == nil {
|
||||
lock = &admissionLocalLock{token: make(chan struct{}, 1)}
|
||||
s.locks[key] = lock
|
||||
}
|
||||
lock.refs++
|
||||
return lock
|
||||
}
|
||||
|
||||
func (s *admissionLocalLockSet) unreference(key string, lock *admissionLocalLock) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
lock.refs--
|
||||
if lock.refs == 0 && s.locks[key] == lock {
|
||||
delete(s.locks, key)
|
||||
}
|
||||
}
|
||||
|
||||
func normalizedAdmissionLockKeys(keys []string) []string {
|
||||
seen := make(map[string]struct{}, len(keys))
|
||||
normalized := make([]string, 0, len(keys))
|
||||
for _, key := range keys {
|
||||
if key == "" {
|
||||
continue
|
||||
}
|
||||
if _, exists := seen[key]; exists {
|
||||
continue
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
normalized = append(normalized, key)
|
||||
}
|
||||
sort.Strings(normalized)
|
||||
return normalized
|
||||
}
|
||||
|
||||
func admissionOperationLockKeys(input TaskAdmissionInput) []string {
|
||||
keys := []string{"task-admission:" + input.TaskID}
|
||||
for _, scope := range normalizedAdmissionScopes(input.Scopes) {
|
||||
keys = append(keys, admissionLockKey(scope))
|
||||
}
|
||||
return normalizedAdmissionLockKeys(keys)
|
||||
}
|
||||
|
||||
func tryAdmissionTransactionLock(ctx context.Context, tx pgx.Tx, key string) error {
|
||||
var locked bool
|
||||
if err := tx.QueryRow(
|
||||
ctx,
|
||||
`SELECT pg_try_advisory_xact_lock(hashtextextended($1, 0))`,
|
||||
key,
|
||||
).Scan(&locked); err != nil {
|
||||
return err
|
||||
}
|
||||
if !locked {
|
||||
return errAdmissionLockBusy
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func retryAdmissionOperation[T any](
|
||||
ctx context.Context,
|
||||
keys []string,
|
||||
operation func() (T, error),
|
||||
) (T, error) {
|
||||
var zero T
|
||||
for attempt := 0; ; attempt++ {
|
||||
release, err := processAdmissionLocks.acquire(ctx, keys)
|
||||
if err != nil {
|
||||
return zero, err
|
||||
}
|
||||
result, operationErr := operation()
|
||||
release()
|
||||
if !errors.Is(operationErr, errAdmissionLockBusy) {
|
||||
return result, operationErr
|
||||
}
|
||||
if err := waitAdmissionLockRetry(ctx, attempt); err != nil {
|
||||
return zero, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func waitAdmissionLockRetry(ctx context.Context, attempt int) error {
|
||||
delay := admissionLockRetryMin
|
||||
for index := 0; index < attempt && delay < admissionLockRetryMax; index++ {
|
||||
delay *= 2
|
||||
}
|
||||
if delay > admissionLockRetryMax {
|
||||
delay = admissionLockRetryMax
|
||||
}
|
||||
timer := time.NewTimer(delay)
|
||||
defer timer.Stop()
|
||||
select {
|
||||
case <-timer.C:
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,221 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/jackc/pgx/v5"
|
||||
)
|
||||
|
||||
func TestAdmissionLocalLockSetSerializesSharedKeys(t *testing.T) {
|
||||
lockSet := newAdmissionLocalLockSet()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
releaseFirst, err := lockSet.acquire(ctx, []string{"scope:b", "scope:a", "scope:a"})
|
||||
if err != nil {
|
||||
t.Fatalf("acquire first lock set: %v", err)
|
||||
}
|
||||
|
||||
const contenders = 64
|
||||
var active atomic.Int32
|
||||
var maxActive atomic.Int32
|
||||
var waitGroup sync.WaitGroup
|
||||
started := make(chan struct{}, contenders)
|
||||
for index := 0; index < contenders; index++ {
|
||||
waitGroup.Add(1)
|
||||
go func() {
|
||||
defer waitGroup.Done()
|
||||
started <- struct{}{}
|
||||
release, acquireErr := lockSet.acquire(ctx, []string{"scope:a", "scope:b"})
|
||||
if acquireErr != nil {
|
||||
t.Errorf("acquire contended lock set: %v", acquireErr)
|
||||
return
|
||||
}
|
||||
current := active.Add(1)
|
||||
for {
|
||||
observed := maxActive.Load()
|
||||
if current <= observed || maxActive.CompareAndSwap(observed, current) {
|
||||
break
|
||||
}
|
||||
}
|
||||
time.Sleep(time.Millisecond)
|
||||
active.Add(-1)
|
||||
release()
|
||||
}()
|
||||
}
|
||||
for index := 0; index < contenders; index++ {
|
||||
<-started
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
if active.Load() != 0 {
|
||||
t.Fatalf("contenders entered while first lock was held: active=%d", active.Load())
|
||||
}
|
||||
releaseFirst()
|
||||
waitGroup.Wait()
|
||||
|
||||
if maxActive.Load() != 1 {
|
||||
t.Fatalf("maximum concurrent holders = %d, want 1", maxActive.Load())
|
||||
}
|
||||
lockSet.mu.Lock()
|
||||
defer lockSet.mu.Unlock()
|
||||
if len(lockSet.locks) != 0 {
|
||||
t.Fatalf("local lock entries leaked: %d", len(lockSet.locks))
|
||||
}
|
||||
}
|
||||
|
||||
func TestRetryAdmissionOperationReleasesLocalLockBeforeBackoff(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
lockKey := "scope:retry-test:" + uuid.NewString()
|
||||
var firstAttempts atomic.Int32
|
||||
firstStarted := make(chan struct{})
|
||||
firstDone := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := retryAdmissionOperation(ctx, []string{lockKey}, func() (struct{}, error) {
|
||||
if firstAttempts.Add(1) == 1 {
|
||||
close(firstStarted)
|
||||
return struct{}{}, errAdmissionLockBusy
|
||||
}
|
||||
return struct{}{}, nil
|
||||
})
|
||||
firstDone <- err
|
||||
}()
|
||||
<-firstStarted
|
||||
|
||||
secondRan := make(chan struct{})
|
||||
if _, err := retryAdmissionOperation(ctx, []string{lockKey}, func() (struct{}, error) {
|
||||
close(secondRan)
|
||||
return struct{}{}, nil
|
||||
}); err != nil {
|
||||
t.Fatalf("second operation: %v", err)
|
||||
}
|
||||
select {
|
||||
case <-secondRan:
|
||||
default:
|
||||
t.Fatal("second operation did not run during first operation backoff")
|
||||
}
|
||||
if err := <-firstDone; err != nil {
|
||||
t.Fatalf("first operation retry: %v", err)
|
||||
}
|
||||
if firstAttempts.Load() != 2 {
|
||||
t.Fatalf("first attempts = %d, want 2", firstAttempts.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestAdmissionDatabaseLockContentionDoesNotOccupyPool(t *testing.T) {
|
||||
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
|
||||
if databaseURL == "" {
|
||||
t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run task admission PostgreSQL integration tests")
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
db, err := ConnectWithMaxConns(ctx, databaseURL, 8)
|
||||
if err != nil {
|
||||
t.Fatalf("connect store: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
var databaseName string
|
||||
if err := db.pool.QueryRow(ctx, `SELECT current_database()`).Scan(&databaseName); err != nil {
|
||||
t.Fatalf("read test database name: %v", err)
|
||||
}
|
||||
if !strings.Contains(strings.ToLower(databaseName), "test") {
|
||||
t.Fatalf("refusing to use non-test database %q", databaseName)
|
||||
}
|
||||
|
||||
lockKey := "task-admission-contention-test:" + uuid.NewString()
|
||||
holder, err := db.pool.Begin(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("begin lock holder: %v", err)
|
||||
}
|
||||
defer rollbackTransaction(holder)
|
||||
if _, err := holder.Exec(
|
||||
ctx,
|
||||
`SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`,
|
||||
lockKey,
|
||||
); err != nil {
|
||||
t.Fatalf("hold advisory lock: %v", err)
|
||||
}
|
||||
|
||||
const contenders = 64
|
||||
var peakAcquired atomic.Int32
|
||||
sampleDone := make(chan struct{})
|
||||
go func() {
|
||||
ticker := time.NewTicker(time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
current := int32(db.pool.Stat().AcquiredConns())
|
||||
for {
|
||||
observed := peakAcquired.Load()
|
||||
if current <= observed || peakAcquired.CompareAndSwap(observed, current) {
|
||||
break
|
||||
}
|
||||
}
|
||||
case <-sampleDone:
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
start := make(chan struct{})
|
||||
results := make(chan error, contenders)
|
||||
var waitGroup sync.WaitGroup
|
||||
for index := 0; index < contenders; index++ {
|
||||
waitGroup.Add(1)
|
||||
go func() {
|
||||
defer waitGroup.Done()
|
||||
<-start
|
||||
_, retryErr := retryAdmissionOperation(ctx, []string{lockKey}, func() (struct{}, error) {
|
||||
tx, beginErr := db.pool.Begin(ctx)
|
||||
if beginErr != nil {
|
||||
return struct{}{}, beginErr
|
||||
}
|
||||
defer rollbackTransaction(tx)
|
||||
if lockErr := tryAdmissionTransactionLock(ctx, tx, lockKey); lockErr != nil {
|
||||
return struct{}{}, lockErr
|
||||
}
|
||||
return struct{}{}, tx.Commit(ctx)
|
||||
})
|
||||
results <- retryErr
|
||||
}()
|
||||
}
|
||||
close(start)
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
if current := db.pool.Stat().AcquiredConns(); current > 2 {
|
||||
t.Fatalf("connections acquired during lock contention = %d, want at most 2", current)
|
||||
}
|
||||
if err := holder.Commit(ctx); err != nil {
|
||||
t.Fatalf("release advisory lock: %v", err)
|
||||
}
|
||||
waitGroup.Wait()
|
||||
close(sampleDone)
|
||||
close(results)
|
||||
for retryErr := range results {
|
||||
if retryErr != nil {
|
||||
t.Fatalf("contended operation: %v", retryErr)
|
||||
}
|
||||
}
|
||||
if peakAcquired.Load() > 2 {
|
||||
t.Fatalf("peak acquired connections = %d, want at most 2", peakAcquired.Load())
|
||||
}
|
||||
|
||||
var advisoryWaiters int
|
||||
if err := db.pool.QueryRow(ctx, `
|
||||
SELECT COUNT(*)
|
||||
FROM pg_stat_activity
|
||||
WHERE datname = current_database()
|
||||
AND wait_event = 'advisory'`).Scan(&advisoryWaiters); err != nil && err != pgx.ErrNoRows {
|
||||
t.Fatalf("count advisory waiters: %v", err)
|
||||
}
|
||||
if advisoryWaiters != 0 {
|
||||
t.Fatalf("advisory waiters after contention = %d, want 0", advisoryWaiters)
|
||||
}
|
||||
}
|
||||
@@ -130,13 +130,23 @@ func (s *Store) QueueTaskAdmissionWithHook(
|
||||
if input.Mode != "async" {
|
||||
return TaskAdmission{}, errors.New("queued admission registration requires asynchronous mode")
|
||||
}
|
||||
return retryAdmissionOperation(ctx, admissionOperationLockKeys(input), func() (TaskAdmission, error) {
|
||||
return s.queueTaskAdmissionOnce(ctx, input, onQueued)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Store) queueTaskAdmissionOnce(
|
||||
ctx context.Context,
|
||||
input TaskAdmissionInput,
|
||||
onQueued func(pgx.Tx) error,
|
||||
) (TaskAdmission, error) {
|
||||
tx, err := s.pool.Begin(ctx)
|
||||
if err != nil {
|
||||
return TaskAdmission{}, err
|
||||
}
|
||||
defer rollbackTransaction(tx)
|
||||
|
||||
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+input.TaskID); err != nil {
|
||||
if err := tryAdmissionTransactionLock(ctx, tx, "task-admission:"+input.TaskID); err != nil {
|
||||
return TaskAdmission{}, err
|
||||
}
|
||||
var taskActive bool
|
||||
@@ -163,7 +173,7 @@ WHERE id = $1::uuid`, input.TaskID).Scan(&taskActive); err != nil {
|
||||
)
|
||||
}
|
||||
for _, scope := range normalizedAdmissionScopes(lockScopes) {
|
||||
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, admissionLockKey(scope)); err != nil {
|
||||
if err := tryAdmissionTransactionLock(ctx, tx, admissionLockKey(scope)); err != nil {
|
||||
return TaskAdmission{}, err
|
||||
}
|
||||
}
|
||||
@@ -228,13 +238,23 @@ func (s *Store) tryTaskAdmission(
|
||||
if err := validateTaskAdmissionInput(input); err != nil {
|
||||
return TaskAdmissionResult{}, err
|
||||
}
|
||||
return retryAdmissionOperation(ctx, admissionOperationLockKeys(input), func() (TaskAdmissionResult, error) {
|
||||
return s.tryTaskAdmissionOnce(ctx, input, onAdmitted)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Store) tryTaskAdmissionOnce(
|
||||
ctx context.Context,
|
||||
input TaskAdmissionInput,
|
||||
onAdmitted func(pgx.Tx) error,
|
||||
) (TaskAdmissionResult, error) {
|
||||
tx, err := s.pool.Begin(ctx)
|
||||
if err != nil {
|
||||
return TaskAdmissionResult{}, err
|
||||
}
|
||||
defer rollbackTransaction(tx)
|
||||
|
||||
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+input.TaskID); err != nil {
|
||||
if err := tryAdmissionTransactionLock(ctx, tx, "task-admission:"+input.TaskID); err != nil {
|
||||
return TaskAdmissionResult{}, err
|
||||
}
|
||||
var taskActive bool
|
||||
@@ -261,7 +281,7 @@ WHERE id = $1::uuid`, input.TaskID).Scan(&taskActive); err != nil {
|
||||
)
|
||||
}
|
||||
for _, scope := range normalizedAdmissionScopes(lockScopes) {
|
||||
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, admissionLockKey(scope)); err != nil {
|
||||
if err := tryAdmissionTransactionLock(ctx, tx, admissionLockKey(scope)); err != nil {
|
||||
return TaskAdmissionResult{}, err
|
||||
}
|
||||
}
|
||||
@@ -452,12 +472,18 @@ func (s *Store) RebindWaitingTaskAdmission(ctx context.Context, input TaskAdmiss
|
||||
if err := validateTaskAdmissionInput(input); err != nil {
|
||||
return TaskAdmission{}, err
|
||||
}
|
||||
return retryAdmissionOperation(ctx, admissionOperationLockKeys(input), func() (TaskAdmission, error) {
|
||||
return s.rebindWaitingTaskAdmissionOnce(ctx, input)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Store) rebindWaitingTaskAdmissionOnce(ctx context.Context, input TaskAdmissionInput) (TaskAdmission, error) {
|
||||
tx, err := s.pool.Begin(ctx)
|
||||
if err != nil {
|
||||
return TaskAdmission{}, err
|
||||
}
|
||||
defer rollbackTransaction(tx)
|
||||
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+input.TaskID); err != nil {
|
||||
if err := tryAdmissionTransactionLock(ctx, tx, "task-admission:"+input.TaskID); err != nil {
|
||||
return TaskAdmission{}, err
|
||||
}
|
||||
current, found, err := loadTaskAdmissionTx(ctx, tx, input.TaskID)
|
||||
@@ -477,7 +503,7 @@ func (s *Store) RebindWaitingTaskAdmission(ctx context.Context, input TaskAdmiss
|
||||
AdmissionScope{ScopeType: "user_group", ScopeKey: current.UserGroupID},
|
||||
)
|
||||
for _, scope := range normalizedAdmissionScopes(lockScopes) {
|
||||
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, admissionLockKey(scope)); err != nil {
|
||||
if err := tryAdmissionTransactionLock(ctx, tx, admissionLockKey(scope)); err != nil {
|
||||
return TaskAdmission{}, err
|
||||
}
|
||||
}
|
||||
@@ -629,12 +655,19 @@ WHERE admission.task_id = owned.task_id
|
||||
}
|
||||
|
||||
func (s *Store) DeleteTaskAdmission(ctx context.Context, taskID string) error {
|
||||
_, err := retryAdmissionOperation(ctx, []string{"task-admission:" + taskID}, func() (struct{}, error) {
|
||||
return struct{}{}, s.deleteTaskAdmissionOnce(ctx, taskID)
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) deleteTaskAdmissionOnce(ctx context.Context, taskID string) error {
|
||||
tx, err := s.pool.Begin(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer rollbackTransaction(tx)
|
||||
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+taskID); err != nil {
|
||||
if err := tryAdmissionTransactionLock(ctx, tx, "task-admission:"+taskID); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := tx.Exec(ctx, `
|
||||
|
||||
Reference in New Issue
Block a user