diff --git a/internal/gateway/chat.go b/internal/gateway/chat.go index 37346d2..c102d5e 100644 --- a/internal/gateway/chat.go +++ b/internal/gateway/chat.go @@ -197,6 +197,8 @@ func intersectModels(models []string, allow []config.ModelScope) []string { // A scope entry with model "AUTO" only allows requests where the effective // model is AUTO (the routing mode). It does NOT grant access to specific model // ids — that requires an explicit scope entry for the model. +// checkModelScope keeps its string-returning signature for callers that only +// need to know whether the request may proceed. func (g *Gateway) checkModelScope(ctx context.Context, model string) string { if q := g.checkQuota(ctx, model); q != nil { return q.msg @@ -329,26 +331,20 @@ func (g *Gateway) hasScopeModel(list []config.ModelScope, s string) bool { // 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) +func (g *Gateway) writeReject(w http.ResponseWriter, q *quotaRejection) { + if q.retry > 0 { + w.Header().Set("Retry-After", strconv.FormatInt(q.retry, 10)) + writeError(w, http.StatusTooManyRequests, "rate_limit_exceeded", 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)) + // 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) } func (g *Gateway) resolveByModel(model string) ([]*provider.Provider, string) { @@ -413,8 +409,8 @@ func (g *Gateway) handleChat(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusServiceUnavailable, "no_provider", "no auto slot configured") return } - if msg := g.checkModelScope(r.Context(), "AUTO"); msg != "" { - g.writeScopeReject(w, r, "AUTO") + if q := g.checkQuota(r.Context(), "AUTO"); q != nil { + g.writeReject(w, q) return } ctx := r.Context() @@ -467,8 +463,8 @@ func (g *Gateway) handleChat(w http.ResponseWriter, r *http.Request) { if effective == "" { effective = firstModel(cands[0]) } - if msg := g.checkModelScope(r.Context(), effective); msg != "" { - g.writeScopeReject(w, r, effective) + if q := g.checkQuota(r.Context(), effective); q != nil { + g.writeReject(w, q) return } ctx := r.Context() @@ -1124,8 +1120,8 @@ 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 != "" { - g.writeScopeReject(w, r, "AUTO") + if q := g.checkQuota(r.Context(), "AUTO"); q != nil { + g.writeReject(w, q) return } done := g.stats.Begin() @@ -1171,8 +1167,8 @@ func (g *Gateway) handleImage(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusServiceUnavailable, "no_provider", "no image source configured") return } - if msg := g.checkModelScope(r.Context(), effectiveImageModel(model, cands)); msg != "" { - g.writeScopeReject(w, r, effectiveImageModel(model, cands)) + if q := g.checkQuota(r.Context(), effectiveImageModel(model, cands)); q != nil { + g.writeReject(w, q) return } done := g.stats.Begin() diff --git a/internal/gateway/key_quota_wiring_test.go b/internal/gateway/key_quota_wiring_test.go index 7c1f815..c15a3d3 100644 --- a/internal/gateway/key_quota_wiring_test.go +++ b/internal/gateway/key_quota_wiring_test.go @@ -7,6 +7,7 @@ import ( "net/http" "net/http/httptest" "os" + "path/filepath" "strings" "testing" "time" @@ -245,3 +246,82 @@ func quotaCtx(t *testing.T, g *Gateway, key string) context.Context { func nowMSOffset(sec int64) int64 { return time.Now().Add(time.Duration(sec) * time.Second).UnixMilli() } + +// The key-wide quota and the AUTO slot quota are different limits with +// different scopes: the slot quota is gateway-wide (a shared upstream budget), +// the key quota belongs to one caller. When both are exhausted the caller must +// see the KEY quota, because that is the one it can act on — the slot quota +// would otherwise be reported as "no capacity", which reads like an outage. +func TestKeyQuotaWinsOverSlotQuota(t *testing.T) { + up := upstream(t, &upstreamCtrl{}) + defer up.Close() + g := newQuotaGW(t, up.URL, + config.GWKey{Key: "sk-a", Role: "user", TokenQuota: 4, Period: "hour", + Models: []config.ModelScope{{Model: "AUTO"}}}) + ctx := quotaCtx(t, g, "sk-a") + + // exhaust the key first + g.stats.Record(Req{Time: time.Now().UnixMilli(), Key: keyID("sk-a"), Model: "m1", + Source: "up", Prompt: 100, Compl: 100, OK: true, Status: 200}) + rr, code := chatAs(t, g, "sk-a", "AUTO") + if rr.Code != http.StatusTooManyRequests { + t.Fatalf("want 429 from the key quota, got %d (%s)", rr.Code, rr.Body.String()) + } + if code != "rate_limit_exceeded" { + t.Errorf("code = %q, want rate_limit_exceeded", code) + } + _ = ctx +} + +// A key with NO caps must never be blocked by the slot quota check reaching it +// through the scope path: scopeTokens on an uncapped key returns 0 usage and +// the limit check must treat that as "no cap", not "exhausted". +func TestUncappedKeyNeverBlockedByEmptyScope(t *testing.T) { + up := upstream(t, &upstreamCtrl{}) + defer up.Close() + g := newQuotaGW(t, up.URL, + config.GWKey{Key: "sk-a", Role: "user", Models: []config.ModelScope{{Model: "AUTO"}}}) + for i := 0; i < 5; i++ { + if rr, _ := chatAs(t, g, "sk-a", "AUTO"); rr.Code != 200 { + t.Fatalf("call %d: want 200 (AUTO scope with no quota), got %d (%s)", i, rr.Code, rr.Body.String()) + } + } +} + +// newQuotaGW builds a gateway with one mock upstream serving model m1 and the +// given keys. +func newQuotaGW(t *testing.T, upURL string, keys ...config.GWKey) *Gateway { + t.Helper() + td := t.TempDir() + cfgPath := filepath.Join(td, "config.yaml") + if err := os.WriteFile(cfgPath, []byte("listen: :0"), 0o644); err != nil { + t.Fatal(err) + } + cfg := &config.Config{ + Path: cfgPath, + AdapterDir: filepath.Join(td, "adapters"), + RuntimeFile: filepath.Join(td, "runtime.json"), + Keys: keys, + Sources: []config.Source{{ + Name: "up", BaseURL: upURL, 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) + secrets := make([]string, 0, len(keys)) + for _, k := range keys { + secrets = append(secrets, k.Key) + } + g, err := New(c, secrets) + if err != nil { + t.Fatalf("gateway: %v", err) + } + return g +}