diff --git a/internal/config/config.go b/internal/config/config.go index 816d6b7..a856d27 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -375,6 +375,82 @@ type GWKey struct { Note string `yaml:"note,omitempty" json:"note,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 + // 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, diff --git a/internal/core/core.go b/internal/core/core.go index 905c836..b791c40 100644 --- a/internal/core/core.go +++ b/internal/core/core.go @@ -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. 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() defer c.mu.Unlock() 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 { return config.GWKey{}, err } + if err := q.Validate(); err != nil { + return config.GWKey{}, err + } rec := config.GWKey{ Key: "sk-gw-" + hex.EncodeToString(key), Role: role, @@ -265,9 +273,8 @@ func (c *Core) CreateKey(name, role string, models []config.ModelScope, note str Note: note, CreatedAt: time.Now().Unix(), } - if rec.Role == "" { - rec.Role = "user" - } + rec.Role = config.NormalizeRole(rec.Role) + rec.ApplyQuota(q) c.cfg.Keys = append(c.cfg.Keys, rec) if err := c.saveConfig(); err != nil { 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. 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() defer c.mu.Unlock() + if quota != nil { + if err := quota.Validate(); err != nil { + return config.GWKey{}, err + } + } for i, k := range c.cfg.Keys { if k.Key == key { 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].Note = note + if quota != nil { + c.cfg.Keys[i].ApplyQuota(*quota) + } if err := c.saveConfig(); err != nil { return config.GWKey{}, err } diff --git a/internal/gateway/apiv1.go b/internal/gateway/apiv1.go index 363f9ab..15393be 100644 --- a/internal/gateway/apiv1.go +++ b/internal/gateway/apiv1.go @@ -116,6 +116,12 @@ func (g *Gateway) apiV1Routes(w http.ResponseWriter, r *http.Request) { "note": k.Note, "created_at": k.CreatedAt, "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 // already has it from creation time or from config.yaml. "key_prefix": maskKey(k.Key), diff --git a/internal/gateway/chat.go b/internal/gateway/chat.go index cb018ee..37346d2 100644 --- a/internal/gateway/chat.go +++ b/internal/gateway/chat.go @@ -8,6 +8,7 @@ import ( "log" "net/http" "regexp" + "strconv" "strings" "sync/atomic" "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 // ids — that requires an explicit scope entry for the model. 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) if allow == nil { - return "" + return nil } for _, sc := range allow { if sc.Model != model { @@ -208,23 +233,75 @@ func (g *Gateway) checkModelScope(ctx context.Context, model string) string { if sc.TokenQuota > 0 { used := g.scopeTokens(ctx, sc) 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 -// window (total for AUTO / per model otherwise). +// checkKeyQuotaRetry enforces the key-wide caps and reports the remaining +// 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 { k := keyID(reqKey(ctx)) - if sc.Model != "" && strings.EqualFold(sc.Model, "AUTO") { - return g.stats.KeyTokens(k) - } 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" / @@ -248,6 +325,32 @@ func (g *Gateway) hasScopeModel(list []config.ModelScope, s string) bool { 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) { if isAuto(model) { return g.core.Registry().Resolve("AUTO"), "" @@ -311,7 +414,7 @@ func (g *Gateway) handleChat(w http.ResponseWriter, r *http.Request) { return } if msg := g.checkModelScope(r.Context(), "AUTO"); msg != "" { - writeError(w, http.StatusForbidden, "model_not_allowed", msg) + g.writeScopeReject(w, r, "AUTO") return } ctx := r.Context() @@ -365,7 +468,7 @@ func (g *Gateway) handleChat(w http.ResponseWriter, r *http.Request) { effective = firstModel(cands[0]) } if msg := g.checkModelScope(r.Context(), effective); msg != "" { - writeError(w, http.StatusForbidden, "model_not_allowed", msg) + g.writeScopeReject(w, r, effective) return } ctx := r.Context() @@ -1022,7 +1125,7 @@ func (g *Gateway) handleImage(w http.ResponseWriter, r *http.Request) { if isAuto(model) { if chain := g.core.AutoImageChain(); chain != nil && len(chain.Tiers) > 0 { if msg := g.checkModelScope(r.Context(), "AUTO"); msg != "" { - writeError(w, http.StatusForbidden, "model_not_allowed", msg) + g.writeScopeReject(w, r, "AUTO") return } done := g.stats.Begin() @@ -1069,7 +1172,7 @@ func (g *Gateway) handleImage(w http.ResponseWriter, r *http.Request) { return } 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 } done := g.stats.Begin() diff --git a/internal/gateway/key_quota_api_test.go b/internal/gateway/key_quota_api_test.go new file mode 100644 index 0000000..fe8d9d4 --- /dev/null +++ b/internal/gateway/key_quota_api_test.go @@ -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 diff --git a/internal/gateway/key_quota_test.go b/internal/gateway/key_quota_test.go new file mode 100644 index 0000000..b15a6c8 --- /dev/null +++ b/internal/gateway/key_quota_test.go @@ -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) + } +} diff --git a/internal/gateway/key_quota_wiring_test.go b/internal/gateway/key_quota_wiring_test.go new file mode 100644 index 0000000..7c1f815 --- /dev/null +++ b/internal/gateway/key_quota_wiring_test.go @@ -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() +} diff --git a/internal/gateway/keys.go b/internal/gateway/keys.go index d95d58e..3c2440b 100644 --- a/internal/gateway/keys.go +++ b/internal/gateway/keys.go @@ -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()}) case http.MethodPost: var body struct { - Name string `json:"name"` - Role string `json:"role"` - Models []config.ModelScope `json:"models"` - Note string `json:"note"` + Name string `json:"name"` + Role string `json:"role"` + Models []config.ModelScope `json:"models"` + 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 { writeError(w, http.StatusBadRequest, "invalid_request", "invalid json: "+err.Error()) return } - if body.Role == "" { - body.Role = "user" + body.Role = config.NormalizeRole(body.Role) + 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" { - 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) + rec, err := g.core.CreateKeyWithQuota(body.Name, body.Role, body.Models, body.Note, q) if err != nil { writeError(w, http.StatusBadRequest, "key_error", err.Error()) return @@ -67,16 +71,33 @@ func (g *Gateway) handleKeysAPI(w http.ResponseWriter, r *http.Request) { return } var body struct { - Name string `json:"name"` - Role string `json:"role"` - Models []config.ModelScope `json:"models"` - Note string `json:"note"` + Name string `json:"name"` + Role string `json:"role"` + Models []config.ModelScope `json:"models"` + 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 { writeError(w, http.StatusBadRequest, "invalid_request", "invalid json: "+err.Error()) 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 { writeError(w, http.StatusBadRequest, "key_error", err.Error()) 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 // themselves; admins can use this as a convenience too). func (g *Gateway) handleKeyMe(w http.ResponseWriter, r *http.Request) { diff --git a/internal/gateway/stats.go b/internal/gateway/stats.go index 5f4db12..0324aee 100644 --- a/internal/gateway/stats.go +++ b/internal/gateway/stats.go @@ -84,6 +84,18 @@ type Stats struct { auditPath string replayPartial bool // aggregates built from a bounded audit tail 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 @@ -116,14 +128,17 @@ func NewStats(maxRecords int) *Stats { maxRecords = defaultRingSize } return &Stats{ - byKey: map[string]*Stat{}, - byModel: map[string]*Stat{}, - bySrc: map[string]*Stat{}, - byKeyModel: map[string]map[string]*Stat{}, - byKeySrc: map[string]map[string]*Stat{}, - byStatus: map[int]*Stat{}, - modelHour: map[string]map[int64]int64{}, - maxRecs: maxRecords, + byKey: map[string]*Stat{}, + byModel: map[string]*Stat{}, + bySrc: map[string]*Stat{}, + byKeyModel: map[string]map[string]*Stat{}, + byKeySrc: map[string]map[string]*Stat{}, + byStatus: map[int]*Stat{}, + modelHour: map[string]map[int64]int64{}, + 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 // bucket. Caller must hold s.mu. func (s *Stats) aggregateLocked(r Req) { + if r.Key != "" { + s.addKeyReqLocked(r.Key, (r.Time/1000)/hourSec, 1) + } inc(s.byKey, r.Key, r) if r.Model != "" { inc(s.byModel, r.Model, r) @@ -362,6 +380,12 @@ func (s *Stats) aggregateLocked(r Req) { s.modelHour[key] = hm } 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 // the longest "month" quota window) 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 ..old once it // exceeds auditRotateBytes and prunes old files beyond auditKeepOld, keeping // the newest ones. Caller must hold s.mu. @@ -464,6 +634,27 @@ func AutoPeriodSeconds(period string, hours int64) int64 { 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 // 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