diff --git a/AGENTS.md b/AGENTS.md index 56c1e59..79b0a00 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 87149d5..a1af862 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) # Triage label cleanup: when a ticket leaves new_status and is # assigned to a human, active triage labels are replaced with stale. 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..2829a78 100644 --- a/executor/feedback.go +++ b/executor/feedback.go @@ -51,10 +51,10 @@ 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 } + p.applyCostCapPRLabel(logger, job.TicketKey, settings, false) defer func() { // On failure post error comment (but do NOT revert status -- 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..089636c 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,162 @@ 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) { + 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 bdf3233..0e15aa0 100644 --- a/executor/pipeline.go +++ b/executor/pipeline.go @@ -138,10 +138,10 @@ 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 } + p.applyCostCapPRLabel(logger, job.TicketKey, settings, false) // --- Clean retry: delete stale branches and workspace --- if job.CleanRetry { diff --git a/models/config.go b/models/config.go index 23b3129..339c19a 100644 --- a/models/config.go +++ b/models/config.go @@ -428,6 +428,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 @@ -441,12 +445,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 6255a44..ebc28ec 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 c8fba2e..e094335 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()) } }) }