将同进程相同准入键的请求先在内存中串行化,跨节点使用 pg_try_advisory_xact_lock 非阻塞竞争并在事务外退避,避免等待 advisory lock 时长期占用 PostgreSQL 连接。\n\n新增 64 路竞争回归测试,验证被占用锁下连接池峰值仅保留持锁者和单个尝试者;256 路 Gemini Base64 端到端压力通过,advisory wait 峰值为 0。
222 lines
5.9 KiB
Go
222 lines
5.9 KiB
Go
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)
|
|
}
|
|
}
|