diff --git a/docs/MODEL_RPM.md b/docs/MODEL_RPM.md new file mode 100644 index 000000000..844b429aa --- /dev/null +++ b/docs/MODEL_RPM.md @@ -0,0 +1,26 @@ +# Local per-model completion limits + +Add `modelRPM` to the existing Zero config file: + +```json +{"modelRPM": {"gpt-4.1": 30, "claude-sonnet-4.5": 15}} +``` + +Keys resolve through the model registry; custom IDs match exactly after trimming +whitespace. Missing keys and zero mean unlimited. Negative limits and conflicting +aliases are configuration errors. Project config may add or tighten a user cap, +but cannot relax it. Changes take effect on the next Zero process start. + +For `zero exec` and interactive agent turns, admission happens immediately before +the wrapped turn session's `Stream`. N admissions fit in any sliding 60-second +window; N+1 returns a local rate-limit error before entering that method. The +error reports when a slot expires. Exec uses its existing provider-error exit +code plus a local-limit hint. There is no sleep or hidden retry. Failed admitted +requests still consume a slot. + +Sessions and model switches share a limiter in one process. Separate processes +have independent windows; restarting clears them. This is not an account-wide +cap. It counts Stream admissions, not all HTTP requests: setup/prewarm, +compaction, provider-internal retries, discovery and direct calls outside that +boundary are not counted. No TPM, cost budgets, persistence, cross-process +coordination or new runtime/Python dependency is introduced. diff --git a/docs/README.md b/docs/README.md index 2923c3078..36d118ba5 100644 --- a/docs/README.md +++ b/docs/README.md @@ -6,6 +6,7 @@ Zero. ## User Docs - [Install](INSTALL.md) +- [Local per-model completion limits](MODEL_RPM.md) - [Update flow](UPDATE.md) - [OAuth logins and subscription-backed providers](oauth-subscriptions.md) diff --git a/internal/cli/app.go b/internal/cli/app.go index 160eabc5b..74e1020e6 100644 --- a/internal/cli/app.go +++ b/internal/cli/app.go @@ -959,7 +959,9 @@ func runInteractiveTUIWithSetup(stderr io.Writer, deps appDeps, permissionMode a // notice when project hooks/plugins were dropped for an untrusted workspace. hookDispatcher, hookSkip := newHookDispatcherWithExtra(workspaceRoot, pluginActivation.hooks, trustRoot, executionRunner) emitTrustNotice(stderr, hookSkip, pluginActivation.trustSkip, mcpSkip) + modelRPM := zeroruntime.NewModelRPMLimiter(resolved.ModelRPM) return deps.runTUI(context.Background(), tui.Options{ + ModelRPM: modelRPM, Cwd: workspaceRoot, Version: version, Theme: theme, @@ -982,7 +984,7 @@ func runInteractiveTUIWithSetup(stderr io.Writer, deps appDeps, permissionMode a if provider == nil { return nil } - if optimized, ok := providers.OptimizedTurnSessions(profile, provider, providers.Options{}); ok { + if optimized, ok := providers.ConfiguredTurnSessions(profile, provider, providers.Options{ModelRPM: modelRPM}); ok { return optimized } return providers.DefaultTurnSessions(profile, provider, providers.Options{}) diff --git a/internal/cli/app_test.go b/internal/cli/app_test.go index 1c12da4c8..a2a87cc88 100644 --- a/internal/cli/app_test.go +++ b/internal/cli/app_test.go @@ -343,8 +343,10 @@ func TestRunNoArgsFallsBackToUsableProviderWhenNoneMarkedActive(t *testing.T) { // config.json has providers configured but none marked active (e.g. a // blank/stale activeProvider field) — Resolve returns the successfully // normalized providers list alongside ErrNoActiveProvider. - return config.ResolvedConfig{Providers: []config.ProviderProfile{usable}}, - fmt.Errorf("%w: active provider %q not found", config.ErrNoActiveProvider, "") + return config.ResolvedConfig{ + Providers: []config.ProviderProfile{usable}, + ModelRPM: map[string]int{"gpt-test": 1}, + }, fmt.Errorf("%w: active provider %q not found", config.ErrNoActiveProvider, "") }, newProvider: func(profile config.ProviderProfile) (zeroruntime.Provider, error) { providerProfile = profile @@ -374,6 +376,10 @@ func TestRunNoArgsFallsBackToUsableProviderWhenNoneMarkedActive(t *testing.T) { } if providerProfile.Name != "work" { t.Fatalf("provider used = %q, want fallback to the usable saved provider %q", providerProfile.Name, "work") + + } + if launchedOptions.ModelRPM == nil { + t.Fatal("ModelRPM = nil, want limiter preserved through provider fallback") } } diff --git a/internal/cli/exec.go b/internal/cli/exec.go index d47964c11..1e4dc1edd 100644 --- a/internal/cli/exec.go +++ b/internal/cli/exec.go @@ -452,20 +452,20 @@ func runExec(args []string, stdout io.Writer, stderr io.Writer, deps appDeps) in // escalation so post-switch turns are attributed to the escalated model. currentModel := resolved.Provider.Model // Optimized OpenAI turn sessions (ZERO_OPENAI_TURN_SESSION, default off). nil - // when gated off or the profile is ineligible: agent.Run then wraps the - // provider in its default adapter. This is the run's STARTING session provider - // and is used whether or not escalation is enabled. - turnSessions, _ := providers.OptimizedTurnSessions(resolved.Provider, provider, providers.Options{}) + // when gated off or ineligible, unless a configured RPM limiter needs to + // wrap the default adapter. The limiter is shared across model switches. + sessionOptions := providers.Options{ModelRPM: zeroruntime.NewModelRPMLimiter(resolved.ModelRPM)} + turnSessions, _ := providers.ConfiguredTurnSessions(resolved.Provider, provider, sessionOptions) // Both switchers come from one shared builder so exec and the interactive TUI // cannot drift on the nil contracts the agent loop depends on. The session - // switcher is nil unless this run STARTED optimized, which is what keeps a - // default-adapter run on the default adapter. + // switcher preserves the starting transport and shares the RPM limiter. var modelSwitcher func(context.Context, string) (agent.Provider, error) var modelSessionSwitcher func(context.Context, string) (zeroruntime.TurnSessionProvider, error) if options.allowEscalation { modelSwitcher, modelSessionSwitcher = providers.EscalationSwitchers( resolved.Provider, provider, deps.newProvider, func(modelID string) { currentModel = modelID }, + sessionOptions, ) } diff --git a/internal/cli/model_rpm_test.go b/internal/cli/model_rpm_test.go new file mode 100644 index 000000000..8f4bf023c --- /dev/null +++ b/internal/cli/model_rpm_test.go @@ -0,0 +1,83 @@ +package cli + +import ( + "bytes" + "context" + "errors" + "github.com/Gitlawb/zero/internal/config" + "github.com/Gitlawb/zero/internal/mcp" + "github.com/Gitlawb/zero/internal/tools" + "github.com/Gitlawb/zero/internal/tui" + "github.com/Gitlawb/zero/internal/zeroruntime" + "strings" + "testing" +) + +type rpmExecProvider struct{ calls int } + +func (p *rpmExecProvider) StreamCompletion(context.Context, zeroruntime.CompletionRequest) (<-chan zeroruntime.StreamEvent, error) { + p.calls++ + ch := make(chan zeroruntime.StreamEvent, 4) + if p.calls == 1 { + ch <- zeroruntime.StreamEvent{Type: zeroruntime.StreamEventToolCallStart, ToolCallID: "fixture", ToolName: "read_file"} + ch <- zeroruntime.StreamEvent{Type: zeroruntime.StreamEventToolCallDelta, ToolCallID: "fixture", ArgumentsFragment: `{"path":"missing-fixture.txt"}`} + ch <- zeroruntime.StreamEvent{Type: zeroruntime.StreamEventToolCallEnd, ToolCallID: "fixture"} + } else { + ch <- zeroruntime.StreamEvent{Type: zeroruntime.StreamEventText, Content: "unexpected second completion"} + } + ch <- zeroruntime.StreamEvent{Type: zeroruntime.StreamEventDone} + close(ch) + return ch, nil +} +func rpmTestDeps(t *testing.T, p *rpmExecProvider) appDeps { + t.Helper() + root := t.TempDir() + for _, k := range []string{"HOME", "USERPROFILE", "APPDATA", "LOCALAPPDATA", "XDG_CONFIG_HOME", "XDG_CACHE_HOME", "XDG_DATA_HOME", "XDG_STATE_HOME"} { + t.Setenv(k, root) + } + return appDeps{getwd: func() (string, error) { return root, nil }, resolveConfig: func(string, config.Overrides) (config.ResolvedConfig, error) { + c := execResolvedConfig() + c.ModelRPM = map[string]int{c.Provider.Model: 1} + return c, nil + }, newProvider: func(config.ProviderProfile) (zeroruntime.Provider, error) { return p, nil }, registerMCPTools: func(context.Context, *tools.Registry, config.MCPConfig, mcp.RegisterOptions) (mcpToolRuntime, error) { + return noopMCPRuntime{}, nil + }} +} +func TestRunExecRPMRefusesSecondCompletion(t *testing.T) { + p := &rpmExecProvider{} + deps := rpmTestDeps(t, p) + var out, err bytes.Buffer + code := runWithDeps([]string{"exec", "fixture"}, &out, &err, deps) + if code != exitProvider || p.calls != 1 || !strings.Contains(err.String(), "Local modelRPM cap reached") { + t.Fatalf("exit=%d calls=%d stdout=%s stderr=%s", code, p.calls, out.String(), err.String()) + } +} +func TestInteractiveRPMFactorySharesWindow(t *testing.T) { + p := &rpmExecProvider{} + deps := rpmTestDeps(t, p) + launched := false + deps.runTUI = func(ctx context.Context, o tui.Options) int { + launched = true + if o.ModelRPM == nil { + t.Fatal("TUI escalation lacks limiter") + } + for i := 0; i < 2; i++ { + ss := o.NewTurnSessionProvider(o.ProviderProfile, p) + s, e := ss.OpenTurnSession(ctx) + if e != nil { + t.Fatal(e) + } + _, e = s.Stream(ctx, zeroruntime.CompletionRequest{}) + var hit *zeroruntime.RPMLimitError + if (i == 0 && e != nil) || (i == 1 && !errors.As(e, &hit)) { + t.Fatalf("run=%d err=%v", i, e) + } + } + return exitSuccess + } + var out, err bytes.Buffer + code := runWithDeps(nil, &out, &err, deps) + if !launched || code != exitSuccess || p.calls != 1 { + t.Fatalf("launched=%v exit=%d calls=%d stderr=%s", launched, code, p.calls, err.String()) + } +} diff --git a/internal/config/model_rpm.go b/internal/config/model_rpm.go new file mode 100644 index 000000000..4c0bf7a26 --- /dev/null +++ b/internal/config/model_rpm.go @@ -0,0 +1,48 @@ +package config + +import ( + "fmt" + "github.com/Gitlawb/zero/internal/modelregistry" + "strings" +) + +func normalizeModelRPM(limits map[string]int) (map[string]int, error) { + if len(limits) == 0 { + return nil, nil + } + registry, err := modelregistry.DefaultRegistry() + if err != nil { + return nil, err + } + out := make(map[string]int, len(limits)) + for key, n := range limits { + model := strings.TrimSpace(key) + if model == "" || n < 0 { + return nil, fmt.Errorf("invalid modelRPM entry %q: model must be non-empty and limit >= 0", key) + } + if id, ok := registry.ResolveID(model); ok { + model = id + } + if old, exists := out[model]; exists && old != n { + return nil, fmt.Errorf("conflicting modelRPM entries resolve to %q", model) + } + out[model] = n + } + return out, nil +} + +// Project config can tighten a user cap but cannot relax it. +func mergeModelRPM(dst *map[string]int, src map[string]int, tighten bool) { + if len(src) == 0 { + return + } + if *dst == nil { + *dst = make(map[string]int) + } + for model, n := range src { + old := (*dst)[model] + if !tighten || old == 0 || (n > 0 && n < old) { + (*dst)[model] = n + } + } +} diff --git a/internal/config/model_rpm_test.go b/internal/config/model_rpm_test.go new file mode 100644 index 000000000..7a8710f95 --- /dev/null +++ b/internal/config/model_rpm_test.go @@ -0,0 +1,54 @@ +package config + +import ( + "encoding/json" + "reflect" + "testing" +) + +func TestModelRPMRoundTrip(t *testing.T) { + var c FileConfig + if e := json.Unmarshal([]byte(`{"modelRPM":{" OPENAI:GPT-4.1 ":2,"custom":0},"future":true}`), &c); e != nil { + t.Fatal(e) + } + want := map[string]int{"gpt-4.1": 2, "custom": 0} + if !reflect.DeepEqual(c.ModelRPM, want) { + t.Fatal(c.ModelRPM) + } + b, e := json.Marshal(c) + if e != nil { + t.Fatal(e) + } + var d FileConfig + if e = json.Unmarshal(b, &d); e != nil { + t.Fatal(e) + } + if !reflect.DeepEqual(d.ModelRPM, want) || string(d.Extra["future"]) != "true" { + t.Fatalf("round trip=%s", b) + } +} +func TestModelRPMInvalid(t *testing.T) { + for _, s := range []string{`{"modelRPM":{"":1}}`, `{"modelRPM":{"a":-1}}`, `{"modelRPM":{"a":1.5}}`, `{"modelRPM":{"a":"2"}}`, `{"modelRPM":{"gpt-4.1":1,"openai:gpt-4.1":2}}`} { + var c FileConfig + if e := json.Unmarshal([]byte(s), &c); e == nil { + t.Fatalf("accepted %s", s) + } + } +} +func TestModelRPMProjectCanOnlyTighten(t *testing.T) { + user := writeConfig(t, `{"modelRPM":{"gpt-4.1":2},"activeProvider":"test","providers":[{"name":"test","provider":"openai","model":"gpt-4.1"}]}`) + for _, n := range []string{"0", "1", "5"} { + project := writeConfig(t, `{"modelRPM":{"openai:gpt-4.1":`+n+`}}`) + c, e := Resolve(ResolveOptions{UserConfigPath: user, ProjectConfigPath: project, Env: map[string]string{}}) + if e != nil { + t.Fatal(e) + } + want := 2 + if n == "1" { + want = 1 + } + if c.ModelRPM["gpt-4.1"] != want { + t.Fatal(c.ModelRPM) + } + } +} diff --git a/internal/config/resolver.go b/internal/config/resolver.go index 7b1b82854..bcc6eba0a 100644 --- a/internal/config/resolver.go +++ b/internal/config/resolver.go @@ -104,6 +104,9 @@ func Resolve(options ResolveOptions) (ResolvedConfig, error) { // trusted global-config merge, which must keep honouring the setting. commandConfig.Sandbox.Enabled = nil commandConfig.CrossSessionInbound = "" + // External provider commands cannot relax the user cap. + mergeModelRPM(&cfg.ModelRPM, commandConfig.ModelRPM, true) + commandConfig.ModelRPM = nil mergeConfig(&cfg, commandConfig) } @@ -158,10 +161,14 @@ func Resolve(options ResolveOptions) (ResolvedConfig, error) { // normalized (but active-less) profile list — keep it so a caller can fall // back to an already-configured usable provider instead of treating this // like a config with nothing set up at all. - return ResolvedConfig{Providers: providers}, err + return ResolvedConfig{ + ModelRPM: cfg.ModelRPM, + Providers: providers, + }, err } return ResolvedConfig{ + ModelRPM: cfg.ModelRPM, ActiveProvider: active.Name, Providers: providers, Provider: active, @@ -230,6 +237,7 @@ func loadConfigFile(path string) (FileConfig, error) { } func mergeConfig(dst *FileConfig, src FileConfig) { + mergeModelRPM(&dst.ModelRPM, src.ModelRPM, false) if activeProvider := strings.TrimSpace(src.ActiveProvider); activeProvider != "" { dst.ActiveProvider = activeProvider } @@ -291,6 +299,7 @@ func mergeConfig(dst *FileConfig, src FileConfig) { } func mergeProjectConfig(dst *FileConfig, src FileConfig) error { + mergeModelRPM(&dst.ModelRPM, src.ModelRPM, true) if activeProvider := strings.TrimSpace(src.ActiveProvider); activeProvider != "" { dst.ActiveProvider = activeProvider } diff --git a/internal/config/resolver_test.go b/internal/config/resolver_test.go index 11ce0025d..76606ad62 100644 --- a/internal/config/resolver_test.go +++ b/internal/config/resolver_test.go @@ -1161,10 +1161,13 @@ func TestResolveKeepsNormalizedProvidersWhenNoneMarkedActive(t *testing.T) { // activeProvider is blank/stale — a caller like the interactive TUI still // needs the normalized list to fall back to an already-usable provider // instead of forcing a full re-onboarding wizard. - path := writeConfig(t, `{"providers":[ - {"name":"work","provider_kind":"openai","apiKey":"sk-test","model":"gpt-test"}, - {"name":"other","provider_kind":"openai","apiKey":"sk-other","model":"gpt-test"} - ]}`) + path := writeConfig(t, `{ + "modelRPM":{"gpt-test":1}, + "providers":[ + {"name":"work","provider_kind":"openai","apiKey":"sk-test","model":"gpt-test"}, + {"name":"other","provider_kind":"openai","apiKey":"sk-other","model":"gpt-test"} + ] +}`) resolved, err := Resolve(ResolveOptions{ProjectConfigPath: path, Env: map[string]string{}}) if !errors.Is(err, ErrNoActiveProvider) { @@ -1173,6 +1176,9 @@ func TestResolveKeepsNormalizedProvidersWhenNoneMarkedActive(t *testing.T) { if len(resolved.Providers) != 2 { t.Fatalf("Providers = %#v, want the 2 normalized profiles preserved despite the error", resolved.Providers) } + if got := resolved.ModelRPM["gpt-test"]; got != 1 { + t.Fatalf("ModelRPM[gpt-test] = %d, want 1 preserved despite ErrNoActiveProvider", got) + } } func TestResolveTrimsProviderProfileAliasesBeforeFallback(t *testing.T) { diff --git a/internal/config/types.go b/internal/config/types.go index 9343c7c94..d7ba030bc 100644 --- a/internal/config/types.go +++ b/internal/config/types.go @@ -351,6 +351,7 @@ func (cfg *ToolsConfig) UnmarshalJSON(data []byte) error { } type FileConfig struct { + ModelRPM map[string]int `json:"modelRPM,omitempty"` ActiveProvider string `json:"activeProvider,omitempty"` Providers []ProviderProfile `json:"providers,omitempty"` MaxTurns int `json:"maxTurns,omitempty"` @@ -375,6 +376,7 @@ type FileConfig struct { func (cfg FileConfig) MarshalJSON() ([]byte, error) { type rawConfig struct { + ModelRPM map[string]int `json:"modelRPM,omitempty"` ActiveProvider string `json:"activeProvider,omitempty"` Providers []ProviderProfile `json:"providers,omitempty"` MaxTurns int `json:"maxTurns,omitempty"` @@ -390,6 +392,7 @@ func (cfg FileConfig) MarshalJSON() ([]byte, error) { CrossSessionInbound string `json:"crossSessionInbound,omitempty"` } raw := rawConfig{ + ModelRPM: cfg.ModelRPM, ActiveProvider: cfg.ActiveProvider, Providers: cfg.Providers, MaxTurns: cfg.MaxTurns, @@ -465,6 +468,7 @@ type Overrides struct { } type ResolvedConfig struct { + ModelRPM map[string]int ActiveProvider string Providers []ProviderProfile Provider ProviderProfile @@ -533,6 +537,7 @@ type MCPOAuthConfig struct { func (cfg *FileConfig) UnmarshalJSON(data []byte) error { type rawConfig struct { + ModelRPM map[string]int `json:"modelRPM"` ActiveProvider string `json:"activeProvider"` Providers []ProviderProfile `json:"providers"` MaxTurns int `json:"maxTurns"` @@ -567,6 +572,11 @@ func (cfg *FileConfig) UnmarshalJSON(data []byte) error { extra = nil } cfg.ActiveProvider = raw.ActiveProvider + limits, err := normalizeModelRPM(raw.ModelRPM) + if err != nil { + return err + } + cfg.ModelRPM = limits cfg.Providers = raw.Providers // A negative maxTurns is unambiguously invalid; without this it would be // silently dropped by the `MaxTurns > 0` merge gates and fall back to the diff --git a/internal/errhint/errhint.go b/internal/errhint/errhint.go index d1d674057..9c4159e45 100644 --- a/internal/errhint/errhint.go +++ b/internal/errhint/errhint.go @@ -24,6 +24,7 @@ const ( Connectivity ModelNotFound ContextOverflow + LocalRPM ) // providerMarkers are the prefixes the provider layer attaches to every @@ -59,6 +60,8 @@ func Classify(err error) Category { return Unknown } switch { + case strings.Contains(m, "rate limit error: local rpm limit for"): + return LocalRPM case containsAny(m, "auth error:", "unauthorized", "api key", "api_key", "invalid_api_key", "authentication", "permission denied", "forbidden") || containsStatusCode(m, "401", "403"): return Auth @@ -89,6 +92,8 @@ func TUIHint(err error) string { return "API key rejected — run /provider to re-check your credentials" case RateLimit: return "Rate limited — wait a moment, or switch model with /model" + case LocalRPM: + return "Local modelRPM cap reached — wait for the window to expire, or switch model with /model" case Connectivity: return "Can't reach the provider — run /doctor --connectivity" case ModelNotFound: @@ -109,6 +114,8 @@ func CLIHint(err error) string { return "API key rejected — run `zero setup`, `zero auth openrouter` for OpenRouter, or set the provider's API key" case RateLimit: return "Rate limited — wait a moment, or switch model with --model" + case LocalRPM: + return "Local modelRPM cap reached — wait for the window to expire, or switch model with --model" case Connectivity: return "Can't reach the provider — run `zero doctor`" case ModelNotFound: diff --git a/internal/errhint/errhint_test.go b/internal/errhint/errhint_test.go index bab3c9901..b51b9cea8 100644 --- a/internal/errhint/errhint_test.go +++ b/internal/errhint/errhint_test.go @@ -17,6 +17,7 @@ func TestClassify(t *testing.T) { {"raw 401", "provider request error: 401 Unauthorized", Auth}, {"invalid api key", "provider request error: invalid_api_key: incorrect key provided", Auth}, {"rate limit prefix", "rate limit error: 429 too many requests", RateLimit}, + {"local RPM", `rate limit error: local RPM limit for "gpt-4.1" (2 requests/60s); retry in 12s`, LocalRPM}, {"overloaded", "provider error: model is overloaded, please retry", RateLimit}, {"resource exhausted gemini", "provider stream error: rpc error: code = ResourceExhausted desc = quota exceeded", RateLimit}, {"context length openai", "provider request error: this model's maximum context length is 128000 tokens", ContextOverflow}, diff --git a/internal/providers/escalation.go b/internal/providers/escalation.go index 317cdba13..f7ac2aa9d 100644 --- a/internal/providers/escalation.go +++ b/internal/providers/escalation.go @@ -22,14 +22,15 @@ import ( // An error is reported to the loop, which records a note and continues on the // current model. // -// The session switcher comes back nil unless the run STARTED optimized. A run -// that began on the default adapter stays on it, so escalation cannot quietly -// change the transport underneath a session. +// The session switcher is installed for an optimized start or configured RPM +// limits. Default-adapter starts retain that transport after escalation. +// The same limiter wraps switched sessions so switching cannot reset a cap. func EscalationSwitchers( profile config.ProviderProfile, provider zeroruntime.Provider, newProvider func(config.ProviderProfile) (zeroruntime.Provider, error), onSwitch func(modelID string), + sessionOptions ...Options, ) ( func(context.Context, string) (zeroruntime.Provider, error), func(context.Context, string) (zeroruntime.TurnSessionProvider, error), @@ -37,6 +38,10 @@ func EscalationSwitchers( if newProvider == nil { return nil, nil } + options := Options{} + if len(sessionOptions) > 0 { + options = sessionOptions[0] + } // The escalated profile is the run's profile with the model replaced, so the // credential, base URL and headers travel with it. Callers pass a newProvider // that already applies the stored key, which is why there is no per-site key @@ -62,8 +67,8 @@ func EscalationSwitchers( return switchedProvider, nil } - turnSessions, _ := OptimizedTurnSessions(profile, provider, Options{}) - if turnSessions == nil { + turnSessions, _ := OptimizedTurnSessions(profile, provider, options) + if turnSessions == nil && options.ModelRPM == nil { return modelSwitcher, nil } @@ -78,12 +83,14 @@ func EscalationSwitchers( if onSwitch != nil { onSwitch(modelID) } - if optimized, ok := OptimizedTurnSessions(switchedProfile, switchedProvider, Options{}); ok { - return optimized, nil + if turnSessions != nil { + if optimized, ok := OptimizedTurnSessions(switchedProfile, switchedProvider, options); ok { + return limitTurnSessions(switchedProfile, optimized, options), nil + } } // Ineligible target: the default adapter, but carrying the switched // model's own capability projection rather than the original's. - return DefaultTurnSessions(switchedProfile, switchedProvider, Options{}), nil + return limitTurnSessions(switchedProfile, DefaultTurnSessions(switchedProfile, switchedProvider, options), options), nil } return modelSwitcher, sessionSwitcher } diff --git a/internal/providers/factory.go b/internal/providers/factory.go index 965945615..5d692baf1 100644 --- a/internal/providers/factory.go +++ b/internal/providers/factory.go @@ -21,6 +21,7 @@ import ( // Options configures provider construction. type Options struct { + ModelRPM *zeroruntime.ModelRPMLimiter UserAgent string HTTPClient *http.Client ModelRegistry *modelregistry.Registry diff --git a/internal/providers/model_rpm_test.go b/internal/providers/model_rpm_test.go new file mode 100644 index 000000000..09497c6bd --- /dev/null +++ b/internal/providers/model_rpm_test.go @@ -0,0 +1,97 @@ +package providers + +import ( + "context" + "errors" + "github.com/Gitlawb/zero/internal/config" + "github.com/Gitlawb/zero/internal/zeroruntime" + "io" + "net/http" + "strings" + "sync/atomic" + "testing" +) + +type rpmRoundTripper struct{ calls atomic.Int64 } + +func (r *rpmRoundTripper) RoundTrip(q *http.Request) (*http.Response, error) { + r.calls.Add(1) + return &http.Response{StatusCode: 200, Header: http.Header{"Content-Type": {"text/event-stream"}}, Body: io.NopCloser(strings.NewReader("data: [DONE]\n\n")), Request: q}, nil +} +func TestConfiguredRPMBlocksNPlusOneBeforeHTTP(t *testing.T) { + for _, opt := range []string{"0", "1"} { + t.Run("optimized="+opt, func(t *testing.T) { + t.Setenv(openaiTurnSessionEnv, opt) + r := &rpmRoundTripper{} + o := Options{HTTPClient: &http.Client{Transport: r}, ModelRPM: zeroruntime.NewModelRPMLimiter(map[string]int{"gpt-4.1": 2})} + profile := config.ProviderProfile{Name: "test", ProviderKind: config.ProviderKindOpenAI, Model: "openai:gpt-4.1", APIKey: "fixture-only"} + p, e := New(profile, o) + if e != nil { + t.Fatal(e) + } + for i := 0; i < 3; i++ { + if i == 2 { + profile.Model = "gpt-4.1" + } + ss, ok := ConfiguredTurnSessions(profile, p, o) + if !ok { + t.Fatal("missing wrapper") + } + s, e := ss.OpenTurnSession(context.Background()) + if e != nil { + t.Fatal(e) + } + ch, e := s.Stream(context.Background(), zeroruntime.CompletionRequest{Messages: []zeroruntime.Message{{Role: zeroruntime.MessageRoleUser, Content: "fixture"}}}) + if i == 2 { + var hit *zeroruntime.RPMLimitError + if !errors.As(e, &hit) || ch != nil { + t.Fatalf("N+1 was not denied: %v", e) + } + } else { + if e != nil { + t.Fatal(e) + } + for range ch { + } + } + if e := s.Close(); e != nil { + t.Fatal(e) + } + } + if r.calls.Load() != 2 { + t.Fatalf("transport=%d want 2", r.calls.Load()) + } + }) + } +} +func TestEscalationSharesRPMWhenStartingUnoptimized(t *testing.T) { + t.Setenv(openaiTurnSessionEnv, "0") + o := Options{ModelRPM: zeroruntime.NewModelRPMLimiter(map[string]int{"gpt-4.1": 1})} + profile := config.ProviderProfile{Name: "test", ProviderKind: config.ProviderKindOpenAI, Model: "gpt-4.1"} + p := escalationStubProvider{} + ss, _ := ConfiguredTurnSessions(profile, p, o) + s, e := ss.OpenTurnSession(context.Background()) + if e != nil { + t.Fatal(e) + } + if _, e = s.Stream(context.Background(), zeroruntime.CompletionRequest{}); e != nil { + t.Fatal(e) + } + _, switcher := EscalationSwitchers(profile, p, func(config.ProviderProfile) (zeroruntime.Provider, error) { return p, nil }, nil, o) + if switcher == nil { + t.Fatal("missing shared session switcher") + } + ss, e = switcher(context.Background(), "openai:gpt-4.1") + if e != nil { + t.Fatal(e) + } + s, e = ss.OpenTurnSession(context.Background()) + if e != nil { + t.Fatal(e) + } + _, e = s.Stream(context.Background(), zeroruntime.CompletionRequest{}) + var hit *zeroruntime.RPMLimitError + if !errors.As(e, &hit) { + t.Fatalf("switch reset cap: %v", e) + } +} diff --git a/internal/providers/turn_session.go b/internal/providers/turn_session.go index 83bba1920..21d8bdca5 100644 --- a/internal/providers/turn_session.go +++ b/internal/providers/turn_session.go @@ -101,3 +101,25 @@ func DefaultTurnSessions(profile config.ProviderProfile, provider zeroruntime.Pr } return zeroruntime.NewProviderTurnSessionProvider(provider, caps) } + +// ConfiguredTurnSessions preserves transport selection and wraps the starting +// session when RPM is configured. Without either feature, the default is nil. +func ConfiguredTurnSessions(profile config.ProviderProfile, p zeroruntime.Provider, o Options) (zeroruntime.TurnSessionProvider, bool) { + s, ok := OptimizedTurnSessions(profile, p, o) + if o.ModelRPM == nil { + return s, ok + } + if !ok { + s = DefaultTurnSessions(profile, p, o) + } + return limitTurnSessions(profile, s, o), true +} +func limitTurnSessions(profile config.ProviderProfile, s zeroruntime.TurnSessionProvider, o Options) zeroruntime.TurnSessionProvider { + model := strings.TrimSpace(profile.Model) + if registry, err := defaultRegistry(o.ModelRegistry); err == nil { + if id, ok := registry.ResolveID(model); ok { + model = id + } + } + return o.ModelRPM.Wrap(model, s) +} diff --git a/internal/tui/model.go b/internal/tui/model.go index 3c69814df..e06de2105 100644 --- a/internal/tui/model.go +++ b/internal/tui/model.go @@ -88,6 +88,7 @@ type model struct { // allowEscalation mirrors Options.AllowEscalation: it gates the per-run model // switchers, and the caller gates the escalate_model tool on the same flag. allowEscalation bool + modelRPM *zeroruntime.ModelRPMLimiter newProvider func(config.ProviderProfile) (zeroruntime.Provider, error) newTurnSessionProvider func(config.ProviderProfile, zeroruntime.Provider) zeroruntime.TurnSessionProvider probeProviderHealth func(context.Context, providerhealth.Options) providerhealth.Result @@ -1022,6 +1023,7 @@ func newModel(ctx context.Context, options Options) model { sandboxSetupCommand: options.SandboxSetupCommand, agentOptions: options.AgentOptions, allowEscalation: options.AllowEscalation, + modelRPM: options.ModelRPM, sessionCompactor: options.SessionCompactor, runtimeMessageSink: options.RuntimeMessageSink, permissionMode: permissionMode, @@ -5594,6 +5596,7 @@ func (m model) runAgentWithOptions(runID int, runCtx context.Context, prompt str // currentModel: every usage event after a real escalation is billed // to the escalated model, not the one the run started on. func(modelID string) { usageModelID = modelID }, + providers.Options{ModelRPM: m.modelRPM}, ) } diff --git a/internal/tui/options.go b/internal/tui/options.go index 638e82085..cd65ef09a 100644 --- a/internal/tui/options.go +++ b/internal/tui/options.go @@ -22,6 +22,7 @@ import ( // Options configures the reusable Zero terminal UI shell. type Options struct { + ModelRPM *zeroruntime.ModelRPMLimiter Cwd string Version string // CLI build version, shown on the home screen; empty hides it UserConfigPath string diff --git a/internal/zeroruntime/rpm.go b/internal/zeroruntime/rpm.go new file mode 100644 index 000000000..916ad117f --- /dev/null +++ b/internal/zeroruntime/rpm.go @@ -0,0 +1,101 @@ +package zeroruntime + +import ( + "context" + "fmt" + "sync" + "time" +) + +// ModelRPMLimiter shares a sliding 60-second window across turn sessions in +// one runtime process. Separate processes do not share state. +type ModelRPMLimiter struct { + mu sync.Mutex + limits map[string]int + windows map[string][]time.Time + now func() time.Time +} + +// NewModelRPMLimiter copies validated canonical-model limits. Zero/unset is off. +func NewModelRPMLimiter(limits map[string]int) *ModelRPMLimiter { + active := make(map[string]int) + for m, n := range limits { + if n > 0 { + active[m] = n + } + } + if len(active) == 0 { + return nil + } + return &ModelRPMLimiter{limits: active, windows: make(map[string][]time.Time), now: time.Now} +} + +// RPMLimitError denotes rejection before the wrapped Stream method is called. +type RPMLimitError struct { + Model string + Limit int + RetryAfter time.Duration +} + +func (e *RPMLimitError) Error() string { + retry := e.RetryAfter.Round(time.Second) + if retry < e.RetryAfter { + retry += time.Second + } + return fmt.Sprintf("rate limit error: local RPM limit for %q (%d requests/60s); retry in %s", e.Model, e.Limit, retry) +} + +// Wrap gates Stream admissions, not setup/prewarm, compaction or transport retries. +func (l *ModelRPMLimiter) Wrap(model string, p TurnSessionProvider) TurnSessionProvider { + if l == nil || l.limits[model] == 0 || p == nil { + return p + } + return rpmProvider{TurnSessionProvider: p, limiter: l, model: model} +} +func (l *ModelRPMLimiter) admit(ctx context.Context, model string) error { + l.mu.Lock() + defer l.mu.Unlock() + if err := ctx.Err(); err != nil { + return err + } + now := l.now() + w := l.windows[model] + first := 0 + for first < len(w) && !w[first].After(now.Add(-time.Minute)) { + first++ + } + w = w[:copy(w, w[first:])] + l.windows[model] = w + if len(w) >= l.limits[model] { + return &RPMLimitError{Model: model, Limit: l.limits[model], RetryAfter: w[0].Add(time.Minute).Sub(now)} + } + l.windows[model] = append(w, now) + return nil +} + +type rpmProvider struct { + TurnSessionProvider + limiter *ModelRPMLimiter + model string +} + +func (p rpmProvider) OpenTurnSession(ctx context.Context) (TurnSession, error) { + s, err := p.TurnSessionProvider.OpenTurnSession(ctx) + if err != nil { + return nil, err + } + return rpmSession{TurnSession: s, limiter: p.limiter, model: p.model}, nil +} + +type rpmSession struct { + TurnSession + limiter *ModelRPMLimiter + model string +} + +func (s rpmSession) Stream(ctx context.Context, r CompletionRequest) (<-chan StreamEvent, error) { + if err := s.limiter.admit(ctx, s.model); err != nil { + return nil, err + } + return s.TurnSession.Stream(ctx, r) +} diff --git a/internal/zeroruntime/rpm_test.go b/internal/zeroruntime/rpm_test.go new file mode 100644 index 000000000..077f9e56a --- /dev/null +++ b/internal/zeroruntime/rpm_test.go @@ -0,0 +1,128 @@ +package zeroruntime + +import ( + "context" + "errors" + "sync" + "sync/atomic" + "testing" + "time" +) + +type rpmTransport struct{ calls atomic.Int64 } + +func (p *rpmTransport) StreamCompletion(context.Context, CompletionRequest) (<-chan StreamEvent, error) { + p.calls.Add(1) + ch := make(chan StreamEvent) + close(ch) + return ch, nil +} +func rpmSessionFor(t *testing.T, l *ModelRPMLimiter, m string, p Provider) TurnSession { + t.Helper() + s, e := l.Wrap(m, NewProviderTurnSessionProvider(p, ProviderCapabilities{Model: m})).OpenTurnSession(context.Background()) + if e != nil { + t.Fatal(e) + } + return s +} +func TestRPMRejectsNPlusOneBeforeTransportAcrossSessions(t *testing.T) { + l := NewModelRPMLimiter(map[string]int{"m": 3}) + clock := time.Unix(0, 0) + l.now = func() time.Time { return clock } + p := &rpmTransport{} + for i := 0; i < 3; i++ { + if _, e := rpmSessionFor(t, l, "m", p).Stream(context.Background(), CompletionRequest{}); e != nil { + t.Fatal(e) + } + } + s := rpmSessionFor(t, l, "m", p) + ch, e := s.Stream(context.Background(), CompletionRequest{}) + var hit *RPMLimitError + if !errors.As(e, &hit) || ch != nil || hit.RetryAfter != time.Minute { + t.Fatalf("N+1: stream=%v error=%v", ch, e) + } + if p.calls.Load() != 3 { + t.Fatalf("transport calls=%d, want 3", p.calls.Load()) + } + clock = clock.Add(time.Minute) + if _, e = s.Stream(context.Background(), CompletionRequest{}); e != nil { + t.Fatal(e) + } +} +func TestRPMSlidingWindow(t *testing.T) { + l := NewModelRPMLimiter(map[string]int{"m": 2}) + clock := time.Unix(59, 0) + l.now = func() time.Time { return clock } + s := rpmSessionFor(t, l, "m", &rpmTransport{}) + for i := 0; i < 2; i++ { + if _, e := s.Stream(context.Background(), CompletionRequest{}); e != nil { + t.Fatal(e) + } + } + clock = time.Unix(60, 0) + _, e := s.Stream(context.Background(), CompletionRequest{}) + var hit *RPMLimitError + if !errors.As(e, &hit) || hit.RetryAfter != 59*time.Second { + t.Fatalf("minute boundary reset window: %v", e) + } +} +func TestRPMConcurrentAdmissions(t *testing.T) { + l := NewModelRPMLimiter(map[string]int{"m": 7}) + l.now = func() time.Time { return time.Unix(0, 0) } + p := &rpmTransport{} + var grants, unexpected atomic.Int64 + var wg sync.WaitGroup + start := make(chan struct{}) + for i := 0; i < 100; i++ { + s := rpmSessionFor(t, l, "m", p) + wg.Add(1) + go func() { + defer wg.Done() + <-start + _, e := s.Stream(context.Background(), CompletionRequest{}) + var hit *RPMLimitError + if e == nil { + grants.Add(1) + } else if !errors.As(e, &hit) { + unexpected.Add(1) + } + }() + } + close(start) + wg.Wait() + if grants.Load() != 7 || p.calls.Load() != 7 || unexpected.Load() != 0 { + t.Fatalf("grants=%d transport=%d unexpected=%d", grants.Load(), p.calls.Load(), unexpected.Load()) + } +} +func TestRPMCancellationFailedAttemptsAndIndependentModels(t *testing.T) { + l := NewModelRPMLimiter(map[string]int{"a": 1, "b": 1}) + p := &recordingProvider{err: errors.New("fixture transport failure")} + s := rpmSessionFor(t, l, "a", p) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if _, e := s.Stream(ctx, CompletionRequest{}); !errors.Is(e, context.Canceled) { + t.Fatal(e) + } + if _, e := s.Stream(context.Background(), CompletionRequest{}); !errors.Is(e, p.err) { + t.Fatal(e) + } + _, e := s.Stream(context.Background(), CompletionRequest{}) + var hit *RPMLimitError + if !errors.As(e, &hit) || len(p.requests) != 1 { + t.Fatalf("failed attempt refunded: %v", e) + } + _, _ = rpmSessionFor(t, l, "b", p).Stream(context.Background(), CompletionRequest{}) + for i := 0; i < 3; i++ { + _, _ = rpmSessionFor(t, l, "unlimited", p).Stream(context.Background(), CompletionRequest{}) + } + if len(p.requests) != 5 { + t.Fatal("other model incorrectly blocked") + } +} +func TestRPMDisabled(t *testing.T) { + l := NewModelRPMLimiter(map[string]int{"m": 0}) + p := &providerTurnSessionProvider{provider: &rpmTransport{}} + if l != nil || l.Wrap("m", p) != p { + t.Fatal("disabled limit altered provider") + } +}