From 90b6a6c148d5af690b2cb086b142651a01ee2971 Mon Sep 17 00:00:00 2001 From: Andy Dalton Date: Mon, 27 Jul 2026 11:34:57 -0400 Subject: [PATCH 1/2] feat: add ai-budget-exceeded PR label when per-ticket cost cap is hit MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit When the per-ticket AI session cost cap is exceeded, the bot now applies an `ai-budget-exceeded` label to all open PRs for the ticket. Previously the only signal was a generic `blocked` Jira label that PR reviewers were unlikely to see. The label reflects live state: re-applied each scan cycle if removed while the condition holds, and automatically cleared when the bot successfully runs again (e.g., after the cap is raised). It participates in the existing PR validation label mutual exclusivity group alongside `ai-validation-failed` and `ai-nonzero-exit`. The label defaults to "ai-budget-exceeded" when not configured — the only PR validation label with a non-empty default. Set `cost_cap_exceeded: ""` in `pr_validation_labels` to disable. The Jira `blocked` label is no longer set for cost-cap-exceeded (it is still set for general pipeline failures). Assisted-by: Claude Opus 4.6 (1M) --- AGENTS.md | 3 +- config.example.yaml | 1 + executor/export_test.go | 5 + executor/feedback.go | 3 +- executor/labels.go | 45 ++++++++ executor/labels_test.go | 185 +++++++++++++++++++++++++++++++ executor/pipeline.go | 3 +- models/config.go | 23 +++- models/config_test.go | 68 ++++++++++++ projectresolver/resolver.go | 8 +- projectresolver/resolver_test.go | 42 ++++++- 11 files changed, 376 insertions(+), 10 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 6303ff2..7b59533 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -110,9 +110,10 @@ When the `merged` label is applied, the scanner also transitions the ticket to t ### PR Validation Labels -Configurable GitHub PR labels (`pr_validation_labels` in project config) applied when the AI session reports a problem. Labels are mutually exclusive: at most one is set on a PR at any time. Empty strings disable the corresponding label. Suggested values: `ai-validation-failed` and `ai-nonzero-exit`. +Configurable GitHub PR labels (`pr_validation_labels` in project config) applied when the AI session reports a problem or the bot cannot proceed. Labels are mutually exclusive: at most one is set on a PR at any time. Empty strings disable the corresponding label. Suggested values: `ai-validation-failed`, `ai-nonzero-exit`, and `ai-budget-exceeded`. - **`validation_failed`**: Applied when the AI session explicitly reports `validation_passed: false`. - **`nonzero_exit`**: Applied when the AI container exits with a non-zero code (and validation was not explicitly reported as failed). +- **`cost_cap_exceeded`**: Applied when the per-ticket AI session cost cap has been reached. **Defaults to `"ai-budget-exceeded"` when not configured** (the only label with a non-empty default). Set to `""` to disable. Reflects live state: re-applied if removed while the condition holds, removed when the bot successfully runs again (e.g., after the cap is raised). Labels are applied when code is pushed (both new-ticket and feedback paths) and cleared when a subsequent push passes validation. When the AI produces no code changes, labels are left unchanged. Label management is best-effort — failures are logged but never block core operations. diff --git a/config.example.yaml b/config.example.yaml index 5d15128..08b1870 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -63,6 +63,7 @@ jira: pr_validation_labels: validation_failed: "ai-validation-failed" # AI explicitly reported validation failure nonzero_exit: "ai-nonzero-exit" # AI container exited with non-zero code + # cost_cap_exceeded: "ai-budget-exceeded" # Per-ticket cost cap reached (default; set to "" to disable) # Status transitions can be configured per ticket type # All ticket types must be explicitly configured diff --git a/executor/export_test.go b/executor/export_test.go index e64b3b1..029e387 100644 --- a/executor/export_test.go +++ b/executor/export_test.go @@ -84,6 +84,11 @@ func ClearPRValidationLabels(p *Pipeline, logger *zap.Logger, owner, repo string p.clearPRValidationLabels(logger, owner, repo, prNumber, vl) } +// ApplyCostCapPRLabel exposes applyCostCapPRLabel for testing. +func ApplyCostCapPRLabel(p *Pipeline, logger *zap.Logger, ticketKey string, settings *models.ProjectSettings, exceeded bool) { + p.applyCostCapPRLabel(logger, ticketKey, settings, exceeded) +} + // TicketCostPath exposes the ticket cost file path constant for tests. const TicketCostPath = ticketCostPath diff --git a/executor/feedback.go b/executor/feedback.go index a1d7841..a4e6400 100644 --- a/executor/feedback.go +++ b/executor/feedback.go @@ -51,8 +51,7 @@ func (p *Pipeline) executeFeedback(ctx context.Context, job *jobmanager.Job) (re logger.Info("Per-ticket cost cap exceeded, skipping feedback", zap.String("ticket", job.TicketKey), zap.Float64("cap_usd", settings.MaxTicketCostUSD)) - allLabels := models.AllPipelineLabels(settings.FailureLabels, settings.LifecycleLabels) - p.setPipelineLabel(logger, job.TicketKey, allLabels, settings.FailureLabels.Blocked) + p.applyCostCapPRLabel(logger, job.TicketKey, settings, true) return result, errTicketCostCapExceeded } diff --git a/executor/labels.go b/executor/labels.go index 82fe01d..3734e9f 100644 --- a/executor/labels.go +++ b/executor/labels.go @@ -152,3 +152,48 @@ func (p *Pipeline) clearPRValidationLabels( } } } + +// applyCostCapPRLabel applies or removes the cost-cap-exceeded label +// on all open PRs for the given ticket across all configured repos. +// When exceeded is true, the label is set and other validation labels +// are removed (mutual exclusivity). When exceeded is false, only the +// cost cap label is removed. All operations are best-effort. +func (p *Pipeline) applyCostCapPRLabel( + logger *zap.Logger, + ticketKey string, + settings *models.ProjectSettings, + exceeded bool, +) { + label := settings.PRValidationLabels.CostCapLabel() + if label == "" { + return + } + + branchName := fmt.Sprintf("%s/%s", p.cfg.BotUsername, ticketKey) + heads := settings.PRHeads(branchName) + + for _, repo := range settings.Repos { + pr, err := p.findPRByHeadsOptional(repo.Owner, repo.Repo, heads) + if err != nil { + logger.Debug("Failed to find PR for cost cap label", + zap.String("repo", repo.Owner+"/"+repo.Repo), + zap.Error(err)) + continue + } + if pr == nil { + continue + } + + if exceeded { + p.setPRValidationLabel(logger, repo.Owner, repo.Repo, + pr.Number, settings.PRValidationLabels, label) + } else { + if err := p.git.RemovePRLabel(repo.Owner, repo.Repo, pr.Number, label); err != nil { + logger.Debug("Failed to remove cost cap PR label", + zap.String("repo", repo.Owner+"/"+repo.Repo), + zap.Int("pr", pr.Number), + zap.Error(err)) + } + } + } +} diff --git a/executor/labels_test.go b/executor/labels_test.go index e132274..e6021cd 100644 --- a/executor/labels_test.go +++ b/executor/labels_test.go @@ -561,6 +561,39 @@ func TestSetPRValidationLabel(t *testing.T) { p := d.pipeline(t) executor.SetPRValidationLabel(p, zap.NewNop(), "org", "repo", 42, vl, "ai-validation-failed") }) + + t.Run("removes cost cap label when setting validation label", func(t *testing.T) { + costCap := models.DefaultCostCapLabel + vlAll := models.PRValidationLabels{ + ValidationFailed: "ai-validation-failed", + NonzeroExit: "ai-nonzero-exit", + CostCapExceeded: &costCap, + } + var added, removed []string + d := newTestDeps(t) + d.git.AddPRLabelFunc = func(_, _ string, _ int, label string) error { added = append(added, label); return nil } + d.git.RemovePRLabelFunc = func(_, _ string, _ int, label string) error { removed = append(removed, label); return nil } + + p := d.pipeline(t) + executor.SetPRValidationLabel(p, zap.NewNop(), "org", "repo", 42, vlAll, "ai-validation-failed") + + if len(added) != 1 || added[0] != "ai-validation-failed" { + t.Errorf("added = %v, want [ai-validation-failed]", added) + } + removedSet := make(map[string]bool, len(removed)) + for _, l := range removed { + removedSet[l] = true + } + if !removedSet["ai-nonzero-exit"] { + t.Error("expected ai-nonzero-exit to be removed") + } + if !removedSet[models.DefaultCostCapLabel] { + t.Errorf("expected %s to be removed", models.DefaultCostCapLabel) + } + if len(removed) != 2 { + t.Errorf("removed = %v, want exactly 2 entries", removed) + } + }) } func TestClearPRValidationLabels(t *testing.T) { @@ -676,3 +709,155 @@ func TestSetPipelineLabel_ErrorsAreSwallowed(t *testing.T) { } }) } + +func TestApplyCostCapPRLabel(t *testing.T) { + costCapLabel := models.DefaultCostCapLabel + + makeSettings := func(repos []models.RepoSettings) *models.ProjectSettings { + return &models.ProjectSettings{ + Repos: repos, + PRValidationLabels: models.PRValidationLabels{ + ValidationFailed: "ai-validation-failed", + NonzeroExit: "ai-nonzero-exit", + CostCapExceeded: &costCapLabel, + }, + } + } + + t.Run("applies label and removes others on exceeded", func(t *testing.T) { + var added, removed []string + d := newTestDeps(t) + d.git.GetPRForBranchFunc = func(_, _, _ string) (*models.PRDetails, error) { + return &models.PRDetails{Number: 42}, nil + } + d.git.AddPRLabelFunc = func(_, _ string, _ int, label string) error { added = append(added, label); return nil } + d.git.RemovePRLabelFunc = func(_, _ string, _ int, label string) error { removed = append(removed, label); return nil } + + settings := makeSettings([]models.RepoSettings{{Owner: "org", Repo: "repo"}}) + p := d.pipeline(t) + executor.ApplyCostCapPRLabel(p, zap.NewNop(), "TEST-1", settings, true) + + if len(added) != 1 || added[0] != models.DefaultCostCapLabel { + t.Errorf("added = %v, want [%s]", added, models.DefaultCostCapLabel) + } + removedSet := make(map[string]bool, len(removed)) + for _, l := range removed { + removedSet[l] = true + } + if !removedSet["ai-validation-failed"] { + t.Error("expected ai-validation-failed to be removed") + } + if !removedSet["ai-nonzero-exit"] { + t.Error("expected ai-nonzero-exit to be removed") + } + }) + + t.Run("removes only cost cap label when not exceeded", func(t *testing.T) { + var removed []string + d := newTestDeps(t) + d.git.GetPRForBranchFunc = func(_, _, _ string) (*models.PRDetails, error) { + return &models.PRDetails{Number: 42}, nil + } + d.git.RemovePRLabelFunc = func(_, _ string, _ int, label string) error { removed = append(removed, label); return nil } + + settings := makeSettings([]models.RepoSettings{{Owner: "org", Repo: "repo"}}) + p := d.pipeline(t) + executor.ApplyCostCapPRLabel(p, zap.NewNop(), "TEST-1", settings, false) + + if len(removed) != 1 || removed[0] != models.DefaultCostCapLabel { + t.Errorf("removed = %v, want [%s]", removed, models.DefaultCostCapLabel) + } + }) + + t.Run("no-op when label is disabled", func(t *testing.T) { + var prLookups int + empty := "" + d := newTestDeps(t) + d.git.GetPRForBranchFunc = func(_, _, _ string) (*models.PRDetails, error) { + prLookups++ + return &models.PRDetails{Number: 42}, nil + } + + settings := &models.ProjectSettings{ + Repos: []models.RepoSettings{{Owner: "org", Repo: "repo"}}, + PRValidationLabels: models.PRValidationLabels{ + CostCapExceeded: &empty, + }, + } + p := d.pipeline(t) + executor.ApplyCostCapPRLabel(p, zap.NewNop(), "TEST-1", settings, true) + + if prLookups != 0 { + t.Errorf("expected no PR lookups when label disabled, got %d", prLookups) + } + }) + + t.Run("skips repos without open PRs", func(t *testing.T) { + var added []string + d := newTestDeps(t) + d.git.GetPRForBranchFunc = func(_, repo, _ string) (*models.PRDetails, error) { + if repo == "repo-with-pr" { + return &models.PRDetails{Number: 10}, nil + } + return nil, nil + } + d.git.AddPRLabelFunc = func(_, repo string, prNum int, label string) error { + added = append(added, fmt.Sprintf("%s#%d:%s", repo, prNum, label)) + return nil + } + d.git.RemovePRLabelFunc = func(_, _ string, _ int, _ string) error { return nil } + + settings := makeSettings([]models.RepoSettings{ + {Owner: "org", Repo: "repo-no-pr"}, + {Owner: "org", Repo: "repo-with-pr"}, + }) + p := d.pipeline(t) + executor.ApplyCostCapPRLabel(p, zap.NewNop(), "TEST-1", settings, true) + + if len(added) != 1 { + t.Fatalf("added = %v, want 1 entry", added) + } + want := fmt.Sprintf("repo-with-pr#10:%s", models.DefaultCostCapLabel) + if added[0] != want { + t.Errorf("added[0] = %q, want %q", added[0], want) + } + }) + + t.Run("multi-repo applies to all PRs", func(t *testing.T) { + var addedRepos []string + d := newTestDeps(t) + d.git.GetPRForBranchFunc = func(_, repo, _ string) (*models.PRDetails, error) { + if repo == "repo1" { + return &models.PRDetails{Number: 10}, nil + } + return &models.PRDetails{Number: 20}, nil + } + d.git.AddPRLabelFunc = func(_, repo string, _ int, _ string) error { + addedRepos = append(addedRepos, repo) + return nil + } + d.git.RemovePRLabelFunc = func(_, _ string, _ int, _ string) error { return nil } + + settings := makeSettings([]models.RepoSettings{ + {Owner: "org", Repo: "repo1"}, + {Owner: "org", Repo: "repo2"}, + }) + p := d.pipeline(t) + executor.ApplyCostCapPRLabel(p, zap.NewNop(), "TEST-1", settings, true) + + if len(addedRepos) != 2 { + t.Fatalf("expected label on 2 repos, got %d", len(addedRepos)) + } + }) + + t.Run("PR lookup errors are swallowed", func(t *testing.T) { + d := newTestDeps(t) + d.git.GetPRForBranchFunc = func(_, _, _ string) (*models.PRDetails, error) { + return nil, fmt.Errorf("API error") + } + + settings := makeSettings([]models.RepoSettings{{Owner: "org", Repo: "repo"}}) + p := d.pipeline(t) + executor.ApplyCostCapPRLabel(p, zap.NewNop(), "TEST-1", settings, true) + }) +} diff --git a/executor/pipeline.go b/executor/pipeline.go index bdf3233..59e06c0 100644 --- a/executor/pipeline.go +++ b/executor/pipeline.go @@ -138,8 +138,7 @@ func (p *Pipeline) executeNewTicket(ctx context.Context, job *jobmanager.Job) (r logger.Info("Per-ticket cost cap exceeded, skipping", zap.String("ticket", job.TicketKey), zap.Float64("cap_usd", settings.MaxTicketCostUSD)) - allLabels := models.AllPipelineLabels(settings.FailureLabels, settings.LifecycleLabels) - p.setPipelineLabel(logger, job.TicketKey, allLabels, settings.FailureLabels.Blocked) + p.applyCostCapPRLabel(logger, job.TicketKey, settings, true) return result, errTicketCostCapExceeded } diff --git a/models/config.go b/models/config.go index fa1c47b..b9843c4 100644 --- a/models/config.go +++ b/models/config.go @@ -423,6 +423,10 @@ func (ll LifecycleLabels) All() []string { return []string{ll.Queued, ll.Review, ll.Merged} } +// DefaultCostCapLabel is the default GitHub PR label applied when the +// per-ticket AI session cost cap is exceeded. +const DefaultCostCapLabel = "ai-budget-exceeded" + // PRValidationLabels holds configurable GitHub PR labels applied when // the AI session's validation or exit code indicates a problem. Labels // are mutually exclusive: at most one is set on a given PR. Empty @@ -436,12 +440,29 @@ type PRValidationLabels struct { // non-zero code (and validation was not explicitly reported // as failed). NonzeroExit string `yaml:"nonzero_exit" mapstructure:"nonzero_exit"` + + // CostCapExceeded is applied when the per-ticket AI session + // cost cap has been reached. Defaults to DefaultCostCapLabel + // when not configured (nil); set to "" to disable. + CostCapExceeded *string `yaml:"cost_cap_exceeded,omitempty" mapstructure:"cost_cap_exceeded"` +} + +// CostCapLabel returns the raw cost cap exceeded label value. +// Returns empty string when the pointer is nil or points to "". +// The default (DefaultCostCapLabel) is applied by the project +// resolver, not here; callers outside the resolver should only +// see post-resolution values. +func (vl PRValidationLabels) CostCapLabel() string { + if vl.CostCapExceeded == nil { + return "" + } + return *vl.CostCapExceeded } // All returns the configured label strings in a fixed order. Empty // strings (disabled labels) are included; callers should skip them. func (vl PRValidationLabels) All() []string { - return []string{vl.ValidationFailed, vl.NonzeroExit} + return []string{vl.ValidationFailed, vl.NonzeroExit, vl.CostCapLabel()} } // AllPipelineLabels returns all configured failure and lifecycle labels diff --git a/models/config_test.go b/models/config_test.go index f75a7c2..2843674 100644 --- a/models/config_test.go +++ b/models/config_test.go @@ -1907,6 +1907,9 @@ workspaces: if vl.NonzeroExit != "" { t.Errorf("NonzeroExit = %q, want empty", vl.NonzeroExit) } + if vl.CostCapExceeded != nil { + t.Errorf("CostCapExceeded = %v, want nil", vl.CostCapExceeded) + } }) t.Run("custom values", func(t *testing.T) { @@ -1940,6 +1943,7 @@ jira: pr_validation_labels: validation_failed: "custom-vf" nonzero_exit: "custom-nze" + cost_cap_exceeded: "custom-cce" github: app_id: 123456 private_key_path: "` + tmpKeyPath + `" @@ -1969,6 +1973,70 @@ workspaces: if vl.NonzeroExit != "custom-nze" { t.Errorf("NonzeroExit = %q, want custom-nze", vl.NonzeroExit) } + if vl.CostCapExceeded == nil || *vl.CostCapExceeded != "custom-cce" { + t.Errorf("CostCapExceeded = %v, want ptr to custom-cce", vl.CostCapExceeded) + } + }) +} + +func TestPRValidationLabels_CostCapLabel(t *testing.T) { + t.Run("nil pointer returns empty", func(t *testing.T) { + vl := PRValidationLabels{CostCapExceeded: nil} + if got := vl.CostCapLabel(); got != "" { + t.Errorf("CostCapLabel() = %q, want empty", got) + } + }) + + t.Run("empty string pointer returns empty", func(t *testing.T) { + empty := "" + vl := PRValidationLabels{CostCapExceeded: &empty} + if got := vl.CostCapLabel(); got != "" { + t.Errorf("CostCapLabel() = %q, want empty", got) + } + }) + + t.Run("non-empty pointer returns value", func(t *testing.T) { + label := "ai-budget-exceeded" + vl := PRValidationLabels{CostCapExceeded: &label} + if got := vl.CostCapLabel(); got != "ai-budget-exceeded" { + t.Errorf("CostCapLabel() = %q, want ai-budget-exceeded", got) + } + }) +} + +func TestPRValidationLabels_All_IncludesCostCap(t *testing.T) { + t.Run("includes cost cap label", func(t *testing.T) { + label := "ai-budget-exceeded" + vl := PRValidationLabels{ + ValidationFailed: "vf", + NonzeroExit: "nze", + CostCapExceeded: &label, + } + got := vl.All() + want := []string{"vf", "nze", "ai-budget-exceeded"} + if len(got) != len(want) { + t.Fatalf("All() length = %d, want %d", len(got), len(want)) + } + for i := range want { + if got[i] != want[i] { + t.Errorf("All()[%d] = %q, want %q", i, got[i], want[i]) + } + } + }) + + t.Run("nil cost cap returns empty in slice", func(t *testing.T) { + vl := PRValidationLabels{ + ValidationFailed: "vf", + NonzeroExit: "nze", + CostCapExceeded: nil, + } + got := vl.All() + if len(got) != 3 { + t.Fatalf("All() length = %d, want 3", len(got)) + } + if got[2] != "" { + t.Errorf("All()[2] = %q, want empty", got[2]) + } }) } diff --git a/projectresolver/resolver.go b/projectresolver/resolver.go index 1605132..bcc7311 100644 --- a/projectresolver/resolver.go +++ b/projectresolver/resolver.go @@ -74,6 +74,12 @@ func (r *ConfigResolver) ResolveProject(workItem models.WorkItem) (*models.Proje maxTicketCost = *pc.MaxTicketCostUSD } + prVL := pc.PRValidationLabels + if prVL.CostCapExceeded == nil { + defaultLabel := models.DefaultCostCapLabel + prVL.CostCapExceeded = &defaultLabel + } + return &models.ProjectSettings{ Repos: repos, RootRepoURL: ws.RootRepo, @@ -86,7 +92,7 @@ func (r *ConfigResolver) ResolveProject(workItem models.WorkItem) (*models.Proje Container: ws.Container, FailureLabels: pc.FailureLabels, LifecycleLabels: pc.LifecycleLabels, - PRValidationLabels: pc.PRValidationLabels, + PRValidationLabels: prVL, MergedStatus: transitions.Merged, ForkMode: pc.ForkMode, GitHubUsername: ghUsername, diff --git a/projectresolver/resolver_test.go b/projectresolver/resolver_test.go index 0722ea7..ccca871 100644 --- a/projectresolver/resolver_test.go +++ b/projectresolver/resolver_test.go @@ -1029,10 +1029,12 @@ func TestResolveProject_FailureLabels(t *testing.T) { func TestResolveProject_PRValidationLabels(t *testing.T) { t.Run("passes through configured labels", func(t *testing.T) { + customCCE := "custom-cce" cfg := minimalConfig() cfg.Jira.Projects[0].PRValidationLabels = models.PRValidationLabels{ ValidationFailed: "custom-vf", NonzeroExit: "custom-nze", + CostCapExceeded: &customCCE, } r, err := projectresolver.NewConfigResolver(cfg) if err != nil { @@ -1054,10 +1056,44 @@ func TestResolveProject_PRValidationLabels(t *testing.T) { if ps.PRValidationLabels.NonzeroExit != "custom-nze" { t.Errorf("NonzeroExit = %q, want %q", ps.PRValidationLabels.NonzeroExit, "custom-nze") } + if ps.PRValidationLabels.CostCapLabel() != "custom-cce" { + t.Errorf("CostCapLabel() = %q, want %q", ps.PRValidationLabels.CostCapLabel(), "custom-cce") + } }) - t.Run("defaults to empty when not configured", func(t *testing.T) { + t.Run("defaults cost_cap_exceeded when not configured", func(t *testing.T) { + cfg := minimalConfig() + r, err := projectresolver.NewConfigResolver(cfg) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + ps, err := r.ResolveProject(models.WorkItem{ + Key: "PROJ-1", + Type: "Bug", + Components: []string{"backend"}, + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if ps.PRValidationLabels.ValidationFailed != "" { + t.Errorf("ValidationFailed = %q, want empty", ps.PRValidationLabels.ValidationFailed) + } + if ps.PRValidationLabels.NonzeroExit != "" { + t.Errorf("NonzeroExit = %q, want empty", ps.PRValidationLabels.NonzeroExit) + } + if ps.PRValidationLabels.CostCapLabel() != "ai-budget-exceeded" { + t.Errorf("CostCapLabel() = %q, want %q", ps.PRValidationLabels.CostCapLabel(), "ai-budget-exceeded") + } + }) + + t.Run("explicit empty disables cost_cap_exceeded", func(t *testing.T) { + empty := "" cfg := minimalConfig() + cfg.Jira.Projects[0].PRValidationLabels = models.PRValidationLabels{ + CostCapExceeded: &empty, + } r, err := projectresolver.NewConfigResolver(cfg) if err != nil { t.Fatalf("unexpected error: %v", err) @@ -1072,8 +1108,8 @@ func TestResolveProject_PRValidationLabels(t *testing.T) { t.Fatalf("unexpected error: %v", err) } - if ps.PRValidationLabels != (models.PRValidationLabels{}) { - t.Errorf("expected zero PRValidationLabels, got %+v", ps.PRValidationLabels) + if ps.PRValidationLabels.CostCapLabel() != "" { + t.Errorf("CostCapLabel() = %q, want empty (disabled)", ps.PRValidationLabels.CostCapLabel()) } }) } From efbd7c483403723ece622495bb93386371862ef3 Mon Sep 17 00:00:00 2001 From: Andy Dalton Date: Mon, 27 Jul 2026 11:50:08 -0400 Subject: [PATCH 2/2] fix: proactively clear cost cap PR label when cap is no longer exceeded Add applyCostCapPRLabel(..., false) immediately after the cost cap check passes in both pipeline.go and feedback.go. Without this, a stale ai-budget-exceeded label could persist on PRs when the feedback path takes an early return (e.g., no new comments) before the downstream validation label logic runs. Also add assertions to the "PR lookup errors are swallowed" test to verify no label operations occur when the PR lookup fails. Assisted-by: Claude Opus 4.6 (1M) --- executor/feedback.go | 1 + executor/labels_test.go | 7 +++++++ executor/pipeline.go | 1 + 3 files changed, 9 insertions(+) diff --git a/executor/feedback.go b/executor/feedback.go index a4e6400..2829a78 100644 --- a/executor/feedback.go +++ b/executor/feedback.go @@ -54,6 +54,7 @@ func (p *Pipeline) executeFeedback(ctx context.Context, job *jobmanager.Job) (re p.applyCostCapPRLabel(logger, job.TicketKey, settings, true) return result, errTicketCostCapExceeded } + p.applyCostCapPRLabel(logger, job.TicketKey, settings, false) defer func() { // On failure post error comment (but do NOT revert status -- diff --git a/executor/labels_test.go b/executor/labels_test.go index e6021cd..089636c 100644 --- a/executor/labels_test.go +++ b/executor/labels_test.go @@ -851,13 +851,20 @@ func TestApplyCostCapPRLabel(t *testing.T) { }) t.Run("PR lookup errors are swallowed", func(t *testing.T) { + var labelCalls int d := newTestDeps(t) d.git.GetPRForBranchFunc = func(_, _, _ string) (*models.PRDetails, error) { return nil, fmt.Errorf("API error") } + d.git.AddPRLabelFunc = func(_, _ string, _ int, _ string) error { labelCalls++; return nil } + d.git.RemovePRLabelFunc = func(_, _ string, _ int, _ string) error { labelCalls++; return nil } settings := makeSettings([]models.RepoSettings{{Owner: "org", Repo: "repo"}}) p := d.pipeline(t) executor.ApplyCostCapPRLabel(p, zap.NewNop(), "TEST-1", settings, true) + + if labelCalls != 0 { + t.Errorf("expected no label calls on lookup error, got %d", labelCalls) + } }) } diff --git a/executor/pipeline.go b/executor/pipeline.go index 59e06c0..0e15aa0 100644 --- a/executor/pipeline.go +++ b/executor/pipeline.go @@ -141,6 +141,7 @@ func (p *Pipeline) executeNewTicket(ctx context.Context, job *jobmanager.Job) (r p.applyCostCapPRLabel(logger, job.TicketKey, settings, true) return result, errTicketCostCapExceeded } + p.applyCostCapPRLabel(logger, job.TicketKey, settings, false) // --- Clean retry: delete stale branches and workspace --- if job.CleanRetry {