mirror of
https://gitcode.com/JianFeeeee/ModelRouter.git
synced 2026-10-03 23:54:06 +00:00
per-key 配额桶是共享 map:每个被记录的请求写它,每个配额检查读它。 单线程单测完全看不到这里的竞态,只有让多个 goroutine 同时读写才有效。 key_quota_concurrency_test.go:64 goroutine 跑 2 秒,并发 Record + KeyWindowTokens + KeyWindowReqs + KeyWindowModelTokens,其中一条路径 在中途 opt in 惰性创建的 pinned 桶(那条路径一次改两个桶 map)。 go test -race 结果:零 DATA RACE,5444 万 token 全部入账。 全仓 -race(./...)亦全绿。
61 lines
1.7 KiB
Go
61 lines
1.7 KiB
Go
package gateway
|
|
|
|
import (
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
// The per-key quota buckets are shared maps mutated on every recorded request
|
|
// and read on every quota check. Single-threaded unit tests cannot see a race
|
|
// there at all, so this drives 64 goroutines through record + read together
|
|
// under `go test -race`. It also opts one model into the lazily-created pinned
|
|
// bucket mid-flight, which is the path that mutates two bucket maps at once.
|
|
func TestConcurrentQuotaAccounting(t *testing.T) {
|
|
s := NewStats(0)
|
|
keys := []string{"k0", "k1", "k2", "k3", "k4", "k5", "k6"}
|
|
models := []string{"m0", "m1", "m2", "AUTO"}
|
|
var wg sync.WaitGroup
|
|
stop := time.Now().Add(2 * time.Second)
|
|
|
|
for w := 0; w < 64; w++ {
|
|
wg.Add(1)
|
|
go func(w int) {
|
|
defer wg.Done()
|
|
k := keys[w%len(keys)]
|
|
m := models[w%len(models)]
|
|
now := time.Now()
|
|
for time.Now().Before(stop) {
|
|
s.Record(Req{Time: now.UnixMilli(), Key: k, Model: m, Source: "src",
|
|
Prompt: 100, Compl: 50, OK: true, Status: 200, Type: "chat"})
|
|
_ = s.KeyWindowTokens(k, 3600)
|
|
_ = s.KeyWindowReqs(k, 3600)
|
|
_ = s.KeyWindowModelTokens(k, m, "", 3600)
|
|
_ = s.KeyWindowModelTokens(k, m, "src", 3600) // opts into the pinned bucket
|
|
}
|
|
}(w)
|
|
}
|
|
wg.Wait()
|
|
|
|
// Every recorded request must be accounted for exactly once.
|
|
s.mu.Lock()
|
|
var total int64
|
|
for _, hm := range s.keyHour {
|
|
for _, v := range hm {
|
|
total += v
|
|
}
|
|
}
|
|
recs := int64(0)
|
|
for _, r := range s.recs {
|
|
recs += r.Prompt + r.Compl
|
|
}
|
|
s.mu.Unlock()
|
|
if total == 0 {
|
|
t.Fatal("no tokens accounted")
|
|
}
|
|
t.Logf("accounted %d tokens across %d keys", total, len(keys))
|
|
if recs == 0 {
|
|
t.Error("no records retained; the concurrent writers lost work")
|
|
}
|
|
}
|