mirror of
https://gitcode.com/JianFeeeee/ModelRouter.git
synced 2026-10-03 23:54:06 +00:00
feat(gateway): per-key 用量配额(token + 请求数)与重置周期
问题:密钥控制只能限制模型范围。实测发现三个缺陷,其中前两个让
per-model token_quota 在真实链路上从未生效:
1. 桶键不含 key。scopeTokens 调 WindowTokens(model, source, win),
桶键是 model / source::model,与调用方无关。实测两把 key 各用
1000 token,窗口报 2000 —— A key 的额度被 B key 消耗。
2. 无 source pin 的桶永远是空的。真实记录 Source 总被填上,桶键存成
"deepseek::m1",而无 pin 的查询找 "m1" —— 读到 0,永远 < quota,
配额形同虚设。实测 WindowTokens("m1","",1h)=0 而 pinned=2000。
3. AUTO scope 走 KeyTokens(key),是全时段累计、永不重置。实测 30 天
前的 200 token 仍计入 1 小时配额(报 210 而非 10)。配了
period: hour 也不会每小时归零。
生产 5 把 user key 全是 token_quota: 0,所以前两条一直没暴露。
改动:
- Stats 新增 per-key 小时桶 keyModelHour(key → model → hour)与
keyHour(key 总量)、keyReqHour(请求数),retention 40 天,与既有
modelHour 对齐以覆盖最长的 month 窗口;LoadAudit 走 aggregateLocked,
所以窗口用量跨重启存活。modelHour 保持 key-blind:它服务的是 AUTO
槽位配额(限制整个网关对某槽位的消耗),语义不同,不应被 per-key
改造污染。
- 每个请求写两份模型桶:裸 model 与 source::model。无 pin 的 scope
条目读前者,有 pin 的读后者。
- GWKey 新增 TokenQuota / ReqQuota / Period / Hours:整钥配额,
跨该 key 所有模型共享一份预算;ReqQuota 覆盖持续请求量(源上的
RPM 只管突发)。
- 配额耗尽返回 429 + Retry-After(rate_limit_exceeded),而不是 403:
403 让客户端以为这把 key 永远不能用该模型,直接放弃;429 + 等待
才能在窗口重置后自动恢复。模型越权仍是 403。
- admin key 永不受配额限制 —— 否则操作者会把自己锁在门外。
- 周期词表在写入时校验,拼错的 period 被拒绝而不是静默当成永不过期
(那与操作者输入的意图正好相反)。
- PUT /api/keys 的配额字段是指针:省略=保留原值,显式 0=解除限制。
否则只改模型范围就会悄悄清空预算。
判据 3 个文件 24 例,9 个变异全部被抓:key 隔离、pin 桶缺失、
AUTO 周期、key-blind 退化、429→403、admin 被限、PUT 清空配额、
Validate 失效、pinned 桶缺失。前三个变异最初漏网 —— 判据只测了
Stats 层没测接线,补了走真实 HTTP 的接线层与 API 层判据后抓住。
端到端验证:真实进程 + 加密配置往返,配额字段与 enc:v1 密钥均正常。
This commit is contained in:
@ -375,6 +375,82 @@ type GWKey struct {
|
|||||||
Note string `yaml:"note,omitempty" json:"note,omitempty"`
|
Note string `yaml:"note,omitempty" json:"note,omitempty"`
|
||||||
CreatedAt int64 `yaml:"created_at,omitempty" json:"created_at,omitempty"`
|
CreatedAt int64 `yaml:"created_at,omitempty" json:"created_at,omitempty"`
|
||||||
Seed bool `yaml:"seed,omitempty" json:"seed,omitempty"` // true if migrated from config gateway_keys
|
Seed bool `yaml:"seed,omitempty" json:"seed,omitempty"` // true if migrated from config gateway_keys
|
||||||
|
// TokenQuota caps this key's TOTAL tokens across every model it may use.
|
||||||
|
// 0 = unlimited. Period/Hours define the reset window, exactly like
|
||||||
|
// ModelScope: "" never resets, "hour"/"week"/"month" fixed windows,
|
||||||
|
// "nhour" uses Hours.
|
||||||
|
TokenQuota int64 `yaml:"token_quota,omitempty" json:"token_quota,omitempty"`
|
||||||
|
Period string `yaml:"period,omitempty" json:"period,omitempty"`
|
||||||
|
Hours int64 `yaml:"hours,omitempty" json:"hours,omitempty"`
|
||||||
|
// ReqQuota caps the number of requests per reset window; 0 = unlimited.
|
||||||
|
// RPM covers short bursts; this covers sustained volume.
|
||||||
|
ReqQuota int64 `yaml:"req_quota,omitempty" json:"req_quota,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// KeyQuota is the set of key-wide caps accepted by the admin API. It is a
|
||||||
|
// separate struct so a partial update can be expressed as a pointer (nil =
|
||||||
|
// "leave the stored caps alone") instead of zero values meaning "clear".
|
||||||
|
type KeyQuota struct {
|
||||||
|
TokenQuota int64 `json:"token_quota"`
|
||||||
|
ReqQuota int64 `json:"req_quota"`
|
||||||
|
Period string `json:"period"`
|
||||||
|
Hours int64 `json:"hours"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ApplyQuota writes the caps onto a key record.
|
||||||
|
func (k *GWKey) ApplyQuota(q KeyQuota) {
|
||||||
|
k.TokenQuota = q.TokenQuota
|
||||||
|
k.ReqQuota = q.ReqQuota
|
||||||
|
k.Period = q.Period
|
||||||
|
k.Hours = q.Hours
|
||||||
|
}
|
||||||
|
|
||||||
|
// NormalizeRole defaults an empty role to "user", so a key can never end up in
|
||||||
|
// a state where no role means "neither admin nor user".
|
||||||
|
func NormalizeRole(role string) string {
|
||||||
|
if role == "admin" {
|
||||||
|
return "admin"
|
||||||
|
}
|
||||||
|
return "user"
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate rejects a quota configuration that could not work as written. A
|
||||||
|
// period is only meaningful when at least one cap is set, and a cap of zero
|
||||||
|
// means "unlimited" rather than "deny everything", so those are the only two
|
||||||
|
// things worth rejecting.
|
||||||
|
func (q KeyQuota) Validate() error {
|
||||||
|
if q.TokenQuota < 0 {
|
||||||
|
return fmt.Errorf("token_quota must be >= 0 (0 = unlimited)")
|
||||||
|
}
|
||||||
|
if q.ReqQuota < 0 {
|
||||||
|
return fmt.Errorf("req_quota must be >= 0 (0 = unlimited)")
|
||||||
|
}
|
||||||
|
if q.Hours < 0 {
|
||||||
|
return fmt.Errorf("hours must be >= 0")
|
||||||
|
}
|
||||||
|
if q.TokenQuota > 0 || q.ReqQuota > 0 {
|
||||||
|
if err := ValidatePeriod(q.Period, q.Hours); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidatePeriod accepts the quota period vocabulary: "" (never resets),
|
||||||
|
// "hour", "week", "month", or "nhour" with hours >= 1. An unknown period is
|
||||||
|
// rejected rather than silently treated as "never resets", which would turn a
|
||||||
|
// typo into an all-time quota — the opposite of what the operator typed.
|
||||||
|
func ValidatePeriod(period string, hours int64) error {
|
||||||
|
switch period {
|
||||||
|
case "", "hour", "week", "month":
|
||||||
|
return nil
|
||||||
|
case "nhour":
|
||||||
|
if hours < 1 {
|
||||||
|
return fmt.Errorf("period %q needs hours >= 1", period)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return fmt.Errorf("period must be one of \"\", hour, week, month, nhour (got %q)", period)
|
||||||
}
|
}
|
||||||
|
|
||||||
// ModelScope is one allowed model for a key, or one AUTO scheduling slot,
|
// ModelScope is one allowed model for a key, or one AUTO scheduling slot,
|
||||||
|
|||||||
@ -250,6 +250,11 @@ func (c *Core) FindKey(key string) (config.GWKey, bool) {
|
|||||||
|
|
||||||
// CreateKey builds a new random gateway key and persists it to config.yaml.
|
// CreateKey builds a new random gateway key and persists it to config.yaml.
|
||||||
func (c *Core) CreateKey(name, role string, models []config.ModelScope, note string) (config.GWKey, error) {
|
func (c *Core) CreateKey(name, role string, models []config.ModelScope, note string) (config.GWKey, error) {
|
||||||
|
return c.CreateKeyWithQuota(name, role, models, note, config.KeyQuota{})
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateKeyWithQuota is CreateKey plus the key-wide token/request caps.
|
||||||
|
func (c *Core) CreateKeyWithQuota(name, role string, models []config.ModelScope, note string, q config.KeyQuota) (config.GWKey, error) {
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
defer c.mu.Unlock()
|
defer c.mu.Unlock()
|
||||||
models = cleanScopes(models)
|
models = cleanScopes(models)
|
||||||
@ -257,6 +262,9 @@ func (c *Core) CreateKey(name, role string, models []config.ModelScope, note str
|
|||||||
if _, err := rand.Read(key); err != nil {
|
if _, err := rand.Read(key); err != nil {
|
||||||
return config.GWKey{}, err
|
return config.GWKey{}, err
|
||||||
}
|
}
|
||||||
|
if err := q.Validate(); err != nil {
|
||||||
|
return config.GWKey{}, err
|
||||||
|
}
|
||||||
rec := config.GWKey{
|
rec := config.GWKey{
|
||||||
Key: "sk-gw-" + hex.EncodeToString(key),
|
Key: "sk-gw-" + hex.EncodeToString(key),
|
||||||
Role: role,
|
Role: role,
|
||||||
@ -265,9 +273,8 @@ func (c *Core) CreateKey(name, role string, models []config.ModelScope, note str
|
|||||||
Note: note,
|
Note: note,
|
||||||
CreatedAt: time.Now().Unix(),
|
CreatedAt: time.Now().Unix(),
|
||||||
}
|
}
|
||||||
if rec.Role == "" {
|
rec.Role = config.NormalizeRole(rec.Role)
|
||||||
rec.Role = "user"
|
rec.ApplyQuota(q)
|
||||||
}
|
|
||||||
c.cfg.Keys = append(c.cfg.Keys, rec)
|
c.cfg.Keys = append(c.cfg.Keys, rec)
|
||||||
if err := c.saveConfig(); err != nil {
|
if err := c.saveConfig(); err != nil {
|
||||||
return config.GWKey{}, err
|
return config.GWKey{}, err
|
||||||
@ -277,8 +284,20 @@ func (c *Core) CreateKey(name, role string, models []config.ModelScope, note str
|
|||||||
|
|
||||||
// UpdateKey mutates a key's name/role/model scope and persists it.
|
// UpdateKey mutates a key's name/role/model scope and persists it.
|
||||||
func (c *Core) UpdateKey(key, name, role string, models []config.ModelScope, note string) (config.GWKey, error) {
|
func (c *Core) UpdateKey(key, name, role string, models []config.ModelScope, note string) (config.GWKey, error) {
|
||||||
|
return c.UpdateKeyWithQuota(key, name, role, models, note, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateKeyWithQuota is UpdateKey plus the key-wide caps. quota == nil leaves
|
||||||
|
// the existing caps untouched, so a caller that only edits the model scope
|
||||||
|
// does not silently clear a key's budget.
|
||||||
|
func (c *Core) UpdateKeyWithQuota(key, name, role string, models []config.ModelScope, note string, quota *config.KeyQuota) (config.GWKey, error) {
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
defer c.mu.Unlock()
|
defer c.mu.Unlock()
|
||||||
|
if quota != nil {
|
||||||
|
if err := quota.Validate(); err != nil {
|
||||||
|
return config.GWKey{}, err
|
||||||
|
}
|
||||||
|
}
|
||||||
for i, k := range c.cfg.Keys {
|
for i, k := range c.cfg.Keys {
|
||||||
if k.Key == key {
|
if k.Key == key {
|
||||||
if name != "" {
|
if name != "" {
|
||||||
@ -293,6 +312,9 @@ func (c *Core) UpdateKey(key, name, role string, models []config.ModelScope, not
|
|||||||
c.cfg.Keys[i].Models = cleanScopes(models)
|
c.cfg.Keys[i].Models = cleanScopes(models)
|
||||||
}
|
}
|
||||||
c.cfg.Keys[i].Note = note
|
c.cfg.Keys[i].Note = note
|
||||||
|
if quota != nil {
|
||||||
|
c.cfg.Keys[i].ApplyQuota(*quota)
|
||||||
|
}
|
||||||
if err := c.saveConfig(); err != nil {
|
if err := c.saveConfig(); err != nil {
|
||||||
return config.GWKey{}, err
|
return config.GWKey{}, err
|
||||||
}
|
}
|
||||||
|
|||||||
@ -116,6 +116,12 @@ func (g *Gateway) apiV1Routes(w http.ResponseWriter, r *http.Request) {
|
|||||||
"note": k.Note,
|
"note": k.Note,
|
||||||
"created_at": k.CreatedAt,
|
"created_at": k.CreatedAt,
|
||||||
"seed": k.Seed,
|
"seed": k.Seed,
|
||||||
|
// key-wide spend caps (0 = unlimited). Echoed so an agent can
|
||||||
|
// see what budget it has without parsing config.yaml.
|
||||||
|
"token_quota": k.TokenQuota,
|
||||||
|
"req_quota": k.ReqQuota,
|
||||||
|
"period": k.Period,
|
||||||
|
"hours": k.Hours,
|
||||||
// The secret itself is never echoed. An operator that needs it
|
// The secret itself is never echoed. An operator that needs it
|
||||||
// already has it from creation time or from config.yaml.
|
// already has it from creation time or from config.yaml.
|
||||||
"key_prefix": maskKey(k.Key),
|
"key_prefix": maskKey(k.Key),
|
||||||
|
|||||||
@ -8,6 +8,7 @@ import (
|
|||||||
"log"
|
"log"
|
||||||
"net/http"
|
"net/http"
|
||||||
"regexp"
|
"regexp"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
@ -197,9 +198,33 @@ func intersectModels(models []string, allow []config.ModelScope) []string {
|
|||||||
// model is AUTO (the routing mode). It does NOT grant access to specific model
|
// model is AUTO (the routing mode). It does NOT grant access to specific model
|
||||||
// ids — that requires an explicit scope entry for the model.
|
// ids — that requires an explicit scope entry for the model.
|
||||||
func (g *Gateway) checkModelScope(ctx context.Context, model string) string {
|
func (g *Gateway) checkModelScope(ctx context.Context, model string) string {
|
||||||
|
if q := g.checkQuota(ctx, model); q != nil {
|
||||||
|
return q.msg
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// quotaRejection is a quota verdict: the message plus whether the client
|
||||||
|
// should retry. A spent quota is a rate limit (429 + Retry-After), not a
|
||||||
|
// permission failure (403): a client that sees 403 gives up on the key, while
|
||||||
|
// one that sees 429 with a retry hint waits and resumes when the window rolls
|
||||||
|
// over.
|
||||||
|
type quotaRejection struct {
|
||||||
|
msg string
|
||||||
|
retry int64 // seconds until the window resets; 0 = unknown
|
||||||
|
}
|
||||||
|
|
||||||
|
func (q *quotaRejection) Error() string { return q.msg }
|
||||||
|
|
||||||
|
// checkQuota is checkKeyScope for callers that need the retry hint. It
|
||||||
|
// separates the quota verdicts (429) from model-permission verdicts (403).
|
||||||
|
func (g *Gateway) checkQuota(ctx context.Context, model string) *quotaRejection {
|
||||||
|
if q := g.checkKeyQuotaRetry(ctx); q != nil {
|
||||||
|
return q
|
||||||
|
}
|
||||||
allow := g.allowedModels(ctx)
|
allow := g.allowedModels(ctx)
|
||||||
if allow == nil {
|
if allow == nil {
|
||||||
return ""
|
return nil
|
||||||
}
|
}
|
||||||
for _, sc := range allow {
|
for _, sc := range allow {
|
||||||
if sc.Model != model {
|
if sc.Model != model {
|
||||||
@ -208,23 +233,75 @@ func (g *Gateway) checkModelScope(ctx context.Context, model string) string {
|
|||||||
if sc.TokenQuota > 0 {
|
if sc.TokenQuota > 0 {
|
||||||
used := g.scopeTokens(ctx, sc)
|
used := g.scopeTokens(ctx, sc)
|
||||||
if used >= sc.TokenQuota {
|
if used >= sc.TokenQuota {
|
||||||
return fmt.Sprintf("token quota exceeded for %q (%d/%d)", model, used, sc.TokenQuota)
|
return "aRejection{
|
||||||
|
msg: fmt.Sprintf("token quota exceeded for %q (%d/%d)", model, used, sc.TokenQuota),
|
||||||
|
retry: AutoSecondsToReset(sc.Period, sc.Hours),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return ""
|
return nil
|
||||||
}
|
}
|
||||||
return fmt.Sprintf("model %q is not allowed for this key", model)
|
return "aRejection{msg: fmt.Sprintf("model %q is not allowed for this key", model)}
|
||||||
}
|
}
|
||||||
|
|
||||||
// scopeTokens returns the tokens a scope entry has consumed within its reset
|
// checkKeyQuotaRetry enforces the key-wide caps and reports the remaining
|
||||||
// window (total for AUTO / per model otherwise).
|
// seconds of the reset window so the caller can answer with 429 + Retry-After.
|
||||||
|
// An admin key is never capped, and a key with no caps set is never rejected.
|
||||||
|
func (g *Gateway) checkKeyQuotaRetry(ctx context.Context) *quotaRejection {
|
||||||
|
rec, ok := g.core.FindKey(reqKey(ctx))
|
||||||
|
if !ok || rec.Role == "admin" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
k := keyID(reqKey(ctx))
|
||||||
|
win := AutoPeriodSeconds(rec.Period, rec.Hours)
|
||||||
|
if rec.TokenQuota > 0 {
|
||||||
|
if used := g.stats.KeyWindowTokens(k, win); used >= rec.TokenQuota {
|
||||||
|
return "aRejection{
|
||||||
|
msg: fmt.Sprintf("key token quota exceeded (%d/%d%s)", used, rec.TokenQuota, quotaWindowSuffix(rec.Period, rec.Hours)),
|
||||||
|
retry: AutoSecondsToReset(rec.Period, rec.Hours),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if rec.ReqQuota > 0 {
|
||||||
|
if used := g.stats.KeyWindowReqs(k, win); used >= rec.ReqQuota {
|
||||||
|
return "aRejection{
|
||||||
|
msg: fmt.Sprintf("key request quota exceeded (%d/%d%s)", used, rec.ReqQuota, quotaWindowSuffix(rec.Period, rec.Hours)),
|
||||||
|
retry: AutoSecondsToReset(rec.Period, rec.Hours),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// quotaWindowSuffix describes a quota's reset window for an error message, so
|
||||||
|
// a rejected caller can tell a permanent block from one that clears in an hour.
|
||||||
|
func quotaWindowSuffix(period string, hours int64) string {
|
||||||
|
switch {
|
||||||
|
case period == "hour":
|
||||||
|
return ", resets hourly"
|
||||||
|
case period == "week":
|
||||||
|
return ", resets weekly"
|
||||||
|
case period == "month":
|
||||||
|
return ", resets monthly"
|
||||||
|
case period == "nhour" && hours > 1:
|
||||||
|
return fmt.Sprintf(", resets every %dh", hours)
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// scopeTokens returns the tokens this key consumed within the scope entry's
|
||||||
|
// reset window, isolated per key. For an AUTO entry the cap covers everything
|
||||||
|
// the key routed through AUTO; for a model entry it covers that model only.
|
||||||
|
//
|
||||||
|
// It reads the per-key hourly buckets rather than the key-blind model
|
||||||
|
// buckets, so one key's usage can never exhaust another's quota.
|
||||||
func (g *Gateway) scopeTokens(ctx context.Context, sc config.ModelScope) int64 {
|
func (g *Gateway) scopeTokens(ctx context.Context, sc config.ModelScope) int64 {
|
||||||
k := keyID(reqKey(ctx))
|
k := keyID(reqKey(ctx))
|
||||||
if sc.Model != "" && strings.EqualFold(sc.Model, "AUTO") {
|
|
||||||
return g.stats.KeyTokens(k)
|
|
||||||
}
|
|
||||||
win := AutoPeriodSeconds(sc.Period, sc.Hours)
|
win := AutoPeriodSeconds(sc.Period, sc.Hours)
|
||||||
return g.stats.WindowTokens(sc.Model, sc.Source, win)
|
if sc.Model != "" && strings.EqualFold(sc.Model, "AUTO") {
|
||||||
|
return g.stats.KeyWindowTokens(k, win)
|
||||||
|
}
|
||||||
|
return g.stats.KeyWindowModelTokens(k, sc.Model, sc.Source, win)
|
||||||
}
|
}
|
||||||
|
|
||||||
// hasScopeModel reports whether a model (possibly with a "source-model" /
|
// hasScopeModel reports whether a model (possibly with a "source-model" /
|
||||||
@ -248,6 +325,32 @@ func (g *Gateway) hasScopeModel(list []config.ModelScope, s string) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// writeScopeReject answers a model-scope or quota rejection. A spent quota is
|
||||||
|
// 429 (rate_limit_exceeded) with Retry-After, so a client waits and resumes
|
||||||
|
// after the reset; a model the key may not use stays 403 (model_not_allowed),
|
||||||
|
// because retrying cannot help.
|
||||||
|
func (g *Gateway) writeScopeReject(w http.ResponseWriter, r *http.Request, model string) {
|
||||||
|
if q := g.checkQuota(r.Context(), model); q != nil {
|
||||||
|
if q.retry > 0 {
|
||||||
|
w.Header().Set("Retry-After", strconv.FormatInt(q.retry, 10))
|
||||||
|
writeError(w, http.StatusTooManyRequests, "rate_limit_exceeded", q.msg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// no window to wait for: the cap is either permanent or key-wide
|
||||||
|
// with no period. "rate_limit_exceeded" still says "come back
|
||||||
|
// after the operator raises the cap", which 403 would not.
|
||||||
|
if strings.Contains(q.msg, "quota exceeded") {
|
||||||
|
writeError(w, http.StatusTooManyRequests, "rate_limit_exceeded", q.msg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
writeError(w, http.StatusForbidden, "model_not_allowed", q.msg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// The scope check already passed; reaching here means the state changed
|
||||||
|
// between the two calls. Fall back to the pre-existing behaviour.
|
||||||
|
writeError(w, http.StatusForbidden, "model_not_allowed", fmt.Sprintf("model %q is not allowed for this key", model))
|
||||||
|
}
|
||||||
|
|
||||||
func (g *Gateway) resolveByModel(model string) ([]*provider.Provider, string) {
|
func (g *Gateway) resolveByModel(model string) ([]*provider.Provider, string) {
|
||||||
if isAuto(model) {
|
if isAuto(model) {
|
||||||
return g.core.Registry().Resolve("AUTO"), ""
|
return g.core.Registry().Resolve("AUTO"), ""
|
||||||
@ -311,7 +414,7 @@ func (g *Gateway) handleChat(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
if msg := g.checkModelScope(r.Context(), "AUTO"); msg != "" {
|
if msg := g.checkModelScope(r.Context(), "AUTO"); msg != "" {
|
||||||
writeError(w, http.StatusForbidden, "model_not_allowed", msg)
|
g.writeScopeReject(w, r, "AUTO")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
ctx := r.Context()
|
ctx := r.Context()
|
||||||
@ -365,7 +468,7 @@ func (g *Gateway) handleChat(w http.ResponseWriter, r *http.Request) {
|
|||||||
effective = firstModel(cands[0])
|
effective = firstModel(cands[0])
|
||||||
}
|
}
|
||||||
if msg := g.checkModelScope(r.Context(), effective); msg != "" {
|
if msg := g.checkModelScope(r.Context(), effective); msg != "" {
|
||||||
writeError(w, http.StatusForbidden, "model_not_allowed", msg)
|
g.writeScopeReject(w, r, effective)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
ctx := r.Context()
|
ctx := r.Context()
|
||||||
@ -1022,7 +1125,7 @@ func (g *Gateway) handleImage(w http.ResponseWriter, r *http.Request) {
|
|||||||
if isAuto(model) {
|
if isAuto(model) {
|
||||||
if chain := g.core.AutoImageChain(); chain != nil && len(chain.Tiers) > 0 {
|
if chain := g.core.AutoImageChain(); chain != nil && len(chain.Tiers) > 0 {
|
||||||
if msg := g.checkModelScope(r.Context(), "AUTO"); msg != "" {
|
if msg := g.checkModelScope(r.Context(), "AUTO"); msg != "" {
|
||||||
writeError(w, http.StatusForbidden, "model_not_allowed", msg)
|
g.writeScopeReject(w, r, "AUTO")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
done := g.stats.Begin()
|
done := g.stats.Begin()
|
||||||
@ -1069,7 +1172,7 @@ func (g *Gateway) handleImage(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
if msg := g.checkModelScope(r.Context(), effectiveImageModel(model, cands)); msg != "" {
|
if msg := g.checkModelScope(r.Context(), effectiveImageModel(model, cands)); msg != "" {
|
||||||
writeError(w, http.StatusForbidden, "model_not_allowed", msg)
|
g.writeScopeReject(w, r, effectiveImageModel(model, cands))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
done := g.stats.Begin()
|
done := g.stats.Begin()
|
||||||
|
|||||||
230
internal/gateway/key_quota_api_test.go
Normal file
230
internal/gateway/key_quota_api_test.go
Normal file
@ -0,0 +1,230 @@
|
|||||||
|
package gateway
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"llmsproxy/internal/config"
|
||||||
|
"llmsproxy/internal/core"
|
||||||
|
)
|
||||||
|
|
||||||
|
// adminGateway builds a gateway whose admin key can call /api/keys.
|
||||||
|
func adminGateway(t *testing.T, keys ...config.GWKey) *Gateway {
|
||||||
|
t.Helper()
|
||||||
|
td := t.TempDir()
|
||||||
|
cfgPath := td + "/config.yaml"
|
||||||
|
if err := os.WriteFile(cfgPath, []byte("listen: :0"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
cfg := &config.Config{
|
||||||
|
Path: cfgPath,
|
||||||
|
AdapterDir: td + "/adapters",
|
||||||
|
RuntimeFile: td + "/runtime.json",
|
||||||
|
Keys: keys,
|
||||||
|
}
|
||||||
|
if err := cfg.ApplyDefaults(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
c, err := core.NewFromConfig(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("core: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(c.Close)
|
||||||
|
g, err := New(c, []string{"sk-admin"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("gateway: %v", err)
|
||||||
|
}
|
||||||
|
return g
|
||||||
|
}
|
||||||
|
|
||||||
|
func adminReq(t *testing.T, g *Gateway, method, path, body string) *httptest.ResponseRecorder {
|
||||||
|
t.Helper()
|
||||||
|
req, _ := http.NewRequest(method, path, strings.NewReader(body))
|
||||||
|
req.Header.Set("Authorization", "Bearer sk-admin")
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
rr := httptest.NewRecorder()
|
||||||
|
g.Handler().ServeHTTP(rr, req)
|
||||||
|
return rr
|
||||||
|
}
|
||||||
|
|
||||||
|
// keyRecord pulls one key's stored record out of the admin list.
|
||||||
|
func keyRecord(t *testing.T, g *Gateway, secret string) config.GWKey {
|
||||||
|
t.Helper()
|
||||||
|
rr := adminReq(t, g, "GET", "/api/keys", "")
|
||||||
|
if rr.Code != 200 {
|
||||||
|
t.Fatalf("GET /api/keys: %d %s", rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
var out struct {
|
||||||
|
Keys []config.GWKey `json:"keys"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(rr.Body.Bytes(), &out); err != nil {
|
||||||
|
t.Fatalf("decode: %v (%s)", err, rr.Body.String())
|
||||||
|
}
|
||||||
|
for _, k := range out.Keys {
|
||||||
|
if k.Key == secret {
|
||||||
|
return k
|
||||||
|
}
|
||||||
|
}
|
||||||
|
t.Fatalf("key %q not found in %s", secret, rr.Body.String())
|
||||||
|
return config.GWKey{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKeyAPIStoresQuota(t *testing.T) {
|
||||||
|
g := adminGateway(t, config.GWKey{Key: "sk-admin", Role: "admin"})
|
||||||
|
rr := adminReq(t, g, "POST", "/api/keys",
|
||||||
|
`{"name":"agent-x","role":"user","token_quota":50000,"req_quota":200,"period":"nhour","hours":6,"models":[{"model":"m1"}]}`)
|
||||||
|
if rr.Code != 200 {
|
||||||
|
t.Fatalf("create: %d %s", rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
var created struct {
|
||||||
|
Key config.GWKey `json:"key"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(rr.Body.Bytes(), &created); err != nil {
|
||||||
|
t.Fatalf("decode: %v", err)
|
||||||
|
}
|
||||||
|
if created.Key.TokenQuota != 50000 || created.Key.ReqQuota != 200 ||
|
||||||
|
created.Key.Period != "nhour" || created.Key.Hours != 6 {
|
||||||
|
t.Fatalf("created key did not carry the caps: %+v", created.Key)
|
||||||
|
}
|
||||||
|
// and it must survive a read-back (persisted, not just echoed)
|
||||||
|
back := keyRecord(t, g, created.Key.Key)
|
||||||
|
if back.TokenQuota != 50000 || back.Period != "nhour" || back.Hours != 6 {
|
||||||
|
t.Errorf("read-back lost the caps: %+v", back)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Editing only the model scope must not silently clear a key's budget: the
|
||||||
|
// caps are pointers precisely so "absent" is not "zero".
|
||||||
|
func TestKeyAPIUpdateKeepsQuotaWhenOmitted(t *testing.T) {
|
||||||
|
g := adminGateway(t, config.GWKey{Key: "sk-admin", Role: "admin"})
|
||||||
|
rr := adminReq(t, g, "POST", "/api/keys",
|
||||||
|
`{"name":"agent-x","role":"user","token_quota":50000,"period":"day-typo-free","models":[{"model":"m1"}]}`)
|
||||||
|
rr = adminReq(t, g, "POST", "/api/keys", `{"name":"y","role":"user","token_quota":50000,"period":"hour","models":[{"model":"m1"}]}`)
|
||||||
|
if rr.Code != 200 {
|
||||||
|
t.Fatalf("setup create: %d %s", rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
var created struct {
|
||||||
|
Key config.GWKey `json:"key"`
|
||||||
|
}
|
||||||
|
_ = json.Unmarshal(rr.Body.Bytes(), &created)
|
||||||
|
|
||||||
|
// a scope-only edit
|
||||||
|
rr = adminReq(t, g, "PUT", "/api/keys/"+created.Key.Key,
|
||||||
|
`{"name":"agent-y","models":[{"model":"m1"},{"model":"m2"}]}`)
|
||||||
|
if rr.Code != 200 {
|
||||||
|
t.Fatalf("update: %d %s", rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
back := keyRecord(t, g, created.Key.Key)
|
||||||
|
if back.TokenQuota != 50000 {
|
||||||
|
t.Errorf("token_quota was cleared by a scope-only edit: %d", back.TokenQuota)
|
||||||
|
}
|
||||||
|
if back.Period != "hour" {
|
||||||
|
t.Errorf("period was cleared by a scope-only edit: %q", back.Period)
|
||||||
|
}
|
||||||
|
if len(back.Models) != 2 {
|
||||||
|
t.Errorf("scope edit did not apply: %+v", back.Models)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sending 0 explicitly must lift the cap, not be treated as "absent".
|
||||||
|
func TestKeyAPIUpdateZeroLiftsCap(t *testing.T) {
|
||||||
|
g := adminGateway(t, config.GWKey{Key: "sk-admin", Role: "admin"})
|
||||||
|
rr := adminReq(t, g, "POST", "/api/keys", `{"name":"z","role":"user","token_quota":1000,"period":"hour"}`)
|
||||||
|
if rr.Code != 200 {
|
||||||
|
t.Fatalf("create: %d %s", rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
var created struct {
|
||||||
|
Key config.GWKey `json:"key"`
|
||||||
|
}
|
||||||
|
_ = json.Unmarshal(rr.Body.Bytes(), &created)
|
||||||
|
|
||||||
|
rr = adminReq(t, g, "PUT", "/api/keys/"+created.Key.Key, `{"token_quota":0}`)
|
||||||
|
if rr.Code != 200 {
|
||||||
|
t.Fatalf("lift: %d %s", rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
if back := keyRecord(t, g, created.Key.Key); back.TokenQuota != 0 {
|
||||||
|
t.Errorf("token_quota = %d, want 0 (cap lifted)", back.TokenQuota)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A misspelled period must be refused, not quietly turned into an all-time
|
||||||
|
// quota — which is the exact opposite of what the operator typed.
|
||||||
|
func TestKeyAPIRejectsBadPeriod(t *testing.T) {
|
||||||
|
g := adminGateway(t, config.GWKey{Key: "sk-admin", Role: "admin"})
|
||||||
|
rr := adminReq(t, g, "POST", "/api/keys", `{"name":"bad","role":"user","token_quota":1000,"period":"houre"}`)
|
||||||
|
if rr.Code != http.StatusBadRequest {
|
||||||
|
t.Fatalf("want 400 for a bad period, got %d %s", rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
if !strings.Contains(rr.Body.String(), "period") {
|
||||||
|
t.Errorf("error should name the period field: %s", rr.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKeyAPIRejectsNegativeQuota(t *testing.T) {
|
||||||
|
g := adminGateway(t, config.GWKey{Key: "sk-admin", Role: "admin"})
|
||||||
|
rr := adminReq(t, g, "POST", "/api/keys", `{"name":"bad","role":"user","token_quota":-5}`)
|
||||||
|
if rr.Code != http.StatusBadRequest {
|
||||||
|
t.Fatalf("want 400 for a negative quota, got %d %s", rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A non-admin key must not be able to set or read another key's budget.
|
||||||
|
func TestKeyAPIQuotaIsAdminOnly(t *testing.T) {
|
||||||
|
g := adminGateway(t,
|
||||||
|
config.GWKey{Key: "sk-admin", Role: "admin"},
|
||||||
|
config.GWKey{Key: "sk-u", Role: "user", TokenQuota: 10, Period: "hour"},
|
||||||
|
)
|
||||||
|
req, _ := http.NewRequest("POST", "/api/keys", strings.NewReader(`{"name":"x","role":"admin","token_quota":0}`))
|
||||||
|
req.Header.Set("Authorization", "Bearer sk-u")
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
rr := httptest.NewRecorder()
|
||||||
|
g.Handler().ServeHTTP(rr, req)
|
||||||
|
if rr.Code != http.StatusForbidden {
|
||||||
|
t.Fatalf("non-admin create: want 403, got %d %s", rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
// /api/v1/keys echoes the caps but never the secret
|
||||||
|
rr = adminReq(t, g, "GET", "/api/v1/keys", "")
|
||||||
|
if rr.Code != 200 {
|
||||||
|
t.Fatalf("GET /api/v1/keys: %d", rr.Code)
|
||||||
|
}
|
||||||
|
if strings.Contains(rr.Body.String(), "sk-u") {
|
||||||
|
t.Error("/api/v1/keys leaked a key secret")
|
||||||
|
}
|
||||||
|
if !strings.Contains(rr.Body.String(), `"token_quota":10`) {
|
||||||
|
t.Errorf("/api/v1/keys should expose the cap: %s", rr.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The AUTO scope entry must honour its reset window: usage that aged out of
|
||||||
|
// the window must not count against a per-key cap.
|
||||||
|
func TestAutoScopeQuotaHonoursWindow(t *testing.T) {
|
||||||
|
g, _ := quotaGateway(t,
|
||||||
|
config.GWKey{Key: "sk-a", Role: "user", Models: []config.ModelScope{
|
||||||
|
{Model: "AUTO", TokenQuota: 1000, Period: "hour"},
|
||||||
|
}},
|
||||||
|
config.GWKey{Key: "sk-b", Role: "user"},
|
||||||
|
)
|
||||||
|
ctx := quotaCtx(t, g, "sk-a")
|
||||||
|
sc := config.ModelScope{Model: "AUTO", TokenQuota: 1000, Period: "hour"}
|
||||||
|
|
||||||
|
// aged-out usage: 2 days old, 5M tokens — must be invisible to a 1h window
|
||||||
|
g.stats.Record(Req{Time: nowMSOffset(-48 * 3600 * 1000), Key: keyID("sk-a"),
|
||||||
|
Model: "m1", Source: "up", Prompt: 2500000, Compl: 2500000, OK: true, Status: 200})
|
||||||
|
if used := g.scopeTokens(ctx, sc); used != 0 {
|
||||||
|
t.Fatalf("AUTO scope saw %d tokens outside its 1h window; the period is being ignored", used)
|
||||||
|
}
|
||||||
|
|
||||||
|
// in-window usage counts
|
||||||
|
g.stats.Record(Req{Time: nowMSOffset(0), Key: keyID("sk-a"),
|
||||||
|
Model: "m1", Source: "up", Prompt: 400, Compl: 400, OK: true, Status: 200})
|
||||||
|
if used := g.scopeTokens(ctx, sc); used != 800 {
|
||||||
|
t.Fatalf("AUTO scope used = %d, want 800", used)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ = fmt.Sprintf
|
||||||
198
internal/gateway/key_quota_test.go
Normal file
198
internal/gateway/key_quota_test.go
Normal file
@ -0,0 +1,198 @@
|
|||||||
|
package gateway
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"llmsproxy/internal/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ---- data layer: per-key isolation ----
|
||||||
|
|
||||||
|
func TestKeyWindowTokensIsolatesKeys(t *testing.T) {
|
||||||
|
s := NewStats(100)
|
||||||
|
now := time.Now().UnixMilli()
|
||||||
|
s.Record(Req{Time: now, Key: "keyA", Model: "m1", Prompt: 500, Compl: 500, OK: true, Status: 200})
|
||||||
|
s.Record(Req{Time: now, Key: "keyB", Model: "m1", Prompt: 500, Compl: 500, OK: true, Status: 200})
|
||||||
|
|
||||||
|
if got := s.KeyWindowModelTokens("keyA", "m1", "", 3600); got != 1000 {
|
||||||
|
t.Errorf("keyA model tokens = %d, want 1000 (its own usage only)", got)
|
||||||
|
}
|
||||||
|
if got := s.KeyWindowModelTokens("keyB", "m1", "", 3600); got != 1000 {
|
||||||
|
t.Errorf("keyB model tokens = %d, want 1000", got)
|
||||||
|
}
|
||||||
|
// the key-blind bucket stays global on purpose: it backs the AUTO slot
|
||||||
|
// quota, which limits the whole gateway, not one key.
|
||||||
|
if got := s.WindowTokens("m1", "", 3600); got != 2000 {
|
||||||
|
t.Errorf("WindowTokens = %d, want 2000 (global per-model total must be unchanged)", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKeyWindowTokensSeparatesSourcePin(t *testing.T) {
|
||||||
|
s := NewStats(100)
|
||||||
|
now := time.Now().UnixMilli()
|
||||||
|
s.Record(Req{Time: now, Key: "keyA", Model: "m1", Source: "srcX", Prompt: 100, Compl: 100, OK: true, Status: 200}) // 200 tok
|
||||||
|
s.Record(Req{Time: now, Key: "keyA", Model: "m1", Source: "srcY", Prompt: 200, Compl: 200, OK: true, Status: 200}) // 400 tok
|
||||||
|
|
||||||
|
// A source pin scopes the cap to that one upstream.
|
||||||
|
if got := s.KeyWindowModelTokens("keyA", "m1", "srcX", 3600); got != 200 {
|
||||||
|
t.Errorf("srcX-pinned tokens = %d, want 200", got)
|
||||||
|
}
|
||||||
|
if got := s.KeyWindowModelTokens("keyA", "m1", "srcY", 3600); got != 400 {
|
||||||
|
t.Errorf("srcY-pinned tokens = %d, want 400", got)
|
||||||
|
}
|
||||||
|
// No pin covers the model on every source: a cap on "this model" must not
|
||||||
|
// stop applying just because the request was served elsewhere. 200+400.
|
||||||
|
if got := s.KeyWindowModelTokens("keyA", "m1", "", 3600); got != 600 {
|
||||||
|
t.Errorf("unpinned model tokens = %d, want 600 (both sources)", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// This is the shape of every real chat record: Source is always populated.
|
||||||
|
// The unpinned bucket must still see the tokens, or a per-model quota without
|
||||||
|
// a source pin reads an empty bucket and never trips.
|
||||||
|
func TestKeyWindowModelTokensUnpinnedSeesSourcedTraffic(t *testing.T) {
|
||||||
|
s := NewStats(100)
|
||||||
|
s.Record(Req{Time: time.Now().UnixMilli(), Key: "keyA", Model: "m1", Source: "deepseek",
|
||||||
|
Prompt: 1000, Compl: 1000, OK: true, Status: 200})
|
||||||
|
|
||||||
|
if got := s.KeyWindowModelTokens("keyA", "m1", "", 3600); got != 2000 {
|
||||||
|
t.Errorf("unpinned tokens = %d, want 2000 — an unpinned per-model quota would never trip otherwise", got)
|
||||||
|
}
|
||||||
|
if got := s.KeyWindowModelTokens("keyA", "m1", "deepseek", 3600); got != 2000 {
|
||||||
|
t.Errorf("pinned tokens = %d, want 2000", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKeyWindowRespectsResetWindow(t *testing.T) {
|
||||||
|
s := NewStats(100)
|
||||||
|
old := time.Now().Add(-30 * 24 * time.Hour).UnixMilli()
|
||||||
|
s.Record(Req{Time: old, Key: "keyA", Model: "m1", Prompt: 100, Compl: 100, OK: true, Status: 200})
|
||||||
|
s.Record(Req{Time: time.Now().UnixMilli(), Key: "keyA", Model: "m1", Prompt: 5, Compl: 5, OK: true, Status: 200})
|
||||||
|
|
||||||
|
if got := s.KeyWindowTokens("keyA", 3600); got != 10 {
|
||||||
|
t.Errorf("1h window = %d, want 10 (30-day-old usage must not count)", got)
|
||||||
|
}
|
||||||
|
if got := s.KeyWindowTokens("keyA", 0); got != 210 {
|
||||||
|
t.Errorf("all-time total = %d, want 210", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKeyWindowReqsCountsEveryRequest(t *testing.T) {
|
||||||
|
s := NewStats(100)
|
||||||
|
now := time.Now().UnixMilli()
|
||||||
|
// a failed request and an image request both count: a client that loops on
|
||||||
|
// failures must still burn its request quota
|
||||||
|
s.Record(Req{Time: now, Key: "keyA", Model: "m1", Type: "chat", OK: true, Status: 200})
|
||||||
|
s.Record(Req{Time: now, Key: "keyA", Model: "m1", Type: "chat", OK: false, Status: 500})
|
||||||
|
s.Record(Req{Time: now, Key: "keyA", Model: "img", Type: "image", OK: true, Status: 200})
|
||||||
|
s.Record(Req{Time: now, Key: "keyB", Model: "m1", Type: "chat", OK: true, Status: 200})
|
||||||
|
|
||||||
|
if got := s.KeyWindowReqs("keyA", 3600); got != 3 {
|
||||||
|
t.Errorf("keyA reqs = %d, want 3 (chat ok + chat fail + image)", got)
|
||||||
|
}
|
||||||
|
if got := s.KeyWindowReqs("keyB", 3600); got != 1 {
|
||||||
|
t.Errorf("keyB reqs = %d, want 1", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKeyWindowTokensZeroTokensNotCounted(t *testing.T) {
|
||||||
|
s := NewStats(100)
|
||||||
|
now := time.Now().UnixMilli()
|
||||||
|
// a request that reported no usage must not create a bucket entry
|
||||||
|
s.Record(Req{Time: now, Key: "keyA", Model: "", Type: "chat", OK: true, Status: 200})
|
||||||
|
if got := s.KeyWindowTokens("keyA", 3600); got != 0 {
|
||||||
|
t.Errorf("tokens = %d, want 0", got)
|
||||||
|
}
|
||||||
|
if got := s.KeyWindowReqs("keyA", 3600); got != 1 {
|
||||||
|
t.Errorf("reqs = %d, want 1 (request still happened)", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- reset window arithmetic ----
|
||||||
|
|
||||||
|
func TestAutoSecondsToReset(t *testing.T) {
|
||||||
|
if got := AutoSecondsToReset("", 0); got != 0 {
|
||||||
|
t.Errorf("no period = %d, want 0 (never resets -> no retry hint)", got)
|
||||||
|
}
|
||||||
|
now := time.Now().Unix()
|
||||||
|
for _, tc := range []struct {
|
||||||
|
period string
|
||||||
|
hours int64
|
||||||
|
want int64
|
||||||
|
}{
|
||||||
|
{"hour", 0, 3600},
|
||||||
|
{"week", 0, 7 * 24 * 3600},
|
||||||
|
{"month", 0, 30 * 24 * 3600},
|
||||||
|
{"nhour", 6, 6 * 3600},
|
||||||
|
} {
|
||||||
|
got := AutoSecondsToReset(tc.period, tc.hours)
|
||||||
|
if got <= 0 || got > tc.want {
|
||||||
|
t.Errorf("AutoSecondsToReset(%q,%d) = %d, want in (0,%d]", tc.period, tc.hours, got, tc.want)
|
||||||
|
}
|
||||||
|
// must never exceed the window itself
|
||||||
|
if got < tc.want-now%3600 {
|
||||||
|
t.Logf("note: %q hint %ds < remaining-window %ds (rounds to hour boundary)", tc.period, got, tc.want-now%3600)
|
||||||
|
}
|
||||||
|
if got > tc.want {
|
||||||
|
t.Errorf("hint %d exceeds window %d", got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
_ = now
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAutoSecondsToResetNeverExceedsWindow(t *testing.T) {
|
||||||
|
// "nhour" with a tiny window must not hand out a longer wait than the
|
||||||
|
// window itself (which would stall a client past its own reset)
|
||||||
|
for hours := int64(1); hours <= 48; hours++ {
|
||||||
|
got := AutoSecondsToReset("nhour", hours)
|
||||||
|
if got <= 0 || got > hours*3600 {
|
||||||
|
t.Errorf("nhour/%d = %d, want in (0,%d]", hours, got, hours*3600)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- config validation ----
|
||||||
|
|
||||||
|
func TestKeyQuotaValidate(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
q config.KeyQuota
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"unlimited", config.KeyQuota{}, false},
|
||||||
|
{"tokens hourly", config.KeyQuota{TokenQuota: 1000, Period: "hour"}, false},
|
||||||
|
{"tokens n-hour", config.KeyQuota{TokenQuota: 1000, Period: "nhour", Hours: 6}, false},
|
||||||
|
{"reqs weekly", config.KeyQuota{ReqQuota: 100, Period: "week"}, false},
|
||||||
|
{"no period never resets", config.KeyQuota{TokenQuota: 1000}, false},
|
||||||
|
{"negative tokens", config.KeyQuota{TokenQuota: -1}, true},
|
||||||
|
{"negative reqs", config.KeyQuota{ReqQuota: -1}, true},
|
||||||
|
{"negative hours", config.KeyQuota{TokenQuota: 5, Hours: -1}, true},
|
||||||
|
{"typo period becomes all-time if accepted", config.KeyQuota{TokenQuota: 5, Period: "houre"}, true},
|
||||||
|
{"nhour without hours", config.KeyQuota{TokenQuota: 5, Period: "nhour"}, true},
|
||||||
|
{"nhour with 0 hours", config.KeyQuota{TokenQuota: 5, Period: "nhour", Hours: 0}, true},
|
||||||
|
{"junk period with only req quota", config.KeyQuota{ReqQuota: 5, Period: "daily"}, true},
|
||||||
|
{"junk period with no caps is irrelevant", config.KeyQuota{Period: "daily"}, false},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
err := tc.q.Validate()
|
||||||
|
if tc.wantErr && err == nil {
|
||||||
|
t.Errorf("%s: expected error, got nil", tc.name)
|
||||||
|
}
|
||||||
|
if !tc.wantErr && err != nil {
|
||||||
|
t.Errorf("%s: unexpected error: %v", tc.name, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeRole(t *testing.T) {
|
||||||
|
if got := config.NormalizeRole(""); got != "user" {
|
||||||
|
t.Errorf("empty role = %q, want user", got)
|
||||||
|
}
|
||||||
|
if got := config.NormalizeRole("admin"); got != "admin" {
|
||||||
|
t.Errorf("admin role = %q, want admin", got)
|
||||||
|
}
|
||||||
|
if got := config.NormalizeRole("root"); got != "user" {
|
||||||
|
t.Errorf("unknown role = %q, want user (never escalate)", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
247
internal/gateway/key_quota_wiring_test.go
Normal file
247
internal/gateway/key_quota_wiring_test.go
Normal file
@ -0,0 +1,247 @@
|
|||||||
|
package gateway
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"llmsproxy/internal/config"
|
||||||
|
"llmsproxy/internal/core"
|
||||||
|
)
|
||||||
|
|
||||||
|
// quotaGateway builds a gateway with two user keys that differ only in their
|
||||||
|
// configured caps, against one mock upstream that reports 4 tokens per call.
|
||||||
|
func quotaGateway(t *testing.T, keyA, keyB config.GWKey) (*Gateway, *upstreamCtrl) {
|
||||||
|
t.Helper()
|
||||||
|
ctrl := &upstreamCtrl{}
|
||||||
|
up := upstream(t, ctrl)
|
||||||
|
t.Cleanup(up.Close)
|
||||||
|
|
||||||
|
td := t.TempDir()
|
||||||
|
cfgPath := td + "/config.yaml"
|
||||||
|
if err := os.WriteFile(cfgPath, []byte("listen: :0"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
cfg := &config.Config{
|
||||||
|
Path: cfgPath,
|
||||||
|
AdapterDir: td + "/adapters",
|
||||||
|
RuntimeFile: td + "/runtime.json",
|
||||||
|
Keys: []config.GWKey{keyA, keyB},
|
||||||
|
Sources: []config.Source{{
|
||||||
|
Name: "up", BaseURL: up.URL, Adapter: "openai",
|
||||||
|
Models: []config.Model{{ID: "m1", Priority: 100}},
|
||||||
|
}},
|
||||||
|
}
|
||||||
|
if err := cfg.ApplyDefaults(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
c, err := core.NewFromConfig(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("core: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(c.Close)
|
||||||
|
g, err := New(c, []string{keyA.Key, keyB.Key})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("gateway: %v", err)
|
||||||
|
}
|
||||||
|
return g, ctrl
|
||||||
|
}
|
||||||
|
|
||||||
|
// chatAs issues a chat completion with a specific gateway key and returns the
|
||||||
|
// recorder plus the decoded error code.
|
||||||
|
func chatAs(t *testing.T, g *Gateway, key, model string) (*httptest.ResponseRecorder, string) {
|
||||||
|
t.Helper()
|
||||||
|
req, _ := http.NewRequest("POST", "/v1/chat/completions",
|
||||||
|
strings.NewReader(fmt.Sprintf(`{"model":%q,"messages":[{"role":"user","content":"hi"}]}`, model)))
|
||||||
|
req.Header.Set("Authorization", "Bearer "+key)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
rr := httptest.NewRecorder()
|
||||||
|
g.Handler().ServeHTTP(rr, req)
|
||||||
|
code := ""
|
||||||
|
if rr.Code != 200 {
|
||||||
|
var e struct {
|
||||||
|
Error struct {
|
||||||
|
Type string `json:"type"`
|
||||||
|
} `json:"error"`
|
||||||
|
}
|
||||||
|
_ = json.Unmarshal(rr.Body.Bytes(), &e)
|
||||||
|
code = e.Error.Type
|
||||||
|
}
|
||||||
|
return rr, code
|
||||||
|
}
|
||||||
|
|
||||||
|
// A key whose token budget is spent must be refused, and the refusal must be a
|
||||||
|
// 429 the client can retry after the window rolls over — not a 403 that reads
|
||||||
|
// as "this key may never use this model".
|
||||||
|
func TestKeyTokenQuotaBlocksWithRetryAfter(t *testing.T) {
|
||||||
|
g, _ := quotaGateway(t,
|
||||||
|
config.GWKey{Key: "sk-a", Role: "user", TokenQuota: 8, Period: "hour",
|
||||||
|
Models: []config.ModelScope{{Model: "m1"}}},
|
||||||
|
config.GWKey{Key: "sk-b", Role: "user", Name: "b"},
|
||||||
|
)
|
||||||
|
// 4 tokens per call, budget 8 -> the third call crosses it
|
||||||
|
for i := 1; i <= 2; i++ {
|
||||||
|
rr, _ := chatAs(t, g, "sk-a", "m1")
|
||||||
|
if rr.Code != 200 {
|
||||||
|
t.Fatalf("call %d: want 200, got %d (%s)", i, rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
rr, code := chatAs(t, g, "sk-a", "m1")
|
||||||
|
if rr.Code != http.StatusTooManyRequests {
|
||||||
|
t.Fatalf("third call: want 429, got %d (%s)", rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
if code != "rate_limit_exceeded" {
|
||||||
|
t.Errorf("error code = %q, want rate_limit_exceeded", code)
|
||||||
|
}
|
||||||
|
if ra := rr.Header().Get("Retry-After"); ra == "" {
|
||||||
|
t.Error("429 must carry Retry-After so a client knows when to come back")
|
||||||
|
} else if n := mustAtoi(t, ra); n <= 0 || n > 3600 {
|
||||||
|
t.Errorf("Retry-After = %q, want a positive value within the 1h window", ra)
|
||||||
|
}
|
||||||
|
if !strings.Contains(rr.Body.String(), "quota") {
|
||||||
|
t.Errorf("body should name the quota: %s", rr.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The quota is per key: exhausting one key's budget must not affect another
|
||||||
|
// key that is allowed the same model.
|
||||||
|
func TestKeyTokenQuotaIsIsolatedPerKey(t *testing.T) {
|
||||||
|
g, _ := quotaGateway(t,
|
||||||
|
config.GWKey{Key: "sk-a", Role: "user", TokenQuota: 4, Period: "hour",
|
||||||
|
Models: []config.ModelScope{{Model: "m1"}}},
|
||||||
|
config.GWKey{Key: "sk-b", Role: "user", TokenQuota: 1000, Period: "hour",
|
||||||
|
Models: []config.ModelScope{{Model: "m1"}}},
|
||||||
|
)
|
||||||
|
rr, _ := chatAs(t, g, "sk-a", "m1")
|
||||||
|
if rr.Code != 200 {
|
||||||
|
t.Fatalf("keyA first call: want 200, got %d", rr.Code)
|
||||||
|
}
|
||||||
|
if rr, code := chatAs(t, g, "sk-a", "m1"); rr.Code != http.StatusTooManyRequests {
|
||||||
|
t.Fatalf("keyA second call: want 429, got %d (%s)", rr.Code, code)
|
||||||
|
}
|
||||||
|
// keyB has its own budget and must still be served
|
||||||
|
if rr, _ := chatAs(t, g, "sk-b", "m1"); rr.Code != 200 {
|
||||||
|
t.Fatalf("keyB must be unaffected by keyA's exhausted quota, got %d", rr.Code)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A per-model quota inside the scope must also be per key: one key burning a
|
||||||
|
// shared model budget must not lock the other key out.
|
||||||
|
func TestScopeModelTokenQuotaIsolatedPerKey(t *testing.T) {
|
||||||
|
g, _ := quotaGateway(t,
|
||||||
|
config.GWKey{Key: "sk-a", Role: "user", Models: []config.ModelScope{
|
||||||
|
{Model: "m1", TokenQuota: 4, Period: "hour"},
|
||||||
|
}},
|
||||||
|
config.GWKey{Key: "sk-b", Role: "user", Models: []config.ModelScope{
|
||||||
|
{Model: "m1", TokenQuota: 4, Period: "hour"},
|
||||||
|
}},
|
||||||
|
)
|
||||||
|
if rr, _ := chatAs(t, g, "sk-a", "m1"); rr.Code != 200 {
|
||||||
|
t.Fatalf("keyA first call: want 200, got %d", rr.Code)
|
||||||
|
}
|
||||||
|
if rr, _ := chatAs(t, g, "sk-a", "m1"); rr.Code != http.StatusTooManyRequests {
|
||||||
|
t.Fatalf("keyA second call: want 429 (its own model quota), got %d", rr.Code)
|
||||||
|
}
|
||||||
|
// keyB shares the model but not the counter — it must be served twice too
|
||||||
|
if rr, _ := chatAs(t, g, "sk-b", "m1"); rr.Code != 200 {
|
||||||
|
t.Fatalf("keyB first call: want 200, got %d", rr.Code)
|
||||||
|
}
|
||||||
|
if rr, _ := chatAs(t, g, "sk-b", "m1"); rr.Code != http.StatusTooManyRequests {
|
||||||
|
t.Fatalf("keyB second call: want 429, got %d", rr.Code)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// An admin key is never capped: a budget set on an admin key must not lock the
|
||||||
|
// operator out of the gateway they administer.
|
||||||
|
func TestAdminKeyIsNeverQuotaCapped(t *testing.T) {
|
||||||
|
g, _ := quotaGateway(t,
|
||||||
|
config.GWKey{Key: "sk-admin", Role: "admin", TokenQuota: 1, Period: "hour", ReqQuota: 1},
|
||||||
|
config.GWKey{Key: "sk-b", Role: "user"},
|
||||||
|
)
|
||||||
|
for i := 1; i <= 3; i++ {
|
||||||
|
if rr, _ := chatAs(t, g, "sk-admin", "m1"); rr.Code != 200 {
|
||||||
|
t.Fatalf("admin call %d: want 200 (admin keys are uncapped), got %d", i, rr.Code)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A request quota caps sustained volume; RPM-style bursts are not its job but
|
||||||
|
// the count must still stop the key.
|
||||||
|
func TestKeyRequestQuotaBlocks(t *testing.T) {
|
||||||
|
g, ctrl := quotaGateway(t,
|
||||||
|
config.GWKey{Key: "sk-a", Role: "user", ReqQuota: 2, Period: "hour",
|
||||||
|
Models: []config.ModelScope{{Model: "m1"}}},
|
||||||
|
config.GWKey{Key: "sk-b", Role: "user"},
|
||||||
|
)
|
||||||
|
for i := 1; i <= 2; i++ {
|
||||||
|
if rr, _ := chatAs(t, g, "sk-a", "m1"); rr.Code != 200 {
|
||||||
|
t.Fatalf("call %d: want 200, got %d", i, rr.Code)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
rr, code := chatAs(t, g, "sk-a", "m1")
|
||||||
|
if rr.Code != http.StatusTooManyRequests || code != "rate_limit_exceeded" {
|
||||||
|
t.Fatalf("third call: want 429/rate_limit_exceeded, got %d/%s", rr.Code, code)
|
||||||
|
}
|
||||||
|
if ctrl.hits != 2 {
|
||||||
|
t.Errorf("upstream saw %d calls, want 2 (a refused request must not reach upstream)", ctrl.hits)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A model outside the scope is still 403, not 429: retrying cannot help.
|
||||||
|
func TestModelOutsideScopeStaysForbidden(t *testing.T) {
|
||||||
|
g, _ := quotaGateway(t,
|
||||||
|
config.GWKey{Key: "sk-a", Role: "user", TokenQuota: 1000, Period: "hour",
|
||||||
|
Models: []config.ModelScope{{Model: "other-model"}}},
|
||||||
|
config.GWKey{Key: "sk-b", Role: "user"},
|
||||||
|
)
|
||||||
|
rr, code := chatAs(t, g, "sk-a", "m1")
|
||||||
|
if rr.Code != http.StatusForbidden {
|
||||||
|
t.Fatalf("want 403, got %d (%s)", rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
if code != "model_not_allowed" {
|
||||||
|
t.Errorf("error code = %q, want model_not_allowed", code)
|
||||||
|
}
|
||||||
|
if rr.Header().Get("Retry-After") != "" {
|
||||||
|
t.Error("a 403 must not advertise a retry time")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// An unlimited key (no caps at all) is never rejected — the default config
|
||||||
|
// must keep working exactly as before.
|
||||||
|
func TestKeyWithoutCapsIsUnlimited(t *testing.T) {
|
||||||
|
g, _ := quotaGateway(t,
|
||||||
|
config.GWKey{Key: "sk-a", Role: "user", Models: []config.ModelScope{{Model: "m1"}}},
|
||||||
|
config.GWKey{Key: "sk-b", Role: "user"},
|
||||||
|
)
|
||||||
|
for i := 1; i <= 5; i++ {
|
||||||
|
if rr, _ := chatAs(t, g, "sk-a", "m1"); rr.Code != 200 {
|
||||||
|
t.Fatalf("call %d: want 200 (no caps = unlimited), got %d", i, rr.Code)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func mustAtoi(t *testing.T, s string) int64 {
|
||||||
|
t.Helper()
|
||||||
|
var n int64
|
||||||
|
if _, err := fmt.Sscanf(s, "%d", &n); err != nil {
|
||||||
|
t.Fatalf("Retry-After %q is not an integer: %v", s, err)
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
// quotaCtx builds a request context carrying a gateway key, as the auth
|
||||||
|
// middleware would.
|
||||||
|
func quotaCtx(t *testing.T, g *Gateway, key string) context.Context {
|
||||||
|
t.Helper()
|
||||||
|
return withAuth(context.Background(), key, "user")
|
||||||
|
}
|
||||||
|
|
||||||
|
func nowMSOffset(sec int64) int64 {
|
||||||
|
return time.Now().Add(time.Duration(sec) * time.Second).UnixMilli()
|
||||||
|
}
|
||||||
@ -39,23 +39,27 @@ func (g *Gateway) handleKeysAPI(w http.ResponseWriter, r *http.Request) {
|
|||||||
writeJSON(w, http.StatusOK, map[string]interface{}{"keys": g.core.ListKeys()})
|
writeJSON(w, http.StatusOK, map[string]interface{}{"keys": g.core.ListKeys()})
|
||||||
case http.MethodPost:
|
case http.MethodPost:
|
||||||
var body struct {
|
var body struct {
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
Role string `json:"role"`
|
Role string `json:"role"`
|
||||||
Models []config.ModelScope `json:"models"`
|
Models []config.ModelScope `json:"models"`
|
||||||
Note string `json:"note"`
|
Note string `json:"note"`
|
||||||
|
TokenQuota *int64 `json:"token_quota"`
|
||||||
|
ReqQuota *int64 `json:"req_quota"`
|
||||||
|
Period *string `json:"period"`
|
||||||
|
Hours *int64 `json:"hours"`
|
||||||
}
|
}
|
||||||
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
||||||
writeError(w, http.StatusBadRequest, "invalid_request", "invalid json: "+err.Error())
|
writeError(w, http.StatusBadRequest, "invalid_request", "invalid json: "+err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if body.Role == "" {
|
body.Role = config.NormalizeRole(body.Role)
|
||||||
body.Role = "user"
|
q := config.KeyQuota{
|
||||||
|
TokenQuota: optInt64(body.TokenQuota),
|
||||||
|
ReqQuota: optInt64(body.ReqQuota),
|
||||||
|
Period: optString(body.Period),
|
||||||
|
Hours: optInt64(body.Hours),
|
||||||
}
|
}
|
||||||
if body.Role != "admin" && body.Role != "user" {
|
rec, err := g.core.CreateKeyWithQuota(body.Name, body.Role, body.Models, body.Note, q)
|
||||||
writeError(w, http.StatusBadRequest, "invalid_request", "role must be admin or user")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
rec, err := g.core.CreateKey(body.Name, body.Role, body.Models, body.Note)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
writeError(w, http.StatusBadRequest, "key_error", err.Error())
|
writeError(w, http.StatusBadRequest, "key_error", err.Error())
|
||||||
return
|
return
|
||||||
@ -67,16 +71,33 @@ func (g *Gateway) handleKeysAPI(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
var body struct {
|
var body struct {
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
Role string `json:"role"`
|
Role string `json:"role"`
|
||||||
Models []config.ModelScope `json:"models"`
|
Models []config.ModelScope `json:"models"`
|
||||||
Note string `json:"note"`
|
Note string `json:"note"`
|
||||||
|
TokenQuota *int64 `json:"token_quota"`
|
||||||
|
ReqQuota *int64 `json:"req_quota"`
|
||||||
|
Period *string `json:"period"`
|
||||||
|
Hours *int64 `json:"hours"`
|
||||||
}
|
}
|
||||||
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
||||||
writeError(w, http.StatusBadRequest, "invalid_request", "invalid json: "+err.Error())
|
writeError(w, http.StatusBadRequest, "invalid_request", "invalid json: "+err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
rec, err := g.core.UpdateKey(path, body.Name, body.Role, body.Models, body.Note)
|
// Quota fields are pointers so "absent" is distinguishable from
|
||||||
|
// "set to 0": omitting them leaves the stored caps alone, while
|
||||||
|
// sending 0 explicitly lifts a cap. Without this, editing only the
|
||||||
|
// model scope would silently clear a key's budget.
|
||||||
|
var q *config.KeyQuota
|
||||||
|
if body.TokenQuota != nil || body.ReqQuota != nil || body.Period != nil || body.Hours != nil {
|
||||||
|
q = &config.KeyQuota{
|
||||||
|
TokenQuota: optInt64(body.TokenQuota),
|
||||||
|
ReqQuota: optInt64(body.ReqQuota),
|
||||||
|
Period: optString(body.Period),
|
||||||
|
Hours: optInt64(body.Hours),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
rec, err := g.core.UpdateKeyWithQuota(path, body.Name, body.Role, body.Models, body.Note, q)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
writeError(w, http.StatusBadRequest, "key_error", err.Error())
|
writeError(w, http.StatusBadRequest, "key_error", err.Error())
|
||||||
return
|
return
|
||||||
@ -106,6 +127,22 @@ func (g *Gateway) handleKeysAPI(w http.ResponseWriter, r *http.Request) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// optInt64 dereferences an optional quota field, treating absent as 0.
|
||||||
|
func optInt64(p *int64) int64 {
|
||||||
|
if p == nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return *p
|
||||||
|
}
|
||||||
|
|
||||||
|
// optString dereferences an optional quota field, treating absent as "".
|
||||||
|
func optString(p *string) string {
|
||||||
|
if p == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return *p
|
||||||
|
}
|
||||||
|
|
||||||
// handleKeyMe returns the authenticated key's own record (users see only
|
// handleKeyMe returns the authenticated key's own record (users see only
|
||||||
// themselves; admins can use this as a convenience too).
|
// themselves; admins can use this as a convenience too).
|
||||||
func (g *Gateway) handleKeyMe(w http.ResponseWriter, r *http.Request) {
|
func (g *Gateway) handleKeyMe(w http.ResponseWriter, r *http.Request) {
|
||||||
|
|||||||
@ -84,6 +84,18 @@ type Stats struct {
|
|||||||
auditPath string
|
auditPath string
|
||||||
replayPartial bool // aggregates built from a bounded audit tail
|
replayPartial bool // aggregates built from a bounded audit tail
|
||||||
modelHour map[string]map[int64]int64 // model -> unix-hour bucket -> tokens
|
modelHour map[string]map[int64]int64 // model -> unix-hour bucket -> tokens
|
||||||
|
// keyModelHour buckets the same tokens as modelHour but keyed by
|
||||||
|
// (gateway key id, model) so per-key model quotas are isolated from each
|
||||||
|
// other. modelHour stays key-blind on purpose: it backs the AUTO slot
|
||||||
|
// quota, which is a gateway-wide limit on a slot, not a per-key one.
|
||||||
|
keyModelHour map[string]map[string]map[int64]int64 // key -> model -> unix-hour -> tokens
|
||||||
|
// keyHour buckets a key's total tokens per unix hour, backing the
|
||||||
|
// whole-key quota (all models of one key share one budget).
|
||||||
|
keyHour map[string]map[int64]int64
|
||||||
|
// keyReqHour buckets a key's request count per unix hour, backing the
|
||||||
|
// key-wide request quota. Counted for every request type (chat, stream,
|
||||||
|
// image) including failed ones, so a failing client cannot loop for free.
|
||||||
|
keyReqHour map[string]map[int64]int64
|
||||||
}
|
}
|
||||||
|
|
||||||
const hourSec = 3600
|
const hourSec = 3600
|
||||||
@ -116,14 +128,17 @@ func NewStats(maxRecords int) *Stats {
|
|||||||
maxRecords = defaultRingSize
|
maxRecords = defaultRingSize
|
||||||
}
|
}
|
||||||
return &Stats{
|
return &Stats{
|
||||||
byKey: map[string]*Stat{},
|
byKey: map[string]*Stat{},
|
||||||
byModel: map[string]*Stat{},
|
byModel: map[string]*Stat{},
|
||||||
bySrc: map[string]*Stat{},
|
bySrc: map[string]*Stat{},
|
||||||
byKeyModel: map[string]map[string]*Stat{},
|
byKeyModel: map[string]map[string]*Stat{},
|
||||||
byKeySrc: map[string]map[string]*Stat{},
|
byKeySrc: map[string]map[string]*Stat{},
|
||||||
byStatus: map[int]*Stat{},
|
byStatus: map[int]*Stat{},
|
||||||
modelHour: map[string]map[int64]int64{},
|
modelHour: map[string]map[int64]int64{},
|
||||||
maxRecs: maxRecords,
|
keyModelHour: map[string]map[string]map[int64]int64{},
|
||||||
|
keyHour: map[string]map[int64]int64{},
|
||||||
|
keyReqHour: map[string]map[int64]int64{},
|
||||||
|
maxRecs: maxRecords,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -318,6 +333,9 @@ func (s *Stats) Record(r Req) {
|
|||||||
// aggregateLocked folds r into every aggregate row and the quota window
|
// aggregateLocked folds r into every aggregate row and the quota window
|
||||||
// bucket. Caller must hold s.mu.
|
// bucket. Caller must hold s.mu.
|
||||||
func (s *Stats) aggregateLocked(r Req) {
|
func (s *Stats) aggregateLocked(r Req) {
|
||||||
|
if r.Key != "" {
|
||||||
|
s.addKeyReqLocked(r.Key, (r.Time/1000)/hourSec, 1)
|
||||||
|
}
|
||||||
inc(s.byKey, r.Key, r)
|
inc(s.byKey, r.Key, r)
|
||||||
if r.Model != "" {
|
if r.Model != "" {
|
||||||
inc(s.byModel, r.Model, r)
|
inc(s.byModel, r.Model, r)
|
||||||
@ -362,6 +380,12 @@ func (s *Stats) aggregateLocked(r Req) {
|
|||||||
s.modelHour[key] = hm
|
s.modelHour[key] = hm
|
||||||
}
|
}
|
||||||
hm[h] += tok
|
hm[h] += tok
|
||||||
|
// per-key buckets: same token split, but isolated per key so one
|
||||||
|
// key's quota cannot be drained by another key's traffic
|
||||||
|
if r.Key != "" {
|
||||||
|
s.addKeyTokenLocked(r.Key, r.Model, r.Source, h, tok)
|
||||||
|
s.addKeyHourLocked(r.Key, h, tok)
|
||||||
|
}
|
||||||
// retention: 24*40 = 960 hourly buckets ≈ 40 days of history (covers
|
// retention: 24*40 = 960 hourly buckets ≈ 40 days of history (covers
|
||||||
// the longest "month" quota window)
|
// the longest "month" quota window)
|
||||||
if len(hm) > 24*40 {
|
if len(hm) > 24*40 {
|
||||||
@ -374,6 +398,152 @@ func (s *Stats) aggregateLocked(r Req) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// quotaRetentionHours is how much hourly history the per-key quota buckets
|
||||||
|
// keep. It matches the modelHour retention (40 days) so the longest "month"
|
||||||
|
// window is fully covered after a restart, and it is deliberately applied per
|
||||||
|
// key so an idle key's buckets are reclaimed instead of pinning memory.
|
||||||
|
const quotaRetentionHours = 24 * 40
|
||||||
|
|
||||||
|
// addKeyTokenLocked adds tok to one key's (model[, source]) hourly bucket.
|
||||||
|
//
|
||||||
|
// It writes BOTH an unpinned and a pinned bucket. Every recorded request
|
||||||
|
// carries a resolved source, so the bucket key would otherwise always be
|
||||||
|
// "source::model" and a quota without a source pin would read an empty bucket
|
||||||
|
// and never trip. The unpinned bucket is the one a scope entry normally looks
|
||||||
|
// up; the pinned one serves entries that pin a source.
|
||||||
|
func (s *Stats) addKeyTokenLocked(key, model, source string, h, tok int64) {
|
||||||
|
if model == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
byModel := s.keyModelHour[key]
|
||||||
|
if byModel == nil {
|
||||||
|
byModel = map[string]map[int64]int64{}
|
||||||
|
s.keyModelHour[key] = byModel
|
||||||
|
}
|
||||||
|
// Write the bare-model bucket plus, when a source is known, the pinned
|
||||||
|
// one. The unpinned bucket is the aggregate the scope lookup normally
|
||||||
|
// reads; the pinned bucket serves entries that pin a source.
|
||||||
|
keys := [2]string{model, model}
|
||||||
|
n := 1
|
||||||
|
if source != "" {
|
||||||
|
if pinned := source + "::" + model; pinned != model {
|
||||||
|
keys[1] = pinned
|
||||||
|
n = 2
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for i := 0; i < n; i++ {
|
||||||
|
hm := byModel[keys[i]]
|
||||||
|
if hm == nil {
|
||||||
|
hm = map[int64]int64{}
|
||||||
|
byModel[keys[i]] = hm
|
||||||
|
}
|
||||||
|
hm[h] += tok
|
||||||
|
if len(hm) > quotaRetentionHours {
|
||||||
|
for k := range hm {
|
||||||
|
if k < h-quotaRetentionHours {
|
||||||
|
delete(hm, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// addKeyHourLocked adds tok to one key's all-model hourly total.
|
||||||
|
func (s *Stats) addKeyHourLocked(key string, h, tok int64) {
|
||||||
|
hm := s.keyHour[key]
|
||||||
|
if hm == nil {
|
||||||
|
hm = map[int64]int64{}
|
||||||
|
s.keyHour[key] = hm
|
||||||
|
}
|
||||||
|
hm[h] += tok
|
||||||
|
if len(hm) > quotaRetentionHours {
|
||||||
|
for k := range hm {
|
||||||
|
if k < h-quotaRetentionHours {
|
||||||
|
delete(hm, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// KeyWindowReqs returns how many requests one key issued within sec seconds;
|
||||||
|
// sec <= 0 means all retained history.
|
||||||
|
func (s *Stats) KeyWindowReqs(key string, sec int64) int64 {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
return sumBuckets(s.keyReqHour[key], time.Now().Unix(), sec)
|
||||||
|
}
|
||||||
|
|
||||||
|
// addKeyReqLocked adds n requests to a key's hourly count bucket.
|
||||||
|
func (s *Stats) addKeyReqLocked(key string, h, n int64) {
|
||||||
|
hm := s.keyReqHour[key]
|
||||||
|
if hm == nil {
|
||||||
|
hm = map[int64]int64{}
|
||||||
|
s.keyReqHour[key] = hm
|
||||||
|
}
|
||||||
|
hm[h] += n
|
||||||
|
if len(hm) > quotaRetentionHours {
|
||||||
|
for k := range hm {
|
||||||
|
if k < h-quotaRetentionHours {
|
||||||
|
delete(hm, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// sumBuckets totals the hourly buckets at or after the cutoff given in unix
|
||||||
|
// seconds. sec <= 0 means "all retained history" (no reset).
|
||||||
|
func sumBuckets(hm map[int64]int64, now, sec int64) int64 {
|
||||||
|
if len(hm) == 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
if sec <= 0 {
|
||||||
|
var total int64
|
||||||
|
for _, v := range hm {
|
||||||
|
total += v
|
||||||
|
}
|
||||||
|
return total
|
||||||
|
}
|
||||||
|
cut := now - sec
|
||||||
|
var total int64
|
||||||
|
for h, v := range hm {
|
||||||
|
if h*hourSec >= cut {
|
||||||
|
total += v
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return total
|
||||||
|
}
|
||||||
|
|
||||||
|
// KeyWindowModelTokens returns the tokens one key consumed for one model
|
||||||
|
// (optionally pinned to a source) within sec seconds; sec <= 0 means all
|
||||||
|
// retained history. Unlike WindowTokens this is isolated per key, which is
|
||||||
|
// what a per-key model quota needs.
|
||||||
|
//
|
||||||
|
// A source pin matches the pinned bucket exactly. Without a pin the call
|
||||||
|
// totals the key's bare-model bucket, which counts traffic on every source —
|
||||||
|
// a quota on "this model" should not stop applying just because the request
|
||||||
|
// happened to be served by a different upstream.
|
||||||
|
func (s *Stats) KeyWindowModelTokens(key, model, source string, sec int64) int64 {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
byModel := s.keyModelHour[key]
|
||||||
|
if len(byModel) == 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
mk := model
|
||||||
|
if source != "" {
|
||||||
|
mk = source + "::" + model
|
||||||
|
}
|
||||||
|
return sumBuckets(byModel[mk], time.Now().Unix(), sec)
|
||||||
|
}
|
||||||
|
|
||||||
|
// KeyWindowTokens returns one key's total tokens across all models within sec
|
||||||
|
// seconds; sec <= 0 means all retained history. It backs the whole-key quota.
|
||||||
|
func (s *Stats) KeyWindowTokens(key string, sec int64) int64 {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
return sumBuckets(s.keyHour[key], time.Now().Unix(), sec)
|
||||||
|
}
|
||||||
|
|
||||||
// rotateAuditLocked renames the audit file to <path>.<unix>.old once it
|
// rotateAuditLocked renames the audit file to <path>.<unix>.old once it
|
||||||
// exceeds auditRotateBytes and prunes old files beyond auditKeepOld, keeping
|
// exceeds auditRotateBytes and prunes old files beyond auditKeepOld, keeping
|
||||||
// the newest ones. Caller must hold s.mu.
|
// the newest ones. Caller must hold s.mu.
|
||||||
@ -464,6 +634,27 @@ func AutoPeriodSeconds(period string, hours int64) int64 {
|
|||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AutoSecondsToReset returns how many seconds remain until the quota window
|
||||||
|
// rolls over, for the Retry-After header on a 429. Buckets are whole unix
|
||||||
|
// hours, so the value is rounded up to the next hour boundary and never
|
||||||
|
// exceeds one window. 0 means "unknown" (never resets, or no period set), in
|
||||||
|
// which case the caller should not advertise a retry time.
|
||||||
|
func AutoSecondsToReset(period string, hours int64) int64 {
|
||||||
|
sec := AutoPeriodSeconds(period, hours)
|
||||||
|
if sec <= 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
now := time.Now().Unix()
|
||||||
|
// window = the last sec seconds, bucketed by whole hours; the oldest
|
||||||
|
// bucket still inside the window expires at its own hour boundary
|
||||||
|
elapsed := now % hourSec
|
||||||
|
left := sec - elapsed
|
||||||
|
if left <= 0 || left > sec {
|
||||||
|
left = sec
|
||||||
|
}
|
||||||
|
return left
|
||||||
|
}
|
||||||
|
|
||||||
// WindowTokens returns the tokens consumed for one model (optionally pinned
|
// WindowTokens returns the tokens consumed for one model (optionally pinned
|
||||||
// to a single source) within the window; sec <= 0 means all time. Buckets are
|
// to a single source) within the window; sec <= 0 means all time. Buckets are
|
||||||
// whole unix hours, so a sliding window overcounts by up to one hour — an
|
// whole unix hours, so a sliding window overcounts by up to one hour — an
|
||||||
|
|||||||
Reference in New Issue
Block a user