166 lines
4.2 KiB
Go
166 lines
4.2 KiB
Go
package repo
|
||
|
||
import (
|
||
"context"
|
||
"database/sql"
|
||
"errors"
|
||
"sync"
|
||
"testing"
|
||
|
||
"github.com/google/uuid"
|
||
)
|
||
|
||
// setupBudgetDB 复用 quota_test.go 的临时库,再建一个会话。
|
||
// 用真实 SQLite 而非 mock:预算的正确性核心是「判断与自增在同一条 UPDATE 里」,
|
||
// 那正是只有真实数据库才能验证的部分。
|
||
func setupBudgetDB(t *testing.T) uuid.UUID {
|
||
t.Helper()
|
||
setupTestDB(t)
|
||
id, err := CreateSession(context.Background(), nil, "bot", "预算测试", "")
|
||
if err != nil {
|
||
t.Fatalf("create session: %v", err)
|
||
}
|
||
return id
|
||
}
|
||
|
||
func TestSessionBudgetZeroMeansUnlimited(t *testing.T) {
|
||
id := setupBudgetDB(t)
|
||
ctx := context.Background()
|
||
|
||
// 默认 0 = 不限:引入预算不该把已在进行的会话卡死
|
||
b, err := GetSessionBudget(ctx, id)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if !b.Unlimited || b.Remaining != -1 {
|
||
t.Fatalf("默认应为不限:%+v", b)
|
||
}
|
||
// 不限时反复占用都成功
|
||
for i := 0; i < 5; i++ {
|
||
if _, err := ConsumeSessionBudget(ctx, id); err != nil {
|
||
t.Fatalf("不限额下第 %d 次占用失败: %v", i+1, err)
|
||
}
|
||
}
|
||
}
|
||
|
||
func TestSessionBudgetExhausts(t *testing.T) {
|
||
id := setupBudgetDB(t)
|
||
ctx := context.Background()
|
||
|
||
if _, err := SetSessionBudget(ctx, id, 2); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
for i := 1; i <= 2; i++ {
|
||
b, err := ConsumeSessionBudget(ctx, id)
|
||
if err != nil {
|
||
t.Fatalf("第 %d 次应成功: %v", i, err)
|
||
}
|
||
if b.Used != i {
|
||
t.Fatalf("第 %d 次后 used = %d", i, b.Used)
|
||
}
|
||
}
|
||
b, err := ConsumeSessionBudget(ctx, id)
|
||
if !errors.Is(err, ErrSessionBudgetExhausted) {
|
||
t.Fatalf("第 3 次应耗尽,得到 err=%v b=%+v", err, b)
|
||
}
|
||
if b.Remaining != 0 {
|
||
t.Fatalf("耗尽后剩余应为 0:%+v", b)
|
||
}
|
||
}
|
||
|
||
// 判断与自增必须在同一条 UPDATE 里,否则并发下会把预算刷穿
|
||
func TestSessionBudgetConcurrentDoesNotOverdraw(t *testing.T) {
|
||
id := setupBudgetDB(t)
|
||
ctx := context.Background()
|
||
const limit = 10
|
||
if _, err := SetSessionBudget(ctx, id, limit); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
|
||
var wg sync.WaitGroup
|
||
var mu sync.Mutex
|
||
ok := 0
|
||
for i := 0; i < 40; i++ {
|
||
wg.Add(1)
|
||
go func() {
|
||
defer wg.Done()
|
||
if _, err := ConsumeSessionBudget(ctx, id); err == nil {
|
||
mu.Lock()
|
||
ok++
|
||
mu.Unlock()
|
||
}
|
||
}()
|
||
}
|
||
wg.Wait()
|
||
|
||
if ok != limit {
|
||
t.Fatalf("40 并发下成功 %d 次,期望恰好 %d 次(预算被刷穿或误拒)", ok, limit)
|
||
}
|
||
b, _ := GetSessionBudget(ctx, id)
|
||
if b.Used != limit {
|
||
t.Fatalf("used = %d,期望 %d", b.Used, limit)
|
||
}
|
||
}
|
||
|
||
func TestSessionBudgetResetAndLowerBelowUsed(t *testing.T) {
|
||
id := setupBudgetDB(t)
|
||
ctx := context.Background()
|
||
|
||
SetSessionBudget(ctx, id, 5)
|
||
for i := 0; i < 3; i++ {
|
||
ConsumeSessionBudget(ctx, id)
|
||
}
|
||
|
||
// 调到低于已用次数 = 「就到这里为止」,是人的合法意图,不该报错
|
||
b, err := SetSessionBudget(ctx, id, 1)
|
||
if err != nil {
|
||
t.Fatalf("下调预算不该失败: %v", err)
|
||
}
|
||
if b.Remaining != 0 {
|
||
t.Fatalf("已用 3 上限 1 时剩余应为 0:%+v", b)
|
||
}
|
||
if _, err := ConsumeSessionBudget(ctx, id); !errors.Is(err, ErrSessionBudgetExhausted) {
|
||
t.Fatal("下调后应立即拦住")
|
||
}
|
||
|
||
// 重置只清已用次数,不动上限
|
||
b, err = ResetSessionBudget(ctx, id)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if b.Used != 0 || b.Max != 1 {
|
||
t.Fatalf("重置后应为 0/1:%+v", b)
|
||
}
|
||
}
|
||
|
||
// 会话预算先扣、全局配额后扣;全局拦下时必须把会话那次退回去,
|
||
// 否则那格白掉了 —— 那次往返实际上没有发生
|
||
func TestRefundSessionBudget(t *testing.T) {
|
||
id := setupBudgetDB(t)
|
||
ctx := context.Background()
|
||
SetSessionBudget(ctx, id, 3)
|
||
ConsumeSessionBudget(ctx, id)
|
||
|
||
RefundSessionBudget(ctx, id)
|
||
b, _ := GetSessionBudget(ctx, id)
|
||
if b.Used != 0 {
|
||
t.Fatalf("退还后 used 应为 0:%+v", b)
|
||
}
|
||
|
||
// 已经是 0 时再退不该变成负数
|
||
RefundSessionBudget(ctx, id)
|
||
b, _ = GetSessionBudget(ctx, id)
|
||
if b.Used != 0 {
|
||
t.Fatalf("重复退还把 used 变成了 %d", b.Used)
|
||
}
|
||
}
|
||
|
||
func TestGetSessionBudgetMissingSession(t *testing.T) {
|
||
setupBudgetDB(t)
|
||
if _, err := GetSessionBudget(context.Background(), uuid.New()); err == nil {
|
||
t.Fatal("不存在的会话应报错")
|
||
} else if errors.Is(err, sql.ErrNoRows) {
|
||
t.Fatal("应包装成可读错误而不是裸 sql.ErrNoRows")
|
||
}
|
||
}
|