diff --git a/.env.template b/.env.template index 0b1234f7e..c2bca0d51 100644 --- a/.env.template +++ b/.env.template @@ -348,6 +348,16 @@ # CIRCUIT_BREAKER_SUCCESS_THRESHOLD=2 # Circuit breaker open-state timeout duration (default: 30s) # CIRCUIT_BREAKER_TIMEOUT=30s +# Open the breaker instantly when an upstream error message matches a rule. +# Rules are named groups; declare as many as you like with one MATCH +# (+ optional TTL) var per group. A zero TTL uses CIRCUIT_BREAKER_TIMEOUT +# above. Unset = feature inert. +# CIRCUIT_BREAKER_TRIP_ON_WEEKLY_QUOTA_MATCH="weekly \(7-day\) usage limit" +# CIRCUIT_BREAKER_TRIP_ON_WEEKLY_QUOTA_TTL=4h +# CIRCUIT_BREAKER_TRIP_ON_QUOTA_EXCEEDED_MATCH="quota exceeded" # no TTL +# Per-provider override: replaces only the global rule of the same name; +# other global rules stay in effect. +# KIMICODE_CIRCUIT_BREAKER_TRIP_ON_WEEKLY_QUOTA_MATCH="..." # ============================================================================= # Admin API & Dashboard Configuration diff --git a/config/config.example.yaml b/config/config.example.yaml index f492be543..8a442e55d 100644 --- a/config/config.example.yaml +++ b/config/config.example.yaml @@ -335,6 +335,20 @@ resilience: failure_threshold: 5 success_threshold: 2 timeout: 30s + # trip_on opens the breaker instantly when an upstream error message + # matches `match` (a Go regexp tested against the error message and code): + # failures are not counted, the breaker simply stays open for `ttl` + # (default: the `timeout` above), then probes recovery as usual. Meant for + # hard quota/billing stops that retries cannot fix. Rules are named groups + # (docker-compose style); they are inert until configured, `trip_on: {}` + # disables them, and an empty match, an invalid regexp, or a negative ttl + # fails config load. Env overrides by group name: + # CIRCUIT_BREAKER_TRIP_ON__MATCH / _TTL (global), + # _CIRCUIT_BREAKER_TRIP_ON__MATCH / _TTL (per provider). + # trip_on: + # insufficient_quota: + # match: "insufficient_quota|exceeded your current quota" + # ttl: 30m guardrails: enabled: false @@ -479,6 +493,10 @@ providers: # max_retries: 5 # circuit_breaker: # enabled: false + # # A trip_on list here replaces the global one entirely; omit to inherit. + # trip_on: + # - match: "billing hard limit reached" + # ttl: 1h anthropic: type: anthropic diff --git a/config/config_helpers_test.go b/config/config_helpers_test.go index 7d94fcc21..bfdd3d15b 100644 --- a/config/config_helpers_test.go +++ b/config/config_helpers_test.go @@ -361,6 +361,20 @@ func TestApplyEnvOverrides(t *testing.T) { assert.Equal(t, 10*time.Second, cfg.Resilience.CircuitBreaker.Timeout) }, }, + { + name: "circuit breaker trip_on override", + envVars: map[string]string{ + "CIRCUIT_BREAKER_TRIP_ON_QUOTA_MATCH": "quota exceeded", + "CIRCUIT_BREAKER_TRIP_ON_QUOTA_TTL": "15m", + "CIRCUIT_BREAKER_TRIP_ON_USAGE_MATCH": "usage limit", + }, + check: func(t *testing.T, cfg *Config) { + assert.Equal(t, TripRuleMap{ + "QUOTA": {Name: "QUOTA", Match: "quota exceeded", TTL: 15 * time.Minute}, + "USAGE": {Name: "USAGE", Match: "usage limit"}, + }, cfg.Resilience.CircuitBreaker.TripOn) + }, + }, } for _, tt := range tests { @@ -375,3 +389,15 @@ func TestApplyEnvOverrides(t *testing.T) { }) } } + +// A malformed trip_on env group fails configuration loading instead of +// silently dropping quota protection. +func TestApplyEnvOverrides_TripOnInvalid(t *testing.T) { + t.Setenv("CIRCUIT_BREAKER_TRIP_ON_BAD_TTL_MATCH", "quota exceeded") + t.Setenv("CIRCUIT_BREAKER_TRIP_ON_BAD_TTL_TTL", "banana") + + cfg := buildDefaultConfig() + err := applyEnvOverrides(cfg) + require.Error(t, err) + assert.Contains(t, err.Error(), "CIRCUIT_BREAKER_TRIP_ON_BAD_TTL_TTL") +} diff --git a/config/env.go b/config/env.go index 35affa592..58978fe11 100644 --- a/config/env.go +++ b/config/env.go @@ -15,6 +15,18 @@ func applyEnvOverrides(cfg *Config) error { if err := applyEnvOverridesValue(reflect.ValueOf(cfg).Elem()); err != nil { return err } + // Trip rules use named env groups (CIRCUIT_BREAKER_TRIP_ON__MATCH / + // _TTL) instead of a single list variable, mirroring the per-provider + // _CIRCUIT_BREAKER_TRIP_ON__* convention. A group + // replaces the config rule of the same name; other rules survive. + groups, err := CollectTripRuleGroups(os.Environ(), "CIRCUIT_BREAKER_TRIP_ON") + if err != nil { + return err + } + if len(groups) > 0 { + cfg.Resilience.CircuitBreaker.TripOn = TripRuleMapFromList( + MergeTripRuleGroups(cfg.Resilience.CircuitBreaker.TripOn.List(), groups)) + } normalizeModelListURL(cfg) applyOfflineMode(cfg) return nil diff --git a/config/resilience.go b/config/resilience.go index 2e1850b6d..be9aefd4c 100644 --- a/config/resilience.go +++ b/config/resilience.go @@ -1,6 +1,11 @@ package config -import "time" +import ( + "fmt" + "slices" + "strings" + "time" +) // RetryConfig holds resolved retry settings for an LLM client. // This is the canonical type shared between config and llmclient. @@ -26,6 +31,147 @@ func DefaultRetryConfig() RetryConfig { } } +// TripRuleConfig opens the circuit breaker instantly when an upstream error +// message matches Match. A zero TTL uses the breaker's open-state timeout. +// TTL encodes as JSON nanoseconds wherever the rule crosses the admin/store +// wire as {match, ttl}. +// +// Name identifies the group the rule belongs to. It comes from the YAML map +// key in `trip_on` blocks and from env group names; it only matters while +// merging config with env overrides and is empty elsewhere. +type TripRuleConfig struct { + Name string `yaml:"-" json:"name,omitempty"` + Match string `yaml:"match" json:"match"` + TTL time.Duration `yaml:"ttl" json:"ttl"` +} + +// TripRuleMap holds named trip rules the way operators configure them: +// docker-compose style groups keyed by name. The yaml shape is +// +// trip_on: +// weekly_quota: +// match: "weekly usage limit" +// ttl: 4h +// +// Unmarshal injects the key into each rule's Name; resolution turns the map +// into TripRuleMap.List() so evaluation order is deterministic. +type TripRuleMap map[string]TripRuleConfig + +// UnmarshalYAML accepts only the named map form. Each rule inherits the key +// as its Name; an empty key is a config error. +func (m *TripRuleMap) UnmarshalYAML(unmarshal func(any) error) error { + var raw map[string]TripRuleConfig + if err := unmarshal(&raw); err != nil { + return fmt.Errorf("trip_on must be a map of name -> {match, ttl}: %w", err) + } + for name, rule := range raw { + if strings.TrimSpace(name) == "" { + return fmt.Errorf("trip_on group names must not be empty") + } + rule.Name = name + raw[name] = rule + } + *m = TripRuleMap(raw) + return nil +} + +// TripRuleMapFromList converts a name-carrying rule list back into a named +// map. A rule with an empty name gets a generated key (`rule_N`) so unnamed +// entries never collide with named ones. +func TripRuleMapFromList(rules []TripRuleConfig) TripRuleMap { + m := make(TripRuleMap, len(rules)) + unnamed := 0 + for _, rule := range rules { + name := rule.Name + if name == "" { + unnamed++ + name = fmt.Sprintf("rule_%d", unnamed) + } + m[name] = rule + } + return m +} + +// List returns the rules with their names injected, ordered by group name so +// "first matching rule wins" stays deterministic no matter how the map was +// assembled (yaml, env, or merged). +func (m TripRuleMap) List() []TripRuleConfig { + rules := make([]TripRuleConfig, 0, len(m)) + for name, rule := range m { + rule.Name = name + rules = append(rules, rule) + } + slices.SortFunc(rules, func(a, b TripRuleConfig) int { return strings.Compare(a.Name, b.Name) }) + return rules +} + +// MergeTripRuleGroups returns the merged rule set: an incoming group replaces +// the existing rule with the same name (case-insensitive); other existing +// rules survive; new groups append. The result is ordered by name. +func MergeTripRuleGroups(existing []TripRuleConfig, groups []TripRuleConfig) []TripRuleConfig { + merged := make(map[string]TripRuleConfig, len(existing)+len(groups)) + for _, rule := range existing { + merged[strings.ToLower(rule.Name)] = rule + } + for _, group := range groups { + merged[strings.ToLower(group.Name)] = group + } + rules := make([]TripRuleConfig, 0, len(merged)) + for _, rule := range merged { + rules = append(rules, rule) + } + slices.SortFunc(rules, func(a, b TripRuleConfig) int { return strings.Compare(a.Name, b.Name) }) + return rules +} + +// CollectTripRuleGroups reads `env` entries for `__MATCH` and +// `__TTL` pairs. A group must carry MATCH; TTL is an optional +// Go duration (zero uses the breaker timeout). Entries with a different +// prefix or no suffix are ignored. Groups are returned ordered by name. +func CollectTripRuleGroups(environ []string, prefix string) ([]TripRuleConfig, error) { + matchKey := prefix + "_" + groups := make(map[string]TripRuleConfig) + for _, entry := range environ { + key, value, ok := strings.Cut(entry, "=") + if !ok || value == "" || !strings.HasPrefix(key, matchKey) { + continue + } + // The attribute is the LAST underscore segment; group names may + // contain underscores themselves. + rest := key[len(matchKey):] + i := strings.LastIndex(rest, "_") + if i < 1 { + continue + } + name, attr := rest[:i], rest[i+1:] + if name == "" { + continue + } + rule := groups[name] + rule.Name = name + switch attr { + case "MATCH": + rule.Match = value + case "TTL": + ttl, err := time.ParseDuration(strings.TrimSpace(value)) + if err != nil { + return nil, fmt.Errorf("%s: invalid ttl %q: %w", key, value, err) + } + rule.TTL = ttl + } + groups[name] = rule + } + rules := make([]TripRuleConfig, 0, len(groups)) + for _, rule := range groups { + if rule.Match == "" { + return nil, fmt.Errorf("%s_%s_MATCH: group declares a TTL but no match pattern", prefix, rule.Name) + } + rules = append(rules, rule) + } + slices.SortFunc(rules, func(a, b TripRuleConfig) int { return strings.Compare(a.Name, b.Name) }) + return rules, nil +} + // CircuitBreakerConfig holds resolved circuit breaker settings. // This is the canonical type shared between config and llmclient. type CircuitBreakerConfig struct { @@ -35,10 +181,11 @@ type CircuitBreakerConfig struct { // Enabled switches the circuit breaker on or off. When false, requests are // never short-circuited regardless of the thresholds below. // Default: true - Enabled bool `yaml:"enabled" env:"CIRCUIT_BREAKER_ENABLED"` - FailureThreshold int `yaml:"failure_threshold" env:"CIRCUIT_BREAKER_FAILURE_THRESHOLD"` - SuccessThreshold int `yaml:"success_threshold" env:"CIRCUIT_BREAKER_SUCCESS_THRESHOLD"` - Timeout time.Duration `yaml:"timeout" env:"CIRCUIT_BREAKER_TIMEOUT"` + Enabled bool `yaml:"enabled" env:"CIRCUIT_BREAKER_ENABLED"` + FailureThreshold int `yaml:"failure_threshold" env:"CIRCUIT_BREAKER_FAILURE_THRESHOLD"` + SuccessThreshold int `yaml:"success_threshold" env:"CIRCUIT_BREAKER_SUCCESS_THRESHOLD"` + Timeout time.Duration `yaml:"timeout" env:"CIRCUIT_BREAKER_TIMEOUT"` + TripOn TripRuleMap `yaml:"trip_on"` } // DefaultCircuitBreakerConfig returns the default circuit breaker settings. @@ -69,12 +216,13 @@ type RawResilienceConfig struct { // RawCircuitBreakerConfig holds optional per-provider circuit breaker overrides from YAML. // Nil fields inherit from the global CircuitBreakerConfig. type RawCircuitBreakerConfig struct { - FailureOnStatuses []string `yaml:"failure_on_statuses"` - Scope *string `yaml:"scope"` - Enabled *bool `yaml:"enabled"` - FailureThreshold *int `yaml:"failure_threshold"` - SuccessThreshold *int `yaml:"success_threshold"` - Timeout *time.Duration `yaml:"timeout"` + FailureOnStatuses []string `yaml:"failure_on_statuses"` + Scope *string `yaml:"scope"` + Enabled *bool `yaml:"enabled"` + FailureThreshold *int `yaml:"failure_threshold"` + SuccessThreshold *int `yaml:"success_threshold"` + Timeout *time.Duration `yaml:"timeout"` + TripOn TripRuleMap `yaml:"trip_on"` } // RawRetryConfig holds optional per-provider retry overrides from YAML. diff --git a/config/resilience_policy.go b/config/resilience_policy.go index aa93181ab..c0d13877b 100644 --- a/config/resilience_policy.go +++ b/config/resilience_policy.go @@ -1,6 +1,9 @@ package config -import "fmt" +import ( + "fmt" + "regexp" +) // ParseResilienceStatuses expands exact HTTP codes and classes. Nil uses the // supplied defaults; an explicit empty list disables status-based matches. @@ -40,6 +43,9 @@ func validateResilienceConfig(global ResilienceConfig, providers map[string]RawP if cb.Scope != nil { r.CircuitBreaker.Scope = *cb.Scope } + if cb.TripOn != nil { + r.CircuitBreaker.TripOn = cb.TripOn + } } if err := ValidateResilience(r); err != nil { return fmt.Errorf("providers.%s.resilience: %w", name, err) @@ -61,6 +67,18 @@ func ValidateResilience(r ResilienceConfig) error { default: return fmt.Errorf("circuit_breaker.scope must be provider or model") } + for _, rule := range r.CircuitBreaker.TripOn.List() { + name := "trip_on[" + rule.Name + "]" + if rule.Match == "" { + return fmt.Errorf("circuit_breaker.%s: match must not be empty", name) + } + if _, err := regexp.Compile(rule.Match); err != nil { + return fmt.Errorf("circuit_breaker.%s: %w", name, err) + } + if rule.TTL < 0 { + return fmt.Errorf("circuit_breaker.%s: ttl must not be negative", name) + } + } return nil } diff --git a/config/resilience_policy_test.go b/config/resilience_policy_test.go index f4d942c66..cf841b583 100644 --- a/config/resilience_policy_test.go +++ b/config/resilience_policy_test.go @@ -3,6 +3,7 @@ package config import ( "strings" "testing" + "time" "github.com/stretchr/testify/require" "go.yaml.in/yaml/v3" @@ -16,6 +17,10 @@ func TestResiliencePolicyLoading(t *testing.T) { {"invalid breaker", "resilience:\n circuit_breaker:\n failure_on_statuses: [oops]\n", "circuit_breaker.failure_on_statuses"}, {"invalid scope", "resilience:\n circuit_breaker:\n scope: global\n", "circuit_breaker.scope"}, {"invalid provider", "providers:\n cloudflare:\n resilience:\n retry:\n retry_on_statuses: [oops]\n", "providers.cloudflare.resilience"}, + {"trip_on", "resilience:\n circuit_breaker:\n trip_on:\n quota:\n match: \"quota exceeded\"\n ttl: 5m\n", ""}, + {"invalid trip_on regexp", "resilience:\n circuit_breaker:\n trip_on:\n bad:\n match: \"[quota\"\n", "circuit_breaker.trip_on[bad]"}, + {"empty trip_on match", "resilience:\n circuit_breaker:\n trip_on:\n empty:\n ttl: 5m\n", "circuit_breaker.trip_on[empty]"}, + {"negative trip_on ttl", "resilience:\n circuit_breaker:\n trip_on:\n neg:\n match: \"quota\"\n ttl: -5m\n", "circuit_breaker.trip_on[neg]"}, } { t.Run(tc.name, func(t *testing.T) { clearProviderEnvVars(t) @@ -149,6 +154,28 @@ func TestNormalizeBreakerScope(t *testing.T) { } } +func TestTripOnLoadingAndInheritance(t *testing.T) { + clearProviderEnvVars(t) + dir := t.TempDir() + body := "resilience:\n circuit_breaker:\n trip_on:\n quota:\n match: \"quota\"\n ttl: 5m\n overloaded:\n match: \"overloaded\"\nproviders:\n cloudflare:\n resilience:\n circuit_breaker:\n trip_on:\n quota:\n match: \"quota\"\n ttl: 1m\n openai: {}\n" + writeConfigYAML(t, dir, body) + t.Chdir(dir) + result, err := Load() + require.NoError(t, err) + + require.Equal(t, TripRuleMap{ + "overloaded": {Name: "overloaded", Match: "overloaded"}, + "quota": {Name: "quota", Match: "quota", TTL: 5 * time.Minute}, + }, result.Config.Resilience.CircuitBreaker.TripOn) + require.Equal(t, TripRuleMap{"quota": {Name: "quota", Match: "quota", TTL: time.Minute}}, + result.RawProviders["cloudflare"].Resilience.CircuitBreaker.TripOn) + // A provider without trip_on keeps no raw override; resolution inherits the + // global list (asserted in internal/providers). + require.Nil(t, result.RawProviders["openai"].Resilience) + // Defaults inject no trip rules; the breaker stays inert when unset. + require.Nil(t, DefaultCircuitBreakerConfig().TripOn) +} + func TestProviderPolicyOverrideValidation(t *testing.T) { for _, tc := range []struct{ name, body, wantError string }{ { @@ -171,6 +198,16 @@ func TestProviderPolicyOverrideValidation(t *testing.T) { "resilience:\n circuit_breaker:\n scope: model\nproviders:\n cloudflare:\n resilience:\n circuit_breaker:\n scope: provider\n", "", }, + { + "invalid provider trip_on regexp", + "providers:\n cloudflare:\n resilience:\n circuit_breaker:\n trip_on:\n bad:\n match: \"[quota\"\n", + "providers.cloudflare.resilience: circuit_breaker.trip_on[bad]", + }, + { + "valid provider trip_on replaces the global list", + "resilience:\n circuit_breaker:\n trip_on:\n global:\n match: \"global\"\nproviders:\n cloudflare:\n resilience:\n circuit_breaker:\n trip_on:\n quota:\n match: \"quota\"\n ttl: 1m\n", + "", + }, } { t.Run(tc.name, func(t *testing.T) { clearProviderEnvVars(t) diff --git a/docs/advanced/resilience.mdx b/docs/advanced/resilience.mdx index ac0e2821f..ec4f845e1 100644 --- a/docs/advanced/resilience.mdx +++ b/docs/advanced/resilience.mdx @@ -36,6 +36,7 @@ you need. | `failure_threshold` | `5` | Consecutive failures before the circuit opens | | `success_threshold` | `2` | Consecutive successes to close it again | | `timeout` | `30s` | How long the circuit stays open before probing | +| `trip_on` | `[]` | Upstream error-message regexps that open the breaker instantly | To disable retries, set `max_retries: 0` — every request gets exactly one attempt. To disable the circuit breaker, set `circuit_breaker.enabled: false` @@ -144,6 +145,111 @@ preserved. If every slot is protected, a new model receives a failover-eligible State resets on restart and nothing is shared between replicas. +## Trip rules + +`trip_on` rules open the circuit breaker **instantly** when an upstream error +message matches a Go regular expression — without counting failures toward +`failure_threshold`. Use them for errors that mean "stop calling this +provider", such as exhausted quota or a disabled billing key, where retrying +only burns requests. + +```yaml +resilience: + circuit_breaker: + timeout: 30s + trip_on: # docker-compose style named groups + insufficient_quota: + match: "insufficient_quota|exceeded your current quota" + ttl: 30m +``` + +- `match` is a Go regexp tested against the provider's error message and + error code. +- `ttl` is how long the breaker stays open. Omit it to reuse `timeout` + (default `30s`). +- The group name only matters for env overrides; resolution evaluates rules + in group-name order. +- Rules are inert until configured: without them the breaker behaves exactly + as described above. `trip_on: {}` disables them explicitly. +- A provider's `trip_on` map **replaces** the global map entirely; omit the + key to inherit the global rules. + +### Environment variables + +The same rules are settable without a config file — one `MATCH` (plus +optional `TTL`) variable per named group: + +```bash +# global +CIRCUIT_BREAKER_TRIP_ON_WEEKLY_QUOTA_MATCH="weekly \(7-day\) usage limit" +CIRCUIT_BREAKER_TRIP_ON_WEEKLY_QUOTA_TTL=4h + +# per provider +KIMICODE_CIRCUIT_BREAKER_TRIP_ON_WEEKLY_QUOTA_MATCH="..." +``` + +Env groups **override config rules by name only** (case-insensitive): a group +named like an existing config rule replaces it; config rules with other names +survive; env-only names append. A malformed group (missing `MATCH`, invalid +duration) fails config load, matching YAML behavior. +- Trip rules match on translated (non-streaming) and streaming errors; raw + passthrough routes rely on status-based breaker behavior instead. +- An empty `match`, a pattern that does not compile, or a negative `ttl` + fails configuration loading with a `circuit_breaker.trip_on[i]` error. + +Override per provider like any other breaker setting: + +```yaml +providers: + anthropic: + type: anthropic + api_key: ${ANTHROPIC_API_KEY} + resilience: + circuit_breaker: + trip_on: + credit_balance: # group name + match: "credit balance is too low" + ttl: 1h # replaces the global trip_on rules +``` + +A tripped breaker behaves like an ordinary open one: requests fail fast with +a `503` before reaching the provider, failover virtual models route around +it, and after `ttl` the half-open probe decides whether traffic resumes. If +the probe sees the same error, the trip window restarts. + +### Resetting a tripped breaker + +Once the quota or key is fixed, you do not have to wait out the `ttl`. Use +the **Reset breaker** button on the provider's dashboard card (enabled while +the breaker is open or half-open), or call the admin API: + +```bash +curl -X POST http://localhost:8080/admin/providers/anthropic/circuit-breaker/reset \ + -H "Authorization: Bearer $GOMODEL_MASTER_KEY" +``` + +`204` resumes traffic immediately. `404` means the provider name is unknown, +`400` that the provider's adapter cannot reset its breaker, and `503` that +the reset hook is unavailable on the gateway. + +### Why did my provider pause? + +A provider card showing `Circuit Open` while requests fail with `503` before +reaching the provider means the breaker is open — either `failure_threshold` +failures accumulated or a `trip_on` rule matched. `GET +/admin/providers/status` reports the live state per provider as +`circuit_state` (`closed`, `open`, or `half-open`; empty until the provider +has served traffic) alongside the effective `resilience.circuit_breaker` +settings, including `trip_on`. + +Providers configured in the dashboard can edit their trip rules directly in +the credential editor ("Quota breaker"). Dashboard-managed providers inherit +the global rules unless they define their own list; clearing all rule rows +restores the inherited list (yaml's explicit `trip_on: []` to disable has no +dashboard-managed equivalent). Providers declared in `config.yaml` or +environment variables show their rules read-only there; change them in the +YAML file. Trip rules have no environment-variable equivalent. + ## Environment Variables These set the global defaults that apply to every provider unless overridden diff --git a/internal/admin/handler.go b/internal/admin/handler.go index 4bbddc57a..84082cf44 100644 --- a/internal/admin/handler.go +++ b/internal/admin/handler.go @@ -58,6 +58,7 @@ type Handler struct { configuredProviders []providers.SanitizedProviderConfig providerCredentials ProviderCredentialsAdmin requestHealth RequestHealthSource + breakerResetter BreakerResetter quotaTemplates bool mutationMu sync.Mutex @@ -137,6 +138,11 @@ type providerStatusItemResponse struct { LastError string `json:"last_error,omitempty"` Config providers.SanitizedProviderConfig `json:"config"` Runtime providers.ProviderRuntimeSnapshot `json:"runtime"` + // CircuitState mirrors the live circuit-breaker state ("open", + // "half-open", ...) from request health; empty until the provider has + // served traffic or no breaker state is tracked. The dashboard drives + // the reset button off this field. + CircuitState string `json:"circuit_state"` // RequestHealth reports windowed real-traffic outcomes (per-model error // counts and the live circuit-breaker state); nil when the provider has // served no recent requests or request-health tracking is not wired. @@ -369,6 +375,19 @@ func WithRequestHealth(source RequestHealthSource) Option { } } +// BreakerResetter force-closes a named provider's circuit breaker(s), +// backing the circuit-breaker reset endpoint. +type BreakerResetter interface { + ResetCircuitBreaker(providerName string) error +} + +// WithBreakerResetter enables POST /admin/providers/{name}/circuit-breaker/reset. +func WithBreakerResetter(resetter BreakerResetter) Option { + return func(h *Handler) { + h.breakerResetter = resetter + } +} + // WithDashboardRuntimeConfig enables the allowlisted dashboard runtime config endpoint. func WithDashboardRuntimeConfig(values DashboardConfigResponse) Option { return func(h *Handler) { diff --git a/internal/admin/handler_provider_breaker_reset_test.go b/internal/admin/handler_provider_breaker_reset_test.go new file mode 100644 index 000000000..c43c0dae3 --- /dev/null +++ b/internal/admin/handler_provider_breaker_reset_test.go @@ -0,0 +1,95 @@ +package admin + +import ( + "errors" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + "github.com/labstack/echo/v5" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/enterpilot/gomodel/internal/echotest" + "github.com/enterpilot/gomodel/internal/providers" +) + +// breakerResetterFake records reset calls and replays a canned error. +type breakerResetterFake struct { + calls []string + err error +} + +func (f *breakerResetterFake) ResetCircuitBreaker(providerName string) error { + f.calls = append(f.calls, providerName) + return f.err +} + +// resetBreakerRequest builds the handler call for POST .../circuit-breaker/reset. +func resetBreakerRequest(t *testing.T, name string, fake *breakerResetterFake) (*Handler, *echo.Context, *httptest.ResponseRecorder) { + t.Helper() + h := NewHandler(nil, nil, WithBreakerResetter(fake)) + c, rec := echotest.Post(t, "/admin/providers/"+name+"/circuit-breaker/reset", nil, + echotest.WithPathValue("name", name)) + return h, c, rec +} + +func TestResetProviderCircuitBreaker_Success(t *testing.T) { + fake := &breakerResetterFake{} + h, c, rec := resetBreakerRequest(t, "openai-main", fake) + + err := h.ResetProviderCircuitBreaker(c) + require.NoError(t, err) + require.Equal(t, http.StatusNoContent, rec.Code, rec.Body.String()) + assert.Empty(t, rec.Body.String()) + assert.Equal(t, []string{"openai-main"}, fake.calls) +} + +func TestResetProviderCircuitBreaker_UnknownProvider(t *testing.T) { + fake := &breakerResetterFake{ + err: fmt.Errorf("%w: nope", providers.ErrProviderNotFound), + } + h, c, rec := resetBreakerRequest(t, "nope", fake) + + err := h.ResetProviderCircuitBreaker(c) + require.NoError(t, err) + require.Equal(t, http.StatusNotFound, rec.Code, rec.Body.String()) + assert.Contains(t, rec.Body.String(), "provider not found") +} + +func TestResetProviderCircuitBreaker_ProviderCannotReset(t *testing.T) { + fake := &breakerResetterFake{ + err: errors.New(`provider "legacy" does not support circuit breaker reset`), + } + h, c, rec := resetBreakerRequest(t, "legacy", fake) + + err := h.ResetProviderCircuitBreaker(c) + require.NoError(t, err) + require.Equal(t, http.StatusBadRequest, rec.Code, rec.Body.String()) + assert.Contains(t, rec.Body.String(), "does not support circuit breaker reset") +} + +func TestResetProviderCircuitBreaker_FeatureUnavailable(t *testing.T) { + h := NewHandler(nil, nil) + c, rec := echotest.Post(t, "/admin/providers/openai-main/circuit-breaker/reset", nil, + echotest.WithPathValue("name", "openai-main")) + + err := h.ResetProviderCircuitBreaker(c) + require.NoError(t, err) + require.Equal(t, http.StatusServiceUnavailable, rec.Code, rec.Body.String()) +} + +func TestResetProviderCircuitBreaker_EmptyName(t *testing.T) { + fake := &breakerResetterFake{} + h := NewHandler(nil, nil, WithBreakerResetter(fake)) + // A whitespace-only name trims to empty: the resetter must not be called. + c, rec := echotest.Post(t, "/admin/providers/tmp/circuit-breaker/reset", nil, + echotest.WithPathValue("name", " ")) + + err := h.ResetProviderCircuitBreaker(c) + require.NoError(t, err) + assert.Equal(t, http.StatusBadRequest, rec.Code, rec.Body.String()) + assert.Contains(t, rec.Body.String(), "provider name is required") + assert.Empty(t, fake.calls, "an empty provider name must not reach the resetter") +} diff --git a/internal/admin/handler_provider_credentials.go b/internal/admin/handler_provider_credentials.go index a630791c1..6ff00f7e7 100644 --- a/internal/admin/handler_provider_credentials.go +++ b/internal/admin/handler_provider_credentials.go @@ -11,6 +11,7 @@ import ( "github.com/labstack/echo/v5" + "github.com/enterpilot/gomodel/config" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/providers" ) @@ -57,7 +58,11 @@ type upsertProviderCredentialRequest struct { ServiceAccountJSONBase64 string `json:"service_account_json_base64,omitempty"` GCPScope string `json:"gcp_scope,omitempty"` Models []string `json:"models,omitempty"` - Enabled *bool `json:"enabled,omitempty"` + // TripOn carries circuit-breaker trip rules as {match, ttl} objects; + // ttl is nanoseconds, config's duration encoding. Plain configuration, + // never redacted. + TripOn []config.TripRuleConfig `json:"trip_on,omitempty"` + Enabled *bool `json:"enabled,omitempty"` } // providerCredentialFieldResponse describes one credential field a provider @@ -88,26 +93,27 @@ type providerCredentialTypeResponse struct { // credential: its definition (secrets redacted) plus whether it is read-only // (config/env-declared). type providerCredentialViewResponse struct { - Name string `json:"name"` - Type string `json:"type"` - APIKeys []string `json:"api_keys,omitempty"` - SessionStickyKeys bool `json:"session_sticky_keys"` - BaseURL string `json:"base_url,omitempty"` - APIVersion string `json:"api_version,omitempty"` - Backend string `json:"backend,omitempty"` - AuthType string `json:"auth_type,omitempty"` - APIMode string `json:"api_mode,omitempty"` - VertexProject string `json:"vertex_project,omitempty"` - VertexLocation string `json:"vertex_location,omitempty"` - ServiceAccountFile string `json:"service_account_file,omitempty"` - ServiceAccountJSON string `json:"service_account_json,omitempty"` - ServiceAccountJSONBase64 string `json:"service_account_json_base64,omitempty"` - GCPScope string `json:"gcp_scope,omitempty"` - Models []string `json:"models,omitempty"` - Enabled bool `json:"enabled"` - Managed bool `json:"managed"` - CreatedAt *time.Time `json:"created_at,omitempty"` - UpdatedAt *time.Time `json:"updated_at,omitempty"` + Name string `json:"name"` + Type string `json:"type"` + APIKeys []string `json:"api_keys,omitempty"` + SessionStickyKeys bool `json:"session_sticky_keys"` + BaseURL string `json:"base_url,omitempty"` + APIVersion string `json:"api_version,omitempty"` + Backend string `json:"backend,omitempty"` + AuthType string `json:"auth_type,omitempty"` + APIMode string `json:"api_mode,omitempty"` + VertexProject string `json:"vertex_project,omitempty"` + VertexLocation string `json:"vertex_location,omitempty"` + ServiceAccountFile string `json:"service_account_file,omitempty"` + ServiceAccountJSON string `json:"service_account_json,omitempty"` + ServiceAccountJSONBase64 string `json:"service_account_json_base64,omitempty"` + GCPScope string `json:"gcp_scope,omitempty"` + Models []string `json:"models,omitempty"` + TripOn []config.TripRuleConfig `json:"trip_on,omitempty"` + Enabled bool `json:"enabled"` + Managed bool `json:"managed"` + CreatedAt *time.Time `json:"created_at,omitempty"` + UpdatedAt *time.Time `json:"updated_at,omitempty"` } // ListProviderCredentials handles GET /admin/provider-credentials. @@ -250,6 +256,12 @@ func (h *Handler) UpsertProviderCredential(c *echo.Context) error { if err != nil { return handleError(c, err) } + // Validate trip rules regardless of enabled state: an invalid regex or + // negative TTL stored in a disabled credential would surface only at + // enable time and break the operator's workflow. + if err := config.ValidateResilience(config.ResilienceConfig{CircuitBreaker: config.CircuitBreakerConfig{TripOn: config.TripRuleMapFromList(cred.TripOn)}}); err != nil { + return handleError(c, core.NewInvalidRequestError("invalid trip_on rules: "+err.Error(), nil)) + } if err := h.providerCredentials.Upsert(c.Request().Context(), cred); err != nil { return handleError(c, providerCredentialWriteError(err)) } @@ -331,6 +343,7 @@ func (h *Handler) buildProviderCredentialUpsert(ctx context.Context, name string ServiceAccountJSONBase64: serviceAccountJSONBase64, GCPScope: strings.TrimSpace(req.GCPScope), Models: req.Models, + TripOn: req.TripOn, Enabled: enabled, } if current != nil { @@ -418,6 +431,7 @@ func (h *Handler) providerCredentialView(cred providers.ManagedProviderCredentia ServiceAccountFile: cred.ServiceAccountFile, GCPScope: cred.GCPScope, Models: cred.Models, + TripOn: cred.TripOn, Enabled: cred.Enabled, Managed: h.providerCredentials.IsManaged(cred.Name), CreatedAt: nonZeroTime(cred.CreatedAt), @@ -445,6 +459,7 @@ func (h *Handler) declaredProviderCredentialView(cfg providers.SanitizedProvider BaseURL: cfg.BaseURL, APIVersion: cfg.APIVersion, Models: cfg.Models, + TripOn: cfg.Resilience.CircuitBreaker.TripOn, SessionStickyKeys: cfg.SessionStickyKeys, Enabled: true, Managed: true, diff --git a/internal/admin/handler_provider_credentials_test.go b/internal/admin/handler_provider_credentials_test.go index d53dbcf45..15fcf9d07 100644 --- a/internal/admin/handler_provider_credentials_test.go +++ b/internal/admin/handler_provider_credentials_test.go @@ -8,13 +8,16 @@ import ( "net/http/httptest" "sort" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/enterpilot/gomodel/config" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/echotest" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/health" ) // providerCredentialsAdminFake is an in-memory ProviderCredentialsAdmin for @@ -449,6 +452,53 @@ func TestProviderCredentialsEndpointsReturn503WhenUnavailable(t *testing.T) { assertUnavailable("DeleteProviderCredential", h.DeleteProviderCredential(deleteCtx), deleteRec) } +// Trip rules are plain configuration, so the upsert stores them, the stored +// view lists them unredacted, and the declared (config.yaml/env) read-only +// view carries the effective rules from the sanitized config. +func TestUpsertProviderCredential_TripRulesRoundTrip(t *testing.T) { + fake := newProviderCredentialsAdminFake() + h := newProviderCredentialsHandler(fake) + + c, rec := echotest.Request(t, http.MethodPut, "/admin/provider-credentials", + `{"name":"my-openai","type":"openai","api_keys":["sk-real"],"trip_on":[{"match":"insufficient_quota","ttl":60000000000}]}`) + err := h.UpsertProviderCredential(c) + require.NoError(t, err) + require.Equal(t, http.StatusOK, rec.Code, rec.Body.String()) + + want := []config.TripRuleConfig{{Match: "insufficient_quota", TTL: 60 * time.Second}} + stored, ok := fake.rows["my-openai"] + require.True(t, ok) + assert.Equal(t, want, stored.TripOn) + + response := echotest.Decode[providerCredentialViewResponse](t, rec) + assert.Equal(t, want, response.TripOn) +} + +func TestListProviderCredentials_DeclaredViewShowsTripRules(t *testing.T) { + fake := newProviderCredentialsAdminFake() + h := newProviderCredentialsHandlerWithConfigured(fake, []providers.SanitizedProviderConfig{ + { + Name: "openai", + Type: "openai", + Resilience: providers.SanitizedResilienceConfig{ + CircuitBreaker: providers.SanitizedCircuitBreakerConfig{ + TripOn: []config.TripRuleConfig{{Match: "rate limit", TTL: 30 * time.Second}}, + }, + }, + }, + }) + + c, rec := echotest.Get(t, "/admin/provider-credentials") + err := h.ListProviderCredentials(c) + require.NoError(t, err) + require.Equal(t, http.StatusOK, rec.Code, rec.Body.String()) + + body := echotest.Decode[[]providerCredentialViewResponse](t, rec) + require.Len(t, body, 1) + assert.True(t, body[0].Managed) + assert.Equal(t, []config.TripRuleConfig{{Match: "rate limit", TTL: 30 * time.Second}}, body[0].TripOn) +} + func TestUpsertProviderCredential_BubblesProviderErrorOnStoreFailure(t *testing.T) { fake := newProviderCredentialsAdminFake() fake.upsertErr = errors.New("disk full") @@ -528,3 +578,88 @@ func TestProviderStatus_ReportsCredentialServiceConfigForRuntimeProviders(t *tes }) } } + +// requestHealthFake replays canned per-provider health snapshots. +type requestHealthFake struct { + snapshot map[string]health.ProviderHealth +} + +func (f requestHealthFake) Snapshot() map[string]health.ProviderHealth { + return f.snapshot +} + +// The status item must surface the live breaker state as a first-class field +// (the dashboard's reset button keys off it) and the effective trip rules as +// plain, unredacted configuration. +func TestProviderStatus_ExposesCircuitStateAndTripRules(t *testing.T) { + tripOn := []config.TripRuleConfig{{Match: "insufficient_quota", TTL: time.Minute}} + fake := newProviderCredentialsAdminFake() + fake.configured = []providers.SanitizedProviderConfig{{ + Name: "dash-openai", + Type: "openai", + Resilience: providers.SanitizedResilienceConfig{ + CircuitBreaker: providers.SanitizedCircuitBreakerConfig{TripOn: tripOn}, + }, + }} + h := NewHandler(nil, nil, + WithProviderCredentials(fake), + WithRequestHealth(requestHealthFake{snapshot: map[string]health.ProviderHealth{ + "dash-openai": {CircuitState: "open"}, + }}), + ) + + c, rec := echotest.Get(t, "/admin/providers/status") + err := h.ProviderStatus(c) + require.NoError(t, err) + require.Equal(t, http.StatusOK, rec.Code, rec.Body.String()) + + body := echotest.Decode[providerStatusResponse](t, rec) + require.Len(t, body.Providers, 1) + + item := body.Providers[0] + assert.Equal(t, "open", item.CircuitState) + assert.Equal(t, tripOn, item.Config.Resilience.CircuitBreaker.TripOn) +} + +// A provider with no traffic yet reports an empty circuit_state rather than a +// made-up state. +func TestProviderStatus_CircuitStateEmptyWithoutTraffic(t *testing.T) { + h := NewHandler(nil, nil, WithConfiguredProviders([]providers.SanitizedProviderConfig{ + {Name: "idle", Type: "openai"}, + })) + + c, rec := echotest.Get(t, "/admin/providers/status") + err := h.ProviderStatus(c) + require.NoError(t, err) + require.Equal(t, http.StatusOK, rec.Code, rec.Body.String()) + + body := echotest.Decode[providerStatusResponse](t, rec) + require.Len(t, body.Providers, 1) + assert.Empty(t, body.Providers[0].CircuitState) +} + +// A disabled credential still validates trip_on: an invalid regex is rejected +// with 400 so operators never store unusable rules that surface only at enable +// time. A negative TTL is rejected the same way. +func TestUpsertProviderCredential_DisabledCredentialWithInvalidTripOnReturns400(t *testing.T) { + fake := newProviderCredentialsAdminFake() + h := newProviderCredentialsHandler(fake) + + t.Run("invalid regex", func(t *testing.T) { + c, rec := echotest.Request(t, http.MethodPut, "/admin/provider-credentials", + `{"name":"x","type":"openai","api_keys":["sk"],"trip_on":[{"match":"(unclosed","ttl":0}],"enabled":false}`) + err := h.UpsertProviderCredential(c) + require.NoError(t, err) + require.Equal(t, http.StatusBadRequest, rec.Code, rec.Body.String()) + assert.Empty(t, fake.rows, "invalid credential must not be stored") + }) + + t.Run("negative ttl", func(t *testing.T) { + c, rec := echotest.Request(t, http.MethodPut, "/admin/provider-credentials", + `{"name":"y","type":"openai","api_keys":["sk"],"trip_on":[{"match":"bad","ttl":-1}],"enabled":false}`) + err := h.UpsertProviderCredential(c) + require.NoError(t, err) + require.Equal(t, http.StatusBadRequest, rec.Code, rec.Body.String()) + assert.Empty(t, fake.rows, "invalid credential must not be stored") + }) +} diff --git a/internal/admin/handler_providers.go b/internal/admin/handler_providers.go index 4a3b628ca..4b9390e68 100644 --- a/internal/admin/handler_providers.go +++ b/internal/admin/handler_providers.go @@ -163,10 +163,50 @@ func buildProviderStatusItem(name string, cfg providers.SanitizedProviderConfig, LastError: lastError, Config: cfg, Runtime: runtime, + CircuitState: circuitStateFor(requestHealth), RequestHealth: requestHealth, } } +// circuitStateFor lifts the live breaker state into a first-class response +// field; empty until the provider has served traffic. +func circuitStateFor(requestHealth *health.ProviderHealth) string { + if requestHealth == nil { + return "" + } + return requestHealth.CircuitState +} + +// ResetProviderCircuitBreaker handles POST /admin/providers/:name/circuit-breaker/reset. +// +// @Summary Force-close a provider's circuit breaker(s) +// @Description Clears an open or half-open breaker (including any quota trip window) so traffic resumes immediately, without a restart. Responds 404 for an unknown provider and 400 for a provider whose adapter cannot reset its breaker. +// @Tags admin +// @Security BearerAuth +// @Param name path string true "Provider name" +// @Success 204 "No Content" +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Failure 404 {object} core.GatewayError +// @Failure 503 {object} core.GatewayError +// @Router /admin/providers/{name}/circuit-breaker/reset [post] +func (h *Handler) ResetProviderCircuitBreaker(c *echo.Context) error { + if h.breakerResetter == nil { + return handleError(c, featureUnavailableError("circuit breaker reset is unavailable")) + } + name := strings.TrimSpace(c.Param("name")) + if name == "" { + return handleError(c, core.NewInvalidRequestError("provider name is required", nil)) + } + if err := h.breakerResetter.ResetCircuitBreaker(name); err != nil { + if errors.Is(err, providers.ErrProviderNotFound) { + return handleError(c, core.NewNotFoundError("provider not found: "+name)) + } + return handleError(c, core.NewInvalidRequestError(err.Error(), err)) + } + return c.NoContent(http.StatusNoContent) +} + // requestHealthFor matches a status row (keyed by trimmed provider name) to // its health snapshot; snapshot keys are trimmed too since llmclient records // the configured name as-is. diff --git a/internal/admin/routes.go b/internal/admin/routes.go index 49c32c4e5..1d93e8600 100644 --- a/internal/admin/routes.go +++ b/internal/admin/routes.go @@ -46,6 +46,7 @@ func (h *Handler) RegisterRoutes(g RouteRegistrar) { g.GET("/audit/conversation", h.AuditConversation) g.GET("/providers/status", h.ProviderStatus, global) + g.POST("/providers/:name/circuit-breaker/reset", h.ResetProviderCircuitBreaker, global) g.POST("/runtime/refresh", h.RefreshRuntime, global) g.GET("/provider-credentials", h.ListProviderCredentials, global) diff --git a/internal/admin/routes_test.go b/internal/admin/routes_test.go index ddc800302..3ecf59354 100644 --- a/internal/admin/routes_test.go +++ b/internal/admin/routes_test.go @@ -54,6 +54,7 @@ func TestRegisterRoutes_RegistersExpectedPaths(t *testing.T) { "GET /admin/audit/conversation", "GET /admin/providers/status", + "POST /admin/providers/:name/circuit-breaker/reset", "POST /admin/runtime/refresh", "GET /admin/provider-credentials", diff --git a/internal/app/init_admin.go b/internal/app/init_admin.go index 5d3aae169..c86b8d787 100644 --- a/internal/app/init_admin.go +++ b/internal/app/init_admin.go @@ -204,6 +204,7 @@ func newAdminHandlers( admin.WithDashboardRuntimeConfig(runtimeConfig), admin.WithLiveBroker(liveBroker), admin.WithRequestHealth(requestHealth), + admin.WithBreakerResetter(providers.NewBreakerResetter(registry)), ) var dashHandler *dashboard.Handler diff --git a/internal/gateway/resilience_failover_test.go b/internal/gateway/resilience_failover_test.go index 05683f762..087493503 100644 --- a/internal/gateway/resilience_failover_test.go +++ b/internal/gateway/resilience_failover_test.go @@ -2,14 +2,18 @@ package gateway import ( "context" + "io" "net/http" - "net/http/httptest" "strings" + "sync/atomic" "testing" "time" + goconfig "github.com/enterpilot/gomodel/config" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -27,21 +31,21 @@ func (p *retryFailoverProvider) ChatCompletion(ctx context.Context, req *core.Ch func TestCloudflareTimeoutRetriesBeforeModelFailover(t *testing.T) { var calls []string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + server, capture := providertest.Server(t, func(w http.ResponseWriter, r *http.Request) { calls = append(calls, r.URL.Path) if r.URL.Path == "/model1" { w.WriteHeader(524) return } _, _ = w.Write([]byte(`{"id":"backup","model":"model2","choices":[{"index":0,"finish_reason":"stop","message":{"role":"assistant","content":"ok"}}]}`)) - })) - defer server.Close() + }) cfg := llmclient.DefaultConfig("cloudflare", server.URL) cfg.Retry.MaxRetries = 2 cfg.Retry.InitialBackoff = time.Nanosecond cfg.CircuitBreaker.Scope = "model" cfg.CircuitBreaker.FailureThreshold = 1 provider := &retryFailoverProvider{client: llmclient.New(cfg, nil)} + _ = capture // uses recorded requests for audit in tests that need it orchestrator := NewInferenceOrchestrator(InferenceConfig{ Provider: provider, FailoverResolver: failoverResolverFunc(func(*core.RequestModelResolution, core.Operation) []core.ModelSelector { @@ -62,3 +66,117 @@ func TestCloudflareTimeoutRetriesBeforeModelFailover(t *testing.T) { got := strings.Join(calls, ",") require.Equal(t, "/model1,/model1,/model1,/model2,/model2", got) } + +// quotaTripProvider is a core.Provider whose ChatCompletion/StreamChatCompletion +// delegate to an llmclient.Client. It satisfies core.Provider and the +// embedded providerTypeResolverStub. +type quotaTripProvider struct { + providerTypeResolverStub + client *llmclient.Client +} + +func (p *quotaTripProvider) ChatCompletion(ctx context.Context, req *core.ChatRequest) (*core.ChatResponse, error) { + var resp core.ChatResponse + err := p.client.Do(ctx, llmclient.Request{Method: "GET", Endpoint: req.Model, Model: req.Model}, &resp) + if err != nil { + return nil, err + } + resp.Provider = req.Provider + return &resp, nil +} + +func (p *quotaTripProvider) StreamChatCompletion(ctx context.Context, req *core.ChatRequest) (io.ReadCloser, error) { + var resp core.ChatResponse + err := p.client.Do(ctx, llmclient.Request{Method: "GET", Endpoint: req.Model, Model: req.Model}, &resp) + if err != nil { + return nil, err + } + return nil, nil +} + +// TestTripRuleQuotaTripsPrimaryInFailoverChain exercises the full trip-on + +// failover + reset cycle: a quota error that matches a trip_on rule opens the +// breaker on the primary provider, the orchestrator falls back to the backup +// target, and a reset allows the primary to be retried. +func TestTripRuleQuotaTripsPrimaryInFailoverChain(t *testing.T) { + t.Parallel() + + var aHits, bHits atomic.Int32 + // Combined server: /model1 returns quota 500 with usage-limit message, + // /model2 returns 200. We use 500 because the default failover policy + // does not treat 403 as retryable, but the trip rule matches the + // message text regardless of status. + server, _ := providertest.Server(t, func(w http.ResponseWriter, r *http.Request) { + path := r.URL.Path + if path == "/model1" { + aHits.Add(1) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte(`{"error":{"message":"You've reached your weekly (7-day) usage limit, please try again after 0:02:52","code":"usage_limit_exceeded"}}`)) + } else { + bHits.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"b-resp","model":"model2","choices":[{"index":0,"finish_reason":"stop","message":{"role":"assistant","content":"ok"}}]}`)) + } + }) + + // Primary provider: quotaTripProvider with trip rules hitting /model1. + primary := "aTripProvider{ + client: llmclient.New(llmclient.Config{ + ProviderName: "primary", + BaseURL: server.URL, + Retry: goconfig.DefaultRetryConfig(), + CircuitBreaker: goconfig.CircuitBreakerConfig{ + Enabled: true, + FailureThreshold: 5, + SuccessThreshold: 1, + Timeout: time.Minute, + Scope: "model", + TripOn: goconfig.TripRuleMap{ + "weekly": {Match: `weekly.*usage limit`, TTL: time.Minute}, + }, + }, + }, nil), + } + + // Orchestrator: primary is the tripClientProvider for /model1; + // failover falls back to the same provider with /model2 (model-scoped + // breaker means /model1 trip doesn't block /model2). + orchestrator := NewInferenceOrchestrator(InferenceConfig{ + Provider: primary, + FailoverResolver: failoverResolverFunc(func(*core.RequestModelResolution, core.Operation) []core.ModelSelector { + return []core.ModelSelector{{Provider: "primary", Model: "/model2"}} + }), + }) + workflow := &core.Workflow{ + Endpoint: core.DescribeEndpoint("POST", "/v1/chat/completions"), + Resolution: &core.RequestModelResolution{ResolvedSelector: core.ModelSelector{Provider: "primary", Model: "/model1"}}, + Policy: &core.ResolvedWorkflowPolicy{Features: core.WorkflowFeatures{Failover: true}}, + } + + // Request 1: trips breaker on model1, failover to model2. + resp, _, err := orchestrator.DispatchChatCompletion(context.Background(), workflow, &core.ChatRequest{Model: "/model1"}) + require.NoError(t, err) + require.Equal(t, "b-resp", resp.ID) + assert.Equal(t, int32(1), aHits.Load(), "model1 should be called once (tripped)") + assert.Equal(t, int32(1), bHits.Load(), "model2 should be called once (failover)") + + // Request 2: model1 breaker is open → tripped → skip. Failover to model2. + resp, _, err = orchestrator.DispatchChatCompletion(context.Background(), workflow, &core.ChatRequest{Model: "/model1"}) + require.NoError(t, err) + require.Equal(t, "b-resp", resp.ID) + assert.Equal(t, int32(1), aHits.Load(), "model1 must not be called again while tripped") + assert.Equal(t, int32(2), bHits.Load(), "model2 serves the second request via failover") + + // Reset breaker via the trip client's ResetBreaker method. + // (In production the admin API calls providers.NewBreakerResetter, which + // type-asserts the adapter and calls ResetBreaker.) + primary.client.ResetBreaker() + + // Request 3: breaker is reset, model1 is contacted again (and trips). + resp, _, err = orchestrator.DispatchChatCompletion(context.Background(), workflow, &core.ChatRequest{Model: "/model1"}) + require.NoError(t, err) + require.Equal(t, "b-resp", resp.ID) + assert.Equal(t, int32(2), aHits.Load(), "model1 must be contacted again after reset") + assert.Equal(t, int32(3), bHits.Load(), "failover serves again as model1 trips once more") +} diff --git a/internal/llmclient/circuit_breaker.go b/internal/llmclient/circuit_breaker.go index 7aa95a084..6320a0318 100644 --- a/internal/llmclient/circuit_breaker.go +++ b/internal/llmclient/circuit_breaker.go @@ -1,10 +1,19 @@ package llmclient import ( + "regexp" "sync" "time" ) +// TripRule opens the breaker instantly when an upstream error message matches +// Pattern, skipping the failure threshold. A zero TTL defers to the breaker's +// open-state timeout at trip time. +type TripRule struct { + Pattern *regexp.Regexp + TTL time.Duration +} + // circuitBreaker implements a circuit breaker pattern with half-open state protection type circuitBreaker struct { mu sync.Mutex @@ -15,7 +24,10 @@ type circuitBreaker struct { successThreshold int timeout time.Duration lastFailure time.Time - halfOpenAllowed bool // Controls single-request probe in half-open state + // quotaUntil keeps the breaker open for a quota trip window even after + // lastFailure ages past timeout. Zero unless a quota trip happened. + quotaUntil time.Time + halfOpenAllowed bool // Controls single-request probe in half-open state } type circuitState int @@ -46,8 +58,8 @@ func (cb *circuitBreaker) acquire() (bool, bool) { case circuitClosed: return true, false case circuitOpen: - // Check if timeout has passed - if time.Since(cb.lastFailure) > cb.timeout { + // Check if timeout has passed and no quota trip window is active + if time.Since(cb.lastFailure) > cb.timeout && time.Now().After(cb.quotaUntil) { cb.state = circuitHalfOpen cb.successes = 0 cb.halfOpenAllowed = true // Allow the first probe request @@ -120,6 +132,34 @@ func (cb *circuitBreaker) RecordFailure() { } } +// RecordQuotaTrip opens the breaker instantly, without counting failures, +// because a matching upstream error proves the provider's quota is gone. The +// breaker stays open for ttl regardless of the open-state timeout; acquire +// honors whichever deadline expires last. +func (cb *circuitBreaker) RecordQuotaTrip(ttl time.Duration) { + cb.mu.Lock() + defer cb.mu.Unlock() + + now := time.Now() + cb.state = circuitOpen + cb.lastFailure = now + cb.quotaUntil = now.Add(ttl) + cb.successes = 0 + cb.halfOpenAllowed = true // Reset for next timeout period +} + +// Reset force-closes the breaker and clears any quota trip window. +func (cb *circuitBreaker) Reset() { + cb.mu.Lock() + defer cb.mu.Unlock() + + cb.state = circuitClosed + cb.failures = 0 + cb.successes = 0 + cb.quotaUntil = time.Time{} + cb.halfOpenAllowed = true +} + // State returns the current circuit state (for testing/monitoring) func (cb *circuitBreaker) State() string { cb.mu.Lock() diff --git a/internal/llmclient/circuit_breaker_trip_test.go b/internal/llmclient/circuit_breaker_trip_test.go new file mode 100644 index 000000000..36c15b507 --- /dev/null +++ b/internal/llmclient/circuit_breaker_trip_test.go @@ -0,0 +1,587 @@ +package llmclient + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "sync/atomic" + "testing" + "time" + + goconfig "github.com/enterpilot/gomodel/config" + "github.com/enterpilot/gomodel/internal/core" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +const quotaErrorBody = `{"error":{"message":"quota exceeded for organization, window refreshes soon","code":"quota_exceeded"}}` + +func newQuotaTripTestClient(t *testing.T, serverURL string, mutate func(cb *goconfig.CircuitBreakerConfig)) *Client { + t.Helper() + + cfg := DefaultConfig("test", serverURL) + cfg.Retry.MaxRetries = 0 + cfg.CircuitBreaker = goconfig.CircuitBreakerConfig{ + Enabled: true, + FailureThreshold: 5, + SuccessThreshold: 1, + Timeout: 20 * time.Millisecond, + TripOn: goconfig.TripRuleMap{ + "quota": {Match: `quota (exceeded|exhausted)`, TTL: 150 * time.Millisecond}, + }, + } + if mutate != nil { + mutate(&cfg.CircuitBreaker) + } + return New(cfg, nil) +} + +// A matching upstream error message opens the breaker on the first failure, +// and the next request is rejected locally without reaching the upstream. +func TestTripRule_MatchingMessageTripsInstantly(t *testing.T) { + t.Parallel() + + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempts.Add(1) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusForbidden) + _, _ = w.Write([]byte(quotaErrorBody)) + })) + defer server.Close() + + // 403 is not a failure status by default, so only the message match can + // explain the trip. + client := newQuotaTripTestClient(t, server.URL, nil) + + err := client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + var gatewayErr *core.GatewayError + require.ErrorAs(t, err, &gatewayErr) + assert.Equal(t, http.StatusForbidden, gatewayErr.StatusCode) + assert.Equal(t, int32(1), attempts.Load()) + assert.Equal(t, "open", client.circuitBreaker.State()) + + err = client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + require.ErrorAs(t, err, &gatewayErr) + assert.Contains(t, gatewayErr.Message, "circuit breaker is open") + assert.Equal(t, int32(1), attempts.Load(), "rejected request must not reach the upstream") +} + +// A matching quota error on a retryable status trips on the first attempt: +// the retry loop stops instead of hammering the quota-exhausted provider +// max_retries+1 times. +func TestTripRule_RetryableQuotaErrorTripsBeforeRetry(t *testing.T) { + t.Parallel() + + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempts.Add(1) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusTooManyRequests) + _, _ = w.Write([]byte(quotaErrorBody)) + })) + defer server.Close() + + cfg := DefaultConfig("test", server.URL) + cfg.Retry.MaxRetries = 2 + cfg.CircuitBreaker = goconfig.CircuitBreakerConfig{ + Enabled: true, + FailureThreshold: 5, + SuccessThreshold: 1, + Timeout: 20 * time.Millisecond, + TripOn: goconfig.TripRuleMap{"quota": {Match: `quota exceeded`, TTL: 150 * time.Millisecond}}, + } + client := New(cfg, nil) + + err := client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + var gatewayErr *core.GatewayError + require.ErrorAs(t, err, &gatewayErr) + assert.Equal(t, http.StatusTooManyRequests, gatewayErr.StatusCode) + assert.Equal(t, int32(1), attempts.Load(), "quota trip must stop the retry loop after the first attempt") + assert.Equal(t, "open", client.circuitBreaker.State()) +} + +// The quota window outlives the breaker timeout: while quotaUntil is in the +// future the breaker stays open even after lastFailure ages past timeout. +// Once the window lapses the half-open probe decides recovery. +func TestTripRule_TTLExpiryAndHalfOpenProbe(t *testing.T) { + t.Parallel() + + t.Run("probe success closes", func(t *testing.T) { + t.Parallel() + + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if attempts.Add(1) == 1 { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusForbidden) + _, _ = w.Write([]byte(quotaErrorBody)) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"message":"ok"}`)) + })) + defer server.Close() + + client := newQuotaTripTestClient(t, server.URL, nil) + + err := client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + assert.Equal(t, "open", client.circuitBreaker.State()) + + // Past the breaker timeout (20ms) but inside the quota window. + time.Sleep(45 * time.Millisecond) + err = client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + var gatewayErr *core.GatewayError + require.ErrorAs(t, err, &gatewayErr) + assert.Contains(t, gatewayErr.Message, "circuit breaker is open") + assert.Equal(t, int32(1), attempts.Load(), "quota window must keep the breaker closed to traffic") + + // Past the quota window: the probe goes upstream and closes the breaker. + time.Sleep(150 * time.Millisecond) + err = client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.NoError(t, err) + assert.Equal(t, int32(2), attempts.Load()) + assert.Equal(t, "closed", client.circuitBreaker.State()) + }) + + t.Run("probe failure reopens", func(t *testing.T) { + t.Parallel() + + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempts.Add(1) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusForbidden) + _, _ = w.Write([]byte(quotaErrorBody)) + })) + defer server.Close() + + client := newQuotaTripTestClient(t, server.URL, nil) + + err := client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + time.Sleep(200 * time.Millisecond) + + // The probe still sees the quota error: the trip window restarts. + err = client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + var gatewayErr *core.GatewayError + require.ErrorAs(t, err, &gatewayErr) + assert.Equal(t, http.StatusForbidden, gatewayErr.StatusCode) + assert.Equal(t, "open", client.circuitBreaker.State()) + + err = client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + require.ErrorAs(t, err, &gatewayErr) + assert.Contains(t, gatewayErr.Message, "circuit breaker is open") + assert.Equal(t, int32(2), attempts.Load()) + }) +} + +// A zero rule TTL defers to the breaker's open-state timeout at trip time. +func TestTripRule_ZeroTTLUsesBreakerTimeout(t *testing.T) { + t.Parallel() + + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if attempts.Add(1) == 1 { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusForbidden) + _, _ = w.Write([]byte(quotaErrorBody)) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"message":"ok"}`)) + })) + defer server.Close() + + client := newQuotaTripTestClient(t, server.URL, func(cb *goconfig.CircuitBreakerConfig) { + cb.Timeout = 50 * time.Millisecond + cb.TripOn = goconfig.TripRuleMap{"quota": {Match: `quota exceeded`}} + }) + + err := client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + assert.Equal(t, "open", client.circuitBreaker.State()) + + err = client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + assert.Equal(t, int32(1), attempts.Load()) + + time.Sleep(80 * time.Millisecond) + err = client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.NoError(t, err) + assert.Equal(t, int32(2), attempts.Load()) + assert.Equal(t, "closed", client.circuitBreaker.State()) +} + +// Without a message match nothing changes: failures still accumulate toward +// the configured threshold instead of tripping instantly. +func TestTripRule_NonMatchingFailureKeepsThresholdSemantics(t *testing.T) { + t.Parallel() + + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempts.Add(1) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte(`{"error":{"message":"internal server error"}}`)) + })) + defer server.Close() + + client := newQuotaTripTestClient(t, server.URL, func(cb *goconfig.CircuitBreakerConfig) { + cb.FailureThreshold = 2 + }) + + err := client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + assert.Equal(t, "closed", client.circuitBreaker.State(), "non-matching failure must not trip instantly") + + err = client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + assert.Equal(t, "open", client.circuitBreaker.State(), "threshold must still open the breaker") + assert.Equal(t, int32(2), attempts.Load()) + + err = client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + var gatewayErr *core.GatewayError + require.ErrorAs(t, err, &gatewayErr) + assert.Contains(t, gatewayErr.Message, "circuit breaker is open") + assert.Equal(t, int32(2), attempts.Load()) +} + +// No rules configured: the feature is inert and even a quota-shaped error +// leaves the breaker closed. +func TestTripRule_NoRulesStaysInert(t *testing.T) { + t.Parallel() + + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempts.Add(1) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusForbidden) + _, _ = w.Write([]byte(quotaErrorBody)) + })) + defer server.Close() + + client := newQuotaTripTestClient(t, server.URL, func(cb *goconfig.CircuitBreakerConfig) { + cb.TripOn = nil + }) + require.Empty(t, client.tripRules) + + for i := range 2 { + err := client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + var gatewayErr *core.GatewayError + require.ErrorAs(t, err, &gatewayErr, "attempt %d", i+1) + assert.Equal(t, http.StatusForbidden, gatewayErr.StatusCode, "attempt %d", i+1) + } + assert.Equal(t, int32(2), attempts.Load()) + assert.Equal(t, "closed", client.circuitBreaker.State()) +} + +// ResetBreaker force-closes the provider-level breaker and every model-scoped +// breaker, ending a quota trip immediately. +func TestResetBreaker_ClearsProviderAndModelBreakers(t *testing.T) { + t.Parallel() + + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if attempts.Add(1) == 1 { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusForbidden) + _, _ = w.Write([]byte(quotaErrorBody)) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"message":"ok"}`)) + })) + defer server.Close() + + client := newQuotaTripTestClient(t, server.URL, func(cb *goconfig.CircuitBreakerConfig) { + cb.Scope = "model" + cb.Timeout = time.Minute + cb.TripOn = goconfig.TripRuleMap{"quota": {Match: `quota exceeded`}} + }) + + req := Request{Method: http.MethodGet, Endpoint: "/test", Model: "m1"} + err := client.Do(context.Background(), req, nil) + require.Error(t, err) + + var modelBreaker *circuitBreaker + for _, entry := range client.modelBreakers { + modelBreaker = entry.breaker + } + require.NotNil(t, modelBreaker, "model-scoped breaker must exist after a request") + assert.Equal(t, "open", modelBreaker.State(), "the model breaker carries the trip") + + // The provider-level breaker serves empty/unknown models; trip it too so + // both maps are exercised. + client.circuitBreaker.RecordQuotaTrip(time.Minute) + require.Equal(t, "open", client.circuitBreaker.State()) + + client.ResetBreaker() + assert.Equal(t, "closed", client.circuitBreaker.State()) + assert.Equal(t, "closed", modelBreaker.State()) + + err = client.Do(context.Background(), req, nil) + require.NoError(t, err) + assert.Equal(t, int32(2), attempts.Load()) +} + +// An invalid rule pattern surfaces through the same configErr channel as the +// other resilience validation errors. +func TestTripRule_InvalidRegexReportsConfigError(t *testing.T) { + t.Parallel() + + client := New(Config{ + ProviderName: "test", + BaseURL: "http://localhost", + CircuitBreaker: goconfig.CircuitBreakerConfig{ + Enabled: true, + TripOn: goconfig.TripRuleMap{"invalid": {Match: "(unclosed"}}, + }, + }, nil) + require.Error(t, client.configErr) + assert.Contains(t, client.configErr.Error(), "trip_on") + + err := client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + assert.Contains(t, err.Error(), "invalid resilience configuration") +} + +func TestQuotaTripTTL(t *testing.T) { + t.Parallel() + + timeout := time.Minute + quotaErr := core.NewProviderError("test", http.StatusForbidden, "quota exceeded", nil) + codedErr := core.NewProviderError("test", http.StatusForbidden, "usage limit reached", nil).WithCode("insufficient_quota") + + tests := []struct { + name string + tripOn []goconfig.TripRuleConfig + err error + wantTTL time.Duration + wantOK bool + }{ + { + name: "no rules is inert", + tripOn: nil, + err: quotaErr, + }, + { + name: "nil error never matches", + tripOn: []goconfig.TripRuleConfig{{Match: "quota"}}, + err: nil, + }, + { + name: "non-gateway error never matches", + tripOn: []goconfig.TripRuleConfig{{Match: "quota"}}, + err: errors.New("quota exceeded"), + }, + { + name: "matching message uses rule TTL", + tripOn: []goconfig.TripRuleConfig{{Match: `quota exceeded`, TTL: 5 * time.Minute}}, + err: quotaErr, + wantTTL: 5 * time.Minute, + wantOK: true, + }, + { + name: "matching code alone matches", + tripOn: []goconfig.TripRuleConfig{{Match: `insufficient_quota$`}}, + err: codedErr, + wantTTL: timeout, + wantOK: true, + }, + { + name: "zero rule TTL substitutes breaker timeout", + tripOn: []goconfig.TripRuleConfig{{Match: `quota exceeded`}}, + err: quotaErr, + wantTTL: timeout, + wantOK: true, + }, + { + name: "first matching rule wins", + tripOn: []goconfig.TripRuleConfig{ + {Match: `no match here`, TTL: time.Second}, + {Match: `quota exceeded`, TTL: 2 * time.Minute}, + }, + err: quotaErr, + wantTTL: 2 * time.Minute, + wantOK: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + client := New(Config{ + ProviderName: "test", + BaseURL: "http://localhost", + CircuitBreaker: goconfig.CircuitBreakerConfig{ + Enabled: true, + Timeout: timeout, + TripOn: goconfig.TripRuleMapFromList(tt.tripOn), + }, + }, nil) + require.NoError(t, client.configErr) + + ttl, ok := client.quotaTripTTL(tt.err) + assert.Equal(t, tt.wantOK, ok) + assert.Equal(t, tt.wantTTL, ttl) + }) + } +} + +// Breaker-level behavior: the quota window and the open-state timeout gate +// the half-open transition independently. +func TestCircuitBreaker_RecordQuotaTrip(t *testing.T) { + t.Parallel() + + t.Run("window outlives timeout", func(t *testing.T) { + t.Parallel() + + cb := newCircuitBreaker(3, 1, 20*time.Millisecond) + cb.RecordQuotaTrip(120 * time.Millisecond) + assert.Equal(t, "open", cb.State()) + + allowed, probe := cb.acquire() + assert.False(t, allowed) + assert.False(t, probe) + + // Past the timeout but inside the quota window: still open. + time.Sleep(45 * time.Millisecond) + allowed, probe = cb.acquire() + assert.False(t, allowed) + assert.False(t, probe) + + // Past the quota window: the half-open probe is granted. + time.Sleep(120 * time.Millisecond) + allowed, probe = cb.acquire() + assert.True(t, allowed) + assert.True(t, probe) + }) + + t.Run("zero ttl keeps zero window", func(t *testing.T) { + t.Parallel() + + cb := newCircuitBreaker(3, 1, 20*time.Millisecond) + cb.RecordQuotaTrip(0) + + allowed, probe := cb.acquire() + assert.False(t, allowed) + assert.False(t, probe) + + time.Sleep(30 * time.Millisecond) + allowed, probe = cb.acquire() + assert.True(t, allowed) + assert.True(t, probe) + }) + + t.Run("zero quotaUntil does not change plain timeout behavior", func(t *testing.T) { + t.Parallel() + + cb := newCircuitBreaker(1, 1, 20*time.Millisecond) + cb.RecordFailure() + assert.Equal(t, "open", cb.State()) + + allowed, probe := cb.acquire() + assert.False(t, allowed) + assert.False(t, probe) + + time.Sleep(30 * time.Millisecond) + allowed, probe = cb.acquire() + assert.True(t, allowed) + assert.True(t, probe) + }) +} + +func TestCircuitBreaker_Reset(t *testing.T) { + t.Parallel() + + cb := newCircuitBreaker(2, 1, time.Hour) + cb.RecordFailure() + cb.RecordQuotaTrip(time.Hour) + require.Equal(t, "open", cb.State()) + + allowed, _ := cb.acquire() + assert.False(t, allowed) + + cb.Reset() + assert.Equal(t, "closed", cb.State()) + + allowed, probe := cb.acquire() + assert.True(t, allowed) + assert.False(t, probe, "a reset breaker admits traffic without consuming a probe slot") + + // The failure count restarts from zero: the next failure does not open. + cb.RecordFailure() + assert.Equal(t, "closed", cb.State()) +} + +// A matching quota error in a streaming response body opens the breaker on +// the first stream establishment, just like the non-streaming Do path. When +// the provider returns a non-200 status the error body is parsed identically +// to the response path and quotaTripTTL fires. +func TestTripRule_DoStreamMatchingErrorTripsInstantly(t *testing.T) { + t.Parallel() + + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempts.Add(1) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusForbidden) + _, _ = w.Write([]byte(quotaErrorBody)) + })) + defer server.Close() + + client := newQuotaTripTestClient(t, server.URL, nil) + + stream, err := client.DoStream(context.Background(), Request{Method: http.MethodPost, Endpoint: "/chat"}) + require.Error(t, err) + require.Nil(t, stream) + assert.Equal(t, int32(1), attempts.Load()) + assert.Equal(t, "open", client.circuitBreaker.State()) + + // The second DoStream call must be rejected locally. + stream, err = client.DoStream(context.Background(), Request{Method: http.MethodPost, Endpoint: "/chat"}) + require.Error(t, err) + require.Nil(t, stream) + assert.Equal(t, int32(1), attempts.Load(), "rejected stream must not reach the upstream") +} + +// Some providers answer 200 with a bare {"error": ...} body. A quota message +// there must trip the breaker exactly like a translated error status. +func TestTripRule_Embedded200QuotaErrorTripsInstantly(t *testing.T) { + t.Parallel() + + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempts.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(quotaErrorBody)) + })) + defer server.Close() + + client := newQuotaTripTestClient(t, server.URL, nil) + + err := client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + var gatewayErr *core.GatewayError + require.ErrorAs(t, err, &gatewayErr) + require.ErrorIs(t, err, core.ErrEmbeddedInSuccess) + assert.Equal(t, "open", client.circuitBreaker.State()) + + err = client.Do(context.Background(), Request{Method: http.MethodGet, Endpoint: "/test"}, nil) + require.Error(t, err) + require.ErrorAs(t, err, &gatewayErr) + assert.Contains(t, gatewayErr.Message, "circuit breaker is open") + assert.Equal(t, int32(1), attempts.Load(), "rejected request must not reach the upstream") +} diff --git a/internal/llmclient/client.go b/internal/llmclient/client.go index 3921a8b17..8abcd4617 100644 --- a/internal/llmclient/client.go +++ b/internal/llmclient/client.go @@ -7,8 +7,10 @@ package llmclient import ( "context" + "fmt" "io" "net/http" + "regexp" "sync" "time" @@ -135,6 +137,7 @@ type Client struct { configErr error retryStatuses map[int]bool failureStatuses map[int]bool + tripRules []TripRule } // New creates a new LLM client with the given configuration @@ -159,6 +162,10 @@ func New(cfg Config, headerSetter HeaderSetter) *Client { if c.configErr != nil { return c } + c.tripRules, c.configErr = compileTripRules(cfg.CircuitBreaker.TripOn.List()) + if c.configErr != nil { + return c + } // The breaker is off when explicitly disabled or when it can never trip. if cfg.CircuitBreaker.Enabled && cfg.CircuitBreaker.FailureThreshold > 0 { @@ -179,6 +186,36 @@ func NewWithHTTPClient(httpClient *http.Client, cfg Config, headerSetter HeaderS return c } +// compileTripRules turns configured error-message matchers into breaker trip +// rules. An empty list yields nil: the feature stays inert without rules. +func compileTripRules(rules []config.TripRuleConfig) ([]TripRule, error) { + if len(rules) == 0 { + return nil, nil + } + compiled := make([]TripRule, 0, len(rules)) + for i, rule := range rules { + pattern, err := regexp.Compile(rule.Match) + if err != nil { + return nil, fmt.Errorf("invalid circuit_breaker.trip_on[%d] pattern: %w", i, err) + } + compiled = append(compiled, TripRule{Pattern: pattern, TTL: rule.TTL}) + } + return compiled, nil +} + +// ResetBreaker force-closes the provider-level breaker and every model-scoped +// breaker, clearing any quota trip window so traffic resumes immediately. +func (c *Client) ResetBreaker() { + if c.circuitBreaker != nil { + c.circuitBreaker.Reset() + } + c.modelBreakersMu.Lock() + defer c.modelBreakersMu.Unlock() + for _, entry := range c.modelBreakers { + entry.breaker.Reset() + } +} + // SetBaseURL updates the base URL (thread-safe) func (c *Client) SetBaseURL(url string) { c.mu.Lock() diff --git a/internal/llmclient/client_do.go b/internal/llmclient/client_do.go index c5070f308..ac74cfe0b 100644 --- a/internal/llmclient/client_do.go +++ b/internal/llmclient/client_do.go @@ -114,6 +114,11 @@ func (c *Client) DoRaw(ctx context.Context, req Request) (*Response, error) { lastErr = attachResponseHeaders(core.ParseProviderError(c.config.ProviderName, resp.StatusCode, resp.Body, nil), resp.Header) lastStatusCode = resp.StatusCode lastErrFromTransport = false + if ttl, ok := c.quotaTripTTL(lastErr); ok { + scope.breaker.RecordQuotaTrip(ttl) + c.completeScope(scope, lastStatusCode, lastErr, lastErr) + return nil, lastErr + } if scope.halfOpenProbe { c.completeScope(scope, lastStatusCode, lastErr, nil) return nil, lastErr @@ -134,6 +139,11 @@ func (c *Client) DoRaw(ctx context.Context, req Request) (*Response, error) { lastErr = attachResponseHeaders(embedded, resp.Header) lastStatusCode = embedded.StatusCode lastErrFromTransport = false + if ttl, ok := c.quotaTripTTL(lastErr); ok { + scope.breaker.RecordQuotaTrip(ttl) + c.completeScope(scope, lastStatusCode, lastErr, lastErr) + return nil, lastErr + } if c.isRetryable(embedded.StatusCode) && !scope.halfOpenProbe { continue } diff --git a/internal/llmclient/client_scope.go b/internal/llmclient/client_scope.go index a7973f930..f6540a13c 100644 --- a/internal/llmclient/client_scope.go +++ b/internal/llmclient/client_scope.go @@ -173,7 +173,7 @@ func (r *firstChunkReadCloser) Read(p []byte) (int, error) { // whether the failure was transport-level) and emits the metrics observation. // Use this whenever a code path returns from one of the public Do* methods. func (c *Client) completeScope(scope requestScope, statusCode int, err, cbErr error) { - c.recordCircuitBreakerCompletion(scope, statusCode, cbErr) + c.recordCircuitBreakerCompletion(scope, statusCode, err, cbErr) c.finishRequest(scope, statusCode, err) } @@ -215,19 +215,32 @@ func (c *Client) waitForRetryAttempt(ctx context.Context, scope requestScope, at return nil } -func (c *Client) recordCircuitBreakerCompletion(scope requestScope, statusCode int, err error) { +// recordCircuitBreakerCompletion turns a finished request into a breaker +// outcome. err is the request's final error, cbErr only the transport-level +// portion: HTTP-status and embedded-200 failures carry their parsed provider +// error in err while cbErr stays nil, so trip rules can match their message. +func (c *Client) recordCircuitBreakerCompletion(scope requestScope, statusCode int, err, cbErr error) { if scope.breaker == nil { return } - if err != nil { + if cbErr != nil { // A caller-side cancellation aborts the transport but proves nothing // about provider health, so it is neither a success nor a failure. // Client deadlines (context.DeadlineExceeded) still count: the // provider failed to answer within the latency budget. - if errors.Is(err, context.Canceled) { + if errors.Is(cbErr, context.Canceled) { c.releaseHalfOpenProbe(scope) return } + } + // A quota trip wins over classification: the message proves the provider's + // quota is gone, so open instantly even when the status alone would count + // as a success. + if ttl, ok := c.quotaTripTTL(err); ok { + scope.breaker.RecordQuotaTrip(ttl) + return + } + if cbErr != nil { scope.breaker.RecordFailure() return } @@ -238,6 +251,33 @@ func (c *Client) recordCircuitBreakerCompletion(scope requestScope, statusCode i scope.breaker.RecordSuccess() } +// quotaTripTTL returns the trip window of the first rule matching the error's +// provider message (and code), substituting the breaker timeout for a zero +// rule TTL. Inert without rules. +func (c *Client) quotaTripTTL(err error) (time.Duration, bool) { + if len(c.tripRules) == 0 || err == nil { + return 0, false + } + var gatewayErr *core.GatewayError + if !errors.As(err, &gatewayErr) { + return 0, false + } + text := gatewayErr.Message + if gatewayErr.Code != nil && *gatewayErr.Code != "" { + text += " " + *gatewayErr.Code + } + for _, rule := range c.tripRules { + if rule.Pattern.MatchString(text) { + ttl := rule.TTL + if ttl == 0 { + ttl = c.config.CircuitBreaker.Timeout + } + return ttl, true + } + } + return 0, false +} + func (c *Client) shouldTripCircuitBreaker(statusCode int) bool { return c.failureStatuses[statusCode] } diff --git a/internal/providers/breaker_reset.go b/internal/providers/breaker_reset.go new file mode 100644 index 000000000..bd023de54 --- /dev/null +++ b/internal/providers/breaker_reset.go @@ -0,0 +1,50 @@ +package providers + +import ( + "errors" + "fmt" + "strings" +) + +// ErrProviderNotFound reports a provider instance name nothing is registered +// under right now. +var ErrProviderNotFound = errors.New("provider not found") + +// BreakerReset is implemented by providers whose underlying client can +// force-close its circuit breaker(s). Resetting is opt-in per adapter: the +// admin reset endpoint type-asserts this interface and reports the providers +// that lack it as not resettable rather than failing. +type BreakerReset interface { + ResetBreaker() +} + +// BreakerResetter force-closes a named provider's circuit breaker(s). +// It is the admin API's seam over the live registry; NewBreakerResetter +// builds the production implementation. +type BreakerResetter interface { + ResetCircuitBreaker(providerName string) error +} + +type registryBreakerResetter struct { + registry *ModelRegistry +} + +// NewBreakerResetter returns a BreakerResetter that resolves provider +// instance names against the live registry. +func NewBreakerResetter(registry *ModelRegistry) BreakerResetter { + return registryBreakerResetter{registry: registry} +} + +func (r registryBreakerResetter) ResetCircuitBreaker(providerName string) error { + providerName = strings.TrimSpace(providerName) + provider := r.registry.ProviderByName(providerName) + if provider == nil { + return fmt.Errorf("%w: %s", ErrProviderNotFound, providerName) + } + resetter, ok := provider.(BreakerReset) + if !ok { + return fmt.Errorf("provider %q does not support circuit breaker reset", providerName) + } + resetter.ResetBreaker() + return nil +} diff --git a/internal/providers/breaker_reset_test.go b/internal/providers/breaker_reset_test.go new file mode 100644 index 000000000..f149151bd --- /dev/null +++ b/internal/providers/breaker_reset_test.go @@ -0,0 +1,58 @@ +package providers + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/enterpilot/gomodel/internal/core" +) + +// resetTrackingProvider adapts a plain provider into one implementing +// BreakerReset, counting resets. +type resetTrackingProvider struct { + core.Provider + resets int +} + +func (p *resetTrackingProvider) ResetBreaker() { + p.resets++ +} + +func TestBreakerResetter_ResetsRegisteredProvider(t *testing.T) { + registry := NewModelRegistry() + provider := &resetTrackingProvider{} + registry.RegisterProviderWithNameAndType(provider, "openai-main", "openai") + + err := NewBreakerResetter(registry).ResetCircuitBreaker("openai-main") + require.NoError(t, err) + assert.Equal(t, 1, provider.resets) +} + +func TestBreakerResetter_TrimsProviderName(t *testing.T) { + registry := NewModelRegistry() + provider := &resetTrackingProvider{} + registry.RegisterProviderWithNameAndType(provider, "openai-main", "openai") + + err := NewBreakerResetter(registry).ResetCircuitBreaker(" openai-main ") + require.NoError(t, err) + assert.Equal(t, 1, provider.resets) +} + +func TestBreakerResetter_UnknownProvider(t *testing.T) { + registry := NewModelRegistry() + + err := NewBreakerResetter(registry).ResetCircuitBreaker("missing") + require.ErrorIs(t, err, ErrProviderNotFound) +} + +func TestBreakerResetter_ProviderWithoutBreakerResetSupport(t *testing.T) { + registry := NewModelRegistry() + registry.RegisterProviderWithNameAndType(®istryMockProvider{name: "legacy"}, "legacy", "test") + + err := NewBreakerResetter(registry).ResetCircuitBreaker("legacy") + require.Error(t, err) + require.NotErrorIs(t, err, ErrProviderNotFound) + assert.Contains(t, err.Error(), "does not support circuit breaker reset") +} diff --git a/internal/providers/config.go b/internal/providers/config.go index e927a4db5..9ae58c1e5 100644 --- a/internal/providers/config.go +++ b/internal/providers/config.go @@ -282,6 +282,9 @@ func buildProviderConfig(raw config.RawProviderConfig, global config.ResilienceC if cb.Timeout != nil { resolved.Resilience.CircuitBreaker.Timeout = *cb.Timeout } + if cb.TripOn != nil { + resolved.Resilience.CircuitBreaker.TripOn = cb.TripOn + } } return resolved diff --git a/internal/providers/config_env.go b/internal/providers/config_env.go index cbaa3df47..487727484 100644 --- a/internal/providers/config_env.go +++ b/internal/providers/config_env.go @@ -9,6 +9,7 @@ import ( "sort" "strconv" "strings" + "time" "unicode" "github.com/enterpilot/gomodel/config" @@ -97,6 +98,10 @@ type providerEnvValues struct { ModelFilterMaxPrice *float64 SessionStickyKeys *bool FairnessFromUserPath *bool + // TripOnGroups holds named trip rules from `_CIRCUIT_BREAKER_ + // TRIP_ON__MATCH|_TTL` vars. A group replaces the config rule of + // the same name at overlay time; other config rules survive. + TripOnGroups map[string]config.TripRuleConfig } // modelFilter assembles the filter this env group declares. @@ -174,10 +179,93 @@ func (v providerEnvValues) empty() bool { strings.TrimSpace(v.InferenceObjective) == "" && v.SessionStickyKeys == nil && v.FairnessFromUserPath == nil && + len(v.TripOnGroups) == 0 && len(v.Models) == 0 && v.modelFilter().Empty() } +// tripOnResilience assembles the raw resilience overlay this env group +// declares, or nil when no trip rules are set. Env groups override config +// rules by name only; rules the env does not name survive. +func (v providerEnvValues) tripOnResilience() *config.RawResilienceConfig { + if len(v.TripOnGroups) == 0 { + return nil + } + rules := make([]config.TripRuleConfig, 0, len(v.TripOnGroups)) + for _, rule := range v.TripOnGroups { + rules = append(rules, rule) + } + slices.SortFunc(rules, func(a, b config.TripRuleConfig) int { + return strings.Compare(a.Name, b.Name) + }) + return &config.RawResilienceConfig{ + CircuitBreaker: &config.RawCircuitBreakerConfig{ + TripOn: config.TripRuleMapFromList(rules), + }, + } +} + +// mergeTripOnResilience applies env trip-rule groups onto the provider's raw +// resilience overlay, replacing same-named config rules and keeping every +// other YAML-declared override intact. +func mergeTripOnResilience(existing *config.RawResilienceConfig, groups []config.TripRuleConfig) *config.RawResilienceConfig { + tripOn := config.TripRuleMapFromList(groups) + if existing == nil { + return &config.RawResilienceConfig{ + CircuitBreaker: &config.RawCircuitBreakerConfig{TripOn: tripOn}, + } + } + res := *existing + if res.CircuitBreaker == nil { + res.CircuitBreaker = &config.RawCircuitBreakerConfig{TripOn: tripOn} + return &res + } + cb := *res.CircuitBreaker + if cb.TripOn == nil { + cb.TripOn = tripOn + } else { + cb.TripOn = config.TripRuleMapFromList(config.MergeTripRuleGroups(cb.TripOn.List(), groups)) + } + res.CircuitBreaker = &cb + return &res +} + +// parseTripRuleEnvKey splits a provider trip-rule env var into the +// provider-name suffix, the group name, and the attribute (MATCH or TTL): +// [_]_CIRCUIT_BREAKER_TRIP_ON__. It mirrors the +// API-key special case: the generic field table cannot match multi-token +// names, so this runs ahead of it. +func parseTripRuleEnvKey(prefix, key string) (suffix, group, attr string, ok bool) { + rest, found := strings.CutPrefix(key, prefix+"_") + if !found { + return "", "", "", false + } + const marker = "CIRCUIT_BREAKER_TRIP_ON_" + i := strings.Index(rest, marker) + if i < 0 { + return "", "", "", false + } + tail := rest[i+len(marker):] + switch { + case strings.HasSuffix(tail, "_MATCH"): + attr = "MATCH" + group = strings.TrimSuffix(tail, "_MATCH") + case strings.HasSuffix(tail, "_TTL"): + attr = "TTL" + group = strings.TrimSuffix(tail, "_TTL") + default: + return "", "", "", false + } + if group == "" { + return "", "", "", false + } + suffix = strings.TrimSuffix(rest[:i], "_") + if suffix != "" && !validProviderEnvSuffix(suffix) { + return "", "", "", false + } + return suffix, group, attr, true +} + func providerEnvSources(providerType string, spec DiscoveryConfig) []providerEnvSource { separator := spec.NameSeparator if separator == "" { @@ -201,6 +289,38 @@ func collectProviderEnvValues(prefix string, spec DiscoveryConfig, environ []str continue } + // Trip-rule groups carry their own key shape — [_]_ + // CIRCUIT_BREAKER_TRIP_ON__MATCH|_TTL — before the generic + // single-field parse. + if suffix, group, attr, isTripRule := parseTripRuleEnvKey(prefix, key); isTripRule { + values := groups[suffix] + rule := values.TripOnGroups[group] + rule.Name = group + switch attr { + case "MATCH": + rule.Match = value + case "TTL": + ttl, err := time.ParseDuration(strings.TrimSpace(value)) + if err != nil { + // Fail closed: a poison rule (empty match) makes resilience + // validation reject the provider, the same stance as the + // NaN price cap, instead of silently dropping quota + // protection because one TTL was mistyped. + slog.Warn("provider trip_on env ttl is malformed and will fail validation", + "env_prefix", prefix, "group", group, "error", err) + rule = config.TripRuleConfig{Name: group} + } else { + rule.TTL = ttl + } + } + if values.TripOnGroups == nil { + values.TripOnGroups = make(map[string]config.TripRuleConfig) + } + values.TripOnGroups[group] = rule + groups[suffix] = values + continue + } + suffix, field, index, ok := parseProviderEnvKey(prefix, key, spec) if !ok { continue @@ -511,6 +631,7 @@ func (v providerEnvValues) rawConfig(providerType string, spec DiscoveryConfig) Models: rawProviderModelsFromIDs(v.Models), ModelFilter: v.modelFilter(), SessionStickyKeys: v.SessionStickyKeys, + Resilience: v.tripOnResilience(), } } @@ -588,6 +709,9 @@ func overlayProviderEnvValues(existing config.RawProviderConfig, values provider if values.ModelFilterMaxPrice != nil { existing.ModelFilter.MaxPricePerMtok = values.ModelFilterMaxPrice } + if len(values.TripOnGroups) > 0 { + existing.Resilience = mergeTripOnResilience(existing.Resilience, values.tripOnResilience().CircuitBreaker.TripOn.List()) + } return existing } @@ -637,10 +761,19 @@ func (v providerEnvValues) withoutFieldsSetBy(existing config.RawProviderConfig) drop("model_filter.include", len(v.ModelFilterInclude) > 0, len(existing.ModelFilter.Include) > 0, func() { v.ModelFilterInclude = nil }) drop("model_filter.exclude", len(v.ModelFilterExclude) > 0, len(existing.ModelFilter.Exclude) > 0, func() { v.ModelFilterExclude = nil }) drop("model_filter.max_price_per_mtok", v.ModelFilterMaxPrice != nil, existing.ModelFilter.MaxPricePerMtok != nil, func() { v.ModelFilterMaxPrice = nil }) + drop("trip_on", len(v.TripOnGroups) > 0, rawProviderHasTripOn(existing), func() { v.TripOnGroups = nil }) return v, ignored } +// rawProviderHasTripOn reports whether the config provider declares trip +// rules in YAML, so a bare _TRIP_ON never borrows onto it. +func rawProviderHasTripOn(cfg config.RawProviderConfig) bool { + return cfg.Resilience != nil && + cfg.Resilience.CircuitBreaker != nil && + cfg.Resilience.CircuitBreaker.TripOn != nil +} + // rawProviderHasResolvedModel reports whether the config provider declares at // least one model ID that is not an unresolved ${VAR} placeholder, so a list // left to the environment does not block the bare _MODELS fill. diff --git a/internal/providers/config_env_test.go b/internal/providers/config_env_test.go index 66e7be1a2..e36855234 100644 --- a/internal/providers/config_env_test.go +++ b/internal/providers/config_env_test.go @@ -5,6 +5,7 @@ import ( "log/slog" "strings" "testing" + "time" "github.com/enterpilot/gomodel/config" "github.com/stretchr/testify/assert" @@ -147,3 +148,94 @@ func TestApplyProviderEnvVars_BareTypeEnvVarsAgainstRenamedProviders(t *testing. }) } } + +// Named trip-rule env groups override config rules by name only; a malformed +// value must fail resilience validation instead of silently dropping quota +// protection. +func TestApplyProviderEnvVars_TripOnGroups(t *testing.T) { + oneMinute := time.Minute + + setEnv := func(t *testing.T) { + t.Helper() + t.Setenv("KIMICODE_CIRCUIT_BREAKER_TRIP_ON_QUOTA_EXCEEDED_MATCH", "quota exceeded") + t.Setenv("KIMICODE_CIRCUIT_BREAKER_TRIP_ON_QUOTA_EXCEEDED_TTL", "15m") + t.Setenv("KIMICODE_CIRCUIT_BREAKER_TRIP_ON_USAGE_LIMIT_MATCH", "usage limit") + } + wantGroups := config.TripRuleMap{ + "QUOTA_EXCEEDED": {Name: "QUOTA_EXCEEDED", Match: "quota exceeded", TTL: 15 * time.Minute}, + "USAGE_LIMIT": {Name: "USAGE_LIMIT", Match: "usage limit"}, + } + + t.Run("bare type env creates provider with parsed trip rules", func(t *testing.T) { + setEnv(t) + got := applyProviderEnvVars(map[string]config.RawProviderConfig{}, map[string]DiscoveryConfig{ + "kimicode": {DefaultBaseURL: "https://api.kimi.com/coding/v1"}, + }) + + p, ok := got["kimicode"] + require.True(t, ok, "kimicode provider missing") + require.NotNil(t, p.Resilience, "resilience overlay missing") + require.NotNil(t, p.Resilience.CircuitBreaker, "circuit_breaker overlay missing") + assert.Equal(t, wantGroups, p.Resilience.CircuitBreaker.TripOn) + }) + + t.Run("env group overrides only the config rule of the same name", func(t *testing.T) { + t.Setenv("KIMICODE_CIRCUIT_BREAKER_TRIP_ON_QUOTA_EXCEEDED_MATCH", "quota exhausted") + existing := map[string]config.RawProviderConfig{ + "kimicode": { + Type: "kimicode", + Resilience: &config.RawResilienceConfig{ + CircuitBreaker: &config.RawCircuitBreakerConfig{ + Timeout: &oneMinute, + TripOn: config.TripRuleMap{ + "quota_exceeded": {Match: "old pattern", TTL: time.Hour}, + "other_rule": {Match: "yaml only", TTL: 2 * time.Hour}, + }, + }, + }, + }, + } + got := applyProviderEnvVars(existing, map[string]DiscoveryConfig{"kimicode": {}}) + + p := got["kimicode"] + require.NotNil(t, p.Resilience) + require.NotNil(t, p.Resilience.CircuitBreaker) + assert.Equal(t, time.Minute, *p.Resilience.CircuitBreaker.Timeout, "YAML timeout must survive the env overlay") + tripOn := p.Resilience.CircuitBreaker.TripOn + assert.Equal(t, config.TripRuleConfig{Name: "QUOTA_EXCEEDED", Match: "quota exhausted"}, tripOn["QUOTA_EXCEEDED"], "same-named group replaces the config rule") + assert.Equal(t, config.TripRuleConfig{Name: "other_rule", Match: "yaml only", TTL: 2 * time.Hour}, tripOn["other_rule"], "unnamed-by-env rule survives") + }) + + t.Run("renamed provider with YAML trip_on ignores the env value", func(t *testing.T) { + t.Setenv("KIMICODE_CIRCUIT_BREAKER_TRIP_ON_QUOTA_MATCH", "quota exceeded") + yamlRules := config.TripRuleMap{"yaml": {Match: "yaml rule"}} + existing := map[string]config.RawProviderConfig{ + "kimi-renamed": { + Type: "kimicode", + Resilience: &config.RawResilienceConfig{ + CircuitBreaker: &config.RawCircuitBreakerConfig{TripOn: yamlRules}, + }, + }, + } + logs := captureSlog(t) + got := applyProviderEnvVars(existing, map[string]DiscoveryConfig{"kimicode": {}}) + + p := got["kimi-renamed"] + assert.Equal(t, yamlRules, p.Resilience.CircuitBreaker.TripOn, "YAML rules must win over env") + assert.Contains(t, logs.String(), "trip_on") + }) + + t.Run("malformed ttl fails resilience validation", func(t *testing.T) { + t.Setenv("KIMICODE_CIRCUIT_BREAKER_TRIP_ON_BAD_MATCH", "quota exceeded") + t.Setenv("KIMICODE_CIRCUIT_BREAKER_TRIP_ON_BAD_TTL", "banana") + got := applyProviderEnvVars(map[string]config.RawProviderConfig{}, map[string]DiscoveryConfig{ + "kimicode": {DefaultBaseURL: "https://api.kimi.com/coding/v1"}, + }) + + p := got["kimicode"] + require.NotNil(t, p.Resilience) + require.Error(t, config.ValidateResilience(config.ResilienceConfig{ + CircuitBreaker: config.CircuitBreakerConfig{TripOn: p.Resilience.CircuitBreaker.TripOn}, + }), "malformed TRIP_ON group must fail validation") + }) +} diff --git a/internal/providers/credentials.go b/internal/providers/credentials.go index cf05ab053..82159a021 100644 --- a/internal/providers/credentials.go +++ b/internal/providers/credentials.go @@ -48,6 +48,11 @@ type ManagedProviderCredential struct { GCPScope string Models []string + // TripOn holds circuit-breaker trip rules for this provider. Rules are + // configuration, not secrets, so they travel unredacted through the + // admin API and are persisted as plain JSON. + TripOn []config.TripRuleConfig + // Enabled controls whether this credential is applied to the running // registry. Disabling one keeps the row (and its keys) on file without // routing traffic to it — the same effect as deleting it, without losing @@ -80,6 +85,14 @@ func (m ManagedProviderCredential) toRawProviderConfig() config.RawProviderConfi GCPScope: m.GCPScope, Models: rawProviderModelsFromIDs(m.Models), } + // Trip rules ride the same per-provider override pipeline declarative + // providers use; only set when present so rows without rules keep the + // resolved global breaker config untouched. + if len(m.TripOn) > 0 { + raw.Resilience = &config.RawResilienceConfig{ + CircuitBreaker: &config.RawCircuitBreakerConfig{TripOn: config.TripRuleMapFromList(m.TripOn)}, + } + } if len(m.APIKeys) > 0 { raw.APIKey = m.APIKeys[0] } diff --git a/internal/providers/credentials_store.go b/internal/providers/credentials_store.go index 4a4b9cfb9..8579f511f 100644 --- a/internal/providers/credentials_store.go +++ b/internal/providers/credentials_store.go @@ -5,6 +5,8 @@ import ( "time" "github.com/goccy/go-json" + + "github.com/enterpilot/gomodel/config" ) func encodeCredentialList(value []string) (string, error) { @@ -32,6 +34,35 @@ func decodeCredentialList(data []byte) ([]string, error) { return value, nil } +// encodeTripRules stores trip rules as a JSON array of {match, ttl}. Nil +// rules encode as an empty array, matching encodeCredentialList. +func encodeTripRules(rules []config.TripRuleConfig) (string, error) { + if rules == nil { + rules = []config.TripRuleConfig{} + } + data, err := json.Marshal(rules) + if err != nil { + return "", err + } + return string(data), nil +} + +// decodeTripRules reads trip rules back; empty input or an empty array reads +// as nil, like never set. +func decodeTripRules(data []byte) ([]config.TripRuleConfig, error) { + if len(data) == 0 { + return nil, nil + } + var rules []config.TripRuleConfig + if err := json.Unmarshal(data, &rules); err != nil { + return nil, err + } + if len(rules) == 0 { + return nil, nil + } + return rules, nil +} + // stampCredentialUpsert sets timestamps: CreatedAt on insert, UpdatedAt always. func stampCredentialUpsert(cred *ManagedProviderCredential) { now := time.Now().UTC() diff --git a/internal/providers/credentials_store_mongodb.go b/internal/providers/credentials_store_mongodb.go index 0b99e478b..abfb8d99c 100644 --- a/internal/providers/credentials_store_mongodb.go +++ b/internal/providers/credentials_store_mongodb.go @@ -9,28 +9,31 @@ import ( "go.mongodb.org/mongo-driver/v2/bson" "go.mongodb.org/mongo-driver/v2/mongo" "go.mongodb.org/mongo-driver/v2/mongo/options" + + "github.com/enterpilot/gomodel/config" ) type mongoCredentialDocument struct { - ID string `bson:"_id"` - Type string `bson:"type"` - APIKeys []string `bson:"api_keys,omitempty"` - BaseURL string `bson:"base_url,omitempty"` - APIVersion string `bson:"api_version,omitempty"` - Backend string `bson:"backend,omitempty"` - AuthType string `bson:"auth_type,omitempty"` - APIMode string `bson:"api_mode,omitempty"` - VertexProject string `bson:"vertex_project,omitempty"` - VertexLocation string `bson:"vertex_location,omitempty"` - ServiceAccountFile string `bson:"service_account_file,omitempty"` - ServiceAccountJSON string `bson:"service_account_json,omitempty"` - ServiceAccountJSONBase64 string `bson:"service_account_json_base64,omitempty"` - GCPScope string `bson:"gcp_scope,omitempty"` - Models []string `bson:"models,omitempty"` - SessionStickyKeys *bool `bson:"session_sticky_keys,omitempty"` - Enabled bool `bson:"enabled"` - CreatedAt time.Time `bson:"created_at"` - UpdatedAt time.Time `bson:"updated_at"` + ID string `bson:"_id"` + Type string `bson:"type"` + APIKeys []string `bson:"api_keys,omitempty"` + BaseURL string `bson:"base_url,omitempty"` + APIVersion string `bson:"api_version,omitempty"` + Backend string `bson:"backend,omitempty"` + AuthType string `bson:"auth_type,omitempty"` + APIMode string `bson:"api_mode,omitempty"` + VertexProject string `bson:"vertex_project,omitempty"` + VertexLocation string `bson:"vertex_location,omitempty"` + ServiceAccountFile string `bson:"service_account_file,omitempty"` + ServiceAccountJSON string `bson:"service_account_json,omitempty"` + ServiceAccountJSONBase64 string `bson:"service_account_json_base64,omitempty"` + GCPScope string `bson:"gcp_scope,omitempty"` + Models []string `bson:"models,omitempty"` + TripOn []config.TripRuleConfig `bson:"trip_on,omitempty"` + SessionStickyKeys *bool `bson:"session_sticky_keys,omitempty"` + Enabled bool `bson:"enabled"` + CreatedAt time.Time `bson:"created_at"` + UpdatedAt time.Time `bson:"updated_at"` } type mongoCredentialIDFilter struct { @@ -114,6 +117,7 @@ func (s *MongoDBCredentialStore) Upsert(ctx context.Context, cred ManagedProvide "service_account_json_base64": cred.ServiceAccountJSONBase64, "gcp_scope": cred.GCPScope, "models": cred.Models, + "trip_on": cred.TripOn, "session_sticky_keys": sessionStickyKeysEnabled(cred.SessionStickyKeys), "enabled": cred.Enabled, "updated_at": cred.UpdatedAt, @@ -170,5 +174,8 @@ func credentialFromMongo(doc mongoCredentialDocument) ManagedProviderCredential if len(doc.Models) > 0 { cred.Models = append([]string(nil), doc.Models...) } + if len(doc.TripOn) > 0 { + cred.TripOn = append([]config.TripRuleConfig(nil), doc.TripOn...) + } return cred } diff --git a/internal/providers/credentials_store_sql.go b/internal/providers/credentials_store_sql.go index 1263dcabd..9ead73237 100644 --- a/internal/providers/credentials_store_sql.go +++ b/internal/providers/credentials_store_sql.go @@ -33,6 +33,7 @@ var credentialSQLSchema = []string{ service_account_json_base64 TEXT NOT NULL DEFAULT '', gcp_scope TEXT NOT NULL DEFAULT '', models TEXT NOT NULL DEFAULT '[]', + trip_on TEXT NOT NULL DEFAULT '[]', enabled ` + sqlx.TypeBool + ` NOT NULL DEFAULT TRUE, created_at ` + sqlx.TypeInt64 + ` NOT NULL, updated_at ` + sqlx.TypeInt64 + ` NOT NULL @@ -42,7 +43,7 @@ var credentialSQLSchema = []string{ const selectCredentialColumns = `name, type, api_keys, base_url, api_version, backend, auth_type, api_mode, ` + `vertex_project, vertex_location, service_account_file, service_account_json, ` + - `service_account_json_base64, gcp_scope, models, session_sticky_keys, enabled, created_at, updated_at` + `service_account_json_base64, gcp_scope, models, trip_on, session_sticky_keys, enabled, created_at, updated_at` // NewSQLCredentialStore creates the provider_credentials table and indexes if // needed. @@ -55,6 +56,7 @@ func NewSQLCredentialStore(ctx context.Context, db sqlx.DB) (*SQLCredentialStore } if err := sqlx.AddColumns(ctx, db, `ALTER TABLE provider_credentials ADD COLUMN session_sticky_keys `+sqlx.TypeBool+` NOT NULL DEFAULT TRUE`, + `ALTER TABLE provider_credentials ADD COLUMN trip_on TEXT NOT NULL DEFAULT '[]'`, ); err != nil { return nil, fmt.Errorf("migrate provider_credentials: %w", err) } @@ -101,13 +103,17 @@ func (s *SQLCredentialStore) Upsert(ctx context.Context, cred ManagedProviderCre if err != nil { return err } + tripOnJSON, err := encodeTripRules(cred.TripOn) + if err != nil { + return err + } _, err = s.db.Exec(ctx, ` INSERT INTO provider_credentials ( name, type, api_keys, base_url, api_version, backend, auth_type, api_mode, vertex_project, vertex_location, service_account_file, service_account_json, - service_account_json_base64, gcp_scope, models, session_sticky_keys, enabled, created_at, updated_at + service_account_json_base64, gcp_scope, models, trip_on, session_sticky_keys, enabled, created_at, updated_at ) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(name) DO UPDATE SET type = excluded.type, api_keys = excluded.api_keys, @@ -123,6 +129,7 @@ func (s *SQLCredentialStore) Upsert(ctx context.Context, cred ManagedProviderCre service_account_json_base64 = excluded.service_account_json_base64, gcp_scope = excluded.gcp_scope, models = excluded.models, + trip_on = excluded.trip_on, session_sticky_keys = excluded.session_sticky_keys, enabled = excluded.enabled, updated_at = excluded.updated_at @@ -142,6 +149,7 @@ func (s *SQLCredentialStore) Upsert(ctx context.Context, cred ManagedProviderCre cred.ServiceAccountJSONBase64, cred.GCPScope, modelsJSON, + tripOnJSON, sessionStickyKeysEnabled(cred.SessionStickyKeys), cred.Enabled, cred.CreatedAt.Unix(), @@ -171,7 +179,7 @@ func (s *SQLCredentialStore) Close() error { func scanSQLCredential(scanner sqlx.Row) (ManagedProviderCredential, error) { var cred ManagedProviderCredential - var apiKeys, models []byte + var apiKeys, models, tripOn []byte var sessionStickyKeys bool var createdAt, updatedAt int64 if err := scanner.Scan( @@ -190,6 +198,7 @@ func scanSQLCredential(scanner sqlx.Row) (ManagedProviderCredential, error) { &cred.ServiceAccountJSONBase64, &cred.GCPScope, &models, + &tripOn, &sessionStickyKeys, &cred.Enabled, &createdAt, @@ -205,6 +214,9 @@ func scanSQLCredential(scanner sqlx.Row) (ManagedProviderCredential, error) { if cred.Models, err = decodeCredentialList(models); err != nil { return ManagedProviderCredential{}, err } + if cred.TripOn, err = decodeTripRules(tripOn); err != nil { + return ManagedProviderCredential{}, err + } cred.CreatedAt = sqlutil.TimeFromUnix(createdAt) cred.UpdatedAt = sqlutil.TimeFromUnix(updatedAt) return cred, nil diff --git a/internal/providers/credentials_store_sql_test.go b/internal/providers/credentials_store_sql_test.go index f31a36831..72265acc4 100644 --- a/internal/providers/credentials_store_sql_test.go +++ b/internal/providers/credentials_store_sql_test.go @@ -67,3 +67,24 @@ func TestSQLCredentialStoreReopenKeepsRows(t *testing.T) { require.False(t, *got.SessionStickyKeys) }) } + +// A row whose trip_on column holds invalid JSON must surface as an error at +// read time, not silently yield an empty rule set. +func TestSQLCredentialStoreCorruptTripOnRowReturnsError(t *testing.T) { + runSQLCredentialStoreTest(t, func(t *testing.T, store *SQLCredentialStore, db sqlx.DB) { + ctx := context.Background() + err := store.Upsert(ctx, ManagedProviderCredential{ + Name: "my-openai", + Type: "openai", + APIKeys: []string{"sk-one"}, + Enabled: true, + }) + require.NoError(t, err) + + _, err = db.Exec(ctx, `UPDATE provider_credentials SET trip_on = '{' WHERE name = 'my-openai'`) + require.NoError(t, err) + + _, err = store.List(ctx) + require.Error(t, err) + }) +} diff --git a/internal/providers/credentials_store_test.go b/internal/providers/credentials_store_test.go index 5e75472aa..e5bd4f905 100644 --- a/internal/providers/credentials_store_test.go +++ b/internal/providers/credentials_store_test.go @@ -3,8 +3,12 @@ package providers import ( "context" "testing" + "time" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/enterpilot/gomodel/config" ) func TestCredentialStore_RoundTrip(t *testing.T) { @@ -29,6 +33,7 @@ func TestCredentialStore_RoundTrip(t *testing.T) { ServiceAccountJSONBase64: "eyJ0eXBlIjoic2VydmljZV9hY2NvdW50In0=", GCPScope: "https://www.googleapis.com/auth/cloud-platform", Models: []string{"gemini-2.5-pro", "gemini-2.5-flash"}, + TripOn: []config.TripRuleConfig{{Match: "quota exceeded", TTL: 30 * time.Second}}, Enabled: true, } require.NoError(t, store.Upsert(ctx, cred)) @@ -89,17 +94,19 @@ func TestCredentialStore_Defaults(t *testing.T) { require.NoError(t, err) require.Nil(t, got.APIKeys) require.Nil(t, got.Models) + require.Nil(t, got.TripOn) require.NotNil(t, got.SessionStickyKeys, "nil SessionStickyKeys means enabled and is stored as such") require.True(t, *got.SessionStickyKeys) require.False(t, got.Enabled) require.Empty(t, got.BaseURL) // Empty slices read back as nil, like never set. - require.NoError(t, store.Upsert(ctx, ManagedProviderCredential{Name: "ollama", Type: "ollama", APIKeys: []string{}, Models: []string{}})) + require.NoError(t, store.Upsert(ctx, ManagedProviderCredential{Name: "ollama", Type: "ollama", APIKeys: []string{}, Models: []string{}, TripOn: []config.TripRuleConfig{}})) got, err = store.Get(ctx, "ollama") require.NoError(t, err) require.Nil(t, got.APIKeys) require.Nil(t, got.Models) + require.Nil(t, got.TripOn) }) } @@ -111,6 +118,7 @@ func TestCredentialStore_UpsertClearsLists(t *testing.T) { Type: "openai", APIKeys: []string{"sk-one"}, Models: []string{"gpt-4o"}, + TripOn: []config.TripRuleConfig{{Match: "quota exceeded"}}, Enabled: true, })) @@ -122,6 +130,7 @@ func TestCredentialStore_UpsertClearsLists(t *testing.T) { require.NoError(t, err) require.Nil(t, got.APIKeys) require.Nil(t, got.Models) + require.Nil(t, got.TripOn) require.False(t, got.Enabled) }) } @@ -170,3 +179,28 @@ func TestCredentialStore_ListOrdersByName(t *testing.T) { require.Equal(t, []string{"alpha", "mid", "zeta"}, names) }) } + +func TestDecodeTripRules(t *testing.T) { + t.Parallel() + + t.Run("nil input reads as never set", func(t *testing.T) { + t.Parallel() + rules, err := decodeTripRules(nil) + require.NoError(t, err) + assert.Nil(t, rules) + }) + + t.Run("empty array reads as never set", func(t *testing.T) { + t.Parallel() + rules, err := decodeTripRules([]byte(`[]`)) + require.NoError(t, err) + assert.Nil(t, rules) + }) + + t.Run("corrupt JSON returns an error", func(t *testing.T) { + t.Parallel() + rules, err := decodeTripRules([]byte(`{`)) + require.Error(t, err) + assert.Nil(t, rules) + }) +} diff --git a/internal/providers/credentials_test.go b/internal/providers/credentials_test.go index 37e7494cb..92dfe04ee 100644 --- a/internal/providers/credentials_test.go +++ b/internal/providers/credentials_test.go @@ -398,3 +398,34 @@ func TestCredentialsService_ConfiguredProvidersCarryGlobalResilience(t *testing. }) } } + +func TestToRawProviderConfig_TripRules(t *testing.T) { + t.Run("with rules, overrides the breaker trip list only", func(t *testing.T) { + cred := ManagedProviderCredential{ + Name: "openai-main", + Type: "openai", + TripOn: []config.TripRuleConfig{ + {Match: "insufficient_quota", TTL: time.Minute}, + }, + } + + raw := cred.toRawProviderConfig() + require.NotNil(t, raw.Resilience) + require.NotNil(t, raw.Resilience.CircuitBreaker) + require.Equal(t, config.TripRuleMap{"rule_1": cred.TripOn[0]}, raw.Resilience.CircuitBreaker.TripOn) + require.Nil(t, raw.Resilience.Retry) + + // The rules ride the same pipeline declarative providers use: other + // breaker settings keep inheriting from the global config. + resolved := buildProviderConfig(raw, config.ResilienceConfig{ + CircuitBreaker: config.DefaultCircuitBreakerConfig(), + }) + require.Equal(t, config.TripRuleMap{"rule_1": cred.TripOn[0]}, resolved.Resilience.CircuitBreaker.TripOn) + require.Equal(t, config.DefaultCircuitBreakerConfig().FailureThreshold, resolved.Resilience.CircuitBreaker.FailureThreshold) + }) + + t.Run("without rules, resilience stays untouched", func(t *testing.T) { + raw := ManagedProviderCredential{Name: "openai-main", Type: "openai"}.toRawProviderConfig() + require.Nil(t, raw.Resilience) + }) +} diff --git a/internal/providers/openai/compatible_provider.go b/internal/providers/openai/compatible_provider.go index 77dcb1c08..de37fa552 100644 --- a/internal/providers/openai/compatible_provider.go +++ b/internal/providers/openai/compatible_provider.go @@ -144,6 +144,12 @@ func (p *CompatibleProvider) SetRequestMutator(mutator RequestMutator) { p.requestMutator = mutator } +// ResetBreaker force-closes the provider-level and model-scoped circuit +// breakers so traffic resumes immediately after a trip, without a restart. +func (p *CompatibleProvider) ResetBreaker() { + p.client.ResetBreaker() +} + func (p *CompatibleProvider) prepareRequest(req llmclient.Request) llmclient.Request { if p.requestMutator != nil { p.requestMutator(&req) diff --git a/internal/providers/openai/compatible_provider_test.go b/internal/providers/openai/compatible_provider_test.go index 1dfb4f773..82048a93d 100644 --- a/internal/providers/openai/compatible_provider_test.go +++ b/internal/providers/openai/compatible_provider_test.go @@ -261,3 +261,17 @@ func TestCompatibleProvider_CreateBatch_InlineRequests(t *testing.T) { }) } } + +func TestCompatibleProvider_ResetBreaker(t *testing.T) { + server, _ := providertest.JSONServer(t, http.StatusOK, providertest.ModelsJSON) + provider := NewCompatibleProviderWithHTTPClient( + "test-key", + server.Client(), + llmclient.Hooks{}, + CompatibleProviderConfig{ProviderName: "upstream-only", BaseURL: server.URL}, + ) + + // Must be safe to call at any time: it force-closes the client-level + // breaker (and every model-scoped one) without touching the upstream. + assert.NotPanics(t, func() { provider.ResetBreaker() }) +} diff --git a/internal/providers/provider_status.go b/internal/providers/provider_status.go index 73e9edad7..9c4c7faa7 100644 --- a/internal/providers/provider_status.go +++ b/internal/providers/provider_status.go @@ -25,6 +25,9 @@ type SanitizedCircuitBreakerConfig struct { FailureThreshold int `json:"failure_threshold"` SuccessThreshold int `json:"success_threshold"` Timeout string `json:"timeout"` + // TripOn lists the trip rules verbatim: rules are declarative + // configuration, not secrets, so no redaction applies. + TripOn []config.TripRuleConfig `json:"trip_on"` } // SanitizedResilienceConfig exposes effective resilience settings. @@ -125,6 +128,7 @@ func SanitizeProviderConfigs(configs map[string]ProviderConfig) []SanitizedProvi FailureThreshold: cfg.Resilience.CircuitBreaker.FailureThreshold, SuccessThreshold: cfg.Resilience.CircuitBreaker.SuccessThreshold, Timeout: cfg.Resilience.CircuitBreaker.Timeout.String(), + TripOn: cfg.Resilience.CircuitBreaker.TripOn.List(), }, }, }) diff --git a/internal/providers/resilience_policy_test.go b/internal/providers/resilience_policy_test.go index 7d044bb37..817ca852b 100644 --- a/internal/providers/resilience_policy_test.go +++ b/internal/providers/resilience_policy_test.go @@ -3,6 +3,7 @@ package providers import ( "encoding/json" "testing" + "time" "github.com/enterpilot/gomodel/config" "github.com/stretchr/testify/require" @@ -25,6 +26,33 @@ func TestProviderResilienceStatusOverrides(t *testing.T) { require.Equal(t, global, buildProviderConfig(config.RawProviderConfig{}, global).Resilience) } +func TestProviderResilienceTripOnOverrides(t *testing.T) { + global := config.ResilienceConfig{Retry: config.DefaultRetryConfig(), CircuitBreaker: config.DefaultCircuitBreakerConfig()} + global.CircuitBreaker.TripOn = config.TripRuleMap{"global": {Match: "global", TTL: time.Minute}} + + // A provider-level trip_on list replaces the global list wholesale. + var raw config.RawProviderConfig + err := yaml.Unmarshal([]byte("type: openai\nresilience:\n circuit_breaker:\n trip_on:\n quota:\n match: \"quota\"\n ttl: 2m\n"), &raw) + require.NoError(t, err) + got := buildProviderConfig(raw, global).Resilience + require.Equal(t, config.TripRuleMap{"quota": {Name: "quota", Match: "quota", TTL: 2 * time.Minute}}, got.CircuitBreaker.TripOn) + + // An explicit empty map replaces the global rules and disables message trips. + var disabled config.RawProviderConfig + err = yaml.Unmarshal([]byte("type: openai\nresilience:\n circuit_breaker:\n trip_on: {}\n"), &disabled) + require.NoError(t, err) + got = buildProviderConfig(disabled, global).Resilience + require.NotNil(t, got.CircuitBreaker.TripOn) + require.Empty(t, got.CircuitBreaker.TripOn) + + // Without a provider trip_on the global list is inherited as-is. + got = buildProviderConfig(config.RawProviderConfig{Type: "openai"}, global).Resilience + require.Equal(t, global.CircuitBreaker.TripOn, got.CircuitBreaker.TripOn) + + // Unset everywhere stays nil so llmclient never trips on messages. + require.Nil(t, buildProviderConfig(config.RawProviderConfig{}, config.ResilienceConfig{}).Resilience.CircuitBreaker.TripOn) +} + func TestSanitizedResiliencePolicies(t *testing.T) { for _, scope := range []string{"model", ""} { t.Run("scope="+scope, func(t *testing.T) { @@ -50,6 +78,9 @@ func TestFactoryRejectsInvalidResilience(t *testing.T) { {Retry: config.RetryConfig{RetryOnStatuses: []string{"bad"}}}, {CircuitBreaker: config.CircuitBreakerConfig{FailureOnStatuses: []string{"600"}}}, {CircuitBreaker: config.CircuitBreakerConfig{Scope: "bad"}}, + {CircuitBreaker: config.CircuitBreakerConfig{TripOn: config.TripRuleMap{"invalid": {Match: "["}}}}, + {CircuitBreaker: config.CircuitBreakerConfig{TripOn: config.TripRuleMap{"neg": {Match: "quota", TTL: -time.Second}}}}, + {CircuitBreaker: config.CircuitBreakerConfig{TripOn: config.TripRuleMap{"empty": {}}}}, } { _, err := factory.Create(ProviderConfig{Resilience: r}) require.Error(t, err) diff --git a/internal/testconventions/assertions_test.go b/internal/testconventions/assertions_test.go index 0f8a2e884..c9e8a38e8 100644 --- a/internal/testconventions/assertions_test.go +++ b/internal/testconventions/assertions_test.go @@ -21,6 +21,7 @@ import ( var skippedDirs = map[string]bool{ ".git": true, ".cache": true, + ".worktrees": true, "node_modules": true, "third_party": true, "vendor": true, diff --git a/web/dashboard/messages/de.json b/web/dashboard/messages/de.json index 480e86f65..29f75b9a3 100644 --- a/web/dashboard/messages/de.json +++ b/web/dashboard/messages/de.json @@ -389,6 +389,11 @@ "overview_api_version": "API-Version", "overview_recent_requests": "Letzte Anfragen", "overview_breaker_state": "Circuit-Breaker-Status", + "overview_reset_breaker": "Breaker zurücksetzen", + "overview_reset_breaker_title": "Circuit Breaker dieses Anbieters sofort schließen", + "overview_reset_breaker_success": "Circuit Breaker für „{name}“ zurückgesetzt.", + "overview_reset_breaker_failed": "Circuit Breaker konnte nicht zurückgesetzt werden.", + "overview_reset_breaker_unavailable": "Der Circuit-Breaker-Reset ist auf dem Gateway nicht verfügbar.", "overview_models_recent_traffic": "Modelle (letzter Traffic)", "overview_configured_models": "Konfigurierte Modelle", "overview_retry": "Retry", @@ -1188,6 +1193,16 @@ "providers_api_key": "API-Schlüssel {number}", "providers_remove_api_key": "API-Schlüssel {number} entfernen", "providers_add_key": "Schlüssel hinzufügen", + "providers_trip_on": "Quota-Breaker", + "providers_trip_on_hint": "Löst den Circuit Breaker sofort aus, wenn ein Fehler auf eine Regel passt; er bleibt für die TTL offen und tastet dann die Erholung ab. Regeln sind optional.", + "providers_trip_on_match": "Treffer {number}", + "providers_trip_on_name": "Regelname {number}", + "providers_trip_on_ttl": "TTL {number}", + "providers_remove_trip_rule": "Regel {number} entfernen", + "providers_add_trip_rule": "Regel hinzufügen", + "providers_trip_on_row_blank": "Leere Zeile entfernen, statt eine Regel leer zu lassen.", + "providers_trip_on_match_required": "Treffer ist erforderlich, wenn eine TTL gesetzt ist.", + "providers_trip_on_ttl_invalid": "Die TTL muss eine Dauer wie 15m, 1h oder 90s sein.", "providers_default": "Anbieter-Standard", "providers_edit_action": "Anbieter {name} bearbeiten", "providers_delete_action": "Anbieter {name} löschen", diff --git a/web/dashboard/messages/en.json b/web/dashboard/messages/en.json index 03ab56496..dba5bec9e 100644 --- a/web/dashboard/messages/en.json +++ b/web/dashboard/messages/en.json @@ -389,6 +389,11 @@ "overview_api_version": "API Version", "overview_recent_requests": "Recent Requests", "overview_breaker_state": "Breaker State", + "overview_reset_breaker": "Reset breaker", + "overview_reset_breaker_title": "Force-close this provider's circuit breaker", + "overview_reset_breaker_success": "Circuit breaker for \"{name}\" reset.", + "overview_reset_breaker_failed": "Failed to reset the circuit breaker.", + "overview_reset_breaker_unavailable": "Circuit breaker reset is unavailable on the gateway.", "overview_models_recent_traffic": "Models (Recent Traffic)", "overview_configured_models": "Configured Models", "overview_retry": "Retry", @@ -1188,6 +1193,16 @@ "providers_api_key": "API key {number}", "providers_remove_api_key": "Remove API key {number}", "providers_add_key": "Add key", + "providers_trip_on": "Quota breaker", + "providers_trip_on_hint": "Trip the circuit breaker instantly when an error matches a rule; it stays open for the TTL, then probes recovery. Rules are optional.", + "providers_trip_on_match": "Match {number}", + "providers_trip_on_name": "Rule name {number}", + "providers_trip_on_ttl": "TTL {number}", + "providers_remove_trip_rule": "Remove rule {number}", + "providers_add_trip_rule": "Add rule", + "providers_trip_on_row_blank": "Remove the empty row instead of leaving a rule blank.", + "providers_trip_on_match_required": "Match is required when a TTL is set.", + "providers_trip_on_ttl_invalid": "TTL must be a duration like 15m, 1h, or 90s.", "providers_default": "Provider default", "providers_edit_action": "Edit provider {name}", "providers_delete_action": "Delete provider {name}", diff --git a/web/dashboard/messages/pl.json b/web/dashboard/messages/pl.json index 8c93c13be..2888c0fe0 100644 --- a/web/dashboard/messages/pl.json +++ b/web/dashboard/messages/pl.json @@ -395,6 +395,11 @@ "overview_api_version": "Wersja API", "overview_recent_requests": "Ostatnie Requesty", "overview_breaker_state": "Stan circuit breakera", + "overview_reset_breaker": "Resetuj breaker", + "overview_reset_breaker_title": "Wymuś zamknięcie circuit breakera tego dostawcy", + "overview_reset_breaker_success": "Circuit breaker dostawcy „{name}” został zresetowany.", + "overview_reset_breaker_failed": "Nie udało się zresetować circuit breakera.", + "overview_reset_breaker_unavailable": "Resetowanie circuit breakera jest niedostępne na bramie.", "overview_models_recent_traffic": "Modele (ostatni ruch)", "overview_configured_models": "Skonfigurowane modele", "overview_retry": "Retry", @@ -1208,6 +1213,16 @@ "providers_api_key": "Klucz API {number}", "providers_remove_api_key": "Usuń klucz API {number}", "providers_add_key": "Dodaj klucz", + "providers_trip_on": "Quota breaker", + "providers_trip_on_hint": "Natychmiast otwiera circuit breaker, gdy błąd pasuje do reguły; pozostaje otwarty przez TTL, po czym sprawdza powrót sprawności. Reguły są opcjonalne.", + "providers_trip_on_match": "Dopasowanie {number}", + "providers_trip_on_name": "Nazwa reguły {number}", + "providers_trip_on_ttl": "TTL {number}", + "providers_remove_trip_rule": "Usuń regułę {number}", + "providers_add_trip_rule": "Dodaj regułę", + "providers_trip_on_row_blank": "Usuń pusty wiersz zamiast zostawiać regułę pustą.", + "providers_trip_on_match_required": "Dopasowanie jest wymagane, gdy ustawiono TTL.", + "providers_trip_on_ttl_invalid": "TTL musi być czasem trwania, np. 15m, 1h lub 90s.", "providers_default": "Domyślne dostawcy", "providers_edit_action": "Edytuj dostawcę {name}", "providers_delete_action": "Usuń dostawcę {name}", diff --git a/web/dashboard/messages/zh-CN.json b/web/dashboard/messages/zh-CN.json index b140fc38f..ef4978d93 100644 --- a/web/dashboard/messages/zh-CN.json +++ b/web/dashboard/messages/zh-CN.json @@ -362,6 +362,11 @@ "overview_api_version": "API 版本", "overview_recent_requests": "最近的请求", "overview_breaker_state": "熔断状态", + "overview_reset_breaker": "重置熔断器", + "overview_reset_breaker_title": "强制关闭该提供商的熔断器", + "overview_reset_breaker_success": "已重置“{name}”的熔断器。", + "overview_reset_breaker_failed": "重置熔断器失败。", + "overview_reset_breaker_unavailable": "网关不支持熔断器重置。", "overview_models_recent_traffic": "模型(近期流量)", "overview_configured_models": "已配置的模型", "overview_retry": "重试", @@ -1098,6 +1103,16 @@ "providers_api_key": "API 密钥 {number}", "providers_remove_api_key": "移除 API 密钥 {number}", "providers_add_key": "添加密钥", + "providers_trip_on": "配额熔断", + "providers_trip_on_hint": "错误匹配到规则时立即触发熔断器,在 TTL 内保持断开,随后试探恢复。规则可选。", + "providers_trip_on_match": "匹配 {number}", + "providers_trip_on_name": "规则名称 {number}", + "providers_trip_on_ttl": "TTL {number}", + "providers_remove_trip_rule": "移除规则 {number}", + "providers_add_trip_rule": "添加规则", + "providers_trip_on_row_blank": "请移除空行,不要留空规则。", + "providers_trip_on_match_required": "设置 TTL 时必须填写匹配。", + "providers_trip_on_ttl_invalid": "TTL 必须是时长格式,如 15m、1h 或 90s。", "providers_default": "供应商默认值", "providers_edit_action": "编辑供应商 {name}", "providers_delete_action": "删除供应商 {name}", diff --git a/web/dashboard/src/lib/api/client.js b/web/dashboard/src/lib/api/client.js index 593a40839..aaf80346c 100644 --- a/web/dashboard/src/lib/api/client.js +++ b/web/dashboard/src/lib/api/client.js @@ -92,6 +92,20 @@ export function sendJSON(path, method, body, options = {}) { // errorMessage extracts a human-readable message from an admin error payload. export { errorMessage, errorPayloadMessage } from "./errors.js"; +// resetCircuitBreaker force-closes one provider's tripped circuit breaker +// (POST /admin/providers/{name}/circuit-breaker/reset; 204 No Content on +// success, 404 for an unknown provider, 400 when the breaker cannot reset). +export function resetCircuitBreaker(providerName) { + return sendJSON( + "/admin/providers/" + + encodeURIComponent(providerName) + + "/circuit-breaker/reset", + "POST", + undefined, + { label: "reset circuit breaker" }, + ); +} + export function isAbortError(error) { return Boolean(error) && (error.name === "AbortError" || error.code === 20); } diff --git a/web/dashboard/src/pages/overview/ProviderStatusCard.svelte b/web/dashboard/src/pages/overview/ProviderStatusCard.svelte index 0a119f1b4..d23e1ead2 100644 --- a/web/dashboard/src/pages/overview/ProviderStatusCard.svelte +++ b/web/dashboard/src/pages/overview/ProviderStatusCard.svelte @@ -18,6 +18,7 @@ providerLastCheckedTime, providerLastCheckedTitle, providerStatusPillTitle, + providerBreakerResettable, } from "./providersLogic.js"; import { ChevronDown } from "lucide"; import * as m from "$lib/paraglide/messages.js"; @@ -25,6 +26,10 @@ let { provider } = $props(); const expanded = $derived(providerStatusState.cardExpanded(provider)); + // Only a tripped breaker (open/half-open) offers the reset; the button + // also waits out any in-flight POST (single-flight guard blocks all cards). + const breakerResettable = $derived(providerBreakerResettable(provider)); + const resetting = $derived(providerStatusState.resettingName); const formatTimestamp = (ts) => timezone.formatTimestamp(ts); @@ -54,10 +59,19 @@ {/if} - {provider.status_label} +
+ + {provider.status_label} +
@@ -152,6 +166,45 @@ gap: 12px; } + .provider-status-head-end { + display: flex; + align-items: center; + gap: 8px; + } + + /* Breaker reset: a quiet action next to the status pill, greyed out unless + the breaker actually needs it. Shares the section toggle's palette. */ + .provider-breaker-reset { + padding: 4px 10px; + background: var(--bg-surface); + border: 1px solid var(--border); + border-radius: 6px; + color: var(--text); + font-size: 12px; + font-family: inherit; + font-weight: 600; + white-space: nowrap; + cursor: pointer; + transition: + background-color 0.18s ease, + border-color 0.18s ease, + color 0.18s ease; + } + + .provider-breaker-reset:hover:enabled { + background: var(--bg-surface-hover); + } + + .provider-breaker-reset:focus-visible { + outline: 2px solid color-mix(in srgb, var(--accent) 28%, transparent); + outline-offset: 2px; + } + + .provider-breaker-reset:disabled { + opacity: 0.45; + cursor: not-allowed; + } + .provider-status-name { display: flex; flex-wrap: wrap; diff --git a/web/dashboard/src/pages/overview/overviewState.svelte.js b/web/dashboard/src/pages/overview/overviewState.svelte.js index 5003d1fa0..91ffa7fdc 100644 --- a/web/dashboard/src/pages/overview/overviewState.svelte.js +++ b/web/dashboard/src/pages/overview/overviewState.svelte.js @@ -2,9 +2,10 @@ // card expand preferences), audit stats, the MCP servers summary, and the // activity-calendar data. -import { getJSON, isAbortError } from "$lib/api/client.js"; +import { getJSON, isAbortError, resetCircuitBreaker, errorPayloadMessage } from "$lib/api/client.js"; import * as m from "$lib/paraglide/messages.js"; import { dateRange } from "$lib/stores/dateRange.svelte.js"; +import { flash } from "$lib/stores/flash.svelte.js"; import { runtimeConfig } from "$lib/stores/runtimeConfig.svelte.js"; import { emptyProviderStatus, @@ -24,6 +25,9 @@ class ProviderStatusState { loadedOnce = $state(false); detailsExpanded = $state(false); cardOverrides = $state({}); + // Provider name whose breaker reset is in flight, so its card's button + // cannot be double-clicked while the POST is pending. + resettingName = $state(""); #controller = null; #pollTimer = null; #prefsLoaded = false; @@ -107,6 +111,45 @@ class ProviderStatusState { } } + // resetBreaker force-closes a tripped circuit breaker (open or half-open) + // and refreshes the provider status so the card's button re-disables on + // fresh state. Idempotent and low-stakes, so no confirmation dialog. + async resetBreaker(provider) { + const name = String((provider && provider.name) || "").trim(); + if (!name || this.resettingName) { + return; + } + this.resettingName = name; + try { + let result; + try { + result = await resetCircuitBreaker(name); + } catch (e) { + console.error("Failed to reset circuit breaker:", e); + flash.error(m.overview_reset_breaker_failed()); + return; + } + if (result.stale) return; + if (result.status === 503) { + flash.error(m.overview_reset_breaker_unavailable()); + return; + } + if (!result.ok) { + // 401 stays silent: the global auth dialog owns it. + if (result.status !== 401) { + flash.error( + errorPayloadMessage(result.data, m.overview_reset_breaker_failed()), + ); + } + return; + } + flash.success(m.overview_reset_breaker_success({ name })); + await this.fetch(); + } finally { + this.resettingName = ""; + } + } + // Re-probe every 3s while any provider is still "Starting", so the cards // settle without a manual refresh. #schedulePoll() { diff --git a/web/dashboard/src/pages/overview/providersLogic.js b/web/dashboard/src/pages/overview/providersLogic.js index 4e61ef867..936f61390 100644 --- a/web/dashboard/src/pages/overview/providersLogic.js +++ b/web/dashboard/src/pages/overview/providersLogic.js @@ -270,6 +270,22 @@ export function providerBreakerState(provider) { return requestHealth ? String(requestHealth.circuit_state || "").trim() : ""; } +// Live breaker state from the status payload's top-level circuit_state +// ("open", "half-open", "closed", or "" before the provider served traffic). +// This is the field the reset button is driven from; it mirrors the +// request-health snapshot the details section renders. +export function providerCircuitState(provider) { + return String((provider && provider.circuit_state) || "").trim(); +} + +// Reset is only meaningful while the breaker blocks traffic or is probing +// recovery: open has tripped, half-open is mid-recovery. A closed breaker +// (or an unconfigured one, which reports no state) needs no reset. +export function providerBreakerResettable(provider) { + const state = providerCircuitState(provider); + return state === "open" || state === "half-open"; +} + export function providerBreakerStateLabel(provider) { const state = providerBreakerState(provider); if (!state) return ""; diff --git a/web/dashboard/src/pages/providers-config/ProviderCredentialEditor.svelte b/web/dashboard/src/pages/providers-config/ProviderCredentialEditor.svelte index 7cf9aa3c4..1417b64be 100644 --- a/web/dashboard/src/pages/providers-config/ProviderCredentialEditor.svelte +++ b/web/dashboard/src/pages/providers-config/ProviderCredentialEditor.svelte @@ -9,11 +9,14 @@ import EnabledToggle from "$lib/components/atoms/EnabledToggle.svelte"; import EditorDialog from "$lib/components/organisms/EditorDialog.svelte"; import ProviderCredentialField from "./ProviderCredentialField.svelte"; + import TableActionButton from "$lib/components/atoms/TableActionButton.svelte"; + import Icon from "$lib/components/atoms/Icon.svelte"; import { providersConfig } from "./providersConfig.svelte.js"; import { providerCredentialTypeOptions, suggestProviderCredentialName, } from "./providersConfigLogic.js"; + import { Plus, Trash2 } from "lucide"; import * as m from "$lib/paraglide/messages.js"; const typeOptions = $derived( @@ -22,6 +25,15 @@ const fields = $derived(providersConfig.formFields); const nameError = $derived(providersConfig.fieldErrors.name || ""); const typeError = $derived(providersConfig.fieldErrors.type || ""); + const tripOnError = $derived(providersConfig.fieldErrors.trip_on || ""); + + // The id lands on the first rule's match input (or the add button while + // there are no rows), so a rejected save can scroll/focus the block the + // same way it does for the schema-driven fields. + const tripOnId = "provider-credential-trip_on"; + const tripOnTargetId = $derived( + providersConfig.form.trip_on.length > 0 ? tripOnId + "-match-0" : tripOnId + "-add", + ); // onTypeChange resets the Name field to a fresh suggestion whenever the // Type selection changes while creating a provider (Type is immutable once @@ -47,7 +59,14 @@ return; } providersConfig.focusField = ""; - const element = document.getElementById("provider-credential-" + target); + // "trip_on" wraps a rule list: focus the derived target element + // (first rule's match input or the add button). + let element; + if (target === "trip_on") { + element = document.getElementById(tripOnTargetId); + } else { + element = document.getElementById("provider-credential-" + target); + } if (element) { element.scrollIntoView({ block: "center" }); element.focus({ preventScroll: true }); @@ -127,6 +146,68 @@ {/each} +
+ +
+ {#each providersConfig.form.trip_on as rule, index (index)} +
+ providersConfig.clearFieldError("trip_on")} + /> + providersConfig.clearFieldError("trip_on")} + /> + providersConfig.clearFieldError("trip_on")} + /> + providersConfig.removeTripRuleRow(index)} + > + + +
+ {/each} +
+
+ +
+ {#if tripOnError} + {tripOnError} + {:else} + {m.providers_trip_on_hint()} + {/if} +
+
{/if} + + diff --git a/web/dashboard/src/pages/providers-config/ProviderCredentialList.svelte b/web/dashboard/src/pages/providers-config/ProviderCredentialList.svelte index 9c0ea6f94..7d9b14158 100644 --- a/web/dashboard/src/pages/providers-config/ProviderCredentialList.svelte +++ b/web/dashboard/src/pages/providers-config/ProviderCredentialList.svelte @@ -9,11 +9,14 @@ import { providerCredentialAuthLabel, providerCredentialModelsLabel, + providerCredentialTripRulesLabel, providerRowsHaveActions, + providerRowsHaveTripRules, } from "./providersConfigLogic.js"; import { Pencil, X } from "lucide"; const showActions = $derived(providerRowsHaveActions(providersConfig.filteredRows)); + const showTripRules = $derived(providerRowsHaveTripRules(providersConfig.filteredRows));
@@ -25,6 +28,9 @@ {m.overview_base_url()} {m.providers_auth()} {m.providers_models()} + {#if showTripRules} + {m.providers_trip_on()} + {/if} {m.providers_enabled()} {m.providers_updated()} {#if showActions} @@ -48,6 +54,10 @@ {row.base_url || "—"} {providerCredentialAuthLabel(row)} {providerCredentialModelsLabel(row)} + {#if showTripRules} + {providerCredentialTripRulesLabel(row) || "—"} + {/if} + 0) { + const minutes = rest % 60; + rest = (rest - minutes) / 60; + out = minutes + "m" + out; + if (rest > 0) { + out = rest + "h" + out; + } + } + } + return (neg ? "-" : "") + out; +} + +// tripRulesToRows converts a view row's trip_on array ({match, ttl} with ttl +// in nanoseconds) into editable rows whose ttl is a Go duration string. +export function tripRulesToRows(tripOn) { + return (Array.isArray(tripOn) ? tripOn : []).map((rule) => ({ + // Optional group name: env overrides match config rules by name. + name: String((rule && rule.name) || ""), + match: String((rule && rule.match) || ""), + // ttl 0 or absent means "use breaker timeout"; show a blank field, not "0s". + ttl: rule && rule.ttl ? formatGoDurationNs(rule.ttl) : "", + })); +} + +// tripRuleRowsToWire converts editor rows into the wire trip_on array +// ({match, ttl} with ttl in nanoseconds), trimming the match and dropping +// rows the operator left completely empty. Returns null when a filled row's +// ttl is not a valid Go duration. An empty list is valid: no rules. +export function tripRuleRowsToWire(rows) { + const wire = []; + for (const row of Array.isArray(rows) ? rows : []) { + const name = String((row && row.name) || "").trim(); + const match = String((row && row.match) || "").trim(); + const ttlText = String((row && row.ttl) || "").trim(); + if (!match && !ttlText && !name) { + continue; + } + // TTL is optional: empty string means use breaker timeout (encode as 0). + const ttl = ttlText ? parseGoDuration(ttlText) : 0; + if (ttl === null) { + return null; + } + const wireRule = { match, ttl }; + if (name) { + wireRule.name = name; + } + wire.push(wireRule); + } + return wire; +} + +// validateTripRuleRows returns the editor message for the first trip-rule +// problem, or "" when the rules can be submitted. +function validateTripRuleRows(rows) { + for (const row of Array.isArray(rows) ? rows : []) { + const match = String((row && row.match) || "").trim(); + const ttlText = String((row && row.ttl) || "").trim(); + if (!match && !ttlText && !String((row && row.name) || "").trim()) { + return m.providers_trip_on_row_blank(); + } + if (!match) { + return m.providers_trip_on_match_required(); + } + // TTL is optional; when omitted the backend uses the breaker's open-state + // timeout. Validate only when the field has content. + if (ttlText && parseGoDuration(ttlText) === null) { + return m.providers_trip_on_ttl_invalid(); + } + } + return ""; +} + +// providerCredentialTripRulesLabel summarizes a row's trip rules for the +// list ("insufficient_quota (15m0s), rate limit (1m0s)"); empty when the row +// declares none. Read-only for every row: rules are declarative +// configuration for config-declared providers and editable in the editor for +// dashboard-managed ones. +export function providerCredentialTripRulesLabel(row) { + return (Array.isArray(row && row.trip_on) ? row.trip_on : []) + .map((rule) => { + // ttl 0 or absent means "use breaker timeout" (operator left it blank); + // omit the suffix instead of showing "0s" or "()". + if (!rule || !rule.ttl) { + return String((rule && rule.match) || ""); + } + return String((rule && rule.match) || "") + " (" + formatGoDurationNs(rule.ttl) + ")"; + }) + .filter(Boolean) + .join(", "); +} + // providerCredentialKeyRowsToArray flattens editor rows back into the wire // array. Values are NOT trimmed: an untouched "***********" mask must be sent // verbatim so the server preserves the stored key at that position. @@ -367,6 +538,7 @@ export function providerCredentialRowToForm(row) { ), gcp_scope: String((row && row.gcp_scope) || ""), models: (Array.isArray(row && row.models) ? row.models : []).join(", "), + trip_on: tripRulesToRows(row && row.trip_on), enabled: !row || row.enabled !== false, }; } @@ -411,6 +583,11 @@ export function validateProviderCredentialForm(form, mode, existingRows, schema) errors[field.name] = message; } } + + const tripError = validateTripRuleRows(form && form.trip_on); + if (tripError) { + errors.trip_on = tripError; + } return errors; } @@ -468,6 +645,11 @@ export function buildProviderCredentialPayload(form, schema) { type: String((form && form.type) || "").trim(), enabled: Boolean(form && form.enabled), }; + // PUT replaces the whole row, so trip_on is always sent — an empty array + // is how clearing every rule works. (Validation rejects invalid ttls + // before the payload is built; a null here means "no valid rules".) + const tripOn = tripRuleRowsToWire(form && form.trip_on); + payload.trip_on = Array.isArray(tripOn) ? tripOn : []; const { primary, advanced } = providerCredentialFormFields(schema); const hasAdvertisedFields = Boolean( schema && Array.isArray(schema.fields) && schema.fields.length > 0, @@ -525,3 +707,12 @@ function providerCredentialPayloadValue(form, name) { export function providerRowsHaveActions(rows) { return (Array.isArray(rows) ? rows : []).some((row) => row && !row.managed); } + +// providerRowsHaveTripRules reports whether any listed provider declares +// quota-breaker trip rules, which is when the list shows the read-only +// rules column at all. +export function providerRowsHaveTripRules(rows) { + return (Array.isArray(rows) ? rows : []).some( + (row) => Array.isArray(row && row.trip_on) && row.trip_on.length > 0, + ); +} diff --git a/web/dashboard/tests/overview-breaker-reset.test.js b/web/dashboard/tests/overview-breaker-reset.test.js new file mode 100644 index 000000000..f50a7b6d1 --- /dev/null +++ b/web/dashboard/tests/overview-breaker-reset.test.js @@ -0,0 +1,107 @@ +// Reset-button tests for the Providers Overview cards. The button's state +// logic is pure (providersLogic.js); the POST flow lives in the Svelte state +// module, which node cannot import (runes), so its wiring is asserted against +// the source the same way mcp-servers.test.js inspects fetchServers. +import test from "node:test"; +import assert from "node:assert/strict"; +import { readFileSync } from "node:fs"; +import { fileURLToPath } from "node:url"; + +import { + providerCircuitState, + providerBreakerResettable, +} from "../src/pages/overview/providersLogic.js"; + +const CLIENT_SOURCE = fileURLToPath( + new URL("../src/lib/api/client.js", import.meta.url), +); +const STATE_SOURCE = fileURLToPath( + new URL("../src/pages/overview/overviewState.svelte.js", import.meta.url), +); + +test("only an open or half-open breaker offers the reset button", () => { + const withState = (state) => ({ name: "openai", circuit_state: state }); + + // Tripped states: resettable. + assert.equal(providerBreakerResettable(withState("open")), true); + assert.equal(providerBreakerResettable(withState("half-open")), true); + + // Healthy, unknown, and never-served providers: greyed out. + assert.equal(providerBreakerResettable(withState("closed")), false); + assert.equal(providerBreakerResettable(withState("")), false); + assert.equal(providerBreakerResettable(withState("tripped")), false); + + // Missing or malformed rows stay disabled rather than throwing. + assert.equal(providerBreakerResettable({}), false); + assert.equal(providerBreakerResettable({ circuit_state: null }), false); + assert.equal(providerBreakerResettable(null), false); + assert.equal(providerBreakerResettable(undefined), false); +}); + +test("providerCircuitState reads the top-level circuit_state field verbatim", () => { + assert.equal(providerCircuitState({ circuit_state: "half-open" }), "half-open"); + assert.equal(providerCircuitState({ circuit_state: " open " }), "open"); + assert.equal(providerCircuitState({ circuit_state: "" }), ""); + assert.equal(providerCircuitState({}), ""); + assert.equal(providerCircuitState(null), ""); +}); + +test("resetCircuitBreaker posts to the provider's circuit-breaker reset endpoint", () => { + const source = readFileSync(CLIENT_SOURCE, "utf8"); + const declaration = source.match( + /export function resetCircuitBreaker\(providerName\) \{[\s\S]*?\n\}/, + ); + assert.ok(declaration, "resetCircuitBreaker declaration missing"); + + const body = declaration[0]; + // The provider name is path-encoded, and the endpoint matches the gateway + // route exactly; the request is a body-less POST. + assert.match(body, /sendJSON\(/); + assert.match(body, /encodeURIComponent\(providerName\)/); + assert.match(body, /\/admin\/providers\/"\s*\+\s*\n?\s*encodeURIComponent\(providerName\)\s*\+\s*\n?\s*"\/circuit-breaker\/reset"/); + assert.match(body, /"POST"/); +}); + +test("resetBreaker posts per provider name, flashes failures, and refreshes status", () => { + const source = readFileSync(STATE_SOURCE, "utf8"); + const declaration = source.match( + /async resetBreaker\(provider\) \{[\s\S]*?\n \}/, + ); + assert.ok(declaration, "resetBreaker declaration missing"); + const body = declaration[0]; + + // The right provider name goes to the API client. + assert.match(body, /const name = String\(\(provider && provider\.name\) \|\| ""\)\.trim\(\)/); + assert.match(body, /resetCircuitBreaker\(name\)/); + // A re-entrant click while a reset is in flight does nothing. + assert.match(body, /this\.resettingName\)/); + // 204 refreshes the provider status data through the existing fetch path, + // so the button re-disables on the fresh circuit_state. + assert.match(body, /await this\.fetch\(\)/); + // Every failure path surfaces through the flash store; success flashes too. + assert.match(body, /flash\.error\(m\.overview_reset_breaker_failed\(\)\)/); + assert.match(body, /flash\.error\(m\.overview_reset_breaker_unavailable\(\)\)/); + assert.match(body, /flash\.success\(m\.overview_reset_breaker_success\(\{ name \}\)\)/); + // The in-flight marker clears even when the POST throws. + assert.match(body, /finally \{\s*this\.resettingName = "";?\s*\}/); +}); + +test("the reset button renders on every card, disabled off the breaker state", () => { + const source = readFileSync( + fileURLToPath( + new URL("../src/pages/overview/ProviderStatusCard.svelte", import.meta.url), + ), + "utf8", + ); + + assert.match(source, /providerBreakerResettable/); + assert.match( + source, + /disabled=\{!breakerResettable \|\| resetting\}/, + ); + assert.match(source, /providerStatusState\.resetBreaker\(provider\)/); + // The reset guard is global: any in-flight reset disables ALL cards' + // buttons (single-flight). It does not compare resettingName to a + // specific provider.name. + assert.match(source, /const resetting = \$derived\(providerStatusState\.resettingName\)/); +}); diff --git a/web/dashboard/tests/providers-config.test.js b/web/dashboard/tests/providers-config.test.js index 26bf5c67a..be890b584 100644 --- a/web/dashboard/tests/providers-config.test.js +++ b/web/dashboard/tests/providers-config.test.js @@ -23,6 +23,12 @@ import { validateProviderCredentialForm, buildProviderCredentialPayload, providerRowsHaveActions, + providerRowsHaveTripRules, + providerCredentialTripRulesLabel, + parseGoDuration, + formatGoDurationNs, + tripRulesToRows, + tripRuleRowsToWire, } from "../src/pages/providers-config/providersConfigLogic.js"; // Schemas shaped like GET /admin/provider-credentials/types serves them. @@ -102,6 +108,7 @@ test("the payload carries the fields the provider type accepts", () => { "models", "name", "session_sticky_keys", + "trip_on", "type", ]); }); @@ -570,3 +577,245 @@ test("providerRowsHaveActions is false when every provider is managed", () => { ); assert.equal(providerRowsHaveActions(undefined), false); }); + +// --- Quota breaker trip rules --- + +test("parseGoDuration mirrors time.ParseDuration for the strings operators type", () => { + assert.equal(parseGoDuration("15m"), 15 * 60 * 1e9); + assert.equal(parseGoDuration("1h"), 3600 * 1e9); + assert.equal(parseGoDuration("1h30m"), 5400 * 1e9); + assert.equal(parseGoDuration("90s"), 90 * 1e9); + assert.equal(parseGoDuration("500ms"), 5e8); + assert.equal(parseGoDuration("2h45m"), 9900 * 1e9); + assert.equal(parseGoDuration("0"), 0); + assert.equal(parseGoDuration("-500ms"), -5e8); + assert.equal(parseGoDuration("1.5h"), 5400 * 1e9); + assert.equal(parseGoDuration(" 15m "), 900 * 1e9); + + // Go's exact strings round-trip through both directions. + assert.equal(parseGoDuration(formatGoDurationNs(900 * 1e9)), 900 * 1e9); + + // Invalid: missing unit, empty, unknown suffix, trailing junk. + assert.equal(parseGoDuration(""), null); + assert.equal(parseGoDuration("15"), null); + assert.equal(parseGoDuration("abc"), null); + assert.equal(parseGoDuration("15x"), null); + assert.equal(parseGoDuration("15m30"), null); + assert.equal(parseGoDuration(null), null); + assert.equal(parseGoDuration(undefined), null); +}); + +test("formatGoDurationNs renders durations the way Go's Duration.String does", () => { + assert.equal(formatGoDurationNs(0), "0s"); + assert.equal(formatGoDurationNs(900 * 1e9), "15m0s"); + assert.equal(formatGoDurationNs(3600 * 1e9), "1h0m0s"); + assert.equal(formatGoDurationNs(5400 * 1e9), "1h30m0s"); + assert.equal(formatGoDurationNs(90 * 1e9), "1m30s"); + assert.equal(formatGoDurationNs(45 * 1e9), "45s"); + assert.equal(formatGoDurationNs(5e8), "500ms"); + assert.equal(formatGoDurationNs(15e5), "1.5ms"); + assert.equal(formatGoDurationNs(15e2), "1.5µs"); + assert.equal(formatGoDurationNs(42), "42ns"); + assert.equal(formatGoDurationNs(-5e8), "-500ms"); + assert.equal(formatGoDurationNs("not a number"), ""); + // A missing ttl decodes to 0 and displays as Go's zero duration. + assert.equal(formatGoDurationNs(null), "0s"); + assert.equal(formatGoDurationNs(undefined), ""); +}); + +test("tripRulesToRows converts the view's nanosecond ttl into duration strings", () => { + assert.deepEqual( + tripRulesToRows([ + { match: "insufficient_quota", ttl: 900 * 1e9 }, + { match: "rate limit", ttl: 60000000000 }, + ]), + [ + { name: "", match: "insufficient_quota", ttl: "15m0s" }, + { name: "", match: "rate limit", ttl: "1m0s" }, + ], + ); + assert.deepEqual(tripRulesToRows(undefined), []); + assert.deepEqual(tripRulesToRows("junk"), []); +}); + +test("tripRulesToRows renders a zero or absent ttl as a blank field, not 0s", () => { + assert.deepEqual( + tripRulesToRows([ + { match: "insufficient_quota", ttl: 0 }, + { match: "no ttl key" }, + ]), + [ + { name: "", match: "insufficient_quota", ttl: "" }, + { name: "", match: "no ttl key", ttl: "" }, + ], + ); +}); + +test("tripRuleRowsToWire builds the {match, ttl} payload with nanosecond ttls", () => { + const wire = tripRuleRowsToWire([ + { match: " insufficient_quota ", ttl: " 15m " }, + { match: "", ttl: "" }, + { match: "rate limit", ttl: "1h30m" }, + ]); + + assert.deepEqual(wire, [ + { match: "insufficient_quota", ttl: 900 * 1e9 }, + { match: "rate limit", ttl: 5400 * 1e9 }, + ]); + + // An empty list is fine: rules are optional. + assert.deepEqual(tripRuleRowsToWire([]), []); + assert.deepEqual(tripRuleRowsToWire(undefined), []); + assert.deepEqual(tripRuleRowsToWire([{ match: "", ttl: "" }]), []); + + // A filled row with an unparsable ttl is reported, not guessed at. + assert.equal(tripRuleRowsToWire([{ match: "quota", ttl: "soon" }]), null); + + // A group name rides along so env overrides can match by name; blanks + // stay omitted. + assert.deepEqual( + tripRuleRowsToWire([{ name: " weekly_quota ", match: "usage limit", ttl: "4h" }]), + [{ name: "weekly_quota", match: "usage limit", ttl: 4 * 3600 * 1e9 }], + ); +}); + +test("validation rejects blank trip-rule rows and half-filled or invalid rules", () => { + const form = (trip_on) => ({ + ...defaultProviderCredentialForm(), + name: "my-openai", + type: "openai", + api_keys: [{ value: "sk-live" }], + trip_on, + }); + + assert.equal(validateProviderCredentialForm(form([]), "create", [], OPENAI_SCHEMA).trip_on, undefined); + assert.equal( + validateProviderCredentialForm(form([{ match: "", ttl: "", name: "" }]), "create", [], OPENAI_SCHEMA).trip_on, + "Remove the empty row instead of leaving a rule blank.", + ); + assert.equal( + validateProviderCredentialForm(form([{ match: "", ttl: "15m" }]), "create", [], OPENAI_SCHEMA).trip_on, + "Match is required when a TTL is set.", + ); + assert.equal( + validateProviderCredentialForm(form([{ match: "quota", ttl: "" }]), "create", [], OPENAI_SCHEMA).trip_on, + undefined, + ); + assert.match( + validateProviderCredentialForm(form([{ match: "quota", ttl: "soon" }]), "create", [], OPENAI_SCHEMA).trip_on, + /duration like 15m/, + ); +}); + +test("trip_on round-trips from a stored row through the form into the PUT payload", () => { + const row = { + name: "my-openai", + type: "openai", + api_keys: ["***********"], + trip_on: [ + { match: "insufficient_quota", ttl: 900 * 1e9 }, + { match: "rate limit", ttl: 60000000000 }, + ], + managed: false, + enabled: true, + }; + + const form = providerCredentialRowToForm(row); + assert.deepEqual(form.trip_on, [ + { name: "", match: "insufficient_quota", ttl: "15m0s" }, + { name: "", match: "rate limit", ttl: "1m0s" }, + ]); + + const body = buildProviderCredentialPayload(form, OPENAI_SCHEMA); + assert.deepEqual(body.trip_on, [ + { match: "insufficient_quota", ttl: 900 * 1e9 }, + { match: "rate limit", ttl: 60000000000 }, + ]); +}); + +test("an empty trip_on list is always in the payload so clearing rules works", () => { + const body = buildProviderCredentialPayload(defaultProviderCredentialForm(), OPENAI_SCHEMA); + assert.deepEqual(body.trip_on, []); +}); + +test("match-only rows are valid and produce ttl: 0 in the wire payload", () => { + const form = (trip_on) => ({ + ...defaultProviderCredentialForm(), + name: "my-openai", + type: "openai", + api_keys: [{ value: "sk-live" }], + trip_on, + }); + + // Match-only row passes validation. + assert.equal( + validateProviderCredentialForm(form([{ match: "quota" }]), "create", [], OPENAI_SCHEMA).trip_on, + undefined, + ); + + // Wire payload carries ttl: 0 for a match-only row. + const wire = tripRuleRowsToWire([{ match: "quota", ttl: "" }]); + assert.deepEqual(wire, [{ match: "quota", ttl: 0 }]); +}); + +test("changing provider type keeps trip rules, like the other identity values", () => { + const typed = { + ...defaultProviderCredentialForm(), + name: "my-openai", + type: "openai", + api_keys: [{ value: "sk-live" }], + trip_on: [{ match: "insufficient_quota", ttl: "15m" }], + vertex_project: "left-over", + }; + + const form = resetProviderCredentialFields( + typed, + providerCredentialFormFields(OPENAI_SCHEMA), + ); + + assert.equal(form.vertex_project, ""); + assert.deepEqual(form.trip_on, [{ match: "insufficient_quota", ttl: "15m" }]); +}); + +test("trip-rule list helpers summarize rows and gate the read-only column", () => { + const quota = 900 * 1e9; + const row = { + name: "dash-openai", + managed: true, + trip_on: [ + { match: "insufficient_quota", ttl: quota }, + { match: "rate limit", ttl: 60 * 1e9 }, + ], + }; + + // Read-only rendering: every row's rules are shown in the list, declared + // or managed. + assert.equal( + providerCredentialTripRulesLabel(row), + "insufficient_quota (15m0s), rate limit (1m0s)", + ); + assert.equal(providerCredentialTripRulesLabel({ trip_on: [] }), ""); + assert.equal(providerCredentialTripRulesLabel({}), ""); + assert.equal(providerCredentialTripRulesLabel(null), ""); + + // ttl 0 means "use breaker timeout"; the label shows match only, no (0s). + assert.equal( + providerCredentialTripRulesLabel({ + trip_on: [{ match: "insufficient_quota", ttl: 0 }], + }), + "insufficient_quota", + ); + + // An absent ttl key must not render an empty () suffix either. + assert.equal( + providerCredentialTripRulesLabel({ + trip_on: [{ match: "insufficient_quota" }], + }), + "insufficient_quota", + ); + + assert.equal(providerRowsHaveTripRules([row, { trip_on: [] }]), true); + assert.equal(providerRowsHaveTripRules([{ trip_on: [] }, {}]), false); + assert.equal(providerRowsHaveTripRules([]), false); + assert.equal(providerRowsHaveTripRules(undefined), false); +});