From 74f9269ee98b9f17d72580714cdd79f400db4df4 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Sun, 2 Aug 2026 18:01:18 +0000 Subject: [PATCH 1/4] Initial plan From d6dd0f49e44c22d78f088b7efe0c87fa977028e2 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Sun, 2 Aug 2026 18:40:53 +0000 Subject: [PATCH 2/4] refactor: reduce function lengths in pkg/workflow and pkg/cli (slices 1-3, in progress) Co-authored-by: pelikhan <4175913+pelikhan@users.noreply.github.com> --- pkg/cli/audit_diff.go | 968 ++++++++++-------- pkg/cli/audit_diff_command.go | 303 +++--- pkg/cli/audit_diff_render.go | 460 ++++----- pkg/sliceutil/sliceutil.go | 7 + pkg/workflow/awf_config.go | 385 +++---- pkg/workflow/awf_helpers.go | 897 ++++++++-------- pkg/workflow/codex_logs.go | 260 +++-- pkg/workflow/compiler_activation_daily_aic.go | 124 ++- pkg/workflow/compiler_activation_job.go | 173 ++-- .../compiler_activation_permissions.go | 89 +- pkg/workflow/compiler_activation_steps.go | 144 +-- pkg/workflow/compiler_aw_context.go | 97 +- pkg/workflow/compiler_custom_jobs.go | 391 ++++--- pkg/workflow/compiler_difc_proxy.go | 94 +- pkg/workflow/compiler_experiments.go | 279 ++--- pkg/workflow/compiler_github_actions_steps.go | 66 +- pkg/workflow/compiler_github_mcp_steps.go | 158 ++- pkg/workflow/compiler_jobs.go | 85 +- pkg/workflow/compiler_main_job.go | 52 +- pkg/workflow/compiler_main_job_helpers.go | 105 +- pkg/workflow/compiler_orchestrator_engine.go | 231 +++-- .../compiler_orchestrator_frontmatter.go | 242 ++--- .../compiler_orchestrator_workflow.go | 261 +++-- pkg/workflow/compiler_safe_output_jobs.go | 481 +++++---- pkg/workflow/compiler_safe_outputs_steps.go | 360 +++---- pkg/workflow/compiler_string_api.go | 186 ++-- pkg/workflow/compiler_unlock_job.go | 126 +-- pkg/workflow/compiler_workflow_call.go | 186 ++-- pkg/workflow/compiler_yaml.go | 164 ++- pkg/workflow/compiler_yaml_normalize.go | 172 ++-- pkg/workflow/copilot_engine_tools.go | 331 +++--- pkg/workflow/gemini_engine.go | 259 ++--- .../permissions_compiler_validator.go | 205 ++-- pkg/workflow/universal_llm_consumer_engine.go | 123 +-- 34 files changed, 4102 insertions(+), 4362 deletions(-) diff --git a/pkg/cli/audit_diff.go b/pkg/cli/audit_diff.go index f002dee2cda..972c246431b 100644 --- a/pkg/cli/audit_diff.go +++ b/pkg/cli/audit_diff.go @@ -66,28 +66,33 @@ type FirewallDiffSummary struct { // Either analysis may be nil, indicating no firewall data for that run. func computeFirewallDiff(run1ID, run2ID int64, run1, run2 *FirewallAnalysis) *FirewallDiff { auditDiffLog.Printf("Computing firewall diff: run1=%d, run2=%d", run1ID, run2ID) - diff := &FirewallDiff{ - Run1ID: run1ID, - Run2ID: run2ID, - } - - // Handle nil cases - run1Stats := make(map[string]DomainRequestStats) - run2Stats := make(map[string]DomainRequestStats) - - if run1 != nil { - run1Stats = run1.RequestsByDomain + ctx := &firewallDiffContext{diff: &FirewallDiff{Run1ID: run1ID, Run2ID: run2ID}} + run1Stats := firewallRequestStats(run1) + run2Stats := firewallRequestStats(run2) + if len(run1Stats) == 0 && len(run2Stats) == 0 { + return ctx.diff } - if run2 != nil { - run2Stats = run2.RequestsByDomain + for _, domain := range collectFirewallDomains(run1Stats, run2Stats) { + stats1, inRun1 := run1Stats[domain] + stats2, inRun2 := run2Stats[domain] + processFirewallDomain(ctx, domain, stats1, inRun1, stats2, inRun2) } + return finalizeFirewallDiff(ctx) +} - // If both are nil/empty, return empty diff - if len(run1Stats) == 0 && len(run2Stats) == 0 { - return diff +type firewallDiffContext struct { + diff *FirewallDiff + anomalyCount int +} + +func firewallRequestStats(analysis *FirewallAnalysis) map[string]DomainRequestStats { + if analysis == nil { + return map[string]DomainRequestStats{} } + return analysis.RequestsByDomain +} - // Collect all domains +func collectFirewallDomains(run1Stats, run2Stats map[string]DomainRequestStats) []string { allDomains := make(map[string]struct{}) for domain := range run1Stats { allDomains[domain] = struct{}{} @@ -95,122 +100,131 @@ func computeFirewallDiff(run1ID, run2ID int64, run1, run2 *FirewallAnalysis) *Fi for domain := range run2Stats { allDomains[domain] = struct{}{} } + return sliceutil.SortedKeys(allDomains) +} - // Sorted domain list for deterministic output - sortedDomains := sliceutil.SortedKeys(allDomains) +func processFirewallDomain(ctx *firewallDiffContext, domain string, stats1 DomainRequestStats, inRun1 bool, stats2 DomainRequestStats, inRun2 bool) { + switch { + case !inRun1 && inRun2: + ctx.diff.NewDomains = append(ctx.diff.NewDomains, buildNewFirewallDomainEntry(ctx, domain, stats2)) + case inRun1 && !inRun2: + ctx.diff.RemovedDomains = append(ctx.diff.RemovedDomains, buildRemovedFirewallDomainEntry(ctx, domain, stats1)) + default: + processSharedFirewallDomain(ctx, domain, stats1, stats2) + } +} - anomalyCount := 0 +func buildNewFirewallDomainEntry(ctx *firewallDiffContext, domain string, stats DomainRequestStats) DomainDiffEntry { + entry := DomainDiffEntry{ + Domain: domain, + DiffEntryBase: DiffEntryBase{Status: "new"}, + Run2Allowed: stats.Allowed, + Run2Blocked: stats.Blocked, + Run2Status: classifyFirewallDomainStatus(stats), + } + if stats.Blocked > 0 { + entry.IsAnomaly = true + entry.AnomalyNote = "new denied domain" + ctx.anomalyCount++ + } + return entry +} - for _, domain := range sortedDomains { - stats1, inRun1 := run1Stats[domain] - stats2, inRun2 := run2Stats[domain] +func buildRemovedFirewallDomainEntry(ctx *firewallDiffContext, domain string, stats DomainRequestStats) DomainDiffEntry { + entry := DomainDiffEntry{ + Domain: domain, + DiffEntryBase: DiffEntryBase{Status: "removed"}, + Run1Allowed: stats.Allowed, + Run1Blocked: stats.Blocked, + Run1Status: classifyFirewallDomainStatus(stats), + } + if stats.Blocked > 0 { + entry.IsAnomaly = true + entry.AnomalyNote = "denied in base run — absent from comparison run" + ctx.anomalyCount++ + } + return entry +} - if !inRun1 && inRun2 { - // New domain in run 2 - entry := DomainDiffEntry{ - Domain: domain, - DiffEntryBase: DiffEntryBase{Status: "new"}, - Run2Allowed: stats2.Allowed, - Run2Blocked: stats2.Blocked, - Run2Status: classifyFirewallDomainStatus(stats2), - } - // Anomaly: new denied domain - if stats2.Blocked > 0 { - entry.IsAnomaly = true - entry.AnomalyNote = "new denied domain" - anomalyCount++ - } - diff.NewDomains = append(diff.NewDomains, entry) - } else if inRun1 && !inRun2 { - // Removed domain - entry := DomainDiffEntry{ - Domain: domain, - DiffEntryBase: DiffEntryBase{Status: "removed"}, - Run1Allowed: stats1.Allowed, - Run1Blocked: stats1.Blocked, - Run1Status: classifyFirewallDomainStatus(stats1), - } - // Anomaly: the removed domain was denied in the base run. This indicates a - // transient firewall block that prevented the agent from reaching an MCP server - // (e.g. awmg-mcpg:8080) — even though the domain is absent from the comparison - // run (and therefore looks "normal"), its prior denial is worth surfacing so - // post-completion relaunch failures are detectable in audit diffs. - if stats1.Blocked > 0 { - entry.IsAnomaly = true - entry.AnomalyNote = "denied in base run — absent from comparison run" - anomalyCount++ - } - diff.RemovedDomains = append(diff.RemovedDomains, entry) - } else { - // Domain exists in both runs - check for changes - status1 := classifyFirewallDomainStatus(stats1) - status2 := classifyFirewallDomainStatus(stats2) - - if status1 != status2 { - // Status changed - entry := DomainDiffEntry{ - Domain: domain, - DiffEntryBase: DiffEntryBase{Status: "status_changed"}, - Run1Allowed: stats1.Allowed, - Run1Blocked: stats1.Blocked, - Run2Allowed: stats2.Allowed, - Run2Blocked: stats2.Blocked, - Run1Status: status1, - Run2Status: status2, - } - // Anomaly: previously denied, now allowed - if status1 == "denied" && status2 == "allowed" { - entry.IsAnomaly = true - entry.AnomalyNote = "previously denied, now allowed" - anomalyCount++ - } - // Anomaly: previously allowed, now denied - if status1 == "allowed" && status2 == "denied" { - entry.IsAnomaly = true - entry.AnomalyNote = "previously allowed, now denied" - anomalyCount++ - } - diff.StatusChanges = append(diff.StatusChanges, entry) - } else { - // Check for significant volume changes (>100% threshold) - total1 := stats1.Allowed + stats1.Blocked - total2 := stats2.Allowed + stats2.Blocked - - if total1 > 0 { - pctChange := (float64(total2-total1) / float64(total1)) * 100 - if math.Abs(pctChange) > volumeChangeThresholdPercent { - entry := DomainDiffEntry{ - Domain: domain, - DiffEntryBase: DiffEntryBase{Status: "volume_changed"}, - Run1Allowed: stats1.Allowed, - Run1Blocked: stats1.Blocked, - Run2Allowed: stats2.Allowed, - Run2Blocked: stats2.Blocked, - Run1Status: status1, - Run2Status: status2, - VolumeChange: formatVolumeChange(total1, total2), - } - diff.VolumeChanges = append(diff.VolumeChanges, entry) - } - } - } - } +func processSharedFirewallDomain(ctx *firewallDiffContext, domain string, stats1, stats2 DomainRequestStats) { + status1 := classifyFirewallDomainStatus(stats1) + status2 := classifyFirewallDomainStatus(stats2) + if maybeAppendFirewallStatusChange(ctx, domain, stats1, stats2, status1, status2) { + return } + maybeAppendFirewallVolumeChange(ctx, domain, stats1, stats2, status1, status2) +} - diff.Summary = FirewallDiffSummary{ - NewDomainCount: len(diff.NewDomains), - RemovedDomainCount: len(diff.RemovedDomains), - StatusChangeCount: len(diff.StatusChanges), - VolumeChangeCount: len(diff.VolumeChanges), - HasAnomalies: anomalyCount > 0, - AnomalyCount: anomalyCount, +func maybeAppendFirewallStatusChange(ctx *firewallDiffContext, domain string, stats1, stats2 DomainRequestStats, status1, status2 string) bool { + if status1 == status2 { + return false + } + entry := DomainDiffEntry{ + Domain: domain, + DiffEntryBase: DiffEntryBase{Status: "status_changed"}, + Run1Allowed: stats1.Allowed, + Run1Blocked: stats1.Blocked, + Run2Allowed: stats2.Allowed, + Run2Blocked: stats2.Blocked, + Run1Status: status1, + Run2Status: status2, + } + markFirewallStatusChangeAnomaly(ctx, &entry, status1, status2) + ctx.diff.StatusChanges = append(ctx.diff.StatusChanges, entry) + return true +} + +func markFirewallStatusChangeAnomaly(ctx *firewallDiffContext, entry *DomainDiffEntry, status1, status2 string) { + switch { + case status1 == "denied" && status2 == "allowed": + entry.IsAnomaly = true + entry.AnomalyNote = "previously denied, now allowed" + ctx.anomalyCount++ + case status1 == "allowed" && status2 == "denied": + entry.IsAnomaly = true + entry.AnomalyNote = "previously allowed, now denied" + ctx.anomalyCount++ } +} + +func maybeAppendFirewallVolumeChange(ctx *firewallDiffContext, domain string, stats1, stats2 DomainRequestStats, status1, status2 string) { + total1 := stats1.Allowed + stats1.Blocked + total2 := stats2.Allowed + stats2.Blocked + if total1 == 0 { + return + } + pctChange := (float64(total2-total1) / float64(total1)) * 100 + if math.Abs(pctChange) <= volumeChangeThresholdPercent { + return + } + ctx.diff.VolumeChanges = append(ctx.diff.VolumeChanges, DomainDiffEntry{ + Domain: domain, + DiffEntryBase: DiffEntryBase{Status: "volume_changed"}, + Run1Allowed: stats1.Allowed, + Run1Blocked: stats1.Blocked, + Run2Allowed: stats2.Allowed, + Run2Blocked: stats2.Blocked, + Run1Status: status1, + Run2Status: status2, + VolumeChange: formatVolumeChange(total1, total2), + }) +} +func finalizeFirewallDiff(ctx *firewallDiffContext) *FirewallDiff { + ctx.diff.Summary = FirewallDiffSummary{ + NewDomainCount: len(ctx.diff.NewDomains), + RemovedDomainCount: len(ctx.diff.RemovedDomains), + StatusChangeCount: len(ctx.diff.StatusChanges), + VolumeChangeCount: len(ctx.diff.VolumeChanges), + HasAnomalies: ctx.anomalyCount > 0, + AnomalyCount: ctx.anomalyCount, + } auditDiffLog.Printf("Firewall diff complete: new=%d, removed=%d, status_changes=%d, volume_changes=%d, anomalies=%d", - len(diff.NewDomains), len(diff.RemovedDomains), len(diff.StatusChanges), len(diff.VolumeChanges), anomalyCount) - return diff + len(ctx.diff.NewDomains), len(ctx.diff.RemovedDomains), len(ctx.diff.StatusChanges), len(ctx.diff.VolumeChanges), ctx.anomalyCount) + return ctx.diff } +// classifyFirewallDomainStatus returns "allowed", "denied", or "mixed" based on request stats // classifyFirewallDomainStatus returns "allowed", "denied", or "mixed" based on request stats func classifyFirewallDomainStatus(stats DomainRequestStats) string { if stats.Allowed > 0 && stats.Blocked == 0 { @@ -413,188 +427,227 @@ func mcpToolKey(serverName, toolName string) string { // computeMCPToolsDiff computes the diff between two runs' MCP tool usage. // run1 is the "before" (baseline) and run2 is the "after" (comparison target). func computeMCPToolsDiff(run1, run2 *MCPToolUsageData) *MCPToolsDiff { - run1Count, run2Count := 0, 0 - if run1 != nil { - run1Count = len(run1.Summary) + run1Tools := mcpToolSummaryMap(run1) + run2Tools := mcpToolSummaryMap(run2) + auditDiffLog.Printf("Computing MCP tools diff: run1_tools=%d, run2_tools=%d", len(run1Tools), len(run2Tools)) + ctx := &mcpToolsDiffContext{diff: &MCPToolsDiff{}} + for _, key := range collectMCPToolKeys(run1Tools, run2Tools) { + processMCPToolDiffKey(ctx, key, run1Tools, run2Tools) + } + ctx.diff.Summary = MCPToolsDiffSummary{ + NewToolCount: len(ctx.diff.NewTools), + RemovedToolCount: len(ctx.diff.RemovedTools), + ChangedToolCount: len(ctx.diff.ChangedTools), + HasAnomalies: ctx.anomalyCount > 0, + AnomalyCount: ctx.anomalyCount, + } + return ctx.diff +} + +type mcpToolsDiffContext struct { + diff *MCPToolsDiff + anomalyCount int +} + +func mcpToolSummaryMap(run *MCPToolUsageData) map[string]MCPToolSummary { + tools := make(map[string]MCPToolSummary) + if run == nil { + return tools } - if run2 != nil { - run2Count = len(run2.Summary) + for _, summary := range run.Summary { + tools[mcpToolKey(summary.ServerName, summary.ToolName)] = summary } - auditDiffLog.Printf("Computing MCP tools diff: run1_tools=%d, run2_tools=%d", run1Count, run2Count) - run1Tools := make(map[string]MCPToolSummary) - run2Tools := make(map[string]MCPToolSummary) + return tools +} - if run1 != nil { - for _, s := range run1.Summary { - run1Tools[mcpToolKey(s.ServerName, s.ToolName)] = s - } +func collectMCPToolKeys(run1Tools, run2Tools map[string]MCPToolSummary) []string { + allKeys := make(map[string]struct{}) + for key := range run1Tools { + allKeys[key] = struct{}{} } - if run2 != nil { - for _, s := range run2.Summary { - run2Tools[mcpToolKey(s.ServerName, s.ToolName)] = s - } + for key := range run2Tools { + allKeys[key] = struct{}{} } + return sliceutil.SortedKeys(allKeys) +} - allKeys := make(map[string]struct{}) - for k := range run1Tools { - allKeys[k] = struct{}{} +func processMCPToolDiffKey(ctx *mcpToolsDiffContext, key string, run1Tools, run2Tools map[string]MCPToolSummary) { + s1, inRun1 := run1Tools[key] + s2, inRun2 := run2Tools[key] + switch { + case !inRun1 && inRun2: + ctx.diff.NewTools = append(ctx.diff.NewTools, buildNewMCPToolDiffEntry(ctx, s2)) + case inRun1 && !inRun2: + ctx.diff.RemovedTools = append(ctx.diff.RemovedTools, buildRemovedMCPToolDiffEntry(s1)) + case s1.CallCount != s2.CallCount || s1.ErrorCount != s2.ErrorCount: + ctx.diff.ChangedTools = append(ctx.diff.ChangedTools, buildChangedMCPToolDiffEntry(ctx, s1, s2)) } - for k := range run2Tools { - allKeys[k] = struct{}{} - } - - sortedKeys := sliceutil.SortedKeys(allKeys) - - diff := &MCPToolsDiff{} - anomalyCount := 0 - - for _, key := range sortedKeys { - s1, inRun1 := run1Tools[key] - s2, inRun2 := run2Tools[key] - - if !inRun1 && inRun2 { - entry := MCPToolDiffEntry{ - ServerName: s2.ServerName, - ToolName: s2.ToolName, - DiffEntryBase: DiffEntryBase{Status: "new"}, - Run2CallCount: s2.CallCount, - Run2ErrorCount: s2.ErrorCount, - } - if s2.ErrorCount > 0 { - entry.IsAnomaly = true - entry.AnomalyNote = "new tool with errors" - anomalyCount++ - } - diff.NewTools = append(diff.NewTools, entry) - } else if inRun1 && !inRun2 { - diff.RemovedTools = append(diff.RemovedTools, MCPToolDiffEntry{ - ServerName: s1.ServerName, - ToolName: s1.ToolName, - DiffEntryBase: DiffEntryBase{Status: "removed"}, - Run1CallCount: s1.CallCount, - Run1ErrorCount: s1.ErrorCount, - }) - } else if s1.CallCount != s2.CallCount || s1.ErrorCount != s2.ErrorCount { - entry := MCPToolDiffEntry{ - ServerName: s1.ServerName, - ToolName: s1.ToolName, - DiffEntryBase: DiffEntryBase{Status: "changed"}, - Run1CallCount: s1.CallCount, - Run2CallCount: s2.CallCount, - Run1ErrorCount: s1.ErrorCount, - Run2ErrorCount: s2.ErrorCount, - CallCountChange: formatCountChange(s1.CallCount, s2.CallCount), - } - if s2.ErrorCount > s1.ErrorCount { - entry.IsAnomaly = true - entry.AnomalyNote = "error count increased" - anomalyCount++ - } - diff.ChangedTools = append(diff.ChangedTools, entry) - } +} + +func buildNewMCPToolDiffEntry(ctx *mcpToolsDiffContext, summary MCPToolSummary) MCPToolDiffEntry { + entry := MCPToolDiffEntry{ + ServerName: summary.ServerName, + ToolName: summary.ToolName, + DiffEntryBase: DiffEntryBase{Status: "new"}, + Run2CallCount: summary.CallCount, + Run2ErrorCount: summary.ErrorCount, } + if summary.ErrorCount > 0 { + entry.IsAnomaly = true + entry.AnomalyNote = "new tool with errors" + ctx.anomalyCount++ + } + return entry +} - diff.Summary = MCPToolsDiffSummary{ - NewToolCount: len(diff.NewTools), - RemovedToolCount: len(diff.RemovedTools), - ChangedToolCount: len(diff.ChangedTools), - HasAnomalies: anomalyCount > 0, - AnomalyCount: anomalyCount, +func buildRemovedMCPToolDiffEntry(summary MCPToolSummary) MCPToolDiffEntry { + return MCPToolDiffEntry{ + ServerName: summary.ServerName, + ToolName: summary.ToolName, + DiffEntryBase: DiffEntryBase{Status: "removed"}, + Run1CallCount: summary.CallCount, + Run1ErrorCount: summary.ErrorCount, } +} - return diff +func buildChangedMCPToolDiffEntry(ctx *mcpToolsDiffContext, before, after MCPToolSummary) MCPToolDiffEntry { + entry := MCPToolDiffEntry{ + ServerName: before.ServerName, + ToolName: before.ToolName, + DiffEntryBase: DiffEntryBase{Status: "changed"}, + Run1CallCount: before.CallCount, + Run2CallCount: after.CallCount, + Run1ErrorCount: before.ErrorCount, + Run2ErrorCount: after.ErrorCount, + CallCountChange: formatCountChange(before.CallCount, after.CallCount), + } + if after.ErrorCount > before.ErrorCount { + entry.IsAnomaly = true + entry.AnomalyNote = "error count increased" + ctx.anomalyCount++ + } + return entry } +// computeRunMetricsDiff computes the diff of run-level metrics between two runs. // computeRunMetricsDiff computes the diff of run-level metrics between two runs. // Returns nil if no meaningful metrics data is available. func computeRunMetricsDiff(summary1, summary2 *RunSummary) *RunMetricsDiff { - var run1Tokens, run2Tokens int - var run1Duration, run2Duration time.Duration - var run1Turns, run2Turns int - var tu1, tu2 *TokenUsageSummary - var rl1, rl2 *GitHubRateLimitUsage - var m1, m2 *LogMetrics + inputs := extractRunMetricsInputs(summary1, summary2) + if !hasRunMetricsData(inputs) { + return nil + } + diff := newRunMetricsDiff(inputs) + applyRunMetricsTokenUsage(diff) + applyRunMetricsDuration(diff, inputs) + applyRunMetricsTokensPerTurn(diff) + diff.TokenUsageDetails = computeTokenUsageDiff(inputs.tokenUsage1, inputs.tokenUsage2) + diff.GitHubRateLimitDetails = computeGitHubRateLimitDiff(inputs.rateLimit1, inputs.rateLimit2) + diff.ToolCallsDiff = computeToolCallsDiff(inputs.metrics1, inputs.metrics2) + auditDiffLog.Printf("Run metrics diff: tokens %d->%d, turns %d->%d, has_token_details=%t, has_rate_limit_details=%t", + inputs.run1Tokens, inputs.run2Tokens, inputs.run1Turns, inputs.run2Turns, inputs.tokenUsage1 != nil || inputs.tokenUsage2 != nil, inputs.rateLimit1 != nil || inputs.rateLimit2 != nil) + return diff +} + +type runMetricsInputs struct { + run1Tokens int + run2Tokens int + run1Turns int + run2Turns int + run1Duration time.Duration + run2Duration time.Duration + tokenUsage1 *TokenUsageSummary + tokenUsage2 *TokenUsageSummary + rateLimit1 *GitHubRateLimitUsage + rateLimit2 *GitHubRateLimitUsage + metrics1 *LogMetrics + metrics2 *LogMetrics +} +func extractRunMetricsInputs(summary1, summary2 *RunSummary) runMetricsInputs { + inputs := runMetricsInputs{} if summary1 != nil { - run1Tokens = summary1.Run.TokenUsage - run1Duration = summary1.Run.Duration - // Run.Turns may be zero on cached-summary paths; Metrics.Turns is authoritative. - run1Turns = summary1.Run.Turns - if run1Turns == 0 && summary1.Metrics.Turns > 0 { - run1Turns = summary1.Metrics.Turns - } - tu1 = summary1.TokenUsage - rl1 = summary1.GitHubRateLimitUsage - m1 = &summary1.Metrics + inputs.run1Tokens = summary1.Run.TokenUsage + inputs.run1Duration = summary1.Run.Duration + inputs.run1Turns = runSummaryTurnCount(summary1) + inputs.tokenUsage1 = summary1.TokenUsage + inputs.rateLimit1 = summary1.GitHubRateLimitUsage + inputs.metrics1 = &summary1.Metrics } if summary2 != nil { - run2Tokens = summary2.Run.TokenUsage - run2Duration = summary2.Run.Duration - // Run.Turns may be zero on cached-summary paths; Metrics.Turns is authoritative. - run2Turns = summary2.Run.Turns - if run2Turns == 0 && summary2.Metrics.Turns > 0 { - run2Turns = summary2.Metrics.Turns - } - tu2 = summary2.TokenUsage - rl2 = summary2.GitHubRateLimitUsage - m2 = &summary2.Metrics - } + inputs.run2Tokens = summary2.Run.TokenUsage + inputs.run2Duration = summary2.Run.Duration + inputs.run2Turns = runSummaryTurnCount(summary2) + inputs.tokenUsage2 = summary2.TokenUsage + inputs.rateLimit2 = summary2.GitHubRateLimitUsage + inputs.metrics2 = &summary2.Metrics + } + return inputs +} - // Skip if there is no meaningful data - hasTokenDetails := tu1 != nil || tu2 != nil - hasRateLimitDetails := rl1 != nil || rl2 != nil - if run1Tokens == 0 && run2Tokens == 0 && run1Duration == 0 && run2Duration == 0 && run1Turns == 0 && run2Turns == 0 && !hasTokenDetails && !hasRateLimitDetails { - return nil +func runSummaryTurnCount(summary *RunSummary) int { + turns := summary.Run.Turns + if turns == 0 && summary.Metrics.Turns > 0 { + return summary.Metrics.Turns } + return turns +} + +func hasRunMetricsData(inputs runMetricsInputs) bool { + return !(inputs.run1Tokens == 0 && inputs.run2Tokens == 0 && + inputs.run1Duration == 0 && inputs.run2Duration == 0 && + inputs.run1Turns == 0 && inputs.run2Turns == 0 && + inputs.tokenUsage1 == nil && inputs.tokenUsage2 == nil && + inputs.rateLimit1 == nil && inputs.rateLimit2 == nil) +} - diff := &RunMetricsDiff{ - Run1TokenUsage: run1Tokens, - Run2TokenUsage: run2Tokens, - Run1Turns: run1Turns, - Run2Turns: run2Turns, - TurnsChange: run2Turns - run1Turns, +func newRunMetricsDiff(inputs runMetricsInputs) *RunMetricsDiff { + return &RunMetricsDiff{ + Run1TokenUsage: inputs.run1Tokens, + Run2TokenUsage: inputs.run2Tokens, + Run1Turns: inputs.run1Turns, + Run2Turns: inputs.run2Turns, + TurnsChange: inputs.run2Turns - inputs.run1Turns, } +} - if run1Tokens > 0 || run2Tokens > 0 { - diff.TokenUsageChange = formatVolumeChange(run1Tokens, run2Tokens) +func applyRunMetricsTokenUsage(diff *RunMetricsDiff) { + if diff.Run1TokenUsage > 0 || diff.Run2TokenUsage > 0 { + diff.TokenUsageChange = formatVolumeChange(diff.Run1TokenUsage, diff.Run2TokenUsage) } +} - if run1Duration > 0 { - diff.Run1Duration = run1Duration.Round(time.Second).String() +func applyRunMetricsDuration(diff *RunMetricsDiff, inputs runMetricsInputs) { + if inputs.run1Duration > 0 { + diff.Run1Duration = inputs.run1Duration.Round(time.Second).String() } - if run2Duration > 0 { - diff.Run2Duration = run2Duration.Round(time.Second).String() + if inputs.run2Duration > 0 { + diff.Run2Duration = inputs.run2Duration.Round(time.Second).String() } - if run1Duration > 0 && run2Duration > 0 { - delta := run2Duration - run1Duration - if delta >= 0 { - diff.DurationChange = "+" + delta.Round(time.Second).String() - } else { - diff.DurationChange = delta.Round(time.Second).String() - } + if inputs.run1Duration == 0 || inputs.run2Duration == 0 { + return + } + delta := inputs.run2Duration - inputs.run1Duration + if delta >= 0 { + diff.DurationChange = "+" + delta.Round(time.Second).String() + return } + diff.DurationChange = delta.Round(time.Second).String() +} - // Compute tokens per turn using engine-level token usage. - run1PerTurn := run1Tokens - run2PerTurn := run2Tokens - if run1Turns > 0 { - diff.Run1TokensPerTurn = run1PerTurn / run1Turns +func applyRunMetricsTokensPerTurn(diff *RunMetricsDiff) { + if diff.Run1Turns > 0 { + diff.Run1TokensPerTurn = diff.Run1TokenUsage / diff.Run1Turns } - if run2Turns > 0 { - diff.Run2TokensPerTurn = run2PerTurn / run2Turns + if diff.Run2Turns > 0 { + diff.Run2TokensPerTurn = diff.Run2TokenUsage / diff.Run2Turns } if diff.Run1TokensPerTurn > 0 || diff.Run2TokensPerTurn > 0 { diff.TokensPerTurnChange = formatVolumeChange(diff.Run1TokensPerTurn, diff.Run2TokensPerTurn) } - - diff.TokenUsageDetails = computeTokenUsageDiff(tu1, tu2) - diff.GitHubRateLimitDetails = computeGitHubRateLimitDiff(rl1, rl2) - diff.ToolCallsDiff = computeToolCallsDiff(m1, m2) - - auditDiffLog.Printf("Run metrics diff: tokens %d->%d, turns %d->%d, has_token_details=%t, has_rate_limit_details=%t", run1Tokens, run2Tokens, run1Turns, run2Turns, hasTokenDetails, hasRateLimitDetails) - return diff } +// isBashTool returns true if the tool name represents a bash/shell invocation. // isBashTool returns true if the tool name represents a bash/shell invocation. // It matches the generic "bash" / "Bash" tool names used by most engines and the // per-command "bash_*" entries generated by the Codex log parser. @@ -606,139 +659,172 @@ func isBashTool(name string) bool { // computeToolCallsDiff diffs engine-level tool calls from two LogMetrics values. // Returns nil when both metrics have no tool call data. func computeToolCallsDiff(m1, m2 *LogMetrics) *ToolCallsDiff { - run1Tools := make(map[string]ToolCallInfo) - run2Tools := make(map[string]ToolCallInfo) - - // aggregateToolCall merges a tool call entry into the map, summing call counts and - // taking the max of size fields to handle duplicate entries across log files. - aggregateToolCall := func(tools map[string]ToolCallInfo, tc ToolCallInfo) { - if existing, ok := tools[tc.Name]; ok { - existing.CallCount += tc.CallCount - if tc.MaxInputSize > existing.MaxInputSize { - existing.MaxInputSize = tc.MaxInputSize - } - if tc.MaxOutputSize > existing.MaxOutputSize { - existing.MaxOutputSize = tc.MaxOutputSize - } - if tc.MaxDuration > existing.MaxDuration { - existing.MaxDuration = tc.MaxDuration - } - tools[tc.Name] = existing - return - } - tools[tc.Name] = tc + run1Tools := aggregateToolCallMap(m1) + run2Tools := aggregateToolCallMap(m2) + if len(run1Tools) == 0 && len(run2Tools) == 0 { + return nil } + state := newToolCallDiffState(run1Tools, run2Tools) + for _, name := range collectToolCallNames(run1Tools, run2Tools) { + processToolCallDiffEntry(state, name, run1Tools, run2Tools) + } + finalizeToolCallDiffState(state) + auditDiffLog.Printf("Tool calls diff: new=%d, removed=%d, changed=%d, run1_total=%d, run2_total=%d", + len(state.diff.NewTools), len(state.diff.RemovedTools), len(state.diff.ChangedTools), state.run1Total, state.run2Total) + return state.diff +} - if m1 != nil { - for _, tc := range m1.ToolCalls { - aggregateToolCall(run1Tools, tc) - } +type toolCallDiffState struct { + diff *ToolCallsDiff + run1Total int + run2Total int + bashRun1 map[string]ToolCallInfo + bashRun2 map[string]ToolCallInfo +} + +func aggregateToolCallMap(metrics *LogMetrics) map[string]ToolCallInfo { + tools := make(map[string]ToolCallInfo) + if metrics == nil { + return tools + } + for _, call := range metrics.ToolCalls { + aggregateToolCallInfo(tools, call) } - if m2 != nil { - for _, tc := range m2.ToolCalls { - aggregateToolCall(run2Tools, tc) + return tools +} + +func aggregateToolCallInfo(tools map[string]ToolCallInfo, call ToolCallInfo) { + if existing, ok := tools[call.Name]; ok { + existing.CallCount += call.CallCount + if call.MaxInputSize > existing.MaxInputSize { + existing.MaxInputSize = call.MaxInputSize } + if call.MaxOutputSize > existing.MaxOutputSize { + existing.MaxOutputSize = call.MaxOutputSize + } + if call.MaxDuration > existing.MaxDuration { + existing.MaxDuration = call.MaxDuration + } + tools[call.Name] = existing + return } + tools[call.Name] = call +} - if len(run1Tools) == 0 && len(run2Tools) == 0 { - return nil +func newToolCallDiffState(run1Tools, run2Tools map[string]ToolCallInfo) *toolCallDiffState { + return &toolCallDiffState{ + diff: &ToolCallsDiff{}, + bashRun1: make(map[string]ToolCallInfo), + bashRun2: make(map[string]ToolCallInfo), } +} +func collectToolCallNames(run1Tools, run2Tools map[string]ToolCallInfo) []string { allNames := make(map[string]struct{}) - for k := range run1Tools { - allNames[k] = struct{}{} + for name := range run1Tools { + allNames[name] = struct{}{} } - for k := range run2Tools { - allNames[k] = struct{}{} + for name := range run2Tools { + allNames[name] = struct{}{} } + return sliceutil.SortedKeys(allNames) +} - sortedNames := sliceutil.SortedKeys(allNames) - - diff := &ToolCallsDiff{} - var run1Total, run2Total int - // Collect bash tools during the main iteration to avoid a second traversal in computeBashCommandsDiff. - bashRun1 := make(map[string]ToolCallInfo) - bashRun2 := make(map[string]ToolCallInfo) +func processToolCallDiffEntry(state *toolCallDiffState, name string, run1Tools, run2Tools map[string]ToolCallInfo) { + tc1, inRun1 := run1Tools[name] + tc2, inRun2 := run2Tools[name] + trackToolCallTotals(state, name, tc1, inRun1, tc2, inRun2) + entry := buildToolCallDiffEntry(name, tc1, inRun1, tc2, inRun2) + appendToolCallDiffEntry(state.diff, entry) + state.diff.AllTools = append(state.diff.AllTools, entry) +} - for _, name := range sortedNames { - tc1, inRun1 := run1Tools[name] - tc2, inRun2 := run2Tools[name] - - if inRun1 { - run1Total += tc1.CallCount - if isBashTool(name) { - bashRun1[name] = tc1 - } +func trackToolCallTotals(state *toolCallDiffState, name string, tc1 ToolCallInfo, inRun1 bool, tc2 ToolCallInfo, inRun2 bool) { + if inRun1 { + state.run1Total += tc1.CallCount + if isBashTool(name) { + state.bashRun1[name] = tc1 } - if inRun2 { - run2Total += tc2.CallCount - if isBashTool(name) { - bashRun2[name] = tc2 - } + } + if inRun2 { + state.run2Total += tc2.CallCount + if isBashTool(name) { + state.bashRun2[name] = tc2 } + } +} - var entry ToolCallDiffEntry - switch { - case !inRun1 && inRun2: - entry = ToolCallDiffEntry{ - Name: name, - DiffEntryBase: DiffEntryBase{Status: "new"}, - Run2CallCount: tc2.CallCount, - Run2MaxInputSize: tc2.MaxInputSize, - Run2MaxOutputSize: tc2.MaxOutputSize, - } - diff.NewTools = append(diff.NewTools, entry) - case inRun1 && !inRun2: - entry = ToolCallDiffEntry{ - Name: name, - DiffEntryBase: DiffEntryBase{Status: "removed"}, - Run1CallCount: tc1.CallCount, - Run1MaxInputSize: tc1.MaxInputSize, - Run1MaxOutputSize: tc1.MaxOutputSize, - } - diff.RemovedTools = append(diff.RemovedTools, entry) - case tc1.CallCount != tc2.CallCount: - entry = ToolCallDiffEntry{ - Name: name, - DiffEntryBase: DiffEntryBase{Status: "changed"}, - Run1CallCount: tc1.CallCount, - Run2CallCount: tc2.CallCount, - CallCountChange: formatCountChange(tc1.CallCount, tc2.CallCount), - Run1MaxInputSize: tc1.MaxInputSize, - Run2MaxInputSize: tc2.MaxInputSize, - Run1MaxOutputSize: tc1.MaxOutputSize, - Run2MaxOutputSize: tc2.MaxOutputSize, - } - diff.ChangedTools = append(diff.ChangedTools, entry) - default: - entry = ToolCallDiffEntry{ - Name: name, - DiffEntryBase: DiffEntryBase{Status: "unchanged"}, - Run1CallCount: tc1.CallCount, - Run2CallCount: tc2.CallCount, - Run1MaxInputSize: tc1.MaxInputSize, - Run2MaxInputSize: tc2.MaxInputSize, - Run1MaxOutputSize: tc1.MaxOutputSize, - Run2MaxOutputSize: tc2.MaxOutputSize, - } - } - diff.AllTools = append(diff.AllTools, entry) +func buildToolCallDiffEntry(name string, tc1 ToolCallInfo, inRun1 bool, tc2 ToolCallInfo, inRun2 bool) ToolCallDiffEntry { + entry := ToolCallDiffEntry{Name: name} + switch { + case !inRun1 && inRun2: + entry.DiffEntryBase = DiffEntryBase{Status: "new"} + entry.Run2CallCount = tc2.CallCount + entry.Run2MaxInputSize = tc2.MaxInputSize + entry.Run2MaxOutputSize = tc2.MaxOutputSize + case inRun1 && !inRun2: + entry.DiffEntryBase = DiffEntryBase{Status: "removed"} + entry.Run1CallCount = tc1.CallCount + entry.Run1MaxInputSize = tc1.MaxInputSize + entry.Run1MaxOutputSize = tc1.MaxOutputSize + case tc1.CallCount != tc2.CallCount: + entry = buildChangedToolCallDiffEntry(name, tc1, tc2) + default: + entry = buildUnchangedToolCallDiffEntry(name, tc1, tc2) + } + return entry +} + +func buildChangedToolCallDiffEntry(name string, before, after ToolCallInfo) ToolCallDiffEntry { + return ToolCallDiffEntry{ + Name: name, + DiffEntryBase: DiffEntryBase{Status: "changed"}, + Run1CallCount: before.CallCount, + Run2CallCount: after.CallCount, + CallCountChange: formatCountChange(before.CallCount, after.CallCount), + Run1MaxInputSize: before.MaxInputSize, + Run2MaxInputSize: after.MaxInputSize, + Run1MaxOutputSize: before.MaxOutputSize, + Run2MaxOutputSize: after.MaxOutputSize, } +} - diff.BashDiff = computeBashCommandsDiff(bashRun1, bashRun2) - diff.Summary = ToolCallsDiffSummary{ - NewToolCount: len(diff.NewTools), - RemovedToolCount: len(diff.RemovedTools), - ChangedToolCount: len(diff.ChangedTools), - Run1TotalCalls: run1Total, - Run2TotalCalls: run2Total, +func buildUnchangedToolCallDiffEntry(name string, before, after ToolCallInfo) ToolCallDiffEntry { + return ToolCallDiffEntry{ + Name: name, + DiffEntryBase: DiffEntryBase{Status: "unchanged"}, + Run1CallCount: before.CallCount, + Run2CallCount: after.CallCount, + Run1MaxInputSize: before.MaxInputSize, + Run2MaxInputSize: after.MaxInputSize, + Run1MaxOutputSize: before.MaxOutputSize, + Run2MaxOutputSize: after.MaxOutputSize, } +} - auditDiffLog.Printf("Tool calls diff: new=%d, removed=%d, changed=%d, run1_total=%d, run2_total=%d", - len(diff.NewTools), len(diff.RemovedTools), len(diff.ChangedTools), run1Total, run2Total) - return diff +func appendToolCallDiffEntry(diff *ToolCallsDiff, entry ToolCallDiffEntry) { + switch entry.Status { + case "new": + diff.NewTools = append(diff.NewTools, entry) + case "removed": + diff.RemovedTools = append(diff.RemovedTools, entry) + case "changed": + diff.ChangedTools = append(diff.ChangedTools, entry) + } } +func finalizeToolCallDiffState(state *toolCallDiffState) { + state.diff.BashDiff = computeBashCommandsDiff(state.bashRun1, state.bashRun2) + state.diff.Summary = ToolCallsDiffSummary{ + NewToolCount: len(state.diff.NewTools), + RemovedToolCount: len(state.diff.RemovedTools), + ChangedToolCount: len(state.diff.ChangedTools), + Run1TotalCalls: state.run1Total, + Run2TotalCalls: state.run2Total, + } +} + +// computeBashCommandsDiff builds bash-specific analysis from pre-filtered bash tool call maps. // computeBashCommandsDiff builds bash-specific analysis from pre-filtered bash tool call maps. // The maps should contain only bash-related entries (generic "bash"/"Bash" and per-command "bash_*"). // Returns nil when no bash tool calls are present in either map. @@ -853,78 +939,92 @@ func computeTokenUsageDiff(tu1, tu2 *TokenUsageSummary) *TokenUsageDiff { if tu1 == nil && tu2 == nil { return nil } + inputs := extractTokenUsageInputs(tu1, tu2) + diff := &TokenUsageDiff{ + Run1InputTokens: inputs.run1Input, + Run2InputTokens: inputs.run2Input, + Run1OutputTokens: inputs.run1Output, + Run2OutputTokens: inputs.run2Output, + Run1CacheReadTokens: inputs.run1CacheRead, + Run2CacheReadTokens: inputs.run2CacheRead, + Run1CacheWriteTokens: inputs.run1CacheWrite, + Run2CacheWriteTokens: inputs.run2CacheWrite, + Run1AIC: inputs.run1AIC, + Run2AIC: inputs.run2AIC, + Run1TotalRequests: inputs.run1Requests, + Run2TotalRequests: inputs.run2Requests, + Run1CacheEfficiency: inputs.run1CacheEff, + Run2CacheEfficiency: inputs.run2CacheEff, + } + applyTokenUsageChanges(diff) + return diff +} - var ( - run1Input, run2Input int - run1Output, run2Output int - run1CacheRead, run2CacheRead int - run1CacheWrite, run2CacheWrite int - run1AIC, run2AIC float64 - run1Requests, run2Requests int - run1CacheEff, run2CacheEff float64 - ) +type tokenUsageInputs struct { + run1Input int + run2Input int + run1Output int + run2Output int + run1CacheRead int + run2CacheRead int + run1CacheWrite int + run2CacheWrite int + run1AIC float64 + run2AIC float64 + run1Requests int + run2Requests int + run1CacheEff float64 + run2CacheEff float64 +} +func extractTokenUsageInputs(tu1, tu2 *TokenUsageSummary) tokenUsageInputs { + inputs := tokenUsageInputs{} if tu1 != nil { - run1Input = tu1.TotalInputTokens - run1Output = tu1.TotalOutputTokens - run1CacheRead = tu1.TotalCacheReadTokens - run1CacheWrite = tu1.TotalCacheWriteTokens - run1AIC = tu1.TotalAIC - run1Requests = tu1.TotalRequests - run1CacheEff = tu1.CacheEfficiency + inputs.run1Input = tu1.TotalInputTokens + inputs.run1Output = tu1.TotalOutputTokens + inputs.run1CacheRead = tu1.TotalCacheReadTokens + inputs.run1CacheWrite = tu1.TotalCacheWriteTokens + inputs.run1AIC = tu1.TotalAIC + inputs.run1Requests = tu1.TotalRequests + inputs.run1CacheEff = tu1.CacheEfficiency } if tu2 != nil { - run2Input = tu2.TotalInputTokens - run2Output = tu2.TotalOutputTokens - run2CacheRead = tu2.TotalCacheReadTokens - run2CacheWrite = tu2.TotalCacheWriteTokens - run2AIC = tu2.TotalAIC - run2Requests = tu2.TotalRequests - run2CacheEff = tu2.CacheEfficiency - } - - diff := &TokenUsageDiff{ - Run1InputTokens: run1Input, - Run2InputTokens: run2Input, - Run1OutputTokens: run1Output, - Run2OutputTokens: run2Output, - Run1CacheReadTokens: run1CacheRead, - Run2CacheReadTokens: run2CacheRead, - Run1CacheWriteTokens: run1CacheWrite, - Run2CacheWriteTokens: run2CacheWrite, - Run1AIC: run1AIC, - Run2AIC: run2AIC, - Run1TotalRequests: run1Requests, - Run2TotalRequests: run2Requests, - Run1CacheEfficiency: run1CacheEff, - Run2CacheEfficiency: run2CacheEff, - } + inputs.run2Input = tu2.TotalInputTokens + inputs.run2Output = tu2.TotalOutputTokens + inputs.run2CacheRead = tu2.TotalCacheReadTokens + inputs.run2CacheWrite = tu2.TotalCacheWriteTokens + inputs.run2AIC = tu2.TotalAIC + inputs.run2Requests = tu2.TotalRequests + inputs.run2CacheEff = tu2.CacheEfficiency + } + return inputs +} - if run1Input > 0 || run2Input > 0 { - diff.InputTokensChange = formatVolumeChange(run1Input, run2Input) +func applyTokenUsageChanges(diff *TokenUsageDiff) { + if diff.Run1InputTokens > 0 || diff.Run2InputTokens > 0 { + diff.InputTokensChange = formatVolumeChange(diff.Run1InputTokens, diff.Run2InputTokens) } - if run1Output > 0 || run2Output > 0 { - diff.OutputTokensChange = formatVolumeChange(run1Output, run2Output) + if diff.Run1OutputTokens > 0 || diff.Run2OutputTokens > 0 { + diff.OutputTokensChange = formatVolumeChange(diff.Run1OutputTokens, diff.Run2OutputTokens) } - if run1CacheRead > 0 || run2CacheRead > 0 { - diff.CacheReadTokensChange = formatVolumeChange(run1CacheRead, run2CacheRead) + if diff.Run1CacheReadTokens > 0 || diff.Run2CacheReadTokens > 0 { + diff.CacheReadTokensChange = formatVolumeChange(diff.Run1CacheReadTokens, diff.Run2CacheReadTokens) } - if run1CacheWrite > 0 || run2CacheWrite > 0 { - diff.CacheWriteTokensChange = formatVolumeChange(run1CacheWrite, run2CacheWrite) + if diff.Run1CacheWriteTokens > 0 || diff.Run2CacheWriteTokens > 0 { + diff.CacheWriteTokensChange = formatVolumeChange(diff.Run1CacheWriteTokens, diff.Run2CacheWriteTokens) } - if run1AIC > 0 || run2AIC > 0 { - diff.AICChange = formatFloatDelta(run1AIC, run2AIC) + if diff.Run1AIC > 0 || diff.Run2AIC > 0 { + diff.AICChange = formatFloatDelta(diff.Run1AIC, diff.Run2AIC) } - if run1Requests > 0 || run2Requests > 0 { - diff.RequestsDelta = formatCountChange(run1Requests, run2Requests) + if diff.Run1TotalRequests > 0 || diff.Run2TotalRequests > 0 { + diff.RequestsDelta = formatCountChange(diff.Run1TotalRequests, diff.Run2TotalRequests) } - if run1CacheEff > 0 || run2CacheEff > 0 { - diff.CacheEfficiencyChange = formatPercentagePointChange(run1CacheEff, run2CacheEff) + if diff.Run1CacheEfficiency > 0 || diff.Run2CacheEfficiency > 0 { + diff.CacheEfficiencyChange = formatPercentagePointChange(diff.Run1CacheEfficiency, diff.Run2CacheEfficiency) } - - return diff } +// loadRunSummaryForDiff loads or builds a RunSummary for a given run for use in diffing. // loadRunSummaryForDiff loads or builds a RunSummary for a given run for use in diffing. // It first tries to load from a cached RunSummary (which includes MCP tool usage and run // metrics); otherwise it downloads artifacts and analyzes firewall logs, returning a partial diff --git a/pkg/cli/audit_diff_command.go b/pkg/cli/audit_diff_command.go index b36ed5e9024..27fffb72ed5 100644 --- a/pkg/cli/audit_diff_command.go +++ b/pkg/cli/audit_diff_command.go @@ -16,7 +16,17 @@ import ( // Deprecated: pass multiple run IDs directly to `audit` instead (e.g. `gh aw audit `). // This subcommand is hidden and kept for backward compatibility only. func NewAuditDiffSubcommand() *cobra.Command { - cmd := &cobra.Command{ + cmd := newAuditDiffSubcommandDefinition() + addOutputFlag(cmd, defaultLogsOutputDir) + addJSONFlag(cmd) + addRepoFlag(cmd) + cmd.Flags().String("format", "pretty", "Output format: pretty, markdown") + cmd.Flags().StringSlice("artifacts", nil, "Artifact sets to download (default: all, because auditing requires comprehensive artifacts for analysis). Valid sets: "+strings.Join(ValidArtifactSetNames(), ", ")) + return cmd +} + +func newAuditDiffSubcommandDefinition() *cobra.Command { + return &cobra.Command{ Use: "diff ...", Short: "[Deprecated] Compare workflow runs (use: gh aw audit )", Hidden: true, @@ -47,168 +57,215 @@ analyzes their data, and produces a diff showing: ` + string(constants.CLIExtensionPrefix) + ` audit diff 12345 12346 --json # JSON for CI integration ` + string(constants.CLIExtensionPrefix) + ` audit diff 12345 12346 --repo owner/repo # Specify repository`, Args: cobra.MinimumNArgs(2), - RunE: func(cmd *cobra.Command, args []string) error { - baseRunID, err := strconv.ParseInt(args[0], 10, 64) - if err != nil { - return fmt.Errorf("invalid base run ID %q: must be a numeric run ID", args[0]) - } - - compareRunIDs := make([]int64, 0, len(args)-1) - seen := make(map[int64]bool) - for _, arg := range args[1:] { - id, err := strconv.ParseInt(arg, 10, 64) - if err != nil { - return fmt.Errorf("invalid run ID %q: must be a numeric run ID", arg) - } - if id == baseRunID { - return fmt.Errorf("comparison run ID %d is the same as the base run ID: cannot diff a run against itself", id) - } - if seen[id] { - return fmt.Errorf("duplicate comparison run ID %d: each run ID must appear only once", id) - } - seen[id] = true - compareRunIDs = append(compareRunIDs, id) - } - - outputDir, _ := cmd.Flags().GetString("output") - verbose, _ := cmd.Flags().GetBool("verbose") - jsonOutput, _ := cmd.Flags().GetBool("json") - format, _ := cmd.Flags().GetString("format") - repoFlag, _ := cmd.Flags().GetString("repo") - artifacts, _ := cmd.Flags().GetStringSlice("artifacts") - - var owner, repo, hostname string - if repoFlag != "" { - parts := strings.SplitN(repoFlag, "/", 2) - if len(parts) != 2 || parts[0] == "" || parts[1] == "" { - return fmt.Errorf("invalid repository format '%s': expected 'owner/repo'", repoFlag) - } - owner = parts[0] - repo = parts[1] - } - - return RunAuditDiff(cmd.Context(), baseRunID, compareRunIDs, AuditOptions{ - Owner: owner, - Repo: repo, - Hostname: hostname, - OutputDir: outputDir, - Verbose: verbose, - JSONOutput: jsonOutput, - Format: format, - ArtifactSets: artifacts, - }) - }, + RunE: runAuditDiffSubcommand, } +} - addOutputFlag(cmd, defaultLogsOutputDir) - addJSONFlag(cmd) - addRepoFlag(cmd) - cmd.Flags().String("format", "pretty", "Output format: pretty, markdown") - cmd.Flags().StringSlice("artifacts", nil, "Artifact sets to download (default: all, because auditing requires comprehensive artifacts for analysis). Valid sets: "+strings.Join(ValidArtifactSetNames(), ", ")) +func runAuditDiffSubcommand(cmd *cobra.Command, args []string) error { + baseRunID, compareRunIDs, err := parseAuditDiffRunIDs(args) + if err != nil { + return err + } + opts, err := auditDiffOptionsFromFlags(cmd) + if err != nil { + return err + } + return RunAuditDiff(cmd.Context(), baseRunID, compareRunIDs, opts) +} - return cmd +func parseAuditDiffRunIDs(args []string) (int64, []int64, error) { + baseRunID, err := strconv.ParseInt(args[0], 10, 64) + if err != nil { + return 0, nil, fmt.Errorf("invalid base run ID %q: must be a numeric run ID", args[0]) + } + compareRunIDs := make([]int64, 0, len(args)-1) + seen := make(map[int64]bool) + for _, arg := range args[1:] { + id, err := strconv.ParseInt(arg, 10, 64) + if err != nil { + return 0, nil, fmt.Errorf("invalid run ID %q: must be a numeric run ID", arg) + } + if err := validateAuditDiffCompareRunID(baseRunID, id, seen); err != nil { + return 0, nil, err + } + seen[id] = true + compareRunIDs = append(compareRunIDs, id) + } + return baseRunID, compareRunIDs, nil +} + +func validateAuditDiffCompareRunID(baseRunID, compareRunID int64, seen map[int64]bool) error { + if compareRunID == baseRunID { + return fmt.Errorf("comparison run ID %d is the same as the base run ID: cannot diff a run against itself", compareRunID) + } + if seen[compareRunID] { + return fmt.Errorf("duplicate comparison run ID %d: each run ID must appear only once", compareRunID) + } + return nil +} + +func auditDiffOptionsFromFlags(cmd *cobra.Command) (AuditOptions, error) { + outputDir, _ := cmd.Flags().GetString("output") + verbose, _ := cmd.Flags().GetBool("verbose") + jsonOutput, _ := cmd.Flags().GetBool("json") + format, _ := cmd.Flags().GetString("format") + artifacts, _ := cmd.Flags().GetStringSlice("artifacts") + owner, repo, err := parseAuditDiffRepoFlag(cmd) + if err != nil { + return AuditOptions{}, err + } + return AuditOptions{ + Owner: owner, + Repo: repo, + OutputDir: outputDir, + Verbose: verbose, + JSONOutput: jsonOutput, + Format: format, + ArtifactSets: artifacts, + }, nil +} + +func parseAuditDiffRepoFlag(cmd *cobra.Command) (string, string, error) { + repoFlag, _ := cmd.Flags().GetString("repo") + if repoFlag == "" { + return "", "", nil + } + parts := strings.SplitN(repoFlag, "/", 2) + if len(parts) != 2 || parts[0] == "" || parts[1] == "" { + return "", "", fmt.Errorf("invalid repository format '%s': expected 'owner/repo'", repoFlag) + } + return parts[0], parts[1], nil } // RunAuditDiff compares behavior between a base workflow run and one or more comparison runs. // The base run is the reference point; each comparison run is diffed against it independently. func RunAuditDiff(ctx context.Context, baseRunID int64, compareRunIDs []int64, opts AuditOptions) error { - owner := opts.Owner - repo := opts.Repo - hostname := opts.Hostname - outputDir := opts.OutputDir - verbose := opts.Verbose - format := opts.Format - artifactSets := opts.ArtifactSets - - auditDiffLog.Printf("Starting audit diff: base=%d, compare=%v", baseRunID, compareRunIDs) - - // Validate and resolve artifact sets into a concrete filter. - if err := ValidateArtifactSets(artifactSets); err != nil { + runtime, err := newAuditDiffRuntime(ctx, opts) + if err != nil { return err } - artifactFilter := ResolveArtifactFilter(artifactSets) - if len(artifactFilter) > 0 { - auditDiffLog.Printf("Artifact filter active: %v", artifactFilter) - if verbose { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Artifact filter: downloading only "+strings.Join(artifactFilter, ", "))) - } + printAuditDiffStart(baseRunID, compareRunIDs) + baseSummary, err := loadAuditDiffBaseSummary(ctx, baseRunID, runtime) + if err != nil { + return err + } + diffs, err := computeAuditDiffs(ctx, baseRunID, compareRunIDs, baseSummary, runtime) + if err != nil { + return err + } + return renderAuditDiffOutput(diffs, runtime.opts) +} + +type auditDiffRuntime struct { + opts AuditOptions + hostname string + artifactFilter []string +} + +func newAuditDiffRuntime(ctx context.Context, opts AuditOptions) (*auditDiffRuntime, error) { + auditDiffLog.Printf("Starting audit diff: base options=%+v", opts) + if err := ValidateArtifactSets(opts.ArtifactSets); err != nil { + return nil, err + } + runtime := &auditDiffRuntime{opts: opts, artifactFilter: ResolveArtifactFilter(opts.ArtifactSets)} + printAuditDiffArtifactFilter(runtime.artifactFilter, opts.Verbose) + runtime.hostname = resolveAuditDiffHostname(opts.Hostname) + if err := ensureAuditDiffContext(ctx); err != nil { + return nil, err } + return runtime, nil +} - // Auto-detect GHES host from git remote if hostname is not provided - if hostname == "" { - hostname = getHostFromOriginRemote() - if hostname != "github.com" { - auditDiffLog.Printf("Auto-detected GHES host from git remote: %s", hostname) - } +func printAuditDiffArtifactFilter(artifactFilter []string, verbose bool) { + if len(artifactFilter) == 0 { + return + } + auditDiffLog.Printf("Artifact filter active: %v", artifactFilter) + if verbose { + fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Artifact filter: downloading only "+strings.Join(artifactFilter, ", "))) } +} - // Check context cancellation +func resolveAuditDiffHostname(hostname string) string { + if hostname != "" { + return hostname + } + hostname = getHostFromOriginRemote() + if hostname != "github.com" { + auditDiffLog.Printf("Auto-detected GHES host from git remote: %s", hostname) + } + return hostname +} + +func ensureAuditDiffContext(ctx context.Context) error { select { case <-ctx.Done(): fmt.Fprintln(os.Stderr, console.FormatWarningMessage("Operation cancelled")) return ctx.Err() default: + return nil } +} +func printAuditDiffStart(baseRunID int64, compareRunIDs []int64) { if len(compareRunIDs) == 1 { fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Comparing workflow runs: Run #%d → Run #%d", baseRunID, compareRunIDs[0]))) - } else { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Comparing workflow runs: Run #%d (base) vs %d comparison runs", baseRunID, len(compareRunIDs)))) + return } + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Comparing workflow runs: Run #%d (base) vs %d comparison runs", baseRunID, len(compareRunIDs)))) +} - // Load base run summary once (shared across all comparisons) +func loadAuditDiffBaseSummary(ctx context.Context, baseRunID int64, runtime *auditDiffRuntime) (*RunSummary, error) { fmt.Fprintln(os.Stderr, console.FormatProgressMessage(fmt.Sprintf("Loading data for base run %d...", baseRunID))) - baseSummary, err := loadRunSummaryForDiff(ctx, baseRunID, outputDir, owner, repo, hostname, verbose, artifactFilter) + summary, err := loadRunSummaryForDiff(ctx, baseRunID, runtime.opts.OutputDir, runtime.opts.Owner, runtime.opts.Repo, runtime.hostname, runtime.opts.Verbose, runtime.artifactFilter) if err != nil { - return fmt.Errorf("failed to load data for base run %d: %w", baseRunID, err) + return nil, fmt.Errorf("failed to load data for base run %d: %w", baseRunID, err) } + return summary, nil +} +func computeAuditDiffs(ctx context.Context, baseRunID int64, compareRunIDs []int64, baseSummary *RunSummary, runtime *auditDiffRuntime) ([]*AuditDiff, error) { diffs := make([]*AuditDiff, 0, len(compareRunIDs)) - for _, compareRunID := range compareRunIDs { - // Check context cancellation between downloads - select { - case <-ctx.Done(): - fmt.Fprintln(os.Stderr, console.FormatWarningMessage("Operation cancelled")) - return ctx.Err() - default: + if err := ensureAuditDiffContext(ctx); err != nil { + return nil, err } - - fmt.Fprintln(os.Stderr, console.FormatProgressMessage(fmt.Sprintf("Loading data for run %d...", compareRunID))) - compareSummary, err := loadRunSummaryForDiff(ctx, compareRunID, outputDir, owner, repo, hostname, verbose, artifactFilter) + compareSummary, err := loadAuditDiffCompareSummary(ctx, compareRunID, runtime) if err != nil { - return fmt.Errorf("failed to load data for run %d: %w", compareRunID, err) + return nil, fmt.Errorf("failed to load data for run %d: %w", compareRunID, err) } + warnAboutAuditDiffFirewallData(baseRunID, compareRunID, baseSummary, compareSummary) + diffs = append(diffs, computeAuditDiff(baseRunID, compareRunID, baseSummary, compareSummary)) + } + return diffs, nil +} - // Warn if no firewall data found for this pair - fw1 := baseSummary.FirewallAnalysis - fw2 := compareSummary.FirewallAnalysis - if fw1 == nil && fw2 == nil { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("No firewall data found for run pair %d→%d. Both runs may predate firewall logging.", baseRunID, compareRunID))) - } else { - if fw1 == nil { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("No firewall data found for base run %d (older run may lack firewall logs)", baseRunID))) - } - if fw2 == nil { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("No firewall data found for run %d", compareRunID))) - } - } +func loadAuditDiffCompareSummary(ctx context.Context, compareRunID int64, runtime *auditDiffRuntime) (*RunSummary, error) { + fmt.Fprintln(os.Stderr, console.FormatProgressMessage(fmt.Sprintf("Loading data for run %d...", compareRunID))) + return loadRunSummaryForDiff(ctx, compareRunID, runtime.opts.OutputDir, runtime.opts.Owner, runtime.opts.Repo, runtime.hostname, runtime.opts.Verbose, runtime.artifactFilter) +} - diff := computeAuditDiff(baseRunID, compareRunID, baseSummary, compareSummary) - diffs = append(diffs, diff) +func warnAboutAuditDiffFirewallData(baseRunID, compareRunID int64, baseSummary, compareSummary *RunSummary) { + fw1 := baseSummary.FirewallAnalysis + fw2 := compareSummary.FirewallAnalysis + switch { + case fw1 == nil && fw2 == nil: + fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("No firewall data found for run pair %d→%d. Both runs may predate firewall logging.", baseRunID, compareRunID))) + case fw1 == nil: + fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("No firewall data found for base run %d (older run may lack firewall logs)", baseRunID))) + case fw2 == nil: + fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("No firewall data found for run %d", compareRunID))) } +} - // Render output - if opts.JSONOutput || format == "json" { +func renderAuditDiffOutput(diffs []*AuditDiff, opts AuditOptions) error { + switch { + case opts.JSONOutput || opts.Format == "json": return renderAuditDiffJSON(diffs) - } - - if format == "markdown" { + case opts.Format == "markdown": renderAuditDiffMarkdown(diffs) - return nil + default: + renderAuditDiffPretty(diffs) } - - // Default: pretty console output - renderAuditDiffPretty(diffs) return nil } diff --git a/pkg/cli/audit_diff_render.go b/pkg/cli/audit_diff_render.go index ce55627060e..505189101c1 100644 --- a/pkg/cli/audit_diff_render.go +++ b/pkg/cli/audit_diff_render.go @@ -76,66 +76,71 @@ func renderSingleAuditDiffPretty(diff *AuditDiff) { auditDiffRenderLog.Printf("Rendering audit diff as pretty output: run1=%d, run2=%d", diff.Run1ID, diff.Run2ID) fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Audit Diff: Run #%d → Run #%d", diff.Run1ID, diff.Run2ID))) fmt.Fprintln(os.Stderr) - if isEmptyAuditDiff(diff) { fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("No behavioral changes detected between the two runs.")) return } + renderAuditDiffPrettySummary(diff) + renderFirewallDiffPrettySection(diff.FirewallDiff) + renderMCPToolsDiffPrettySection(diff.MCPToolsDiff) + renderRunMetricsDiffPrettySection(diff.Run1ID, diff.Run2ID, diff.RunMetricsDiff) +} + +func renderAuditDiffPrettySummary(diff *AuditDiff) { + summaryParts, anomalyCount := collectAuditDiffPrettySummary(diff) + if len(summaryParts) > 0 { + fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Changes: "+strings.Join(summaryParts, " | "))) + } + if anomalyCount > 0 { + fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("⚠️ %d anomalies detected", anomalyCount))) + } + fmt.Fprintln(os.Stderr) +} - // Collect top-level summary across all sections +func collectAuditDiffPrettySummary(diff *AuditDiff) ([]string, int) { var summaryParts []string anomalyCount := 0 + summaryParts, anomalyCount = appendFirewallPrettySummary(summaryParts, anomalyCount, diff.FirewallDiff) + return appendMCPPrettySummary(summaryParts, anomalyCount, diff.MCPToolsDiff) +} - if diff.FirewallDiff != nil && !isEmptyFirewallDiff(diff.FirewallDiff) { - fwParts := []string{} - if len(diff.FirewallDiff.NewDomains) > 0 { - fwParts = append(fwParts, fmt.Sprintf("%d new domains", len(diff.FirewallDiff.NewDomains))) - } - if len(diff.FirewallDiff.RemovedDomains) > 0 { - fwParts = append(fwParts, fmt.Sprintf("%d removed domains", len(diff.FirewallDiff.RemovedDomains))) - } - if len(diff.FirewallDiff.StatusChanges) > 0 { - fwParts = append(fwParts, fmt.Sprintf("%d status changes", len(diff.FirewallDiff.StatusChanges))) - } - if len(diff.FirewallDiff.VolumeChanges) > 0 { - fwParts = append(fwParts, fmt.Sprintf("%d volume changes", len(diff.FirewallDiff.VolumeChanges))) - } - if len(fwParts) > 0 { - summaryParts = append(summaryParts, "Firewall: "+strings.Join(fwParts, ", ")) - } - anomalyCount += diff.FirewallDiff.Summary.AnomalyCount +func appendFirewallPrettySummary(summaryParts []string, anomalyCount int, diff *FirewallDiff) ([]string, int) { + if diff == nil || isEmptyFirewallDiff(diff) { + return summaryParts, anomalyCount } - - if diff.MCPToolsDiff != nil && !isEmptyMCPToolsDiff(diff.MCPToolsDiff) { - mcpParts := []string{} - if diff.MCPToolsDiff.Summary.NewToolCount > 0 { - mcpParts = append(mcpParts, fmt.Sprintf("%d new tools", diff.MCPToolsDiff.Summary.NewToolCount)) - } - if diff.MCPToolsDiff.Summary.RemovedToolCount > 0 { - mcpParts = append(mcpParts, fmt.Sprintf("%d removed tools", diff.MCPToolsDiff.Summary.RemovedToolCount)) - } - if diff.MCPToolsDiff.Summary.ChangedToolCount > 0 { - mcpParts = append(mcpParts, fmt.Sprintf("%d changed tools", diff.MCPToolsDiff.Summary.ChangedToolCount)) - } - if len(mcpParts) > 0 { - summaryParts = append(summaryParts, "MCP tools: "+strings.Join(mcpParts, ", ")) - } - anomalyCount += diff.MCPToolsDiff.Summary.AnomalyCount + parts := make([]string, 0, 4) + parts = appendCountSummary(parts, len(diff.NewDomains), "new domains") + parts = appendCountSummary(parts, len(diff.RemovedDomains), "removed domains") + parts = appendCountSummary(parts, len(diff.StatusChanges), "status changes") + parts = appendCountSummary(parts, len(diff.VolumeChanges), "volume changes") + if len(parts) > 0 { + summaryParts = append(summaryParts, "Firewall: "+strings.Join(parts, ", ")) } + return summaryParts, anomalyCount + diff.Summary.AnomalyCount +} - if len(summaryParts) > 0 { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Changes: "+strings.Join(summaryParts, " | "))) +func appendMCPPrettySummary(summaryParts []string, anomalyCount int, diff *MCPToolsDiff) ([]string, int) { + if diff == nil || isEmptyMCPToolsDiff(diff) { + return summaryParts, anomalyCount } - if anomalyCount > 0 { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("⚠️ %d anomalies detected", anomalyCount))) + parts := make([]string, 0, 3) + parts = appendCountSummary(parts, diff.Summary.NewToolCount, "new tools") + parts = appendCountSummary(parts, diff.Summary.RemovedToolCount, "removed tools") + parts = appendCountSummary(parts, diff.Summary.ChangedToolCount, "changed tools") + if len(parts) > 0 { + summaryParts = append(summaryParts, "MCP tools: "+strings.Join(parts, ", ")) } - fmt.Fprintln(os.Stderr) + return summaryParts, anomalyCount + diff.Summary.AnomalyCount +} - renderFirewallDiffPrettySection(diff.FirewallDiff) - renderMCPToolsDiffPrettySection(diff.MCPToolsDiff) - renderRunMetricsDiffPrettySection(diff.Run1ID, diff.Run2ID, diff.RunMetricsDiff) +func appendCountSummary(parts []string, count int, label string) []string { + if count > 0 { + parts = append(parts, fmt.Sprintf("%d %s", count, label)) + } + return parts } +// renderFirewallDiffMarkdownSection renders the firewall diff sub-section as markdown // renderFirewallDiffMarkdownSection renders the firewall diff sub-section as markdown func renderFirewallDiffMarkdownSection(diff *FirewallDiff) { if diff == nil || isEmptyFirewallDiff(diff) { @@ -303,287 +308,198 @@ func renderFirewallDiffPrettySection(diff *FirewallDiff) { if diff == nil || isEmptyFirewallDiff(diff) { return } - fmt.Fprintln(os.Stderr, console.FormatSectionHeader("Firewall Changes")) fmt.Fprintln(os.Stderr) + renderFirewallNewDomainsTable(diff.NewDomains) + renderFirewallRemovedDomainsTable(diff.RemovedDomains) + renderFirewallStatusChangesTable(diff.StatusChanges) + renderFirewallVolumeChangesTable(diff.VolumeChanges) +} - if len(diff.NewDomains) > 0 { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("New Domains (%d)", len(diff.NewDomains)))) - config := console.TableConfig{ - Headers: []string{"Domain", "Status", "Requests", "Anomaly"}, - Rows: make([][]string, 0, len(diff.NewDomains)), - } - for _, entry := range diff.NewDomains { - total := entry.Run2Allowed + entry.Run2Blocked - anomalyNote := formatAnomalyNote(entry.IsAnomaly, entry.AnomalyNote) - config.Rows = append(config.Rows, []string{ - entry.Domain, - firewallStatusEmoji(entry.Run2Status) + " " + entry.Run2Status, - strconv.Itoa(total), - anomalyNote, - }) - } - fmt.Fprint(os.Stderr, console.RenderTable(config)) +func renderFirewallNewDomainsTable(entries []DomainDiffEntry) { + if len(entries) == 0 { + return } + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("New Domains (%d)", len(entries)))) + config := console.TableConfig{Headers: []string{"Domain", "Status", "Requests", "Anomaly"}, Rows: make([][]string, 0, len(entries))} + for _, entry := range entries { + total := entry.Run2Allowed + entry.Run2Blocked + config.Rows = append(config.Rows, []string{entry.Domain, firewallStatusEmoji(entry.Run2Status) + " " + entry.Run2Status, strconv.Itoa(total), formatAnomalyNote(entry.IsAnomaly, entry.AnomalyNote)}) + } + fmt.Fprint(os.Stderr, console.RenderTable(config)) +} - if len(diff.RemovedDomains) > 0 { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Removed Domains (%d)", len(diff.RemovedDomains)))) - config := console.TableConfig{ - Headers: []string{"Domain", "Previous Status", "Previous Requests"}, - Rows: make([][]string, 0, len(diff.RemovedDomains)), - } - for _, entry := range diff.RemovedDomains { - total := entry.Run1Allowed + entry.Run1Blocked - config.Rows = append(config.Rows, []string{ - entry.Domain, - firewallStatusEmoji(entry.Run1Status) + " " + entry.Run1Status, - strconv.Itoa(total), - }) - } - fmt.Fprint(os.Stderr, console.RenderTable(config)) +func renderFirewallRemovedDomainsTable(entries []DomainDiffEntry) { + if len(entries) == 0 { + return } + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Removed Domains (%d)", len(entries)))) + config := console.TableConfig{Headers: []string{"Domain", "Previous Status", "Previous Requests"}, Rows: make([][]string, 0, len(entries))} + for _, entry := range entries { + total := entry.Run1Allowed + entry.Run1Blocked + config.Rows = append(config.Rows, []string{entry.Domain, firewallStatusEmoji(entry.Run1Status) + " " + entry.Run1Status, strconv.Itoa(total)}) + } + fmt.Fprint(os.Stderr, console.RenderTable(config)) +} - if len(diff.StatusChanges) > 0 { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Status Changes (%d)", len(diff.StatusChanges)))) - config := console.TableConfig{ - Headers: []string{"Domain", "Before", "After", "Anomaly"}, - Rows: make([][]string, 0, len(diff.StatusChanges)), - } - for _, entry := range diff.StatusChanges { - anomalyNote := formatAnomalyNote(entry.IsAnomaly, entry.AnomalyNote) - config.Rows = append(config.Rows, []string{ - entry.Domain, - firewallStatusEmoji(entry.Run1Status) + " " + entry.Run1Status, - firewallStatusEmoji(entry.Run2Status) + " " + entry.Run2Status, - anomalyNote, - }) - } - fmt.Fprint(os.Stderr, console.RenderTable(config)) +func renderFirewallStatusChangesTable(entries []DomainDiffEntry) { + if len(entries) == 0 { + return + } + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Status Changes (%d)", len(entries)))) + config := console.TableConfig{Headers: []string{"Domain", "Before", "After", "Anomaly"}, Rows: make([][]string, 0, len(entries))} + for _, entry := range entries { + config.Rows = append(config.Rows, []string{entry.Domain, firewallStatusEmoji(entry.Run1Status) + " " + entry.Run1Status, firewallStatusEmoji(entry.Run2Status) + " " + entry.Run2Status, formatAnomalyNote(entry.IsAnomaly, entry.AnomalyNote)}) } + fmt.Fprint(os.Stderr, console.RenderTable(config)) +} - if len(diff.VolumeChanges) > 0 { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Volume Changes")) - config := console.TableConfig{ - Headers: []string{"Domain", "Requests (before)", "Requests (after)", "Change"}, - Rows: make([][]string, 0, len(diff.VolumeChanges)), - } - for _, entry := range diff.VolumeChanges { - total1 := entry.Run1Allowed + entry.Run1Blocked - total2 := entry.Run2Allowed + entry.Run2Blocked - config.Rows = append(config.Rows, []string{ - entry.Domain, - strconv.Itoa(total1), - strconv.Itoa(total2), - entry.VolumeChange, - }) - } - fmt.Fprint(os.Stderr, console.RenderTable(config)) +func renderFirewallVolumeChangesTable(entries []DomainDiffEntry) { + if len(entries) == 0 { + return + } + fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Volume Changes")) + config := console.TableConfig{Headers: []string{"Domain", "Requests (before)", "Requests (after)", "Change"}, Rows: make([][]string, 0, len(entries))} + for _, entry := range entries { + total1 := entry.Run1Allowed + entry.Run1Blocked + total2 := entry.Run2Allowed + entry.Run2Blocked + config.Rows = append(config.Rows, []string{entry.Domain, strconv.Itoa(total1), strconv.Itoa(total2), entry.VolumeChange}) } + fmt.Fprint(os.Stderr, console.RenderTable(config)) } +// renderMCPToolsDiffPrettySection renders the MCP tools diff as a pretty console sub-section // renderMCPToolsDiffPrettySection renders the MCP tools diff as a pretty console sub-section func renderMCPToolsDiffPrettySection(diff *MCPToolsDiff) { if diff == nil || isEmptyMCPToolsDiff(diff) { return } - fmt.Fprintln(os.Stderr, console.FormatSectionHeader("MCP Tool Changes")) fmt.Fprintln(os.Stderr) + renderMCPNewToolsTable(diff.NewTools) + renderMCPRemovedToolsTable(diff.RemovedTools) + renderMCPChangedToolsTable(diff.ChangedTools) +} - if len(diff.NewTools) > 0 { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("New Tools (%d)", len(diff.NewTools)))) - config := console.TableConfig{ - Headers: []string{"Server", "Tool", "Calls", "Anomaly"}, - Rows: make([][]string, 0, len(diff.NewTools)), - } - for _, entry := range diff.NewTools { - anomalyNote := formatAnomalyNote(entry.IsAnomaly, entry.AnomalyNote) - config.Rows = append(config.Rows, []string{ - entry.ServerName, - entry.ToolName, - strconv.Itoa(entry.Run2CallCount), - anomalyNote, - }) - } - fmt.Fprint(os.Stderr, console.RenderTable(config)) +func renderMCPNewToolsTable(entries []MCPToolDiffEntry) { + if len(entries) == 0 { + return } + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("New Tools (%d)", len(entries)))) + config := console.TableConfig{Headers: []string{"Server", "Tool", "Calls", "Anomaly"}, Rows: make([][]string, 0, len(entries))} + for _, entry := range entries { + config.Rows = append(config.Rows, []string{entry.ServerName, entry.ToolName, strconv.Itoa(entry.Run2CallCount), formatAnomalyNote(entry.IsAnomaly, entry.AnomalyNote)}) + } + fmt.Fprint(os.Stderr, console.RenderTable(config)) +} - if len(diff.RemovedTools) > 0 { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Removed Tools (%d)", len(diff.RemovedTools)))) - config := console.TableConfig{ - Headers: []string{"Server", "Tool", "Previous Calls"}, - Rows: make([][]string, 0, len(diff.RemovedTools)), - } - for _, entry := range diff.RemovedTools { - config.Rows = append(config.Rows, []string{ - entry.ServerName, - entry.ToolName, - strconv.Itoa(entry.Run1CallCount), - }) - } - fmt.Fprint(os.Stderr, console.RenderTable(config)) +func renderMCPRemovedToolsTable(entries []MCPToolDiffEntry) { + if len(entries) == 0 { + return + } + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Removed Tools (%d)", len(entries)))) + config := console.TableConfig{Headers: []string{"Server", "Tool", "Previous Calls"}, Rows: make([][]string, 0, len(entries))} + for _, entry := range entries { + config.Rows = append(config.Rows, []string{entry.ServerName, entry.ToolName, strconv.Itoa(entry.Run1CallCount)}) } + fmt.Fprint(os.Stderr, console.RenderTable(config)) +} - if len(diff.ChangedTools) > 0 { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Changed Tools (%d)", len(diff.ChangedTools)))) - config := console.TableConfig{ - Headers: []string{"Server", "Tool", "Calls (before)", "Calls (after)", "Change", "Errors (before)", "Errors (after)", "Anomaly"}, - Rows: make([][]string, 0, len(diff.ChangedTools)), - } - for _, entry := range diff.ChangedTools { - anomalyNote := formatAnomalyNote(entry.IsAnomaly, entry.AnomalyNote) - config.Rows = append(config.Rows, []string{ - entry.ServerName, - entry.ToolName, - strconv.Itoa(entry.Run1CallCount), - strconv.Itoa(entry.Run2CallCount), - entry.CallCountChange, - strconv.Itoa(entry.Run1ErrorCount), - strconv.Itoa(entry.Run2ErrorCount), - anomalyNote, - }) - } - fmt.Fprint(os.Stderr, console.RenderTable(config)) +func renderMCPChangedToolsTable(entries []MCPToolDiffEntry) { + if len(entries) == 0 { + return + } + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Changed Tools (%d)", len(entries)))) + config := console.TableConfig{Headers: []string{"Server", "Tool", "Calls (before)", "Calls (after)", "Change", "Errors (before)", "Errors (after)", "Anomaly"}, Rows: make([][]string, 0, len(entries))} + for _, entry := range entries { + config.Rows = append(config.Rows, []string{entry.ServerName, entry.ToolName, strconv.Itoa(entry.Run1CallCount), strconv.Itoa(entry.Run2CallCount), entry.CallCountChange, strconv.Itoa(entry.Run1ErrorCount), strconv.Itoa(entry.Run2ErrorCount), formatAnomalyNote(entry.IsAnomaly, entry.AnomalyNote)}) } + fmt.Fprint(os.Stderr, console.RenderTable(config)) } +// renderRunMetricsDiffPrettySection renders the run metrics diff as a pretty console sub-section // renderRunMetricsDiffPrettySection renders the run metrics diff as a pretty console sub-section func renderRunMetricsDiffPrettySection(run1ID, run2ID int64, diff *RunMetricsDiff) { if diff == nil { return } - fmt.Fprintln(os.Stderr, console.FormatSectionHeader(fmt.Sprintf("Run Metrics (Run #%d → Run #%d)", run1ID, run2ID))) fmt.Fprintln(os.Stderr) - - config := console.TableConfig{ - Headers: []string{"Metric", fmt.Sprintf("Run #%d", run1ID), fmt.Sprintf("Run #%d", run2ID), "Change"}, - Rows: make([][]string, 0), + if config := buildRunMetricsPrettyTable(run1ID, run2ID, diff); len(config.Rows) > 0 { + fmt.Fprint(os.Stderr, console.RenderTable(config)) } + renderRunMetricsDetailSections(run1ID, run2ID, diff) +} - if diff.Run1TokenUsage > 0 || diff.Run2TokenUsage > 0 { - config.Rows = append(config.Rows, []string{ - "Token usage", - strconv.Itoa(diff.Run1TokenUsage), - strconv.Itoa(diff.Run2TokenUsage), - diff.TokenUsageChange, - }) - } - if diff.Run1Duration != "" || diff.Run2Duration != "" { - config.Rows = append(config.Rows, []string{ - "Duration", - diff.Run1Duration, - diff.Run2Duration, - diff.DurationChange, - }) - } - if diff.Run1Turns > 0 || diff.Run2Turns > 0 { - config.Rows = append(config.Rows, []string{ - "Turns", - strconv.Itoa(diff.Run1Turns), - strconv.Itoa(diff.Run2Turns), - fmt.Sprintf("%+d", diff.TurnsChange), - }) - } - if diff.Run1TokensPerTurn > 0 || diff.Run2TokensPerTurn > 0 { - config.Rows = append(config.Rows, []string{ - "Tokens / turn", - strconv.Itoa(diff.Run1TokensPerTurn), - strconv.Itoa(diff.Run2TokensPerTurn), - diff.TokensPerTurnChange, - }) - } +func buildRunMetricsPrettyTable(run1ID, run2ID int64, diff *RunMetricsDiff) console.TableConfig { + config := console.TableConfig{Headers: []string{"Metric", fmt.Sprintf("Run #%d", run1ID), fmt.Sprintf("Run #%d", run2ID), "Change"}, Rows: make([][]string, 0)} + appendRunMetricsPrettyRow(&config, diff.Run1TokenUsage > 0 || diff.Run2TokenUsage > 0, []string{"Token usage", strconv.Itoa(diff.Run1TokenUsage), strconv.Itoa(diff.Run2TokenUsage), diff.TokenUsageChange}) + appendRunMetricsPrettyRow(&config, diff.Run1Duration != "" || diff.Run2Duration != "", []string{"Duration", diff.Run1Duration, diff.Run2Duration, diff.DurationChange}) + appendRunMetricsPrettyRow(&config, diff.Run1Turns > 0 || diff.Run2Turns > 0, []string{"Turns", strconv.Itoa(diff.Run1Turns), strconv.Itoa(diff.Run2Turns), fmt.Sprintf("%+d", diff.TurnsChange)}) + appendRunMetricsPrettyRow(&config, diff.Run1TokensPerTurn > 0 || diff.Run2TokensPerTurn > 0, []string{"Tokens / turn", strconv.Itoa(diff.Run1TokensPerTurn), strconv.Itoa(diff.Run2TokensPerTurn), diff.TokensPerTurnChange}) + return config +} - if len(config.Rows) > 0 { - fmt.Fprint(os.Stderr, console.RenderTable(config)) +func appendRunMetricsPrettyRow(config *console.TableConfig, include bool, row []string) { + if include { + config.Rows = append(config.Rows, row) } +} - if diff.TokenUsageDetails != nil { - fmt.Fprintln(os.Stderr) - renderTokenUsageDiffPrettySection(run1ID, run2ID, diff.TokenUsageDetails) +func renderRunMetricsDetailSections(run1ID, run2ID int64, diff *RunMetricsDiff) { + renderRunMetricsTokenUsageSection(run1ID, run2ID, diff.TokenUsageDetails) + renderRunMetricsRateLimitSection(run1ID, run2ID, diff.GitHubRateLimitDetails) + renderRunMetricsToolCallsSection(run1ID, run2ID, diff.ToolCallsDiff) +} + +func renderRunMetricsTokenUsageSection(run1ID, run2ID int64, diff *TokenUsageDiff) { + if diff == nil { + return } - if diff.GitHubRateLimitDetails != nil { - fmt.Fprintln(os.Stderr) - renderGitHubRateLimitDiffPrettySection(run1ID, run2ID, diff.GitHubRateLimitDetails) + fmt.Fprintln(os.Stderr) + renderTokenUsageDiffPrettySection(run1ID, run2ID, diff) +} + +func renderRunMetricsRateLimitSection(run1ID, run2ID int64, diff *GitHubRateLimitDiff) { + if diff == nil { + return } - if diff.ToolCallsDiff != nil { - fmt.Fprintln(os.Stderr) - renderToolCallsDiffPrettySection(run1ID, run2ID, diff.ToolCallsDiff) + fmt.Fprintln(os.Stderr) + renderGitHubRateLimitDiffPrettySection(run1ID, run2ID, diff) +} + +func renderRunMetricsToolCallsSection(run1ID, run2ID int64, diff *ToolCallsDiff) { + if diff == nil { + return } + fmt.Fprintln(os.Stderr) + renderToolCallsDiffPrettySection(run1ID, run2ID, diff) } +// renderTokenUsageDiffPrettySection renders detailed token usage as a pretty console sub-section // renderTokenUsageDiffPrettySection renders detailed token usage as a pretty console sub-section func renderTokenUsageDiffPrettySection(run1ID, run2ID int64, diff *TokenUsageDiff) { fmt.Fprintln(os.Stderr, console.FormatSectionHeader("Token Usage Details")) fmt.Fprintln(os.Stderr) - - config := console.TableConfig{ - Headers: []string{"Token Type", fmt.Sprintf("Run #%d", run1ID), fmt.Sprintf("Run #%d", run2ID), "Change"}, - Rows: make([][]string, 0), - } - - if diff.Run1InputTokens > 0 || diff.Run2InputTokens > 0 { - config.Rows = append(config.Rows, []string{ - "Input", - strconv.Itoa(diff.Run1InputTokens), - strconv.Itoa(diff.Run2InputTokens), - diff.InputTokensChange, - }) - } - if diff.Run1OutputTokens > 0 || diff.Run2OutputTokens > 0 { - config.Rows = append(config.Rows, []string{ - "Output", - strconv.Itoa(diff.Run1OutputTokens), - strconv.Itoa(diff.Run2OutputTokens), - diff.OutputTokensChange, - }) - } - if diff.Run1CacheReadTokens > 0 || diff.Run2CacheReadTokens > 0 { - config.Rows = append(config.Rows, []string{ - "Cache read", - strconv.Itoa(diff.Run1CacheReadTokens), - strconv.Itoa(diff.Run2CacheReadTokens), - diff.CacheReadTokensChange, - }) - } - if diff.Run1CacheWriteTokens > 0 || diff.Run2CacheWriteTokens > 0 { - config.Rows = append(config.Rows, []string{ - "Cache write", - strconv.Itoa(diff.Run1CacheWriteTokens), - strconv.Itoa(diff.Run2CacheWriteTokens), - diff.CacheWriteTokensChange, - }) - } - if diff.Run1AIC > 0 || diff.Run2AIC > 0 { - config.Rows = append(config.Rows, []string{ - "AI Credits", - fmt.Sprintf("%.3f", diff.Run1AIC), - fmt.Sprintf("%.3f", diff.Run2AIC), - diff.AICChange, - }) - } - if diff.Run1TotalRequests > 0 || diff.Run2TotalRequests > 0 { - config.Rows = append(config.Rows, []string{ - "API requests", - strconv.Itoa(diff.Run1TotalRequests), - strconv.Itoa(diff.Run2TotalRequests), - diff.RequestsDelta, - }) - } - if diff.Run1CacheEfficiency > 0 || diff.Run2CacheEfficiency > 0 { - config.Rows = append(config.Rows, []string{ - "Cache efficiency", - fmt.Sprintf("%.1f%%", diff.Run1CacheEfficiency*100), - fmt.Sprintf("%.1f%%", diff.Run2CacheEfficiency*100), - diff.CacheEfficiencyChange, - }) - } - + config := buildTokenUsagePrettyTable(run1ID, run2ID, diff) if len(config.Rows) > 0 { fmt.Fprint(os.Stderr, console.RenderTable(config)) } } +func buildTokenUsagePrettyTable(run1ID, run2ID int64, diff *TokenUsageDiff) console.TableConfig { + config := console.TableConfig{Headers: []string{"Token Type", fmt.Sprintf("Run #%d", run1ID), fmt.Sprintf("Run #%d", run2ID), "Change"}, Rows: make([][]string, 0)} + appendRunMetricsPrettyRow(&config, diff.Run1InputTokens > 0 || diff.Run2InputTokens > 0, []string{"Input", strconv.Itoa(diff.Run1InputTokens), strconv.Itoa(diff.Run2InputTokens), diff.InputTokensChange}) + appendRunMetricsPrettyRow(&config, diff.Run1OutputTokens > 0 || diff.Run2OutputTokens > 0, []string{"Output", strconv.Itoa(diff.Run1OutputTokens), strconv.Itoa(diff.Run2OutputTokens), diff.OutputTokensChange}) + appendRunMetricsPrettyRow(&config, diff.Run1CacheReadTokens > 0 || diff.Run2CacheReadTokens > 0, []string{"Cache read", strconv.Itoa(diff.Run1CacheReadTokens), strconv.Itoa(diff.Run2CacheReadTokens), diff.CacheReadTokensChange}) + appendRunMetricsPrettyRow(&config, diff.Run1CacheWriteTokens > 0 || diff.Run2CacheWriteTokens > 0, []string{"Cache write", strconv.Itoa(diff.Run1CacheWriteTokens), strconv.Itoa(diff.Run2CacheWriteTokens), diff.CacheWriteTokensChange}) + appendRunMetricsPrettyRow(&config, diff.Run1AIC > 0 || diff.Run2AIC > 0, []string{"AI Credits", fmt.Sprintf("%.3f", diff.Run1AIC), fmt.Sprintf("%.3f", diff.Run2AIC), diff.AICChange}) + appendRunMetricsPrettyRow(&config, diff.Run1TotalRequests > 0 || diff.Run2TotalRequests > 0, []string{"API requests", strconv.Itoa(diff.Run1TotalRequests), strconv.Itoa(diff.Run2TotalRequests), diff.RequestsDelta}) + appendRunMetricsPrettyRow(&config, diff.Run1CacheEfficiency > 0 || diff.Run2CacheEfficiency > 0, []string{"Cache efficiency", fmt.Sprintf("%.1f%%", diff.Run1CacheEfficiency*100), fmt.Sprintf("%.1f%%", diff.Run2CacheEfficiency*100), diff.CacheEfficiencyChange}) + return config +} + +// renderGitHubRateLimitDiffMarkdownSection renders the GitHub API rate limit diff as markdown // renderGitHubRateLimitDiffMarkdownSection renders the GitHub API rate limit diff as markdown func renderGitHubRateLimitDiffMarkdownSection(run1ID, run2ID int64, diff *GitHubRateLimitDiff) { fmt.Fprintln(os.Stdout, "#### GitHub API Usage") diff --git a/pkg/sliceutil/sliceutil.go b/pkg/sliceutil/sliceutil.go index d03321ad8c5..45ceea7a910 100644 --- a/pkg/sliceutil/sliceutil.go +++ b/pkg/sliceutil/sliceutil.go @@ -64,6 +64,13 @@ func SortedKeys[K cmp.Ordered, V any](m map[K]V) []K { return slices.Sorted(maps.Keys(m)) } +// SortedStrings returns a sorted copy of the provided string slice. +func SortedStrings(values []string) []string { + sorted := append([]string(nil), values...) + slices.Sort(sorted) + return sorted +} + // Any returns true if at least one element in the slice satisfies the predicate. // Returns false for nil or empty slices. // This is a pure function that does not modify the input slice. diff --git a/pkg/workflow/awf_config.go b/pkg/workflow/awf_config.go index 78d67898cb8..383153004c1 100644 --- a/pkg/workflow/awf_config.go +++ b/pkg/workflow/awf_config.go @@ -469,110 +469,130 @@ func buildAWFConfigSchemaURL(firewallConfig *FirewallConfig) string { // by the AWF --config flag. See BuildAWFCommand for how this is wired together. func BuildAWFConfigJSON(config AWFCommandConfig) (string, error) { awfConfigLog.Printf("Building AWF config JSON: engine=%s, allowed_domains=%q", config.EngineName, config.AllowedDomains) - - // Resolve firewall config once — used for both the schema URL and the container image tag. firewallConfig := getFirewallConfig(config.WorkflowData) + awfConfig := AWFConfigFile{Schema: buildAWFConfigSchemaURL(firewallConfig)} + applyAWFRunnerConfig(&awfConfig, config.WorkflowData) + awfConfig.Network = buildAWFNetworkConfig(config) + applyAWFPlatformConfig(&awfConfig, config.WorkflowData) + awfConfig.APIProxy = buildAWFAPIProxyConfig(config, firewallConfig) + applyAWFContainerConfig(&awfConfig, config, firewallConfig) + applyAWFLoggingConfig(&awfConfig, config.WorkflowData) + applyAWFBoundedQueriesConfig(&awfConfig, config.WorkflowData, firewallConfig) + return marshalAndValidateAWFConfigJSON(awfConfig, config.WorkflowData) +} - awfConfig := AWFConfigFile{ - Schema: buildAWFConfigSchemaURL(firewallConfig), - } +type awfAPIProxyRuntimeOptions struct { + maxAICredits int64 + maxRuns int + maxTurnCacheMisses int + enableTokenSteering bool +} - // ── Runner section ────────────────────────────────────────────────────── - if topology := getRunnerTopology(config.WorkflowData); topology != "" { +func applyAWFRunnerConfig(awfConfig *AWFConfigFile, workflowData *WorkflowData) { + if topology := getRunnerTopology(workflowData); topology != "" { awfConfig.Runner = &AWFRunnerConfig{Topology: topology} awfConfigLog.Printf("Runner section: topology=%s", topology) } +} - // ── Network section ────────────────────────────────────────────────────── - if config.AllowedDomains != "" { - allowList := splitDomainList(config.AllowedDomains) - awfConfig.Network = &AWFNetworkConfig{ - AllowDomains: allowList, - } - awfConfigLog.Printf("Network section: %d allowed domains", len(allowList)) - - // Blocked domains (if configured in the workflow) - if config.WorkflowData != nil { - blockedDomainsStr := formatBlockedDomains(config.WorkflowData.NetworkPermissions) - if blockedDomainsStr != "" { - blockList := splitDomainList(blockedDomainsStr) - awfConfig.Network.BlockDomains = blockList - awfConfigLog.Printf("Network section: %d blocked domains", len(blockList)) - } - } - } - +func buildAWFNetworkConfig(config AWFCommandConfig) *AWFNetworkConfig { + network := buildAWFAllowedAndBlockedDomains(config) if isAWFNetworkIsolationEnabled(config.WorkflowData) { - if awfConfig.Network == nil { - awfConfig.Network = &AWFNetworkConfig{} - } - awfConfig.Network.Isolation = true - awfConfig.Network.TopologyAttach = buildAWFTopologyAttachList(config.WorkflowData) - awfConfigLog.Printf("Network section: isolation enabled with %d topology attachments", len(awfConfig.Network.TopologyAttach)) + network = ensureAWFNetworkConfig(network) + network.Isolation = true + network.TopologyAttach = buildAWFTopologyAttachList(config.WorkflowData) + awfConfigLog.Printf("Network section: isolation enabled with %d topology attachments", len(network.TopologyAttach)) } - - // docker-sbx: the sbx microVM resolves host services via host.docker.internal - // (the Docker bridge gateway, 172.17.0.1). Allow this domain so AWF's network - // policy permits connections from the microVM to the api-proxy, MCP gateway, and - // Squid proxy that are all published on the host bridge. if isDockerSbxRuntime(config.WorkflowData) { - if awfConfig.Network == nil { - awfConfig.Network = &AWFNetworkConfig{} - } + network = ensureAWFNetworkConfig(network) const hostDockerInternal = "host.docker.internal" - if !slices.Contains(awfConfig.Network.AllowDomains, hostDockerInternal) { - awfConfig.Network.AllowDomains = append(awfConfig.Network.AllowDomains, hostDockerInternal) + if !slices.Contains(network.AllowDomains, hostDockerInternal) { + network.AllowDomains = append(network.AllowDomains, hostDockerInternal) awfConfigLog.Printf("Network section: added %s for docker-sbx microVM routing", hostDockerInternal) } } + return network +} - if platformType := extractPlatformType(config.WorkflowData); platformType != "" { - awfConfig.Platform = &AWFPlatformConfig{Type: platformType} - awfConfigLog.Printf("Platform section: type=%s", platformType) +func buildAWFAllowedAndBlockedDomains(config AWFCommandConfig) *AWFNetworkConfig { + if config.AllowedDomains == "" { + return nil + } + network := &AWFNetworkConfig{AllowDomains: splitDomainList(config.AllowedDomains)} + awfConfigLog.Printf("Network section: %d allowed domains", len(network.AllowDomains)) + if config.WorkflowData == nil { + return network } + blockedDomainsStr := formatBlockedDomains(config.WorkflowData.NetworkPermissions) + if blockedDomainsStr != "" { + network.BlockDomains = splitDomainList(blockedDomainsStr) + awfConfigLog.Printf("Network section: %d blocked domains", len(network.BlockDomains)) + } + return network +} - // ── API proxy section ───────────────────────────────────────────────────── - // maxAICredits is taken from frontmatter/imports only; when unset (0) the - // runtime value is resolved from vars.GH_AW_DEFAULT_MAX_AI_CREDITS via a - // GitHub Actions expression injected directly into the JSON string in - // BuildAWFCommand (see injectMaxAICreditsExpression in awf_helpers.go). - maxAICredits := int64(0) - maxRuns := constants.DefaultMaxRuns - // GetMaxTurnCacheMisses handles nil receiver and env-var fallback, so pre-init - // via the nil receiver avoids a redundant os.Getenv when EngineConfig is set. - maxTurnCacheMisses := (*EngineConfig)(nil).GetMaxTurnCacheMisses() - if config.WorkflowData != nil && config.WorkflowData.EngineConfig != nil { - if config.WorkflowData.EngineConfig.MaxAICredits != 0 { - maxAICredits = config.WorkflowData.EngineConfig.MaxAICredits - } - maxRuns = config.WorkflowData.EngineConfig.GetMaxRuns() - maxTurnCacheMisses = config.WorkflowData.EngineConfig.GetMaxTurnCacheMisses() +func ensureAWFNetworkConfig(network *AWFNetworkConfig) *AWFNetworkConfig { + if network != nil { + return network } + return &AWFNetworkConfig{} +} - // Token steering is enabled by default. Setting max-ai-credits to a negative - // value (-1) omits that budget from the AWF config and disables token steering. - // When maxAICredits is 0 (runtime default), token steering stays enabled here. - enableTokenSteering := maxAICredits >= 0 - if maxAICredits < 0 { - // Negative signals "disabled" — omit the budget from the AWF config. - maxAICredits = 0 +func applyAWFPlatformConfig(awfConfig *AWFConfigFile, workflowData *WorkflowData) { + if platformType := extractPlatformType(workflowData); platformType != "" { + awfConfig.Platform = &AWFPlatformConfig{Type: platformType} + awfConfigLog.Printf("Platform section: type=%s", platformType) } +} +func buildAWFAPIProxyConfig(config AWFCommandConfig, firewallConfig *FirewallConfig) *AWFAPIProxyConfig { + options := resolveAWFAPIProxyRuntimeOptions(config.WorkflowData) apiProxy := &AWFAPIProxyConfig{ Enabled: true, - MaxRuns: maxRuns, - MaxTurnCacheMisses: maxTurnCacheMisses, - MaxAICredits: maxAICredits, - EnableTokenSteering: enableTokenSteering && awfSupportsTokenSteering(firewallConfig), + MaxRuns: options.maxRuns, + MaxTurnCacheMisses: options.maxTurnCacheMisses, + MaxAICredits: options.maxAICredits, + EnableTokenSteering: options.enableTokenSteering && awfSupportsTokenSteering(firewallConfig), + } + logAWFTokenSteeringDecision(options.enableTokenSteering, firewallConfig) + applyAWFAPIProxyModelFallback(apiProxy, config.WorkflowData) + applyAWFAPIProxyPricing(apiProxy, config.WorkflowData) + applyAWFAPIProxyTargets(apiProxy, config) + applyAWFAPIProxyProviders(apiProxy, config.WorkflowData, firewallConfig) + applyAWFAPIProxyModels(apiProxy, config.WorkflowData) + return apiProxy +} + +func resolveAWFAPIProxyRuntimeOptions(workflowData *WorkflowData) awfAPIProxyRuntimeOptions { + options := awfAPIProxyRuntimeOptions{ + maxRuns: constants.DefaultMaxRuns, + maxTurnCacheMisses: (*EngineConfig)(nil).GetMaxTurnCacheMisses(), + enableTokenSteering: true, + } + if workflowData != nil && workflowData.EngineConfig != nil { + options.maxAICredits = workflowData.EngineConfig.MaxAICredits + options.maxRuns = workflowData.EngineConfig.GetMaxRuns() + options.maxTurnCacheMisses = workflowData.EngineConfig.GetMaxTurnCacheMisses() + } + if options.maxAICredits < 0 { + options.maxAICredits = 0 + options.enableTokenSteering = false } + return options +} +func logAWFTokenSteeringDecision(enableTokenSteering bool, firewallConfig *FirewallConfig) { if !enableTokenSteering { awfConfigLog.Printf("Skipping apiProxy.enableTokenSteering: max-ai-credits is negative (disabled)") - } else if !awfSupportsTokenSteering(firewallConfig) { + return + } + if !awfSupportsTokenSteering(firewallConfig) { awfConfigLog.Printf("Skipping apiProxy.enableTokenSteering: AWF version %q requires at least %s", getAWFImageTag(firewallConfig), constants.AWFTokenSteeringMinVersion) } +} - if mf := extractModelFallback(config.WorkflowData); mf != nil { +func applyAWFAPIProxyModelFallback(apiProxy *AWFAPIProxyConfig, workflowData *WorkflowData) { + if mf := extractModelFallback(workflowData); mf != nil { apiProxy.ModelFallback = mf enabledDisplay := "" if mf.Enabled != nil { @@ -580,28 +600,45 @@ func BuildAWFConfigJSON(config AWFCommandConfig) (string, error) { } awfConfigLog.Printf("API proxy: modelFallback configured: enabled=%s", enabledDisplay) } +} - if pricing := extractDefaultAiCreditsPricing(config.WorkflowData); pricing != nil { +func applyAWFAPIProxyPricing(apiProxy *AWFAPIProxyConfig, workflowData *WorkflowData) { + if pricing := extractDefaultAiCreditsPricing(workflowData); pricing != nil { apiProxy.DefaultAiCreditsPricing = pricing awfConfigLog.Printf("API proxy: defaultAiCreditsPricing configured: input=%g, output=%g", pricing.Input, pricing.Output) } +} + +func applyAWFAPIProxyTargets(apiProxy *AWFAPIProxyConfig, config AWFCommandConfig) { + targets := buildAWFAPITargets(config.WorkflowData, config.EngineName) + if len(targets) == 0 { + return + } + apiProxy.Targets = targets + awfConfigLog.Printf("API proxy: %d custom targets configured", len(targets)) +} +func buildAWFAPITargets(workflowData *WorkflowData, engineName string) map[string]*AWFAPITargetConfig { targets := map[string]*AWFAPITargetConfig{} + addAWFTargetHost(targets, "openai", extractAPITargetHost(workflowData, "OPENAI_BASE_URL")) + addAWFTargetHost(targets, "anthropic", extractAPITargetHost(workflowData, "ANTHROPIC_BASE_URL")) + applyAWFTargetAuthHeaders(targets, workflowData) + applyAWFCopilotTarget(targets, workflowData) + applyAWFGeminiTarget(targets, workflowData, engineName) + return targets +} - if openaiTarget := extractAPITargetHost(config.WorkflowData, "OPENAI_BASE_URL"); openaiTarget != "" { - targets["openai"] = &AWFAPITargetConfig{Host: openaiTarget} - awfConfigLog.Printf("API proxy: custom openai target=%s", openaiTarget) - } - if anthropicTarget := extractAPITargetHost(config.WorkflowData, "ANTHROPIC_BASE_URL"); anthropicTarget != "" { - targets["anthropic"] = &AWFAPITargetConfig{Host: anthropicTarget} - awfConfigLog.Printf("API proxy: custom anthropic target=%s", anthropicTarget) +func addAWFTargetHost(targets map[string]*AWFAPITargetConfig, provider string, host string) { + if host == "" { + return } + targets[provider] = &AWFAPITargetConfig{Host: host} + awfConfigLog.Printf("API proxy: custom %s target=%s", provider, host) +} - // Apply authHeader overrides from sandbox.agent.targets frontmatter. - // These are independent of the host/env-var settings: authHeader can be set - // even when no custom host is configured. +func applyAWFTargetAuthHeaders(targets map[string]*AWFAPITargetConfig, workflowData *WorkflowData) { for _, provider := range []string{"openai", "anthropic"} { - authHeader := extractAPITargetAuthHeader(config.WorkflowData, provider) + authHeader := extractAPITargetAuthHeader(workflowData, provider) if authHeader == "" { continue } @@ -612,53 +649,58 @@ func BuildAWFConfigJSON(config AWFCommandConfig) (string, error) { } awfConfigLog.Printf("API proxy: custom %s authHeader=%s", provider, authHeader) } - if copilotTarget := GetCopilotAPITarget(config.WorkflowData); copilotTarget != "" { +} + +func applyAWFCopilotTarget(targets map[string]*AWFAPITargetConfig, workflowData *WorkflowData) { + if copilotTarget := GetCopilotAPITarget(workflowData); copilotTarget != "" { targets["copilot"] = &AWFAPITargetConfig{Host: copilotTarget} awfConfigLog.Printf("API proxy: custom copilot target=%s", copilotTarget) } + copilotFrontmatter := extractCopilotTargetConfig(workflowData) + if copilotFrontmatter == nil { + return + } + target := ensureAWFAPITarget(targets, "copilot") + if copilotFrontmatter.AuthHeader != "" { + target.AuthHeader = copilotFrontmatter.AuthHeader + awfConfigLog.Printf("API proxy: copilot authHeader=%s", copilotFrontmatter.AuthHeader) + } + if len(copilotFrontmatter.ExtraHeaders) > 0 { + target.ExtraHeaders = copilotFrontmatter.ExtraHeaders + awfConfigLog.Printf("API proxy: copilot extraHeaders configured (%d header(s))", len(copilotFrontmatter.ExtraHeaders)) + } + if len(copilotFrontmatter.ExtraBodyFields) > 0 { + target.ExtraBodyFields = copilotFrontmatter.ExtraBodyFields + awfConfigLog.Printf("API proxy: copilot extraBodyFields configured (%d field(s))", len(copilotFrontmatter.ExtraBodyFields)) + } + if copilotFrontmatter.SessionId != "" { + target.SessionId = copilotFrontmatter.SessionId + awfConfigLog.Printf("API proxy: copilot sessionId configured") + } +} - // Apply BYOK supplemental fields from sandbox.agent.targets.copilot frontmatter. - // extraHeaders, extraBodyFields, and sessionId are Copilot-specific and map to - // AWF_BYOK_EXTRA_HEADERS, AWF_BYOK_EXTRA_BODY_FIELDS, and AWF_PROVIDER_SESSION_ID. - if copilotFrontmatter := extractCopilotTargetConfig(config.WorkflowData); copilotFrontmatter != nil { - existing, ok := targets["copilot"] - if !ok { - existing = &AWFAPITargetConfig{} - targets["copilot"] = existing - } - if copilotFrontmatter.AuthHeader != "" { - existing.AuthHeader = copilotFrontmatter.AuthHeader - awfConfigLog.Printf("API proxy: copilot authHeader=%s", copilotFrontmatter.AuthHeader) - } - if len(copilotFrontmatter.ExtraHeaders) > 0 { - existing.ExtraHeaders = copilotFrontmatter.ExtraHeaders - awfConfigLog.Printf("API proxy: copilot extraHeaders configured (%d header(s))", len(copilotFrontmatter.ExtraHeaders)) - } - if len(copilotFrontmatter.ExtraBodyFields) > 0 { - existing.ExtraBodyFields = copilotFrontmatter.ExtraBodyFields - awfConfigLog.Printf("API proxy: copilot extraBodyFields configured (%d field(s))", len(copilotFrontmatter.ExtraBodyFields)) - } - if copilotFrontmatter.SessionId != "" { - existing.SessionId = copilotFrontmatter.SessionId - awfConfigLog.Printf("API proxy: copilot sessionId configured") - } +func ensureAWFAPITarget(targets map[string]*AWFAPITargetConfig, provider string) *AWFAPITargetConfig { + if target, ok := targets[provider]; ok { + return target } - if antigravityTarget := GetAntigravityAPITarget(config.WorkflowData, config.EngineName); antigravityTarget != "" { - // Route the Antigravity-resolved API target through the "gemini" provider key - // to match AWF's supported target providers. + targets[provider] = &AWFAPITargetConfig{} + return targets[provider] +} + +func applyAWFGeminiTarget(targets map[string]*AWFAPITargetConfig, workflowData *WorkflowData, engineName string) { + if antigravityTarget := GetAntigravityAPITarget(workflowData, engineName); antigravityTarget != "" { awfConfigLog.Printf("API proxy: mapped antigravity target to gemini provider target=%s", antigravityTarget) targets["gemini"] = &AWFAPITargetConfig{Host: antigravityTarget} - } else if geminiTarget := GetGeminiAPITarget(config.WorkflowData, config.EngineName); geminiTarget != "" { + return + } + if geminiTarget := GetGeminiAPITarget(workflowData, engineName); geminiTarget != "" { awfConfigLog.Printf("API proxy: custom gemini target=%s", geminiTarget) targets["gemini"] = &AWFAPITargetConfig{Host: geminiTarget} } +} - if len(targets) > 0 { - apiProxy.Targets = targets - awfConfigLog.Printf("API proxy: %d custom targets configured", len(targets)) - } - - if providers := extractModelCostProviders(config.WorkflowData); len(providers) > 0 { +func applyAWFAPIProxyProviders(apiProxy *AWFAPIProxyConfig, workflowData *WorkflowData, firewallConfig *FirewallConfig) { + if providers := extractModelCostProviders(workflowData); len(providers) > 0 { if awfSupportsAPIProxyProviders(firewallConfig) { apiProxy.Providers = providers awfConfigLog.Printf("API proxy: %d model-cost provider override(s) configured", len(providers)) @@ -666,13 +708,14 @@ func BuildAWFConfigJSON(config AWFCommandConfig) (string, error) { awfConfigLog.Printf("Skipping apiProxy.providers: AWF version %q requires at least %s", getAWFImageTag(firewallConfig), constants.AWFAPIProxyProvidersMinVersion) } } +} - // ── Models section (nested under apiProxy per AWF config schema) ────────── - if config.WorkflowData != nil && len(config.WorkflowData.ModelMappings) > 0 { - apiProxy.Models = config.WorkflowData.ModelMappings - awfConfigLog.Printf("Models section: %d alias entries", len(config.WorkflowData.ModelMappings)) +func applyAWFAPIProxyModels(apiProxy *AWFAPIProxyConfig, workflowData *WorkflowData) { + if workflowData != nil && len(workflowData.ModelMappings) > 0 { + apiProxy.Models = workflowData.ModelMappings + awfConfigLog.Printf("Models section: %d alias entries", len(workflowData.ModelMappings)) } - allowedModels, disallowedModels := resolveModelPolicyForAWFConfig(config.WorkflowData) + allowedModels, disallowedModels := resolveModelPolicyForAWFConfig(workflowData) if len(allowedModels) > 0 { apiProxy.AllowedModels = allowedModels awfConfigLog.Printf("Models policy: %d allowed model pattern(s)", len(allowedModels)) @@ -681,65 +724,52 @@ func BuildAWFConfigJSON(config AWFCommandConfig) (string, error) { apiProxy.DisallowedModels = disallowedModels awfConfigLog.Printf("Models policy: %d disallowed model pattern(s)", len(disallowedModels)) } +} - awfConfig.APIProxy = apiProxy - - // ── Container section ───────────────────────────────────────────────────── +func applyAWFContainerConfig(awfConfig *AWFConfigFile, config AWFCommandConfig, firewallConfig *FirewallConfig) { awfImageTag := buildAWFImageTagWithDigests(getAWFImageTag(firewallConfig), config.WorkflowData) - agentRuntime := getAgentContainerRuntime(config.WorkflowData) - agentTimeout := 0 - if isDockerSbxRuntime(config.WorkflowData) { - agentTimeout = resolveAWFContainerAgentTimeoutMinutes(config.WorkflowData) + agentRuntime, agentTimeout := resolveAWFContainerRuntime(config.WorkflowData, firewallConfig) + if awfImageTag == "" && !isArcDindTopology(config.WorkflowData) && agentRuntime == "" && agentTimeout == 0 { + return } - // containerRuntime is only emitted when the effective AWF version supports it. - // Gate here to avoid sending an unrecognised field to older AWF binaries. - if !awfSupportsContainerRuntime(firewallConfig) { - if agentRuntime != "" { - awfConfigLog.Printf("Skipping containerRuntime: AWF version %q requires at least %s (gh-aw-firewall#6093)", getAWFImageTag(firewallConfig), constants.AWFContainerRuntimeMinVersion) - } - agentRuntime = "" + awfConfig.Container = &AWFContainerConfig{ImageTag: awfImageTag, AgentTimeout: agentTimeout, ContainerRuntime: agentRuntime} + if awfImageTag != "" { + awfConfigLog.Printf("Container section: image_tag=%s", awfImageTag) } - if awfImageTag != "" || isArcDindTopology(config.WorkflowData) || agentRuntime != "" || agentTimeout > 0 { - container := &AWFContainerConfig{ - ImageTag: awfImageTag, - AgentTimeout: agentTimeout, - ContainerRuntime: agentRuntime, - } - // NOTE: dockerHostPathPrefix is intentionally NOT set for arc-dind topology. - // With sysroot-stage active, the Docker daemon can access all needed paths: - // - Workspace & RUNNER_TEMP: on the shared work volume (/home/runner/_work/) - // - System binaries: provided by the sysroot named volume (not bind mounts) - // - Kernel VFS (/dev, /sys): daemon's own kernel - // Setting a prefix would incorrectly translate the workspace mount source to - // a non-existent path (e.g. /prefix/home/runner/_work/repo → empty dir), - // causing the agent to see an empty workspace. See gh-aw#34896. - awfConfig.Container = container - if awfImageTag != "" { - awfConfigLog.Printf("Container section: image_tag=%s", awfImageTag) - } - if agentRuntime != "" { - awfConfigLog.Printf("Container section: containerRuntime=%s", agentRuntime) - } - if agentTimeout > 0 { - awfConfigLog.Printf("Container section: agentTimeout=%d", agentTimeout) - } + if agentRuntime != "" { + awfConfigLog.Printf("Container section: containerRuntime=%s", agentRuntime) } + if agentTimeout > 0 { + awfConfigLog.Printf("Container section: agentTimeout=%d", agentTimeout) + } +} - // ── Logging section ────────────────────────────────────────────────────── - // Logging paths are set in config. For ARC/DinD, the config file is written at runtime, - // so ${RUNNER_TEMP} can be preserved for shell expansion before AWF reads the JSON. - awfConfig.Logging = &AWFLoggingConfig{ - ProxyLogsDir: string(constants.AWFProxyLogsDir), - AuditDir: string(constants.AWFAuditDir), +func resolveAWFContainerRuntime(workflowData *WorkflowData, firewallConfig *FirewallConfig) (string, int) { + agentRuntime := getAgentContainerRuntime(workflowData) + agentTimeout := 0 + if isDockerSbxRuntime(workflowData) { + agentTimeout = resolveAWFContainerAgentTimeoutMinutes(workflowData) + } + if awfSupportsContainerRuntime(firewallConfig) { + return agentRuntime, agentTimeout + } + if agentRuntime != "" { + awfConfigLog.Printf("Skipping containerRuntime: AWF version %q requires at least %s (gh-aw-firewall#6093)", getAWFImageTag(firewallConfig), constants.AWFContainerRuntimeMinVersion) } - if isArcDindTopology(config.WorkflowData) { + return "", agentTimeout +} + +func applyAWFLoggingConfig(awfConfig *AWFConfigFile, workflowData *WorkflowData) { + awfConfig.Logging = &AWFLoggingConfig{ProxyLogsDir: string(constants.AWFProxyLogsDir), AuditDir: string(constants.AWFAuditDir)} + if isArcDindTopology(workflowData) { awfConfig.Logging.ProxyLogsDir = awfArcDindProxyLogsDirExpr awfConfig.Logging.AuditDir = awfArcDindAuditDirExpr } awfConfigLog.Printf("Logging section: proxyLogsDir=%s, auditDir=%s", awfConfig.Logging.ProxyLogsDir, awfConfig.Logging.AuditDir) +} - // ── Bounded queries section ────────────────────────────────────────────── - if bq := extractBoundedQueriesConfig(config.WorkflowData); bq != nil { +func applyAWFBoundedQueriesConfig(awfConfig *AWFConfigFile, workflowData *WorkflowData, firewallConfig *FirewallConfig) { + if bq := extractBoundedQueriesConfig(workflowData); bq != nil { if awfSupportsBoundedQueries(firewallConfig) { awfConfig.BoundedQueries = bq awfConfigLog.Printf("Bounded queries section: %d private repo(s)", len(bq.PrivateRepos)) @@ -747,20 +777,19 @@ func BuildAWFConfigJSON(config AWFCommandConfig) (string, error) { awfConfigLog.Printf("Skipping boundedQueries: AWF version %q requires at least %s", getAWFImageTag(firewallConfig), constants.AWFBoundedQueriesMinVersion) } } +} +func marshalAndValidateAWFConfigJSON(awfConfig AWFConfigFile, workflowData *WorkflowData) (string, error) { jsonStr, err := jsonutil.MarshalCompactNoHTMLEscape(awfConfig) if err != nil { return "", fmt.Errorf("failed to marshal AWF config to JSON: %w", err) } - awfConfigLog.Printf("AWF config JSON generated: %d bytes", len(jsonStr)) - - if config.WorkflowData != nil && config.WorkflowData.ValidateAWFConfig { + if workflowData != nil && workflowData.ValidateAWFConfig { if err := validateAWFConfigJSON(jsonStr); err != nil { return "", fmt.Errorf("generated AWF config failed schema validation: %w", err) } } - return jsonStr, nil } diff --git a/pkg/workflow/awf_helpers.go b/pkg/workflow/awf_helpers.go index 796d2644312..cbd4b2747d8 100644 --- a/pkg/workflow/awf_helpers.go +++ b/pkg/workflow/awf_helpers.go @@ -217,6 +217,12 @@ func buildWorkflowCallNetworkAllowedUpdateScript() (string, error) { shellEscapeArg(string(ecosystemJSON))), nil } +type awfArcDindRuntimeConfig struct { + dockerHostProbe string + prefixProbe string + dockerHostRef string +} + // BuildAWFCommand builds a complete AWF command with all arguments. // This consolidates the AWF command building logic that was duplicated across // Copilot, Claude, and Codex engines. @@ -229,270 +235,239 @@ func buildWorkflowCallNetworkAllowedUpdateScript() (string, error) { func BuildAWFCommand(config AWFCommandConfig) string { awfHelpersLog.Printf("Building AWF command for engine: %s", config.EngineName) isArcDind := isArcDindTopology(config.WorkflowData) - - // Get AWF command prefix (custom or standard) - awfCommand := GetAWFCommandPrefix(config.WorkflowData) - - // Build AWF arguments. The returned list contains only args that are safe to pass - // through shellJoinArgs. Expandable-var args (--container-workdir "${GITHUB_WORKSPACE}" - // and --mount "${RUNNER_TEMP}/...") are appended raw below so that shell variable - // expansion is not suppressed by single-quoting. - awfArgs := BuildAWFArgs(config) firewallConfig := getFirewallConfig(config.WorkflowData) + arcDind := buildAWFArcDindRuntimeConfig(config, firewallConfig) + expandableArgs, dockerHostProbe := buildAWFExpandableArgs(isArcDind, arcDind.dockerHostProbe) + arcDind.dockerHostProbe = dockerHostProbe + configFileSetup, expandableArgs := buildAWFConfigFileSetup(config, expandableArgs) + expandableArgs = addAWFUploadArtifactMount(expandableArgs, config.WorkflowData) + expandableArgs = addAWFServicePortArgs(expandableArgs, config.WorkflowData) + command := buildCompleteAWFCommand( + config, + GetAWFCommandPrefix(config.WorkflowData), + BuildAWFArgs(config), + expandableArgs, + configFileSetup, + arcDind, + buildModelsJSONPathExportScript(isArcDind), + ) + awfHelpersLog.Print("Successfully built AWF command") + return command +} - // Auto-detect ARC/DinD split daemon topology at runtime: probe DOCKER_HOST for a - // tcp:// scheme and pass it through to AWF via --docker-host. - // All behaviors avoid requiring workflow-authored sandbox.agent.args for standard ARC DinD setups. - // When AWF also supports chroot config (v0.27.1+), the Python patch body is embedded inside - // the same if-block so the script only contains one DOCKER_HOST condition check. - arcDindPrefixProbe := "" - arcDindDockerHostProbe := fmt.Sprintf(`%s="" +func buildAWFArcDindRuntimeConfig(config AWFCommandConfig, firewallConfig *FirewallConfig) awfArcDindRuntimeConfig { + runtimeConfig := awfArcDindRuntimeConfig{ + dockerHostProbe: fmt.Sprintf(`%s="" if [[ "${DOCKER_HOST:-}" =~ %s ]]; then %s="${DOCKER_HOST}" -fi`, - awfDockerHostVarName, - awfArcDindDockerHostRegex, - awfDockerHostVarName, - ) - arcDindDockerHostRef := fmt.Sprintf("${%s:+--docker-host \"$%s\"}", awfDockerHostVarName, awfDockerHostVarName) - if awfSupportsDockerHostPathPrefix(firewallConfig) { - chrootPatchBody := "" - if awfSupportsChrootConfig(firewallConfig) { - if config.WorkflowData != nil && config.WorkflowData.IsDetectionRun { - chrootPatchBody = "\n" + buildArcDindChrootConfigPatchBodyBash() - } else { - chrootPatchBody = "\n" + buildArcDindChrootConfigPatchBody() - } - } - // NOTE: --docker-host-path-prefix is intentionally NOT passed. With sysroot-stage - // active, all bind-mount source paths are on the shared work volume and visible to - // the Docker daemon without translation. The prefix caused AWF to translate - // GITHUB_WORKSPACE to a non-existent path, resulting in an empty workspace (gh-aw#34896). - // The probe block is preserved for the chroot config patch which still requires the - // DOCKER_HOST guard. - if chrootPatchBody != "" { - arcDindPrefixProbe = fmt.Sprintf(`if [[ "${DOCKER_HOST:-}" =~ %s ]]; then%s -fi`, - awfArcDindDockerHostRegex, - chrootPatchBody) - } +fi`, awfDockerHostVarName, awfArcDindDockerHostRegex, awfDockerHostVarName), + dockerHostRef: fmt.Sprintf("${%s:+--docker-host \"$%s\"}", awfDockerHostVarName, awfDockerHostVarName), } - toolCacheMountProbe := fmt.Sprintf(`%s="" + if !awfSupportsDockerHostPathPrefix(firewallConfig) { + return runtimeConfig + } + chrootPatchBody := buildAWFArcDindChrootPatchBody(config.WorkflowData, firewallConfig) + if chrootPatchBody == "" { + return runtimeConfig + } + runtimeConfig.prefixProbe = fmt.Sprintf(`if [[ "${DOCKER_HOST:-}" =~ %s ]]; then%s +fi`, awfArcDindDockerHostRegex, chrootPatchBody) + return runtimeConfig +} + +func buildAWFArcDindChrootPatchBody(workflowData *WorkflowData, firewallConfig *FirewallConfig) string { + if !awfSupportsChrootConfig(firewallConfig) { + return "" + } + if workflowData != nil && workflowData.IsDetectionRun { + return "\n" + buildArcDindChrootConfigPatchBodyBash() + } + return "\n" + buildArcDindChrootConfigPatchBody() +} + +func buildAWFToolCacheMountSupport() (string, string) { + probe := fmt.Sprintf(`%s="" GH_AW_TOOL_CACHE="${RUNNER_TOOL_CACHE:?RUNNER_TOOL_CACHE must be set}" if [ -d "$GH_AW_TOOL_CACHE" ]; then if [[ "$GH_AW_TOOL_CACHE" != /opt/* ]]; then %s="$GH_AW_TOOL_CACHE:$GH_AW_TOOL_CACHE:ro" fi -fi`, - awfToolCacheMountVarName, - awfToolCacheMountVarName, - ) - toolCacheMountRef := fmt.Sprintf("${%s:+--mount \"$%s\"}", awfToolCacheMountVarName, awfToolCacheMountVarName) +fi`, awfToolCacheMountVarName, awfToolCacheMountVarName) + ref := fmt.Sprintf("${%s:+--mount \"$%s\"}", awfToolCacheMountVarName, awfToolCacheMountVarName) + return probe, ref +} - // Build the expandable args string for args that need shell variable expansion. - // These MUST be appended as raw (unescaped) strings because single-quoting would - // prevent the runner's shell from expanding ${GITHUB_WORKSPACE} and ${RUNNER_TEMP}. +func buildAWFExpandableArgs(isArcDind bool, dockerHostProbe string) (string, string) { ghAwDir := constants.GhAwRootDirShell expandableArgs := fmt.Sprintf( `--container-workdir "${GITHUB_WORKSPACE}" --mount "%s:%s:ro" --mount "%s:/host%s:ro"`, ghAwDir, ghAwDir, ghAwDir, ghAwDir, ) - if isArcDind { - expandableArgs += fmt.Sprintf( - ` --mount "%s:%s:rw" --mount "%s:%s:rw"`, - awfArcDindHomePathExpr, awfArcDindHomePathExpr, - awfArcDindRootPathExpr+"/sandbox/agent", awfArcDindRootPathExpr+"/sandbox/agent", - ) - // Explicitly mount the workspace so AWF can see it without path-prefix translation. - // GITHUB_WORKSPACE is on the shared work volume, so the Docker daemon can access it. - expandableArgs += ` --mount "${GITHUB_WORKSPACE}:${GITHUB_WORKSPACE}:rw"` - // Pre-create the rw mount source directories. AWF validates that mount source - // paths exist before starting containers, so these must be created on the host - // before the AWF invocation. The parent ${RUNNER_TEMP}/gh-aw/ already exists - // (created by actions/setup), but the subdirectories may not. - arcDindDockerHostProbe += fmt.Sprintf("\nmkdir -p \"%s\" \"%s\"", - awfArcDindHomePathExpr, - awfArcDindRootPathExpr+"/sandbox/agent", - ) - // Copy prompt files to daemon-visible path. On ARC/DinD, /tmp/gh-aw/ is NOT - // accessible to the Docker daemon. The activation job writes prompts to - // /tmp/gh-aw/aw-prompts/, so we copy them to ${RUNNER_TEMP}/gh-aw/aw-prompts/. - arcDindDockerHostProbe += fmt.Sprintf("\nif [ -d /tmp/gh-aw/aw-prompts ]; then cp -a /tmp/gh-aw/aw-prompts \"%s/aw-prompts\"; fi", - awfArcDindRootPathExpr, - ) - } - - // Generate a JSON config file and reference it via --config "${RUNNER_TEMP}/gh-aw/awf-config.json". - // This replaces several verbose CLI flags (--allow-domains, --enable-api-proxy, --image-tag, - // API targets) with a structured JSON file that is easier to audit and extend. - // - // The config file is written at runtime (inside the run: step) immediately before the AWF - // invocation, using printf to a fixed path inside the pre-existing ${RUNNER_TEMP}/gh-aw/ - // directory that is already set up by actions/setup. - var configFileSetup string + if !isArcDind { + return expandableArgs, dockerHostProbe + } + expandableArgs += fmt.Sprintf( + ` --mount "%s:%s:rw" --mount "%s:%s:rw"`, + awfArcDindHomePathExpr, awfArcDindHomePathExpr, + awfArcDindRootPathExpr+"/sandbox/agent", awfArcDindRootPathExpr+"/sandbox/agent", + ) + expandableArgs += ` --mount "${GITHUB_WORKSPACE}:${GITHUB_WORKSPACE}:rw"` + dockerHostProbe += fmt.Sprintf("\nmkdir -p \"%s\" \"%s\"", + awfArcDindHomePathExpr, + awfArcDindRootPathExpr+"/sandbox/agent", + ) + dockerHostProbe += fmt.Sprintf("\nif [ -d /tmp/gh-aw/aw-prompts ]; then cp -a /tmp/gh-aw/aw-prompts \"%s/aw-prompts\"; fi", + awfArcDindRootPathExpr, + ) + return expandableArgs, dockerHostProbe +} + +func buildAWFConfigFileSetup(config AWFCommandConfig, expandableArgs string) (string, string) { awfConfigJSON, err := BuildAWFConfigJSON(config) if err != nil { awfHelpersLog.Printf("Warning: failed to build AWF config JSON: %v", err) - } else { - // When max-ai-credits is not set by frontmatter/imports, export a local shell - // variable (GH_AW_MAX_AI_CREDITS) holding a GitHub Actions runtime expression, - // then inject a reference to that variable (${GH_AW_MAX_AI_CREDITS}) into the - // "maxAiCredits" field of the apiProxy JSON object. GitHub Actions evaluates - // the ${{ }} expression before the shell runs, so the variable is set to the - // resolved integer by the time printf writes the config file. - // - // Standard agent runs use vars.GH_AW_DEFAULT_MAX_AI_CREDITS with built-in - // fallback 1000. Threat-detection runs use - // vars.GH_AW_DEFAULT_DETECTION_MAX_AI_CREDITS with built-in fallback 400. - // Evals runs use vars.GH_AW_DEFAULT_EVALS_MAX_AI_CREDITS with built-in - // fallback 400 to align with detection budgets. - // EngineConfig.MaxAICredits is 0 when no compile-time value was set - // (neither frontmatter nor detection-engine config provided one). - // In that case, emit a runtime expression that lets the org variable - // or the built-in default resolve the budget at action run time. - // For detection runs, use the detection-specific variable/fallback; - // for standard agent runs, use the main-agent variable/fallback. - var maxAICreditsExportLine string - if config.WorkflowData == nil || config.WorkflowData.EngineConfig == nil || config.WorkflowData.EngineConfig.MaxAICredits == 0 { - defaultMaxAICredits := strconv.FormatInt(constants.DefaultMaxAICredits, 10) - if config.WorkflowData != nil { - switch { - case config.WorkflowData.IsEvalsRun: - defaultMaxAICredits = strconv.FormatInt(constants.DefaultDetectionMaxAICredits, 10) - case config.WorkflowData.IsDetectionRun: - defaultMaxAICredits = strconv.FormatInt(constants.DefaultDetectionMaxAICredits, 10) - } - } - awfConfigJSON = injectMaxAICreditsExpression(awfConfigJSON, fmt.Sprintf("${%s}", awfMaxAICreditsVarName)) - if config.ResolveMaxAICreditsFromEnv { - maxAICreditsExportLine = fmt.Sprintf(`%s="${%s:-%s}"`, awfMaxAICreditsVarName, awfMaxAICreditsVarName, defaultMaxAICredits) - } else { - expr := compilerenv.BuildDefaultMaxAICreditsExpression(defaultMaxAICredits) - if config.WorkflowData != nil { - switch { - case config.WorkflowData.IsEvalsRun: - expr = compilerenv.BuildDefaultEvalsMaxAICreditsExpression(defaultMaxAICredits) - case config.WorkflowData.IsDetectionRun: - expr = compilerenv.BuildDefaultDetectionMaxAICreditsExpression(defaultMaxAICredits) - } - } - maxAICreditsExportLine = fmt.Sprintf(`%s="%s"`, awfMaxAICreditsVarName, expr) - } - awfHelpersLog.Printf("Injected maxAiCredits local var reference into AWF config JSON") - } - // Write the config JSON to ${RUNNER_TEMP}/gh-aw/awf-config.json before AWF runs. - // When the generated JSON contains compiler-owned runtime variables such as - // ${GH_AW_MAX_AI_CREDITS} or ${RUNNER_TEMP}, use shellEscapeArgWithVarsPreserved - // which always uses double-quote wrapping: it escapes bare $ signs (e.g. - // "$schema" → "\$schema") while preserving both ${{ }} GitHub Actions expressions - // (e.g. in AllowedDomains) and approved shell variable references so bash expands - // them to runtime-resolved values. When no such variables are injected, - // shellEscapeArg handles escaping normally. - // Also copy it to /tmp/gh-aw/awf-config.json for the unified agent artifact upload. - var printfArg string - preservedVars := make([]string, 0, 2) - if maxAICreditsExportLine != "" { - preservedVars = append(preservedVars, awfMaxAICreditsVarName) - } - if strings.Contains(awfConfigJSON, awfArcDindRootPathExpr) { - preservedVars = append(preservedVars, "RUNNER_TEMP") - } - if len(preservedVars) > 0 { - printfArg = shellEscapeArgWithVarsPreserved(awfConfigJSON, preservedVars...) - } else { - printfArg = shellEscapeArg(awfConfigJSON) - } - // SC2016 ("Expressions don't expand in single quotes") is only triggered when - // printfArg is single-quoted (no runtime variables injected). Double-quoted args - // already escape bare $ signs as \$schema, so shellcheck does not warn there. - var printfLine string - if strings.HasPrefix(printfArg, "'") { - printfLine = "# shellcheck disable=SC2016\nprintf '%%s\\n' %s > %q" - } else { - printfLine = "printf '%%s\\n' %s > %q" - } - configFileSetup = fmt.Sprintf(printfLine, printfArg, awfConfigRuntimePathExpr) - if maxAICreditsExportLine != "" { - configFileSetup = maxAICreditsExportLine + "\n" + configFileSetup - } - if shouldUseWorkflowCallNetworkAllowedInput(config.WorkflowData) { - updateScript, updateErr := buildWorkflowCallNetworkAllowedUpdateScript() - if updateErr != nil { - awfHelpersLog.Printf("Warning: failed to build workflow_call network_allowed updater: %v", updateErr) - } else { - configFileSetup += "\n" + updateScript - } - } - configFileSetup += fmt.Sprintf("\ncp %q %s", awfConfigRuntimePathExpr, constants.AWFConfigFilePath) - // Add --config as the first expandable arg so it appears before --container-workdir. - expandableArgs = fmt.Sprintf("--config %q ", awfConfigRuntimePathExpr) + expandableArgs - awfHelpersLog.Print("Using AWF config file (--config flag)") - } - modelsJSONPathExport := buildModelsJSONPathExportScript(isArcDind) - - // When upload_artifact is configured, add a read-write mount for the staging directory - // so the model can copy files there from inside the container. The parent ${RUNNER_TEMP}/gh-aw - // is mounted :ro above; this child mount overrides access for the staging subdirectory only. - // The staging directory must already exist on the host (created in Generate Safe Outputs Config step). - if config.WorkflowData != nil && config.WorkflowData.SafeOutputs != nil && config.WorkflowData.SafeOutputs.UploadArtifact != nil { - stagingDir := SafeOutputsUploadArtifactsDir - expandableArgs += fmt.Sprintf(` --mount "%s:%s:rw"`, stagingDir, stagingDir) - awfHelpersLog.Print("Added read-write mount for upload_artifact staging directory") - } - - // Add --allow-host-service-ports for services with port mappings. - // This flag requires --legacy-security since it grants host network access. - // This is appended as a raw (expandable) arg because the value contains - // ${{ job.services..ports[''] }} expressions that include single quotes. - // These expressions are resolved by the GitHub Actions runner before shell execution, - // so they must not be shell-escaped. - agentCfg := getAgentConfig(config.WorkflowData) - isLegacyMode := agentCfg != nil && agentCfg.LegacySecurity - if config.WorkflowData != nil && config.WorkflowData.ServicePortExpressions != "" && isLegacyMode { - expandableArgs += fmt.Sprintf(` --allow-host-service-ports "%s"`, config.WorkflowData.ServicePortExpressions) - awfHelpersLog.Printf("Added --allow-host-service-ports with %s", config.WorkflowData.ServicePortExpressions) - } else if config.WorkflowData != nil && config.WorkflowData.ServicePortExpressions != "" { - awfHelpersLog.Print("Skipping --allow-host-service-ports: requires legacy-security mode") - } - - engineCommand := config.EngineCommand - if isArcDind { - engineCommand = rewriteArcDindEngineCommand(engineCommand) + return "", expandableArgs } + maxAICreditsExportLine, awfConfigJSON := buildAWFConfigRuntimeBudgetSetup(config, awfConfigJSON) + configFileSetup := buildAWFConfigPrintfScript(awfConfigJSON, maxAICreditsExportLine) + configFileSetup = appendAWFWorkflowCallNetworkUpdater(configFileSetup, config.WorkflowData) + configFileSetup += fmt.Sprintf("\ncp %q %s", awfConfigRuntimePathExpr, constants.AWFConfigFilePath) + awfHelpersLog.Print("Using AWF config file (--config flag)") + return configFileSetup, fmt.Sprintf("--config %q ", awfConfigRuntimePathExpr) + expandableArgs +} - // Wrap engine command in shell (command already includes any internal setup like npm PATH) - shellWrappedCommand := WrapCommandInShell(engineCommand) +func buildAWFConfigRuntimeBudgetSetup(config AWFCommandConfig, awfConfigJSON string) (string, string) { + if config.WorkflowData != nil && config.WorkflowData.EngineConfig != nil && config.WorkflowData.EngineConfig.MaxAICredits != 0 { + return "", awfConfigJSON + } + defaultMaxAICredits := resolveAWFDefaultMaxAICredits(config.WorkflowData) + awfConfigJSON = injectMaxAICreditsExpression(awfConfigJSON, fmt.Sprintf("${%s}", awfMaxAICreditsVarName)) + awfHelpersLog.Printf("Injected maxAiCredits local var reference into AWF config JSON") + if config.ResolveMaxAICreditsFromEnv { + return fmt.Sprintf(`%s="${%s:-%s}"`, awfMaxAICreditsVarName, awfMaxAICreditsVarName, defaultMaxAICredits), awfConfigJSON + } + return fmt.Sprintf(`%s="%s"`, awfMaxAICreditsVarName, buildAWFDefaultMaxAICreditsExpression(config.WorkflowData, defaultMaxAICredits)), awfConfigJSON +} - // Pre-create the agent stdio log file with restrictive permissions (0600) before - // starting the AWF container. tee would otherwise create it with the default - // umask (0644), leaving secrets (e.g. MCP gateway tokens) world-readable on the - // runner host until the secret-redaction step runs. - preCreateLog := fmt.Sprintf("(umask 177 && touch %s)", shellEscapeArg(config.LogFile)) +func resolveAWFDefaultMaxAICredits(workflowData *WorkflowData) string { + switch { + case workflowData != nil && workflowData.IsEvalsRun: + return strconv.FormatInt(constants.DefaultDetectionMaxAICredits, 10) + case workflowData != nil && workflowData.IsDetectionRun: + return strconv.FormatInt(constants.DefaultDetectionMaxAICredits, 10) + default: + return strconv.FormatInt(constants.DefaultMaxAICredits, 10) + } +} + +func buildAWFDefaultMaxAICreditsExpression(workflowData *WorkflowData, defaultMaxAICredits string) string { + switch { + case workflowData != nil && workflowData.IsEvalsRun: + return compilerenv.BuildDefaultEvalsMaxAICreditsExpression(defaultMaxAICredits) + case workflowData != nil && workflowData.IsDetectionRun: + return compilerenv.BuildDefaultDetectionMaxAICreditsExpression(defaultMaxAICredits) + default: + return compilerenv.BuildDefaultMaxAICreditsExpression(defaultMaxAICredits) + } +} + +func buildAWFConfigPrintfScript(awfConfigJSON string, maxAICreditsExportLine string) string { + printfArg := buildAWFConfigPrintfArg(awfConfigJSON, maxAICreditsExportLine != "") + printfLine := "printf '%%s\\n' %s > %q" + if strings.HasPrefix(printfArg, "'") { + printfLine = "# shellcheck disable=SC2016\nprintf '%%s\\n' %s > %q" + } + configFileSetup := fmt.Sprintf(printfLine, printfArg, awfConfigRuntimePathExpr) + if maxAICreditsExportLine == "" { + return configFileSetup + } + return maxAICreditsExportLine + "\n" + configFileSetup +} + +func buildAWFConfigPrintfArg(awfConfigJSON string, hasRuntimeBudget bool) string { + preservedVars := make([]string, 0, 2) + if hasRuntimeBudget { + preservedVars = append(preservedVars, awfMaxAICreditsVarName) + } + if strings.Contains(awfConfigJSON, awfArcDindRootPathExpr) { + preservedVars = append(preservedVars, "RUNNER_TEMP") + } + if len(preservedVars) == 0 { + return shellEscapeArg(awfConfigJSON) + } + return shellEscapeArgWithVarsPreserved(awfConfigJSON, preservedVars...) +} + +func appendAWFWorkflowCallNetworkUpdater(configFileSetup string, workflowData *WorkflowData) string { + if !shouldUseWorkflowCallNetworkAllowedInput(workflowData) { + return configFileSetup + } + updateScript, err := buildWorkflowCallNetworkAllowedUpdateScript() + if err != nil { + awfHelpersLog.Printf("Warning: failed to build workflow_call network_allowed updater: %v", err) + return configFileSetup + } + return configFileSetup + "\n" + updateScript +} + +func addAWFUploadArtifactMount(expandableArgs string, workflowData *WorkflowData) string { + if workflowData == nil || workflowData.SafeOutputs == nil || workflowData.SafeOutputs.UploadArtifact == nil { + return expandableArgs + } + stagingDir := SafeOutputsUploadArtifactsDir + awfHelpersLog.Print("Added read-write mount for upload_artifact staging directory") + return expandableArgs + fmt.Sprintf(` --mount "%s:%s:rw"`, stagingDir, stagingDir) +} - // Capture the epoch-millisecond timestamp at the very start of the Execute Agent CLI - // step on the host, before the AWF container launches. sendJobConclusionSpan reads - // this file to set the dedicated gh-aw..agent span start time, which excludes - // pre-agent overhead such as workspace audit and CLI proxy startup. +func addAWFServicePortArgs(expandableArgs string, workflowData *WorkflowData) string { + if workflowData == nil || workflowData.ServicePortExpressions == "" { + return expandableArgs + } + agentCfg := getAgentConfig(workflowData) + if agentCfg != nil && agentCfg.LegacySecurity { + awfHelpersLog.Printf("Added --allow-host-service-ports with %s", workflowData.ServicePortExpressions) + return expandableArgs + fmt.Sprintf(` --allow-host-service-ports "%s"`, workflowData.ServicePortExpressions) + } + awfHelpersLog.Print("Skipping --allow-host-service-ports: requires legacy-security mode") + return expandableArgs +} + +func buildAWFEngineCommand(engineCommand string, isArcDind bool) string { + if isArcDind { + return rewriteArcDindEngineCommand(engineCommand) + } + return engineCommand +} + +func buildCompleteAWFCommand( + config AWFCommandConfig, + awfCommand string, + awfArgs []string, + expandableArgs string, + configFileSetup string, + arcDind awfArcDindRuntimeConfig, + modelsJSONPathExport string, +) string { + toolCacheMountProbe, toolCacheMountRef := buildAWFToolCacheMountSupport() + shellWrappedCommand := WrapCommandInShell(buildAWFEngineCommand(config.EngineCommand, isArcDindTopology(config.WorkflowData))) + preCreateLog := fmt.Sprintf("(umask 177 && touch %s)", shellEscapeArg(config.LogFile)) writeAgentCLIStartMs := "printf '%s' \"$(date +%s%3N)\" > " + shellEscapeArg(AgentCLIStartMsPath) + joinedArgs := shellJoinArgs(awfArgs) + logFileArg := shellEscapeArg(config.LogFile) + switch { + case config.PathSetup != "" && configFileSetup != "": + return formatAWFCommandWithPathSetupAndConfig(writeAgentCLIStartMs, config, preCreateLog, configFileSetup, modelsJSONPathExport, arcDind, toolCacheMountProbe, awfCommand, expandableArgs, toolCacheMountRef, joinedArgs, shellWrappedCommand, logFileArg) + case config.PathSetup != "": + return formatAWFCommandWithPathSetup(writeAgentCLIStartMs, config, preCreateLog, modelsJSONPathExport, arcDind, toolCacheMountProbe, awfCommand, expandableArgs, toolCacheMountRef, joinedArgs, shellWrappedCommand, logFileArg) + case configFileSetup != "": + return formatAWFCommandWithConfig(writeAgentCLIStartMs, preCreateLog, configFileSetup, modelsJSONPathExport, arcDind, toolCacheMountProbe, awfCommand, expandableArgs, toolCacheMountRef, joinedArgs, shellWrappedCommand, logFileArg) + default: + return formatAWFCommand(writeAgentCLIStartMs, preCreateLog, modelsJSONPathExport, arcDind, toolCacheMountProbe, awfCommand, expandableArgs, toolCacheMountRef, joinedArgs, shellWrappedCommand, logFileArg) + } +} - // Build the complete command with proper formatting. - // configFileSetup (if non-empty) writes the AWF config JSON immediately before the - // AWF invocation so the file is present when AWF parses --config. - // - // shellcheck directive rationale: - // - SC1003 is expected because this generated block intentionally contains GitHub - // expression literals (for example ${{ job.services..ports[''] }}) - // that include single quotes and must survive into runtime unchanged. - // - SC2086 is expected because a subset of AWF arguments are intentionally emitted - // as expandable shell fragments (for example ${GH_AW_TOOL_CACHE_MOUNT:+...} and - // ${GH_AW_DOCKER_HOST:+...}). These fragments are produced by trusted - // compiler-owned probes above and are not user-provided free-form shell input. - // - // We keep normal quoting for all user-controlled values via shellEscapeArg/shellJoinArgs - // and scope this suppression to the generated AWF invocation line only. - var command string - if config.PathSetup != "" && configFileSetup != "" { - command = fmt.Sprintf(`set -o pipefail +func formatAWFCommandWithPathSetupAndConfig(writeAgentCLIStartMs string, config AWFCommandConfig, preCreateLog string, configFileSetup string, modelsJSONPathExport string, arcDind awfArcDindRuntimeConfig, toolCacheMountProbe string, awfCommand string, expandableArgs string, toolCacheMountRef string, joinedArgs string, shellWrappedCommand string, logFileArg string) string { + return fmt.Sprintf(`set -o pipefail %s %s %s @@ -504,25 +479,13 @@ fi`, %s %s %s %s %s %s \ -- %s 2>&1 | tee -a %s`, - writeAgentCLIStartMs, - config.PathSetup, - preCreateLog, - configFileSetup, - modelsJSONPathExport, - arcDindDockerHostProbe, - arcDindPrefixProbe, - toolCacheMountProbe, - awfShellcheckDirective, - awfCommand, - expandableArgs, - toolCacheMountRef, - arcDindDockerHostRef, - shellJoinArgs(awfArgs), - shellWrappedCommand, - shellEscapeArg(config.LogFile)) - } else if config.PathSetup != "" { - // Include path setup before AWF command (runs on host before AWF) - command = fmt.Sprintf(`set -o pipefail + writeAgentCLIStartMs, config.PathSetup, preCreateLog, configFileSetup, modelsJSONPathExport, + arcDind.dockerHostProbe, arcDind.prefixProbe, toolCacheMountProbe, awfShellcheckDirective, + awfCommand, expandableArgs, toolCacheMountRef, arcDind.dockerHostRef, joinedArgs, shellWrappedCommand, logFileArg) +} + +func formatAWFCommandWithPathSetup(writeAgentCLIStartMs string, config AWFCommandConfig, preCreateLog string, modelsJSONPathExport string, arcDind awfArcDindRuntimeConfig, toolCacheMountProbe string, awfCommand string, expandableArgs string, toolCacheMountRef string, joinedArgs string, shellWrappedCommand string, logFileArg string) string { + return fmt.Sprintf(`set -o pipefail %s %s %s @@ -533,23 +496,13 @@ fi`, %s %s %s %s %s %s \ -- %s 2>&1 | tee -a %s`, - writeAgentCLIStartMs, - config.PathSetup, - preCreateLog, - modelsJSONPathExport, - arcDindDockerHostProbe, - arcDindPrefixProbe, - toolCacheMountProbe, - awfShellcheckDirective, - awfCommand, - expandableArgs, - toolCacheMountRef, - arcDindDockerHostRef, - shellJoinArgs(awfArgs), - shellWrappedCommand, - shellEscapeArg(config.LogFile)) - } else if configFileSetup != "" { - command = fmt.Sprintf(`set -o pipefail + writeAgentCLIStartMs, config.PathSetup, preCreateLog, modelsJSONPathExport, arcDind.dockerHostProbe, + arcDind.prefixProbe, toolCacheMountProbe, awfShellcheckDirective, awfCommand, expandableArgs, + toolCacheMountRef, arcDind.dockerHostRef, joinedArgs, shellWrappedCommand, logFileArg) +} + +func formatAWFCommandWithConfig(writeAgentCLIStartMs string, preCreateLog string, configFileSetup string, modelsJSONPathExport string, arcDind awfArcDindRuntimeConfig, toolCacheMountProbe string, awfCommand string, expandableArgs string, toolCacheMountRef string, joinedArgs string, shellWrappedCommand string, logFileArg string) string { + return fmt.Sprintf(`set -o pipefail %s %s %s @@ -560,23 +513,13 @@ fi`, %s %s %s %s %s %s \ -- %s 2>&1 | tee -a %s`, - writeAgentCLIStartMs, - preCreateLog, - configFileSetup, - modelsJSONPathExport, - arcDindDockerHostProbe, - arcDindPrefixProbe, - toolCacheMountProbe, - awfShellcheckDirective, - awfCommand, - expandableArgs, - toolCacheMountRef, - arcDindDockerHostRef, - shellJoinArgs(awfArgs), - shellWrappedCommand, - shellEscapeArg(config.LogFile)) - } else { - command = fmt.Sprintf(`set -o pipefail + writeAgentCLIStartMs, preCreateLog, configFileSetup, modelsJSONPathExport, arcDind.dockerHostProbe, + arcDind.prefixProbe, toolCacheMountProbe, awfShellcheckDirective, awfCommand, expandableArgs, + toolCacheMountRef, arcDind.dockerHostRef, joinedArgs, shellWrappedCommand, logFileArg) +} + +func formatAWFCommand(writeAgentCLIStartMs string, preCreateLog string, modelsJSONPathExport string, arcDind awfArcDindRuntimeConfig, toolCacheMountProbe string, awfCommand string, expandableArgs string, toolCacheMountRef string, joinedArgs string, shellWrappedCommand string, logFileArg string) string { + return fmt.Sprintf(`set -o pipefail %s %s %s @@ -586,24 +529,9 @@ fi`, %s %s %s %s %s %s \ -- %s 2>&1 | tee -a %s`, - writeAgentCLIStartMs, - preCreateLog, - modelsJSONPathExport, - arcDindDockerHostProbe, - arcDindPrefixProbe, - toolCacheMountProbe, - awfShellcheckDirective, - awfCommand, - expandableArgs, - toolCacheMountRef, - arcDindDockerHostRef, - shellJoinArgs(awfArgs), - shellWrappedCommand, - shellEscapeArg(config.LogFile)) - } - - awfHelpersLog.Print("Successfully built AWF command") - return command + writeAgentCLIStartMs, preCreateLog, modelsJSONPathExport, arcDind.dockerHostProbe, + arcDind.prefixProbe, toolCacheMountProbe, awfShellcheckDirective, awfCommand, expandableArgs, + toolCacheMountRef, arcDind.dockerHostRef, joinedArgs, shellWrappedCommand, logFileArg) } // BuildAWFArgs constructs common AWF arguments from configuration. @@ -629,68 +557,64 @@ fi`, // --container-workdir and --mount are handled by BuildAWFCommand) func BuildAWFArgs(config AWFCommandConfig) []string { awfHelpersLog.Printf("Building AWF args for engine: %s", config.EngineName) - firewallConfig := getFirewallConfig(config.WorkflowData) agentConfig := getAgentConfig(config.WorkflowData) + awfArgs := appendAWFBaseArgs(nil, config, firewallConfig) + awfArgs = appendAWFEnvArgs(awfArgs, config, firewallConfig) + awfArgs = appendAWFMountArgs(awfArgs, agentConfig) + awfArgs = appendAWFLoggingArgs(awfArgs, config, firewallConfig) + awfArgs = appendAWFLegacySecurityArgs(awfArgs, config, firewallConfig, agentConfig) + awfArgs = appendAWFSkipPullAndCLIProxyArgs(awfArgs, config, firewallConfig) + awfArgs = appendAWFAPIBasePathArgs(awfArgs, config.WorkflowData) + awfArgs = append(awfArgs, getSSLBumpArgs(firewallConfig)...) + awfArgs = appendAWFCustomArgs(awfArgs, firewallConfig, agentConfig) + awfHelpersLog.Printf("Built %d AWF arguments", len(awfArgs)) + return awfArgs +} - var awfArgs []string - - // Add TTY flag if needed (Claude requires this), except for docker-sbx where - // sbx exec --tty can terminate long-running Claude sessions prematurely. +func appendAWFBaseArgs(awfArgs []string, config AWFCommandConfig, firewallConfig *FirewallConfig) []string { if config.UsesTTY && !isDockerSbxRuntime(config.WorkflowData) { awfArgs = append(awfArgs, "--tty") } - - // docker-sbx: tell AWF to launch the agent inside a Docker sbx microVM instead - // of as a standard Docker Compose service. Guard on the effective AWF version so - // older binaries do not receive an unknown flag. - if isDockerSbxRuntime(config.WorkflowData) && awfSupportsContainerRuntime(firewallConfig) { - awfArgs = append(awfArgs, "--container-runtime", "sbx") + if !isDockerSbxRuntime(config.WorkflowData) { + return awfArgs + } + if awfSupportsContainerRuntime(firewallConfig) { awfHelpersLog.Print("Added --container-runtime sbx for docker-sbx microVM runtime") - } else if isDockerSbxRuntime(config.WorkflowData) { - awfHelpersLog.Printf("Skipping --container-runtime sbx: AWF version %q is older than required minimum %s", getAWFImageTag(firewallConfig), constants.AWFContainerRuntimeMinVersion) - } - - // Pass all environment variables to the container, but exclude every variable whose - // step-env value comes from a GitHub Actions secret. AWF's API proxy (--enable-api-proxy) - // handles authentication for these tokens transparently, so the container does not need - // the raw values. Excluding them via --exclude-env prevents a prompt-injected agent from - // exfiltrating tokens through bash tools such as `env` or `printenv`. - // The caller computes ExcludeEnvVarNames from ComputeAWFExcludeEnvVarNames() so that every - // secret-bearing variable is covered — not just a hardcoded subset. - // --exclude-env requires AWF v0.25.3+; skip the flags for workflows that pin an older version. + return append(awfArgs, "--container-runtime", "sbx") + } + awfHelpersLog.Printf("Skipping --container-runtime sbx: AWF version %q is older than required minimum %s", getAWFImageTag(firewallConfig), constants.AWFContainerRuntimeMinVersion) + return awfArgs +} + +func appendAWFEnvArgs(awfArgs []string, config AWFCommandConfig, firewallConfig *FirewallConfig) []string { awfArgs = append(awfArgs, "--env-all") - if awfSupportsExcludeEnv(firewallConfig) { - // Sort for deterministic output in compiled lock files. - sortedExclude := make([]string, len(config.ExcludeEnvVarNames)) - copy(sortedExclude, config.ExcludeEnvVarNames) - sort.Strings(sortedExclude) - for _, excludedVar := range sortedExclude { - awfArgs = append(awfArgs, "--exclude-env", excludedVar) - } - } else { + if !awfSupportsExcludeEnv(firewallConfig) { awfHelpersLog.Printf("Skipping --exclude-env: AWF version %q is older than minimum %s", getAWFImageTag(firewallConfig), constants.AWFExcludeEnvMinVersion) + return awfArgs } + sortedExclude := append([]string(nil), config.ExcludeEnvVarNames...) + sort.Strings(sortedExclude) + for _, excludedVar := range sortedExclude { + awfArgs = append(awfArgs, "--exclude-env", excludedVar) + } + return awfArgs +} - // Note: --container-workdir "${GITHUB_WORKSPACE}" and --mount "${RUNNER_TEMP}/gh-aw:..." - // are intentionally NOT added here. They contain shell variable references that require - // double-quote expansion. These args are appended raw in BuildAWFCommand to ensure - // ${GITHUB_WORKSPACE} and ${RUNNER_TEMP} are expanded by the runner's shell. - - // Add custom mounts from agent config if specified - if agentConfig != nil && len(agentConfig.Mounts) > 0 { - // Sort mounts for consistent output - sortedMounts := make([]string, len(agentConfig.Mounts)) - copy(sortedMounts, agentConfig.Mounts) - sort.Strings(sortedMounts) - - for _, mount := range sortedMounts { - awfArgs = append(awfArgs, "--mount", mount) - } - awfHelpersLog.Printf("Added %d custom mounts from agent config", len(sortedMounts)) +func appendAWFMountArgs(awfArgs []string, agentConfig *AgentSandboxConfig) []string { + if agentConfig == nil || len(agentConfig.Mounts) == 0 { + return awfArgs } + sortedMounts := append([]string(nil), agentConfig.Mounts...) + sort.Strings(sortedMounts) + for _, mount := range sortedMounts { + awfArgs = append(awfArgs, "--mount", mount) + } + awfHelpersLog.Printf("Added %d custom mounts from agent config", len(sortedMounts)) + return awfArgs +} - // Set log level +func appendAWFLoggingArgs(awfArgs []string, config AWFCommandConfig, firewallConfig *FirewallConfig) []string { awfLogLevel := string(constants.AWFDefaultLogLevel) if firewallConfig != nil && firewallConfig.LogLevel != "" { awfLogLevel = firewallConfig.LogLevel @@ -700,101 +624,87 @@ func BuildAWFArgs(config AWFCommandConfig) []string { awfArgs = append(awfArgs, "--diagnostic-logs") awfHelpersLog.Print("Added --diagnostic-logs because awf-diagnostic-logs feature flag is enabled") } + return awfArgs +} - // Legacy security mode: emit --legacy-security, --enable-host-access, and --allow-host-ports - isLegacy := agentConfig != nil && agentConfig.LegacySecurity - if isLegacy { - if awfSupportsLegacySecurity(firewallConfig) { - awfArgs = append(awfArgs, "--legacy-security") - awfHelpersLog.Print("Added --legacy-security (legacy-security: enable in frontmatter)") - } else { - // AWF versions older than v0.27.32 don't support --legacy-security; - // they run in legacy mode by default so the flag is unnecessary. - awfHelpersLog.Printf("Skipping --legacy-security: AWF version %q is older than minimum %s (legacy mode is the default for older versions)", getAWFImageTag(firewallConfig), constants.AWFLegacySecurityMinVersion) - } - - awfArgs = append(awfArgs, "--enable-host-access") - awfHelpersLog.Print("Added --enable-host-access for legacy security mode") - - if awfSupportsAllowHostPorts(firewallConfig) { - mcpGatewayPort := int(DefaultMCPGatewayPort) - if config.WorkflowData != nil && config.WorkflowData.SandboxConfig != nil && - config.WorkflowData.SandboxConfig.MCP != nil && config.WorkflowData.SandboxConfig.MCP.Port > 0 { - mcpGatewayPort = config.WorkflowData.SandboxConfig.MCP.Port - } - hostPorts := fmt.Sprintf("80,443,%d", mcpGatewayPort) - awfArgs = append(awfArgs, "--allow-host-ports", hostPorts) - awfHelpersLog.Printf("Added --allow-host-ports %s for legacy security mode", hostPorts) - } - } else { +func appendAWFLegacySecurityArgs(awfArgs []string, config AWFCommandConfig, firewallConfig *FirewallConfig, agentConfig *AgentSandboxConfig) []string { + if agentConfig == nil || !agentConfig.LegacySecurity { awfHelpersLog.Print("Strict security: skipping host-access flags (default)") + return awfArgs } + awfArgs = appendAWFLegacySecurityModeArg(awfArgs, firewallConfig) + awfArgs = append(awfArgs, "--enable-host-access") + awfHelpersLog.Print("Added --enable-host-access for legacy security mode") + return appendAWFAllowHostPortsArg(awfArgs, config.WorkflowData, firewallConfig) +} - // Skip pulling images since they are pre-downloaded - awfArgs = append(awfArgs, "--skip-pull") - awfHelpersLog.Print("Using --skip-pull since images are pre-downloaded") - - // Enable CLI proxy sidecar when GitHub mode is gh-proxy. - // Start the difc-proxy on the host and tell AWF where to connect - // (firewall v0.25.17+). - if isGitHubCLIModeEnabled(config.WorkflowData) { - if awfSupportsCliProxy(firewallConfig) { - difcProxyHost := "host.docker.internal:18443" - if isAWFNetworkIsolationEnabled(config.WorkflowData) { - difcProxyHost = "awmg-cli-proxy:18443" - } - awfArgs = append(awfArgs, "--difc-proxy-host", difcProxyHost) - awfArgs = append(awfArgs, "--difc-proxy-ca-cert", constants.TmpDIFCProxyTLSCACert) - awfHelpersLog.Print("Added --difc-proxy-host and --difc-proxy-ca-cert for CLI proxy sidecar") - } else { - awfHelpersLog.Printf("Skipping CLI proxy flags: AWF version %q is older than minimum %s", getAWFImageTag(firewallConfig), constants.AWFCliProxyMinVersion) - } +func appendAWFLegacySecurityModeArg(awfArgs []string, firewallConfig *FirewallConfig) []string { + if awfSupportsLegacySecurity(firewallConfig) { + awfHelpersLog.Print("Added --legacy-security (legacy-security: enable in frontmatter)") + return append(awfArgs, "--legacy-security") } + awfHelpersLog.Printf("Skipping --legacy-security: AWF version %q is older than minimum %s (legacy mode is the default for older versions)", getAWFImageTag(firewallConfig), constants.AWFLegacySecurityMinVersion) + return awfArgs +} - // Pass base path if URL contains a path component - // This is required for endpoints with path prefixes (e.g., Databricks /serving-endpoints, - // Azure OpenAI /openai/deployments/, corporate LLM routers with path-based routing) - // Base paths remain as CLI flags — they are not yet represented in the config file schema. - openaiBasePath := extractAPIBasePath(config.WorkflowData, "OPENAI_BASE_URL") - if openaiBasePath != "" { - awfArgs = append(awfArgs, "--openai-api-base-path", openaiBasePath) - awfHelpersLog.Printf("Added --openai-api-base-path=%s", openaiBasePath) +func appendAWFAllowHostPortsArg(awfArgs []string, workflowData *WorkflowData, firewallConfig *FirewallConfig) []string { + if !awfSupportsAllowHostPorts(firewallConfig) { + return awfArgs } - - anthropicBasePath := extractAPIBasePath(config.WorkflowData, "ANTHROPIC_BASE_URL") - if anthropicBasePath != "" { - awfArgs = append(awfArgs, "--anthropic-api-base-path", anthropicBasePath) - awfHelpersLog.Printf("Added --anthropic-api-base-path=%s", anthropicBasePath) + mcpGatewayPort := int(DefaultMCPGatewayPort) + if workflowData != nil && workflowData.SandboxConfig != nil && workflowData.SandboxConfig.MCP != nil && workflowData.SandboxConfig.MCP.Port > 0 { + mcpGatewayPort = workflowData.SandboxConfig.MCP.Port } + hostPorts := fmt.Sprintf("80,443,%d", mcpGatewayPort) + awfHelpersLog.Printf("Added --allow-host-ports %s for legacy security mode", hostPorts) + return append(awfArgs, "--allow-host-ports", hostPorts) +} - geminiBasePath := extractAPIBasePath(config.WorkflowData, "GEMINI_API_BASE_URL") - if geminiBasePath != "" { - awfArgs = append(awfArgs, "--gemini-api-base-path", geminiBasePath) - awfHelpersLog.Printf("Added --gemini-api-base-path=%s", geminiBasePath) +func appendAWFSkipPullAndCLIProxyArgs(awfArgs []string, config AWFCommandConfig, firewallConfig *FirewallConfig) []string { + awfArgs = append(awfArgs, "--skip-pull") + awfHelpersLog.Print("Using --skip-pull since images are pre-downloaded") + if !isGitHubCLIModeEnabled(config.WorkflowData) { + return awfArgs + } + if !awfSupportsCliProxy(firewallConfig) { + awfHelpersLog.Printf("Skipping CLI proxy flags: AWF version %q is older than minimum %s", getAWFImageTag(firewallConfig), constants.AWFCliProxyMinVersion) + return awfArgs + } + difcProxyHost := "host.docker.internal:18443" + if isAWFNetworkIsolationEnabled(config.WorkflowData) { + difcProxyHost = "awmg-cli-proxy:18443" } + awfHelpersLog.Print("Added --difc-proxy-host and --difc-proxy-ca-cert for CLI proxy sidecar") + return append(awfArgs, "--difc-proxy-host", difcProxyHost, "--difc-proxy-ca-cert", constants.TmpDIFCProxyTLSCACert) +} + +func appendAWFAPIBasePathArgs(awfArgs []string, workflowData *WorkflowData) []string { + awfArgs = appendAWFAPIBasePathArg(awfArgs, workflowData, "OPENAI_BASE_URL", "--openai-api-base-path") + awfArgs = appendAWFAPIBasePathArg(awfArgs, workflowData, "ANTHROPIC_BASE_URL", "--anthropic-api-base-path") + return appendAWFAPIBasePathArg(awfArgs, workflowData, "GEMINI_API_BASE_URL", "--gemini-api-base-path") +} - // Add SSL Bump support for HTTPS content inspection (v0.9.0+) - sslBumpArgs := getSSLBumpArgs(firewallConfig) - awfArgs = append(awfArgs, sslBumpArgs...) +func appendAWFAPIBasePathArg(awfArgs []string, workflowData *WorkflowData, envVar string, flagName string) []string { + basePath := extractAPIBasePath(workflowData, envVar) + if basePath == "" { + return awfArgs + } + awfHelpersLog.Printf("Added %s=%s", flagName, basePath) + return append(awfArgs, flagName, basePath) +} - // Add custom args if specified in firewall config +func appendAWFCustomArgs(awfArgs []string, firewallConfig *FirewallConfig, agentConfig *AgentSandboxConfig) []string { if firewallConfig != nil && len(firewallConfig.Args) > 0 { awfArgs = append(awfArgs, firewallConfig.Args...) } - - // Add custom args from agent config if specified if agentConfig != nil && len(agentConfig.Args) > 0 { awfArgs = append(awfArgs, agentConfig.Args...) awfHelpersLog.Printf("Added %d custom args from agent config", len(agentConfig.Args)) } - - // Pass memory limit to AWF container if specified in agent config if agentConfig != nil && agentConfig.Memory != "" { awfArgs = append(awfArgs, "--memory-limit", agentConfig.Memory) awfHelpersLog.Printf("Set AWF memory limit to %s", agentConfig.Memory) } - - awfHelpersLog.Printf("Built %d AWF arguments", len(awfArgs)) return awfArgs } @@ -928,84 +838,85 @@ func WrapCommandInShell(command string) string { // - agent.env var names whose values contain ${{ secrets.* }} or a job-output expression // - names listed in the frontmatter excluded-env field (unconditionally) func ComputeAWFExcludeEnvVarNames(workflowData *WorkflowData, coreSecretVarNames []string) []string { - seen := make(map[string]struct { - }) - var names []string - - addUnique := func(name string) { - if !setutil.Contains(seen, name) { - seen[name] = struct { - }{} - names = append(names, name) - } + collector := newAWFExcludeEnvCollector() + collector.addAll(coreSecretVarNames) + collector.addRuntimeSecrets(workflowData) + collector.addMCPScriptEnv(workflowData) + collector.addWorkflowEnv(workflowData) + if isGitHubCLIModeEnabled(workflowData) { + collector.add("GH_TOKEN") } - - // Core secret vars for this engine (always contain secret references). - for _, name := range coreSecretVarNames { - addUnique(name) + if workflowData != nil { + collector.addAll(workflowData.ExcludedEnv) } + awfHelpersLog.Printf("Computed %d AWF env vars to exclude", len(collector.names)) + return collector.names +} - // MCP gateway API key is always a secret when MCP servers are present. - if HasMCPServers(workflowData) { - addUnique("MCP_GATEWAY_API_KEY") - } +type awfExcludeEnvCollector struct { + seen map[string]struct{} + names []string +} - // GitHub MCP server token is always a secret when the GitHub tool is present. - if hasGitHubTool(workflowData.ParsedTools) { - addUnique("GITHUB_MCP_SERVER_TOKEN") - } +func newAWFExcludeEnvCollector() *awfExcludeEnvCollector { + return &awfExcludeEnvCollector{seen: make(map[string]struct{})} +} - // HTTP MCP header secrets: values are always ${{ secrets.* }} references. - for varName := range collectHTTPMCPHeaderSecrets(workflowData.Tools) { - addUnique(varName) - } - - // mcp-scripts env vars: only add those whose configured values contain a secret reference - // or a job-output expression (e.g. ${{ needs.fetch_token.outputs.token }}). - // (Non-secret vars like GH_DEBUG: "1" must NOT be excluded.) - if workflowData.MCPScripts != nil { - for _, toolConfig := range workflowData.MCPScripts.Tools { - for envName, envValue := range toolConfig.Env { - if strings.Contains(envValue, "${{ secrets.") || ContainsJobOutputExpr(envValue) { - addUnique(envName) - } - } - } +func (c *awfExcludeEnvCollector) add(name string) { + if !setutil.Contains(c.seen, name) { + c.seen[name] = struct{}{} + c.names = append(c.names, name) } +} - // engine.env vars that contain a secret reference or a job-output expression. - if workflowData.EngineConfig != nil { - for varName, varValue := range workflowData.EngineConfig.Env { - if strings.Contains(varValue, "${{ secrets.") || ContainsJobOutputExpr(varValue) { - addUnique(varName) - } - } +func (c *awfExcludeEnvCollector) addAll(names []string) { + for _, name := range names { + c.add(name) } +} - // agent.env vars that contain a secret reference or a job-output expression. - agentConfig := getAgentConfig(workflowData) - if agentConfig != nil { - for varName, varValue := range agentConfig.Env { - if strings.Contains(varValue, "${{ secrets.") || ContainsJobOutputExpr(varValue) { - addUnique(varName) - } +func (c *awfExcludeEnvCollector) addSecretLikeEnv(env map[string]string) { + for envName, envValue := range env { + if strings.Contains(envValue, "${{ secrets.") || ContainsJobOutputExpr(envValue) { + c.add(envName) } } +} - // GH_TOKEN when GitHub mode is gh-proxy: the token is passed in the AWF step env for the - // host difc-proxy but must be excluded from the agent container. - if isGitHubCLIModeEnabled(workflowData) { - addUnique("GH_TOKEN") +func (c *awfExcludeEnvCollector) addRuntimeSecrets(workflowData *WorkflowData) { + if workflowData == nil { + return } + if HasMCPServers(workflowData) { + c.add("MCP_GATEWAY_API_KEY") + } + if hasGitHubTool(workflowData.ParsedTools) { + c.add("GITHUB_MCP_SERVER_TOKEN") + } + for varName := range collectHTTPMCPHeaderSecrets(workflowData.Tools) { + c.add(varName) + } +} - // Explicitly excluded env vars from the frontmatter excluded-env field. - // These are always excluded regardless of their value content. - for _, name := range workflowData.ExcludedEnv { - addUnique(name) +func (c *awfExcludeEnvCollector) addMCPScriptEnv(workflowData *WorkflowData) { + if workflowData == nil || workflowData.MCPScripts == nil { + return } + for _, toolConfig := range workflowData.MCPScripts.Tools { + c.addSecretLikeEnv(toolConfig.Env) + } +} - awfHelpersLog.Printf("Computed %d AWF env vars to exclude", len(names)) - return names +func (c *awfExcludeEnvCollector) addWorkflowEnv(workflowData *WorkflowData) { + if workflowData == nil { + return + } + if workflowData.EngineConfig != nil { + c.addSecretLikeEnv(workflowData.EngineConfig.Env) + } + if agentConfig := getAgentConfig(workflowData); agentConfig != nil { + c.addSecretLikeEnv(agentConfig.Env) + } } // addCliProxyGHTokenToEnv adds GH_TOKEN to the AWF step environment when GitHub diff --git a/pkg/workflow/codex_logs.go b/pkg/workflow/codex_logs.go index b130bc9fc3e..e06c7ebc281 100644 --- a/pkg/workflow/codex_logs.go +++ b/pkg/workflow/codex_logs.go @@ -12,77 +12,35 @@ import ( var codexLogsLog = logger.New("workflow:codex_logs") +type codexLogParseState struct { + currentSequence []string + inThinking bool + lastToolName string + toolCallMap map[string]*ToolCallInfo + tokenUsage int + turns int +} + // ParseLogMetrics implements engine-specific log parsing for Codex func (e *CodexEngine) ParseLogMetrics(logContent string, verbose bool) LogMetrics { codexLogsLog.Printf("Parsing Codex log metrics: log_size=%d bytes, lines=%d", len(logContent), strings.Count(logContent, "\n")+1) - var metrics LogMetrics - var totalTokenUsage int - lines := strings.Split(logContent, "\n") - turns := 0 - inThinkingSection := false - toolCallMap := make(map[string]*ToolCallInfo) // Track tool calls - var currentSequence []string // Track tool sequence - var lastToolName string // Track most recent tool for output size extraction + state := codexLogParseState{ + toolCallMap: make(map[string]*ToolCallInfo), + } for i := range lines { - line := lines[i] - - // Skip empty lines - if strings.TrimSpace(line) == "" { - continue - } - - // Detect thinking sections as indicators of turns - // Support both old format: "] thinking" and new Rust format: "thinking" (standalone line) - trimmedLine := strings.TrimSpace(line) - if strings.Contains(line, "] thinking") || trimmedLine == "thinking" { - if !inThinkingSection { - turns++ - inThinkingSection = true - // Start of a new thinking section, save previous sequence if any - if len(currentSequence) > 0 { - metrics.ToolSequences = append(metrics.ToolSequences, currentSequence) - currentSequence = []string{} - } - } - } else if strings.Contains(line, "] tool") || strings.Contains(line, "] exec") || strings.Contains(line, "] codex") || - strings.HasPrefix(trimmedLine, "tool ") || strings.HasPrefix(trimmedLine, "exec ") { - inThinkingSection = false - } - - // Extract tool calls from Codex logs and add to sequence - if toolName := e.parseCodexToolCallsWithSequence(line, toolCallMap); toolName != "" { - currentSequence = append(currentSequence, toolName) - lastToolName = toolName - } - - // Extract output size from success/failure lines followed by JSON blocks - if outputSize := e.extractOutputSizeFromResult(line, lines, i); outputSize > 0 && lastToolName != "" { - if toolInfo, exists := toolCallMap[lastToolName]; exists { - if outputSize > toolInfo.MaxOutputSize { - toolInfo.MaxOutputSize = outputSize - codexLogsLog.Printf("Updated %s MaxOutputSize to %d characters", lastToolName, outputSize) - } - } - } - - // Extract Codex-specific token usage (always sum for Codex) - if tokenUsage := e.extractCodexTokenUsage(line); tokenUsage > 0 { - totalTokenUsage += tokenUsage - } - - // Basic processing - error/warning counting moved to end of function + e.processCodexLogLine(lines[i], lines, i, &state) } - // Finalize metrics using shared helper + var metrics LogMetrics FinalizeToolMetrics(FinalizeToolMetricsOptions{ Metrics: &metrics, - ToolCallMap: toolCallMap, - CurrentSequence: currentSequence, - Turns: turns, - TokenUsage: totalTokenUsage, + ToolCallMap: state.toolCallMap, + CurrentSequence: state.currentSequence, + Turns: state.turns, + TokenUsage: state.tokenUsage, }) codexLogsLog.Printf("Parsed Codex metrics: turns=%d, token_usage=%d, tool_calls=%d", @@ -91,109 +49,137 @@ func (e *CodexEngine) ParseLogMetrics(logContent string, verbose bool) LogMetric return metrics } +func (e *CodexEngine) processCodexLogLine(line string, lines []string, index int, state *codexLogParseState) { + if strings.TrimSpace(line) == "" { + return + } + e.updateCodexThinkingState(line, state) + if toolName := e.parseCodexToolCallsWithSequence(line, state.toolCallMap); toolName != "" { + state.currentSequence = append(state.currentSequence, toolName) + state.lastToolName = toolName + } + e.updateCodexOutputSize(line, lines, index, state) + state.tokenUsage += e.extractCodexTokenUsage(line) +} + +func (e *CodexEngine) updateCodexThinkingState(line string, state *codexLogParseState) { + trimmedLine := strings.TrimSpace(line) + if strings.Contains(line, "] thinking") || trimmedLine == "thinking" { + if state.inThinking { + return + } + state.turns++ + state.inThinking = true + if len(state.currentSequence) > 0 { + state.currentSequence = append([]string{}, state.currentSequence...) + } + return + } + if strings.Contains(line, "] tool") || strings.Contains(line, "] exec") || strings.Contains(line, "] codex") || + strings.HasPrefix(trimmedLine, "tool ") || strings.HasPrefix(trimmedLine, "exec ") { + if len(state.currentSequence) > 0 && state.inThinking { + state.currentSequence = append([]string{}, state.currentSequence...) + } + state.inThinking = false + } +} + +func (e *CodexEngine) updateCodexOutputSize(line string, lines []string, index int, state *codexLogParseState) { + outputSize := e.extractOutputSizeFromResult(line, lines, index) + if outputSize == 0 || state.lastToolName == "" { + return + } + toolInfo, exists := state.toolCallMap[state.lastToolName] + if !exists || outputSize <= toolInfo.MaxOutputSize { + return + } + toolInfo.MaxOutputSize = outputSize + codexLogsLog.Printf("Updated %s MaxOutputSize to %d characters", state.lastToolName, outputSize) +} + // parseCodexToolCallsWithSequence extracts tool call information from Codex log lines and returns tool name func (e *CodexEngine) parseCodexToolCallsWithSequence(line string, toolCallMap map[string]*ToolCallInfo) string { trimmedLine := strings.TrimSpace(line) + if toolName := parseCodexToolName(line, trimmedLine); toolName != "" { + prettifiedName := prettifyCodexToolName(toolName) + incrementToolCallCount(toolCallMap, prettifiedName) + return prettifiedName + } + if execCommand := parseCodexExecCommand(line, trimmedLine); execCommand != "" { + uniqueBashName := "bash_" + ShortenCommand(execCommand) + incrementToolCallCount(toolCallMap, uniqueBashName) + return uniqueBashName + } + e.updateCodexToolDuration(line, toolCallMap) + return "" +} - // Parse tool calls: "] tool provider.method(...)" (old format) - // or "tool provider.method(...)" (new Rust format) - var toolName string - - // Try old format first: "] tool provider.method(...)" +func parseCodexToolName(line, trimmedLine string) string { if strings.Contains(line, "] tool ") && strings.Contains(line, "(") { if match := codexToolCallOldFormat.FindStringSubmatch(line); len(match) > 1 { - toolName = strings.TrimSpace(match[1]) + return strings.TrimSpace(match[1]) } } - - // Try new Rust format: "tool provider.method(...)" - if toolName == "" && strings.HasPrefix(trimmedLine, "tool ") && strings.Contains(trimmedLine, "(") { + if strings.HasPrefix(trimmedLine, "tool ") && strings.Contains(trimmedLine, "(") { if match := codexToolCallNewFormat.FindStringSubmatch(trimmedLine); len(match) > 1 { - toolName = strings.TrimSpace(match[1]) + return strings.TrimSpace(match[1]) } } + return "" +} - if toolName != "" { - prettifiedName := PrettifyToolName(toolName) - - // For Codex, format provider.method as provider_method (avoiding colons) - if strings.Contains(toolName, ".") { - parts := strings.Split(toolName, ".") - if len(parts) >= 2 { - provider := parts[0] - method := strings.Join(parts[1:], "_") - prettifiedName = fmt.Sprintf("%s_%s", provider, method) - } - } - - // Initialize or update tool call info - if toolInfo, exists := toolCallMap[prettifiedName]; exists { - toolInfo.CallCount++ - } else { - toolCallMap[prettifiedName] = &ToolCallInfo{ - Name: prettifiedName, - CallCount: 1, - MaxOutputSize: 0, // Will be updated when output is extracted from result lines - MaxDuration: 0, // Will be updated when duration is found - } - } - +func prettifyCodexToolName(toolName string) string { + prettifiedName := PrettifyToolName(toolName) + if !strings.Contains(toolName, ".") { return prettifiedName } + parts := strings.Split(toolName, ".") + if len(parts) < 2 { + return prettifiedName + } + return fmt.Sprintf("%s_%s", parts[0], strings.Join(parts[1:], "_")) +} - // Parse exec commands: "] exec command" (old format) - // or "exec command in" (new Rust format) - treat as bash calls - var execCommand string - - // Try old format: "] exec command in" +func parseCodexExecCommand(line, trimmedLine string) string { if strings.Contains(line, "] exec ") { if match := codexExecCommandOldFormat.FindStringSubmatch(line); len(match) > 1 { - execCommand = strings.TrimSpace(match[1]) + return strings.TrimSpace(match[1]) } } - - // Try new Rust format: "exec command in" - if execCommand == "" && strings.HasPrefix(trimmedLine, "exec ") { + if strings.HasPrefix(trimmedLine, "exec ") { if match := codexExecCommandNewFormat.FindStringSubmatch(trimmedLine); len(match) > 1 { - execCommand = strings.TrimSpace(match[1]) + return strings.TrimSpace(match[1]) } } + return "" +} - if execCommand != "" { - // Create unique bash entry with command info, avoiding colons - uniqueBashName := "bash_" + ShortenCommand(execCommand) - - // Initialize or update tool call info - if toolInfo, exists := toolCallMap[uniqueBashName]; exists { - toolInfo.CallCount++ - } else { - toolCallMap[uniqueBashName] = &ToolCallInfo{ - Name: uniqueBashName, - CallCount: 1, - MaxOutputSize: 0, - MaxDuration: 0, // Will be updated when duration is found - } - } - - return uniqueBashName +func incrementToolCallCount(toolCallMap map[string]*ToolCallInfo, name string) { + if toolInfo, exists := toolCallMap[name]; exists { + toolInfo.CallCount++ + return } - - // Parse duration from success/failure lines: "] success in 0.2s" or "] failure in 1.5s" - if strings.Contains(line, "success in") || strings.Contains(line, "failure in") || strings.Contains(line, "failed in") { - // Extract duration pattern like "in 0.2s", "in 1.5s" - if match := codexDurationPattern.FindStringSubmatch(line); len(match) > 1 { - if durationSeconds, err := strconv.ParseFloat(match[1], 64); err == nil { - duration := time.Duration(durationSeconds * float64(time.Second)) - - // Find the most recent tool call to associate with this duration - // Since we don't have direct association, we'll update the most recent entry - // This is a limitation of the log format, but it's the best we can do - e.updateMostRecentToolWithDuration(toolCallMap, duration) - } - } + toolCallMap[name] = &ToolCallInfo{ + Name: name, + CallCount: 1, + MaxOutputSize: 0, + MaxDuration: 0, } +} - return "" // No tool call found +func (e *CodexEngine) updateCodexToolDuration(line string, toolCallMap map[string]*ToolCallInfo) { + if !strings.Contains(line, "success in") && !strings.Contains(line, "failure in") && !strings.Contains(line, "failed in") { + return + } + match := codexDurationPattern.FindStringSubmatch(line) + if len(match) <= 1 { + return + } + durationSeconds, err := strconv.ParseFloat(match[1], 64) + if err != nil { + return + } + e.updateMostRecentToolWithDuration(toolCallMap, time.Duration(durationSeconds*float64(time.Second))) } // updateMostRecentToolWithDuration updates the tool with maximum duration diff --git a/pkg/workflow/compiler_activation_daily_aic.go b/pkg/workflow/compiler_activation_daily_aic.go index 99534aa05a5..0b0e4ca0814 100644 --- a/pkg/workflow/compiler_activation_daily_aic.go +++ b/pkg/workflow/compiler_activation_daily_aic.go @@ -99,69 +99,81 @@ func (c *Compiler) resolveDailyAICToken(data *WorkflowData) string { func (c *Compiler) buildActivationDailyAICGuardrailStep(data *WorkflowData) []string { compilerActivationJobLog.Printf("Building daily AIC guardrail step: dedicated_app=%t, cache_enabled=%t", data.MaxDailyAICreditsGitHubApp != nil, data.WorkflowID != "") var steps []string - // When a dedicated GitHub App is configured for the daily AIC guardrail, mint - // its token first so the subsequent steps can reference it. + steps = append(steps, c.buildDailyAICTokenAndCacheSteps(data)...) + steps = append(steps, c.buildDailyAICGuardrailCheckStep(data)...) + return steps +} + +func (c *Compiler) buildDailyAICTokenAndCacheSteps(data *WorkflowData) []string { + var steps []string if data.MaxDailyAICreditsGitHubApp != nil { compilerActivationJobLog.Print("Prepending dedicated daily-AIC app-token mint step") steps = append(steps, c.buildDailyAICAppTokenMintStep(data.MaxDailyAICreditsGitHubApp)...) } - // Prepend cache restore step so cached AIC values from prior runs are available - // when the guardrail script runs, allowing it to skip artifact downloads. if data.WorkflowID != "" { - sanitized := SanitizeWorkflowIDForCacheKey(data.WorkflowID) - cacheKeyPrefix := fmt.Sprintf("agentic-workflow-usage-%s-", sanitized) - steps = append(steps, " - name: Restore daily AIC usage cache\n") - steps = append(steps, " id: restore-daily-aic-cache\n") - steps = append(steps, fmt.Sprintf(" if: %s\n", maxDailyAICreditsConfiguredIfExpr)) - steps = append(steps, " continue-on-error: true\n") - steps = append(steps, fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/cache/restore", data))) - steps = append(steps, " with:\n") - steps = append(steps, fmt.Sprintf(" key: %s${{ github.run_id }}\n", cacheKeyPrefix)) - steps = append(steps, fmt.Sprintf(" restore-keys: %s\n", cacheKeyPrefix)) - steps = append(steps, " path: /tmp/gh-aw/agentic-workflow-usage-cache.jsonl\n") - // Artifact-based fallback for cross-branch cache misses. - // GitHub Actions actions/cache is branch-scoped: caches written by the conclusion job - // on one PR branch are invisible to the activation job running on a different PR branch. - // This step downloads the most recent aic-usage-cache artifact uploaded by a prior - // conclusion job so that the guardrail script can skip per-run artifact downloads. - // Cache-miss detection is performed inside restore_aic_usage_cache_fallback.cjs using - // the cache restore outputs forwarded via env vars. - steps = append(steps, " - name: Restore daily AIC usage cache (artifact fallback)\n") - steps = append(steps, " id: restore-daily-aic-cache-fallback\n") - steps = append(steps, fmt.Sprintf(" if: %s\n", maxDailyAICreditsConfiguredIfExpr)) - steps = append(steps, " continue-on-error: true\n") - steps = append(steps, fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/github-script", data))) - steps = append(steps, " env:\n") - steps = append(steps, " GH_AW_RESTORE_DAILY_AIC_CACHE_HIT: ${{ steps.restore-daily-aic-cache.outputs.cache-hit }}\n") - steps = append(steps, " GH_AW_RESTORE_DAILY_AIC_CACHE_MATCHED_KEY: ${{ steps.restore-daily-aic-cache.outputs.cache-matched-key }}\n") - steps = append(steps, " with:\n") - steps = append(steps, fmt.Sprintf(" github-token: %s\n", c.resolveDailyAICToken(data))) - steps = append(steps, " script: |\n") - steps = append(steps, " const { setupGlobals } = require('"+SetupActionDestination+"/setup_globals.cjs');\n") - steps = append(steps, " setupGlobals(core, github, context, exec, io, getOctokit);\n") - steps = append(steps, " const { main } = require('"+SetupActionDestination+"/restore_aic_usage_cache_fallback.cjs');\n") - steps = append(steps, " await main();\n") + steps = append(steps, c.buildDailyAICCacheRestoreSteps(data)...) + } + return steps +} + +func (c *Compiler) buildDailyAICCacheRestoreSteps(data *WorkflowData) []string { + sanitized := SanitizeWorkflowIDForCacheKey(data.WorkflowID) + cacheKeyPrefix := fmt.Sprintf("agentic-workflow-usage-%s-", sanitized) + token := c.resolveDailyAICToken(data) + return []string{ + " - name: Restore daily AIC usage cache\n", + " id: restore-daily-aic-cache\n", + fmt.Sprintf(" if: %s\n", maxDailyAICreditsConfiguredIfExpr), + " continue-on-error: true\n", + fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/cache/restore", data)), + " with:\n", + fmt.Sprintf(" key: %s${{ github.run_id }}\n", cacheKeyPrefix), + fmt.Sprintf(" restore-keys: %s\n", cacheKeyPrefix), + " path: /tmp/gh-aw/agentic-workflow-usage-cache.jsonl\n", + " - name: Restore daily AIC usage cache (artifact fallback)\n", + " id: restore-daily-aic-cache-fallback\n", + fmt.Sprintf(" if: %s\n", maxDailyAICreditsConfiguredIfExpr), + " continue-on-error: true\n", + fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/github-script", data)), + " env:\n", + " GH_AW_RESTORE_DAILY_AIC_CACHE_HIT: ${{ steps.restore-daily-aic-cache.outputs.cache-hit }}\n", + " GH_AW_RESTORE_DAILY_AIC_CACHE_MATCHED_KEY: ${{ steps.restore-daily-aic-cache.outputs.cache-matched-key }}\n", + " with:\n", + fmt.Sprintf(" github-token: %s\n", token), + " script: |\n", + " const { setupGlobals } = require('" + SetupActionDestination + "/setup_globals.cjs');\n", + " setupGlobals(core, github, context, exec, io, getOctokit);\n", + " const { main } = require('" + SetupActionDestination + "/restore_aic_usage_cache_fallback.cjs');\n", + " await main();\n", + } +} + +func (c *Compiler) buildDailyAICGuardrailCheckStep(data *WorkflowData) []string { + token := c.resolveDailyAICToken(data) + steps := []string{ + " - name: Check daily workflow token guardrail\n", + " id: daily-effective-workflow-guardrail\n", + fmt.Sprintf(" if: %s\n", maxDailyAICreditsConfiguredIfExpr), + fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/github-script", data)), + " env:\n", + fmt.Sprintf(" GH_AW_WORKFLOW_NAME: %q\n", data.Name), + fmt.Sprintf(" GH_AW_WORKFLOW_ID: %q\n", data.WorkflowID), + " GH_AW_RUN_URL: ${{ github.server_url }}/${{ github.repository }}/actions/runs/${{ github.run_id }}\n", + " GH_AW_WORKFLOW_DISPATCH_AW_CONTEXT: ${{ github.event.inputs.aw_context || '' }}\n", + fmt.Sprintf(" GH_AW_HAS_SLASH_COMMAND: %q\n", strconv.FormatBool(len(data.Command) > 0)), + fmt.Sprintf(" GH_AW_HAS_LABEL_COMMAND: %q\n", strconv.FormatBool(len(data.LabelCommand) > 0)), + fmt.Sprintf(" GH_AW_GITHUB_TOKEN: %s\n", token), } - steps = append(steps, " - name: Check daily workflow token guardrail\n") - steps = append(steps, " id: daily-effective-workflow-guardrail\n") - steps = append(steps, fmt.Sprintf(" if: %s\n", maxDailyAICreditsConfiguredIfExpr)) - steps = append(steps, fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/github-script", data))) - steps = append(steps, " env:\n") - steps = append(steps, fmt.Sprintf(" GH_AW_WORKFLOW_NAME: %q\n", data.Name)) - steps = append(steps, fmt.Sprintf(" GH_AW_WORKFLOW_ID: %q\n", data.WorkflowID)) - steps = append(steps, " GH_AW_RUN_URL: ${{ github.server_url }}/${{ github.repository }}/actions/runs/${{ github.run_id }}\n") - steps = append(steps, " GH_AW_WORKFLOW_DISPATCH_AW_CONTEXT: ${{ github.event.inputs.aw_context || '' }}\n") - steps = append(steps, fmt.Sprintf(" GH_AW_HAS_SLASH_COMMAND: %q\n", strconv.FormatBool(len(data.Command) > 0))) - steps = append(steps, fmt.Sprintf(" GH_AW_HAS_LABEL_COMMAND: %q\n", strconv.FormatBool(len(data.LabelCommand) > 0))) - steps = append(steps, fmt.Sprintf(" GH_AW_GITHUB_TOKEN: %s\n", c.resolveDailyAICToken(data))) steps = append(steps, buildTemplatableIntEnvVar(maxDailyAICreditsEnvVar, data.MaxDailyAICredits)...) - steps = append(steps, " with:\n") - steps = append(steps, fmt.Sprintf(" github-token: %s\n", c.resolveDailyAICToken(data))) - steps = append(steps, " script: |\n") - steps = append(steps, " const { setupGlobals } = require('"+SetupActionDestination+"/setup_globals.cjs');\n") - steps = append(steps, " setupGlobals(core, github, context, exec, io, getOctokit);\n") - steps = append(steps, " const { main } = require('"+SetupActionDestination+"/check_daily_aic_workflow_guardrail.cjs');\n") - steps = append(steps, " await main();\n") + steps = append(steps, + " with:\n", + fmt.Sprintf(" github-token: %s\n", token), + " script: |\n", + " const { setupGlobals } = require('"+SetupActionDestination+"/setup_globals.cjs');\n", + " setupGlobals(core, github, context, exec, io, getOctokit);\n", + " const { main } = require('"+SetupActionDestination+"/check_daily_aic_workflow_guardrail.cjs');\n", + " await main();\n", + ) return steps } diff --git a/pkg/workflow/compiler_activation_job.go b/pkg/workflow/compiler_activation_job.go index 504b9a95cbd..c7b940d0312 100644 --- a/pkg/workflow/compiler_activation_job.go +++ b/pkg/workflow/compiler_activation_job.go @@ -54,23 +54,7 @@ func (c *Compiler) buildActivationJob(data *WorkflowData, preActivationJobCreate return nil, fmt.Errorf("failed to add activation command and label outputs: %w", err) } ctx.steps = append(ctx.steps, buildRuntimeFeaturesSummaryStep()...) - - // Generate experiment selection steps when experiments are declared in the frontmatter. - // These steps run before the prompt is built so that experiments.name expressions - // can be resolved by the substitute_placeholders step. - if experimentSteps := c.generateExperimentSteps(data); len(experimentSteps) > 0 { - compilerActivationJobLog.Printf("Adding %d experiment step(s) for %d experiment(s)", len(experimentSteps), len(data.Experiments)) - ctx.steps = append(ctx.steps, experimentSteps...) - // Expose the combined experiment JSON as a job output so downstream jobs can access - // the variant assignments via needs.activation.outputs.experiments. - ctx.outputs["experiments"] = "${{ steps.pick-experiment.outputs.experiments }}" - // Also expose each experiment variant individually so downstream jobs can reference - // needs.activation.outputs. in timeout-minutes or other expressions. - for _, name := range sortedExperimentNames(data.Experiments) { - ctx.outputs[name] = "${{ steps.pick-experiment.outputs." + name + " }}" - } - } - + c.addActivationExperimentOutputs(ctx) c.configureActivationNeedsAndCondition(ctx) compilerActivationJobLog.Print("Generating prompt in activation job") c.generatePromptInActivationJob(&ctx.steps, data, preActivationJobCreated, ctx.customJobsBeforeActivation) @@ -79,15 +63,34 @@ func (c *Compiler) buildActivationJob(data *WorkflowData, preActivationJobCreate ctx.steps = append(ctx.steps, " - run: echo \"Activation success\"\n") } - if c.actionMode.IsScript() { - ctx.steps = append(ctx.steps, c.generateScriptModeCleanupStep()) - } - permissions, err := c.buildActivationPermissions(ctx) if err != nil { return nil, fmt.Errorf("failed to build activation permissions: %w", err) } + c.addActivationScriptCleanupIfNeeded(ctx) + return c.finalizeActivationJob(data, workflowRunRepoSafety, ctx, permissions), nil +} + +func (c *Compiler) addActivationExperimentOutputs(ctx *activationJobBuildContext) { + experimentSteps := c.generateExperimentSteps(ctx.data) + if len(experimentSteps) == 0 { + return + } + compilerActivationJobLog.Printf("Adding %d experiment step(s) for %d experiment(s)", len(experimentSteps), len(ctx.data.Experiments)) + ctx.steps = append(ctx.steps, experimentSteps...) + ctx.outputs["experiments"] = "${{ steps.pick-experiment.outputs.experiments }}" + for _, name := range sortedExperimentNames(ctx.data.Experiments) { + ctx.outputs[name] = "${{ steps.pick-experiment.outputs." + name + " }}" + } +} + +func (c *Compiler) addActivationScriptCleanupIfNeeded(ctx *activationJobBuildContext) { + if c.actionMode.IsScript() { + ctx.steps = append(ctx.steps, c.generateScriptModeCleanupStep()) + } +} +func (c *Compiler) finalizeActivationJob(data *WorkflowData, workflowRunRepoSafety string, ctx *activationJobBuildContext, permissions string) *Job { return &Job{ Name: string(constants.ActivationJobName), If: ctx.activationCondition, @@ -99,7 +102,7 @@ func (c *Compiler) buildActivationJob(data *WorkflowData, preActivationJobCreate Steps: ctx.steps, Outputs: ctx.outputs, Needs: ctx.activationNeeds, - }, nil + } } func addActivationInteractionPermissions( @@ -392,46 +395,48 @@ func (c *Compiler) generateResolveHostRepoStep(data *WorkflowData) string { // runs before the agent job and needs independent access to workflow files for runtime imports during // prompt generation. func (c *Compiler) generateCheckoutGitHubFolderForActivation(data *WorkflowData) []string { - // Check if action-tag is specified - if so, skip checkout - if data != nil && data.Features != nil { - if actionTagVal, exists := data.Features["action-tag"]; exists { - if actionTagStr, ok := actionTagVal.(string); ok && actionTagStr != "" { - // action-tag is set, no checkout needed - compilerActivationJobLog.Print("Skipping .github checkout in activation: action-tag specified") - return nil - } - } + if shouldSkipActivationGitHubFolderCheckout(data) { + compilerActivationJobLog.Print("Skipping .github checkout in activation: action-tag specified") + return nil } + extraPaths := c.activationSparseCheckoutExtraPaths(data) + activationToken := c.resolveActivationToken(data) + if data != nil && hasWorkflowCallTrigger(data.On) && !data.InlinedImports { + return c.generateWorkflowCallActivationCheckout(activationToken, extraPaths) + } + compilerActivationJobLog.Print("Adding .github, .agents, and engine-specific dirs to sparse checkout for activation job") + return NewCheckoutManager(nil).GenerateGitHubFolderCheckoutStep("", "", activationToken, c.getActionPin, extraPaths...) +} - // Note: We don't check data.Permissions for contents read access here because - // the activation job ALWAYS gets contents:read added to its permissions (see buildActivationJob - // around line 720). The workflow's original permissions may not include contents:read, - // but the activation job will always have it for GitHub API access and runtime imports. - // The agent job uses only the user-specified permissions (no automatic contents:read augmentation). - - // For workflow_call triggers, checkout the callee (platform) repository using the target_repo - // and target_checkout_ref outputs from the resolve-host-repo step. That step uses - // job.workflow_repository and job.workflow_sha to identify the platform repo and pin to the - // exact commit, correctly handling all relay patterns including cross-repo and cross-org scenarios. - // (target_checkout_ref carries the SHA; target_ref carries the dispatch-compatible branch/tag ref.) - // - // Skip when inlined-imports is enabled: content is embedded at compile time and no - // runtime-import macros are used, so the callee's .md files are not needed at runtime. - // In dev mode, actions/setup is referenced via a local workspace path (./actions/setup), - // so it must be included in the sparse-checkout to preserve it for the post step. - // In release/script/action modes, the action is in the runner cache and not the workspace. +func shouldSkipActivationGitHubFolderCheckout(data *WorkflowData) bool { + if data == nil || data.Features == nil { + return false + } + actionTagStr, _ := data.Features["action-tag"].(string) + return actionTagStr != "" +} + +func (c *Compiler) activationSparseCheckoutExtraPaths(data *WorkflowData) []string { var extraPaths []string if c.actionMode.IsDev() { compilerActivationJobLog.Print("Dev mode: adding actions/setup to sparse-checkout to preserve local action post step") extraPaths = append(extraPaths, "actions/setup") } + extraPaths = append(extraPaths, activationEngineSpecificSparseCheckoutDirs(data)...) + repoRoot := c.gitRoot + if repoRoot == "" { + if cwd, err := os.Getwd(); err == nil { + repoRoot = cwd + } + } + extraPaths = resolveSymlinkExtraPaths(repoRoot, extraPaths) + compilerActivationJobLog.Printf("Adding %d engine-specific dirs to sparse-checkout: %v", len(extraPaths), extraPaths) + return extraPaths +} - // Add engine-specific agent config directories to the sparse checkout. - // .github and .agents are already included in GenerateGitHubFolderCheckoutStep's hardcoded list. - // Root instruction files (AGENTS.md, CLAUDE.md, GEMINI.md) are excluded — they are not needed - // during activation and are omitted to keep the shallow checkout minimal. - defaultSparseCheckoutDirs := map[string]struct { - }{".github": {}, ".agents": {}} +func activationEngineSpecificSparseCheckoutDirs(data *WorkflowData) []string { + defaultSparseCheckoutDirs := map[string]struct{}{".github": {}, ".agents": {}} + var extraPaths []string registry := GetGlobalEngineRegistry() for _, folder := range registry.GetAllAgentManifestFolders() { if !setutil.Contains(defaultSparseCheckoutDirs, folder) { @@ -443,50 +448,26 @@ func (c *Compiler) generateCheckoutGitHubFolderForActivation(data *WorkflowData) extraPaths = append(extraPaths, folder) } } - compilerActivationJobLog.Printf("Adding %d engine-specific dirs to sparse-checkout: %v", len(extraPaths), extraPaths) - - // Detect symlinks for well-known .github sub-paths and add their resolved targets - // so that sparse checkout fetches the target directory, not just the symlink blob. - // Use c.gitRoot so detection works regardless of the process CWD. - repoRoot := c.gitRoot - if repoRoot == "" { - if cwd, err := os.Getwd(); err == nil { - repoRoot = cwd - } - } - extraPaths = resolveSymlinkExtraPaths(repoRoot, extraPaths) + return extraPaths +} +func (c *Compiler) generateWorkflowCallActivationCheckout(activationToken string, extraPaths []string) []string { + compilerActivationJobLog.Print("Adding cross-repo-aware .github checkout for workflow_call trigger") cm := NewCheckoutManager(nil) - activationToken := c.resolveActivationToken(data) - if data != nil && hasWorkflowCallTrigger(data.On) && !data.InlinedImports { - compilerActivationJobLog.Print("Adding cross-repo-aware .github checkout for workflow_call trigger") - cm.SetCrossRepoTargetRepo("${{ steps.resolve-host-repo.outputs.target_repo }}") - cm.SetCrossRepoTargetRef("${{ steps.resolve-host-repo.outputs.target_checkout_ref }}") - checkoutSteps := cm.GenerateGitHubFolderCheckoutStep( - cm.GetCrossRepoTargetRepo(), - cm.GetCrossRepoTargetRef(), - activationToken, - c.getActionPin, - extraPaths..., - ) - // When no custom token is configured, GITHUB_TOKEN is scoped to the calling - // repository and cannot read a private callee repository in cross-repo invocations - // (e.g. nbcnews/tvOS-App calling nbcnews/.github). Add an if: condition so the - // checkout is only attempted for same-repo invocations where GITHUB_TOKEN works. - // For cross-repo scenarios, users can enable the checkout by configuring - // activation-github-token or activation-github-app in the workflow frontmatter. - if activationToken == "${{ secrets.GITHUB_TOKEN }}" { - compilerActivationJobLog.Print("No custom activation token — restricting cross-repo checkout to same-repo invocations") - checkoutSteps = addSameRepoIfConditionToSteps(checkoutSteps) - } - return checkoutSteps - } - - // For activation job, sparse checkout .github, .agents, and engine-specific config directories - // (plus actions/setup in dev mode). Root instruction files are excluded as they are not needed - // during activation. sparse-checkout-cone-mode: true ensures subdirectories are recursively included. - compilerActivationJobLog.Print("Adding .github, .agents, and engine-specific dirs to sparse checkout for activation job") - return cm.GenerateGitHubFolderCheckoutStep("", "", activationToken, c.getActionPin, extraPaths...) + cm.SetCrossRepoTargetRepo("${{ steps.resolve-host-repo.outputs.target_repo }}") + cm.SetCrossRepoTargetRef("${{ steps.resolve-host-repo.outputs.target_checkout_ref }}") + checkoutSteps := cm.GenerateGitHubFolderCheckoutStep( + cm.GetCrossRepoTargetRepo(), + cm.GetCrossRepoTargetRef(), + activationToken, + c.getActionPin, + extraPaths..., + ) + if activationToken == "${{ secrets.GITHUB_TOKEN }}" { + compilerActivationJobLog.Print("No custom activation token — restricting cross-repo checkout to same-repo invocations") + return addSameRepoIfConditionToSteps(checkoutSteps) + } + return checkoutSteps } func localSkillSparseCheckoutTopLevelDirs(data *WorkflowData) []string { diff --git a/pkg/workflow/compiler_activation_permissions.go b/pkg/workflow/compiler_activation_permissions.go index b5599ed47bd..d69faff4a3d 100644 --- a/pkg/workflow/compiler_activation_permissions.go +++ b/pkg/workflow/compiler_activation_permissions.go @@ -38,61 +38,41 @@ func activationJobNeedsAppToken(ctx *activationJobBuildContext) bool { func buildActivationAppTokenPermissions(ctx *activationJobBuildContext) *Permissions { appPerms := NewPermissions() - addActivationInteractionPermissions( - appPerms, - activationInteractionPermissionsOptions{ - onSection: ctx.data.On, - hasReaction: ctx.hasReaction, - reactionIncludesIssues: ctx.reactionIssues, - reactionIncludesPullRequests: ctx.reactionPullRequests, - reactionIncludesDiscussions: ctx.reactionDiscussions, - hasStatusComment: ctx.hasStatusComment, - statusCommentIncludesIssues: ctx.statusCommentIssues, - statusCommentIncludesPullRequests: ctx.statusCommentPRs, - statusCommentIncludesDiscussions: ctx.statusCommentDiscussions, - }, - ) + addActivationAppTokenInteractionPermissions(appPerms, ctx) + addActivationAppTokenOperationalPermissions(appPerms, ctx) + addActivationAppTokenInferredPermissions(appPerms, ctx.activationInferredPerms) + return appPerms +} + +func addActivationAppTokenInteractionPermissions(appPerms *Permissions, ctx *activationJobBuildContext) { + options := buildActivationInteractionPermissionOptions(ctx, ctx.data.On) + addActivationInteractionPermissions(appPerms, options) if ctx.data.CommandCentralized && (ctx.hasReaction || ctx.hasStatusComment) { - syntheticOn := buildCentralizedCommandOnSection(ctx.data.CommandEvents) - if syntheticOn != "" { - addActivationInteractionPermissions( - appPerms, - activationInteractionPermissionsOptions{ - onSection: syntheticOn, - hasReaction: ctx.hasReaction, - reactionIncludesIssues: ctx.reactionIssues, - reactionIncludesPullRequests: ctx.reactionPullRequests, - reactionIncludesDiscussions: ctx.reactionDiscussions, - hasStatusComment: ctx.hasStatusComment, - statusCommentIncludesIssues: ctx.statusCommentIssues, - statusCommentIncludesPullRequests: ctx.statusCommentPRs, - statusCommentIncludesDiscussions: ctx.statusCommentDiscussions, - }, - ) + if syntheticOn := buildCentralizedCommandOnSection(ctx.data.CommandEvents); syntheticOn != "" { + addActivationInteractionPermissions(appPerms, buildActivationInteractionPermissionOptions(ctx, syntheticOn)) } } if hasWorkflowCallTrigger(ctx.data.On) && (ctx.hasReaction || ctx.hasStatusComment) { - addActivationInteractionPermissions( - appPerms, - activationInteractionPermissionsOptions{ - hasReaction: ctx.hasReaction, - reactionIncludesIssues: ctx.reactionIssues, - reactionIncludesPullRequests: ctx.reactionPullRequests, - reactionIncludesDiscussions: ctx.reactionDiscussions, - hasStatusComment: ctx.hasStatusComment, - statusCommentIncludesIssues: ctx.statusCommentIssues, - statusCommentIncludesPullRequests: ctx.statusCommentPRs, - statusCommentIncludesDiscussions: ctx.statusCommentDiscussions, - }, - ) + options.onSection = "" + addActivationInteractionPermissions(appPerms, options) + } +} + +func buildActivationInteractionPermissionOptions(ctx *activationJobBuildContext, onSection string) activationInteractionPermissionsOptions { + return activationInteractionPermissionsOptions{ + onSection: onSection, + hasReaction: ctx.hasReaction, + reactionIncludesIssues: ctx.reactionIssues, + reactionIncludesPullRequests: ctx.reactionPullRequests, + reactionIncludesDiscussions: ctx.reactionDiscussions, + hasStatusComment: ctx.hasStatusComment, + statusCommentIncludesIssues: ctx.statusCommentIssues, + statusCommentIncludesPullRequests: ctx.statusCommentPRs, + statusCommentIncludesDiscussions: ctx.statusCommentDiscussions, } - // Keep this aligned with addActivationLabelPermissions: app-token scopes are - // computed separately from GITHUB_TOKEN scopes because app-token permissions - // only apply to steps using the minted app token, while label permissions in - // addActivationLabelPermissions are only for GITHUB_TOKEN execution paths. - // This intentionally mirrors addActivationLabelPermissions without the - // ActivationGitHubApp == nil guard because this function runs only when - // activationJobNeedsAppToken confirms app-token minting is enabled. +} + +func addActivationAppTokenOperationalPermissions(appPerms *Permissions, ctx *activationJobBuildContext) { if ctx.shouldRemoveLabel { if slices.Contains(ctx.filteredLabelEvents, "issues") || slices.Contains(ctx.filteredLabelEvents, "pull_request") { appPerms.Set(PermissionIssues, PermissionWrite) @@ -107,15 +87,14 @@ func buildActivationAppTokenPermissions(ctx *activationJobBuildContext) *Permiss if hasMaxDailyAICGuardrail(ctx.data) { appPerms.Set(PermissionActions, PermissionRead) } - // Add GitHub App-only permissions inferred from activation job gh CLI commands so the - // minted App token includes the scopes those commands require (e.g. codespaces: read - // for `gh codespace list`). Only App-only scopes are passed here. - for scope, level := range ctx.activationInferredPerms { +} + +func addActivationAppTokenInferredPermissions(appPerms *Permissions, inferred map[PermissionScope]PermissionLevel) { + for scope, level := range inferred { if IsGitHubAppOnlyScope(scope) { appPerms.Set(scope, level) } } - return appPerms } // buildActivationPermissions builds activation job permissions from workflow features and selected interactions. diff --git a/pkg/workflow/compiler_activation_steps.go b/pkg/workflow/compiler_activation_steps.go index 18032b918ec..d9c29174c92 100644 --- a/pkg/workflow/compiler_activation_steps.go +++ b/pkg/workflow/compiler_activation_steps.go @@ -171,78 +171,104 @@ func (c *Compiler) addActivationVersionCheckStep(ctx *activationJobBuildContext) } func (c *Compiler) addActivationSkillInstallSteps(ctx *activationJobBuildContext) { - skillRefs := append([]SkillReference(nil), ctx.data.SkillReferences...) - if len(skillRefs) == 0 && len(ctx.data.Skills) > 0 { - skillRefs = make([]SkillReference, 0, len(ctx.data.Skills)) - for _, skill := range ctx.data.Skills { - if strings.TrimSpace(skill) == "" { - continue - } - skillRefs = append(skillRefs, SkillReference{Skill: skill}) - } - } + skillRefs := activationSkillReferences(ctx.data) if len(skillRefs) == 0 { return } engineID := resolveActivationEngineID(ctx.data) - skillDir := GetEngineSkillDir(engineID) + skillDir, skillInstallAgentName := activationSkillInstallMetadata(engineID) + ctx.steps = append(ctx.steps, buildActivationSkillInstallPrereqSteps()...) + for i, skillRef := range skillRefs { + ctx.steps = append(ctx.steps, c.buildActivationSkillInstallStep(ctx, i+1, skillRef, engineID, skillDir, skillInstallAgentName)...) + } + ctx.steps = append(ctx.steps, buildActivationSkillFailureCollectionSteps(ctx.data)...) + ctx.outputs["skill_install_failure_count"] = "${{ steps.collect-skill-install-failures.outputs.failure_count || '0' }}" + ctx.outputs["skill_install_errors"] = "${{ steps.collect-skill-install-failures.outputs.errors || '' }}" +} + +func activationSkillReferences(data *WorkflowData) []SkillReference { + skillRefs := append([]SkillReference(nil), data.SkillReferences...) + if len(skillRefs) > 0 || len(data.Skills) == 0 { + return skillRefs + } + skillRefs = make([]SkillReference, 0, len(data.Skills)) + for _, skill := range data.Skills { + if strings.TrimSpace(skill) != "" { + skillRefs = append(skillRefs, SkillReference{Skill: skill}) + } + } + return skillRefs +} + +func activationSkillInstallMetadata(engineID string) (string, string) { skillInstallAgentName := "" if engine, err := GetGlobalEngineRegistry().GetEngine(strings.ToLower(engineID)); err == nil { skillInstallAgentName = engine.GetGHSkillAgentName() } + return GetEngineSkillDir(engineID), skillInstallAgentName +} - ctx.steps = append(ctx.steps, " - name: Upgrade gh CLI for frontmatter skills\n") - ctx.steps = append(ctx.steps, fmt.Sprintf(" run: bash \"${RUNNER_TEMP}/gh-aw/actions/ensure_gh_cli_min_version.sh\" \"%s\"\n", constants.GhSkillsMinVersion)) - - for i, skillRef := range skillRefs { - tokenExpr := c.resolveActivationToken(ctx.data) - if skillRef.GitHubToken != "" { - tokenExpr = skillRef.GitHubToken - } - if skillRef.GitHubApp != nil { - stepNumber := i + 1 - stepID := fmt.Sprintf("frontmatter-skill-app-token-%d", stepNumber) - ctx.steps = append(ctx.steps, c.buildGitHubAppTokenMintStepWithMeta( - skillRef.GitHubApp, - nil, - "", - "", - fmt.Sprintf("Generate GitHub App token for frontmatter skill %d", stepNumber), - stepID, - )...) - stepTokenExpr := fmt.Sprintf("${{ steps.%s.outputs.token }}", stepID) - if skillRef.GitHubApp.shouldIgnoreMissingKey() { - tokenExpr = combineTokenExpressions(stepTokenExpr, c.resolveActivationToken(ctx.data)) - } else { - tokenExpr = stepTokenExpr - } - } - ctx.steps = append(ctx.steps, fmt.Sprintf(" - name: Install frontmatter skill %d\n", i+1)) - ctx.steps = append(ctx.steps, " env:\n") - ctx.steps = append(ctx.steps, fmt.Sprintf(" GH_TOKEN: %s\n", tokenExpr)) - ctx.steps = append(ctx.steps, formatYAMLEnv(" ", "GH_AW_INFO_ENGINE_ID", engineID)) - ctx.steps = append(ctx.steps, formatYAMLEnv(" ", "GH_AW_GH_SKILL_AGENT_NAME", skillInstallAgentName)) - ctx.steps = append(ctx.steps, formatYAMLEnv(" ", "GH_AW_SKILL_DIR", skillDir)) - ctx.steps = append(ctx.steps, formatYAMLEnv(" ", "GH_AW_FRONTMATTER_SKILLS", skillRef.Skill)) - ctx.steps = append(ctx.steps, fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/github-script", ctx.data))) - ctx.steps = append(ctx.steps, " with:\n") - ctx.steps = append(ctx.steps, " script: |\n") - ctx.steps = append(ctx.steps, generateGitHubScriptWithRequire("install_frontmatter_skills.cjs")) +func buildActivationSkillInstallPrereqSteps() []string { + return []string{ + " - name: Upgrade gh CLI for frontmatter skills\n", + fmt.Sprintf(" run: bash \"${RUNNER_TEMP}/gh-aw/actions/ensure_gh_cli_min_version.sh\" \"%s\"\n", constants.GhSkillsMinVersion), } +} - // Collect skill install failures written by each install step into a shared file. - // Runs with if: always() so failures are captured even if a prior step was unexpectedly hard-failed. - ctx.steps = append(ctx.steps, " - name: Collect skill install failures\n") - ctx.steps = append(ctx.steps, " id: collect-skill-install-failures\n") - ctx.steps = append(ctx.steps, " if: always()\n") - ctx.steps = append(ctx.steps, fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/github-script", ctx.data))) - ctx.steps = append(ctx.steps, " with:\n") - ctx.steps = append(ctx.steps, " script: |\n") - ctx.steps = append(ctx.steps, generateGitHubScriptWithRequire("collect_skill_install_failures.cjs")) +func (c *Compiler) buildActivationSkillInstallStep(ctx *activationJobBuildContext, stepNumber int, skillRef SkillReference, engineID string, skillDir string, skillInstallAgentName string) []string { + tokenExpr, mintSteps := c.resolveActivationSkillToken(ctx, stepNumber, skillRef) + steps := append([]string{}, mintSteps...) + steps = append(steps, + fmt.Sprintf(" - name: Install frontmatter skill %d\n", stepNumber), + " env:\n", + fmt.Sprintf(" GH_TOKEN: %s\n", tokenExpr), + formatYAMLEnv(" ", "GH_AW_INFO_ENGINE_ID", engineID), + formatYAMLEnv(" ", "GH_AW_GH_SKILL_AGENT_NAME", skillInstallAgentName), + formatYAMLEnv(" ", "GH_AW_SKILL_DIR", skillDir), + formatYAMLEnv(" ", "GH_AW_FRONTMATTER_SKILLS", skillRef.Skill), + fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/github-script", ctx.data)), + " with:\n", + " script: |\n", + generateGitHubScriptWithRequire("install_frontmatter_skills.cjs"), + ) + return steps +} - ctx.outputs["skill_install_failure_count"] = "${{ steps.collect-skill-install-failures.outputs.failure_count || '0' }}" - ctx.outputs["skill_install_errors"] = "${{ steps.collect-skill-install-failures.outputs.errors || '' }}" +func (c *Compiler) resolveActivationSkillToken(ctx *activationJobBuildContext, stepNumber int, skillRef SkillReference) (string, []string) { + tokenExpr := c.resolveActivationToken(ctx.data) + if skillRef.GitHubToken != "" { + tokenExpr = skillRef.GitHubToken + } + if skillRef.GitHubApp == nil { + return tokenExpr, nil + } + stepID := fmt.Sprintf("frontmatter-skill-app-token-%d", stepNumber) + mintSteps := c.buildGitHubAppTokenMintStepWithMeta( + skillRef.GitHubApp, + nil, + "", + "", + fmt.Sprintf("Generate GitHub App token for frontmatter skill %d", stepNumber), + stepID, + ) + stepTokenExpr := fmt.Sprintf("${{ steps.%s.outputs.token }}", stepID) + if skillRef.GitHubApp.shouldIgnoreMissingKey() { + return combineTokenExpressions(stepTokenExpr, c.resolveActivationToken(ctx.data)), mintSteps + } + return stepTokenExpr, mintSteps +} + +func buildActivationSkillFailureCollectionSteps(data *WorkflowData) []string { + return []string{ + " - name: Collect skill install failures\n", + " id: collect-skill-install-failures\n", + " if: always()\n", + fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/github-script", data)), + " with:\n", + " script: |\n", + generateGitHubScriptWithRequire("collect_skill_install_failures.cjs"), + } } func (c *Compiler) addActivationTextOutputStep(ctx *activationJobBuildContext) error { diff --git a/pkg/workflow/compiler_aw_context.go b/pkg/workflow/compiler_aw_context.go index 0e57b7f7c46..ad83af0cc89 100644 --- a/pkg/workflow/compiler_aw_context.go +++ b/pkg/workflow/compiler_aw_context.go @@ -57,32 +57,36 @@ func injectInputIntoTrigger(onSection string, triggerName string, inputName stri awContextLog.Printf("Injecting %s input into %s trigger", inputName, triggerName) lines := strings.Split(onSection, "\n") + triggerLineIdx, triggerIndent := findBareTriggerLine(lines, triggerName) + if triggerLineIdx == -1 { + awContextLog.Printf("No bare %s: line found, skipping %s injection", triggerName, inputName) + return onSection + } + awContextLog.Printf("Found %s at line %d (indent=%d), injecting %s", triggerName, triggerLineIdx, triggerIndent, inputName) + inputsLineIdx := findTriggerInputsLine(lines, triggerLineIdx, triggerIndent) + if hasExistingInjectedInput(lines, inputsLineIdx, inputName) { + awContextLog.Printf("%s already injected into %s, skipping", inputName, triggerName) + return onSection + } + inputLines := buildInputLines(triggerIndent) + return strings.Join(injectTriggerInputLines(lines, triggerLineIdx, triggerIndent, inputsLineIdx, inputLines), "\n") +} - // Find the trigger line (bare — no sub-value on same line) - triggerLineIdx := -1 - triggerIndent := 0 +func findBareTriggerLine(lines []string, triggerName string) (int, int) { for i, line := range lines { stripped := strings.TrimLeft(line, " \t") rest, found := strings.CutPrefix(stripped, triggerName+":") if found { rest = strings.TrimSpace(rest) if rest == "" || rest == "null" || rest == "~" { - triggerLineIdx = i - triggerIndent = len(line) - len(stripped) - break + return i, len(line) - len(stripped) } } } + return -1, 0 +} - if triggerLineIdx == -1 { - awContextLog.Printf("No bare %s: line found, skipping %s injection", triggerName, inputName) - return onSection - } - awContextLog.Printf("Found %s at line %d (indent=%d), injecting %s", triggerName, triggerLineIdx, triggerIndent, inputName) - - // Look for an "inputs:" key directly inside the trigger block. - // Only the first non-empty, non-comment line after the trigger matters. - inputsLineIdx := -1 +func findTriggerInputsLine(lines []string, triggerLineIdx int, triggerIndent int) int { for i := triggerLineIdx + 1; i < len(lines); i++ { stripped := strings.TrimLeft(lines[i], " \t") if stripped == "" || strings.HasPrefix(stripped, "#") { @@ -90,55 +94,60 @@ func injectInputIntoTrigger(onSection string, triggerName string, inputName stri } lineIndent := len(lines[i]) - len(stripped) if lineIndent <= triggerIndent { - break // left workflow_dispatch block entirely + break } if strings.HasPrefix(stripped, "inputs:") { - inputsLineIdx = i + return i } - break // only inspect the first substantive child key + break } + return -1 +} - if inputsLineIdx != -1 { - inputsIndent := len(lines[inputsLineIdx]) - len(strings.TrimLeft(lines[inputsLineIdx], " \t")) - for i := inputsLineIdx + 1; i < len(lines); i++ { - stripped := strings.TrimLeft(lines[i], " \t") - if stripped == "" || strings.HasPrefix(stripped, "#") { - continue - } - lineIndent := len(lines[i]) - len(stripped) - if lineIndent <= inputsIndent { - break - } - if strings.HasPrefix(stripped, inputName+":") { - awContextLog.Printf("%s already injected into %s, skipping", inputName, triggerName) - return onSection - } +func hasExistingInjectedInput(lines []string, inputsLineIdx int, inputName string) bool { + if inputsLineIdx == -1 { + return false + } + inputsIndent := len(lines[inputsLineIdx]) - len(strings.TrimLeft(lines[inputsLineIdx], " \t")) + for i := inputsLineIdx + 1; i < len(lines); i++ { + stripped := strings.TrimLeft(lines[i], " \t") + if stripped == "" || strings.HasPrefix(stripped, "#") { + continue + } + if len(lines[i])-len(stripped) <= inputsIndent { + break + } + if strings.HasPrefix(stripped, inputName+":") { + return true } } + return false +} - inputLines := buildInputLines(triggerIndent) - +func injectTriggerInputLines(lines []string, triggerLineIdx int, triggerIndent int, inputsLineIdx int, inputLines []string) []string { result := make([]string, 0, safeAllocationCapacity(len(lines), len(inputLines), 1)) for i, line := range lines { - // When the trigger line contains an explicit null/~ value, - // replace it with a bare trigger so sub-keys can follow. - if i == triggerLineIdx && (strings.HasSuffix(strings.TrimSpace(line), " null") || - strings.HasSuffix(strings.TrimSpace(line), " ~")) { - stripped := strings.TrimLeft(line, " \t") - line = strings.Repeat(" ", triggerIndent) + strings.SplitN(stripped, ":", 2)[0] + ":" + if i == triggerLineIdx { + line = normalizeNullTriggerLine(line, triggerIndent) } result = append(result, line) - if inputsLineIdx != -1 && i == inputsLineIdx { result = append(result, inputLines...) } else if inputsLineIdx == -1 && i == triggerLineIdx { - // Trigger is bare — add inputs: + the requested internal input. result = append(result, strings.Repeat(" ", triggerIndent+2)+"inputs:") result = append(result, inputLines...) } } + return result +} - return strings.Join(result, "\n") +func normalizeNullTriggerLine(line string, triggerIndent int) string { + trimmed := strings.TrimSpace(line) + if !strings.HasSuffix(trimmed, " null") && !strings.HasSuffix(trimmed, " ~") { + return line + } + stripped := strings.TrimLeft(line, " \t") + return strings.Repeat(" ", triggerIndent) + strings.SplitN(stripped, ":", 2)[0] + ":" } // buildAwContextInputLines returns the indented YAML lines for the aw_context input diff --git a/pkg/workflow/compiler_custom_jobs.go b/pkg/workflow/compiler_custom_jobs.go index 5bc343abff2..5c86d0ad8f4 100644 --- a/pkg/workflow/compiler_custom_jobs.go +++ b/pkg/workflow/compiler_custom_jobs.go @@ -178,51 +178,20 @@ func (c *Compiler) extractCustomJobCoreProperties(job *Job, jobName string, conf if _, hasInputs := configMap["inputs"]; hasInputs { return fmt.Errorf("jobs.%s.inputs: inputs are not supported on jobs; use 'env' to pass values to job steps", jobName) } - if err := c.extractCustomJobRunsOn(job, jobName, configMap); err != nil { return err } - - if ifCond, hasIf := configMap["if"]; hasIf { - if ifStr, ok := ifCond.(string); ok { - job.If = c.extractExpressionFromIfString(ifStr) - } - } - - if permissions, hasPermissions := configMap["permissions"]; hasPermissions { - formattedPerms := NewPermissionsParserFromValue(permissions).ToPermissions().RenderToYAML() - if formattedPerms != "" { - job.Permissions = formattedPerms - } - } - - if strategy, hasStrategy := configMap["strategy"]; hasStrategy { - if strategyMap, ok := strategy.(map[string]any); ok { - formattedStrategy, err := formatIndentedYAMLField("strategy", strategyMap, false) - if err != nil { - return fmt.Errorf("failed to convert strategy to YAML for job '%s': %w", jobName, err) - } - job.Strategy = formattedStrategy - } - } - - // Extract name (display name) for custom jobs - if name, hasName := configMap["name"]; hasName { - if nameStr, ok := name.(string); ok { - job.DisplayName = nameStr - } + applyCustomJobConditionalFields(c, job, configMap) + if err := extractCustomJobStrategy(job, jobName, configMap); err != nil { + return err } - if err := extractCustomJobTimeoutMinutes(job, jobName, configMap); err != nil { return err } - if err := extractCustomJobConcurrency(job, jobName, configMap); err != nil { return err } - extractCustomJobEnv(job, configMap) - if err := extractCustomJobContainer(job, jobName, configMap); err != nil { return err } @@ -238,6 +207,41 @@ func (c *Compiler) extractCustomJobCoreProperties(job *Job, jobName string, conf return nil } +func applyCustomJobConditionalFields(c *Compiler, job *Job, configMap map[string]any) { + if ifCond, hasIf := configMap["if"]; hasIf { + if ifStr, ok := ifCond.(string); ok { + job.If = c.extractExpressionFromIfString(ifStr) + } + } + if permissions, hasPermissions := configMap["permissions"]; hasPermissions { + if formattedPerms := NewPermissionsParserFromValue(permissions).ToPermissions().RenderToYAML(); formattedPerms != "" { + job.Permissions = formattedPerms + } + } + if name, hasName := configMap["name"]; hasName { + if nameStr, ok := name.(string); ok { + job.DisplayName = nameStr + } + } +} + +func extractCustomJobStrategy(job *Job, jobName string, configMap map[string]any) error { + strategy, hasStrategy := configMap["strategy"] + if !hasStrategy { + return nil + } + strategyMap, ok := strategy.(map[string]any) + if !ok { + return nil + } + formattedStrategy, err := formatIndentedYAMLField("strategy", strategyMap, false) + if err != nil { + return fmt.Errorf("failed to convert strategy to YAML for job '%s': %w", jobName, err) + } + job.Strategy = formattedStrategy + return nil +} + func (c *Compiler) extractCustomJobRunsOn(job *Job, jobName string, configMap map[string]any) error { runsOn, hasRunsOn := configMap["runs-on"] if !hasRunsOn { @@ -481,94 +485,86 @@ func configureCustomReusableWorkflow(job *Job, jobName string, usesStr string, c } func (c *Compiler) configureCustomJobSteps(job *Job, jobName string, configMap map[string]any, data *WorkflowData) error { + c.ensureCustomJobRunsOn(job, data) + setupSteps, preSteps, regularSteps, hasInjectedSteps, err := c.extractCustomJobStepGroups(jobName, configMap, data) + if err != nil { + return err + } + restoreMemCfg, err := extractRestoreMemoryConfig(configMap, jobName, data) + if err != nil { + return err + } + hasRestoreMemory := restoreMemCfg != nil + applyRestoreMemoryEnv(job, restoreMemCfg, data) + if hasInjectedSteps || hasRestoreMemory { + injectedSteps, err := c.buildConfiguredCustomJobSteps(jobName, data, restoreMemCfg, setupSteps, preSteps, regularSteps) + if err != nil { + return err + } + job.Steps = append(job.Steps, injectedSteps...) + } + return nil +} + +func (c *Compiler) ensureCustomJobRunsOn(job *Job, data *WorkflowData) { if job.RunsOn == "" { job.RunsOn = c.indentYAMLLines(data.RunsOn, " ") if job.RunsOn == "" { job.RunsOn = "runs-on: ubuntu-latest" } } +} - // Add basic steps if specified (only for non-reusable workflow jobs). - // `setup-steps` and `pre-steps` stay distinct so setup-steps can remain the - // first injected steps in the job, followed by compiler scaffolding, - // `pre-steps`, and the regular `steps` list. - var setupSteps []string - var preSteps []string - var regularSteps []string - _, hasSetupStepsField := configMap["setup-steps"] - _, hasPreStepsField := configMap["pre-steps"] - _, hasStepsField := configMap["steps"] - - if hasSetupStepsField { - var err error - setupSteps, err = c.extractPinnedJobSteps("setup-steps", jobName, configMap, data) - if err != nil { - return fmt.Errorf("failed to process setup-steps for job '%s': %w", jobName, err) - } +func (c *Compiler) extractCustomJobStepGroups(jobName string, configMap map[string]any, data *WorkflowData) ([]string, []string, []string, bool, error) { + setupSteps, hasSetupStepsField, err := c.extractPinnedOptionalJobSteps("setup-steps", jobName, configMap, data) + if err != nil { + return nil, nil, nil, false, fmt.Errorf("failed to process setup-steps for job '%s': %w", jobName, err) } - if hasPreStepsField { - var err error - preSteps, err = c.extractPinnedJobSteps("pre-steps", jobName, configMap, data) - if err != nil { - return fmt.Errorf("failed to process pre-steps for job '%s': %w", jobName, err) - } + preSteps, hasPreStepsField, err := c.extractPinnedOptionalJobSteps("pre-steps", jobName, configMap, data) + if err != nil { + return nil, nil, nil, false, fmt.Errorf("failed to process pre-steps for job '%s': %w", jobName, err) } - if hasStepsField { - var err error - regularSteps, err = c.extractPinnedJobSteps("steps", jobName, configMap, data) - if err != nil { - return fmt.Errorf("failed to process steps for job '%s': %w", jobName, err) - } + regularSteps, hasStepsField, err := c.extractPinnedOptionalJobSteps("steps", jobName, configMap, data) + if err != nil { + return nil, nil, nil, false, fmt.Errorf("failed to process steps for job '%s': %w", jobName, err) } + return setupSteps, preSteps, regularSteps, hasSetupStepsField || hasPreStepsField || hasStepsField, nil +} - // Parse restore-memory configuration. - // restore-memory injects read-only memory restore steps into the custom job. - // No write-back or commit steps are ever emitted for memory in custom jobs. - restoreMemCfg, err := extractRestoreMemoryConfig(configMap, jobName, data) - if err != nil { - return err +func (c *Compiler) extractPinnedOptionalJobSteps(fieldName string, jobName string, configMap map[string]any, data *WorkflowData) ([]string, bool, error) { + if _, hasField := configMap[fieldName]; !hasField { + return nil, false, nil } + steps, err := c.extractPinnedJobSteps(fieldName, jobName, configMap, data) + return steps, true, err +} - hasRestoreMemory := restoreMemCfg != nil +func applyRestoreMemoryEnv(job *Job, restoreMemCfg *restoreMemoryConfig, data *WorkflowData) { + if restoreMemCfg == nil || !restoreMemCfg.CacheMemory || data.WorkflowID == "" { + return + } + if job.Env == nil { + job.Env = make(map[string]string) + } + if _, alreadySet := job.Env["GH_AW_WORKFLOW_ID_SANITIZED"]; !alreadySet { + job.Env["GH_AW_WORKFLOW_ID_SANITIZED"] = SanitizeWorkflowIDForCacheKey(data.WorkflowID) + } +} - // When cache-memory restore is requested, inject GH_AW_WORKFLOW_ID_SANITIZED so that - // restore keys match those used by the agent job. Only set it when the user has not - // already provided the variable in their job's env: block. - if hasRestoreMemory && restoreMemCfg.CacheMemory && data.WorkflowID != "" { - sanitized := SanitizeWorkflowIDForCacheKey(data.WorkflowID) - if job.Env == nil { - job.Env = make(map[string]string) - } - if _, alreadySet := job.Env["GH_AW_WORKFLOW_ID_SANITIZED"]; !alreadySet { - job.Env["GH_AW_WORKFLOW_ID_SANITIZED"] = sanitized - } - } - - if hasSetupStepsField || hasPreStepsField || hasStepsField || hasRestoreMemory { - job.Steps = append(job.Steps, setupSteps...) - // Prepend GH_HOST configuration step for GHES/GHEC compatibility. - // Custom frontmatter jobs run as independent GitHub Actions jobs that - // don't inherit GITHUB_ENV from the agent job, so the gh CLI won't - // know which host to target without this step. - job.Steps = append(job.Steps, generateGHESHostConfigurationStep()) - - // Inject gh-aw setup + memory restore steps when restore-memory is requested. - // Setup lines come first (they install scripts needed by repo/comment memory). - // Memory lines follow immediately after (restore/clone/prepare steps). - if hasRestoreMemory { - memorySetupLines, memoryRestoreLines, memErr := c.buildRestoreMemorySteps(restoreMemCfg, jobName, data) - if memErr != nil { - return memErr - } - job.Steps = append(job.Steps, memorySetupLines...) - job.Steps = append(job.Steps, memoryRestoreLines...) +func (c *Compiler) buildConfiguredCustomJobSteps(jobName string, data *WorkflowData, restoreMemCfg *restoreMemoryConfig, setupSteps []string, preSteps []string, regularSteps []string) ([]string, error) { + steps := append([]string{}, setupSteps...) + steps = append(steps, generateGHESHostConfigurationStep()) + if restoreMemCfg != nil { + memorySetupLines, memoryRestoreLines, err := c.buildRestoreMemorySteps(restoreMemCfg, jobName, data) + if err != nil { + return nil, err } - - job.Steps = append(job.Steps, preSteps...) - job.Steps = append(job.Steps, regularSteps...) + steps = append(steps, memorySetupLines...) + steps = append(steps, memoryRestoreLines...) } - - return nil + steps = append(steps, preSteps...) + steps = append(steps, regularSteps...) + return steps, nil } func formatIndentedYAMLField(fieldName string, value any, trimTrailingNewline bool) (string, error) { @@ -690,65 +686,67 @@ func (c *Compiler) applyBuiltinJobNeedsAugmentations(data *WorkflowData) error { if data == nil || data.Jobs == nil { return nil } - allJobs := c.jobManager.GetAllJobs() for configuredJobName, rawConfig := range data.Jobs { - targetJobName := normalizeBuiltinJobAlias(configuredJobName) - if !isBuiltinJobName(targetJobName) { - continue - } - - configMap, ok := rawConfig.(map[string]any) - if !ok { - return fmt.Errorf("jobs.%s must be an object, got %T", configuredJobName, rawConfig) - } - - augmentedNeeds, err := extractBuiltinJobNeedsAugmentation(configuredJobName, configMap) - if err != nil { + if err := c.applyBuiltinJobNeedsAugmentation(configuredJobName, rawConfig, allJobs); err != nil { return err } - if len(augmentedNeeds) == 0 { - continue - } - - targetJob, exists := c.jobManager.GetJob(targetJobName) - if !exists { - return fmt.Errorf("jobs.%s.needs: cannot augment %q because this workflow does not generate that job", configuredJobName, targetJobName) - } + } + return nil +} - normalizedNeeds := make([]string, 0, len(augmentedNeeds)) - for _, rawNeed := range augmentedNeeds { - need := normalizeBuiltinJobAlias(rawNeed) - if need == targetJobName { - return fmt.Errorf("jobs.%s.needs: %q cannot depend on itself", configuredJobName, rawNeed) - } - if _, known := allJobs[need]; !known { - return fmt.Errorf("jobs.%s.needs: unknown job %q", configuredJobName, rawNeed) - } - normalizedNeeds = append(normalizedNeeds, need) - } +func (c *Compiler) applyBuiltinJobNeedsAugmentation(configuredJobName string, rawConfig any, allJobs map[string]*Job) error { + targetJobName := normalizeBuiltinJobAlias(configuredJobName) + if !isBuiltinJobName(targetJobName) { + return nil + } + configMap, ok := rawConfig.(map[string]any) + if !ok { + return fmt.Errorf("jobs.%s must be an object, got %T", configuredJobName, rawConfig) + } + augmentedNeeds, err := extractBuiltinJobNeedsAugmentation(configuredJobName, configMap) + if err != nil || len(augmentedNeeds) == 0 { + return err + } + targetJob, exists := c.jobManager.GetJob(targetJobName) + if !exists { + return fmt.Errorf("jobs.%s.needs: cannot augment %q because this workflow does not generate that job", configuredJobName, targetJobName) + } + normalizedNeeds, err := normalizeBuiltinAugmentedNeeds(configuredJobName, targetJobName, augmentedNeeds, allJobs) + if err != nil { + return err + } + targetJob.Needs = mergeJobNeeds(targetJob.Needs, normalizedNeeds) + compilerJobsLog.Printf("Applied jobs.%s.needs augmentation to %q: %v", configuredJobName, targetJobName, normalizedNeeds) + return nil +} - seen := make(map[string]struct{}, len(targetJob.Needs)+len(normalizedNeeds)) - mergedNeeds := make([]string, 0, len(targetJob.Needs)+len(normalizedNeeds)) - for _, need := range targetJob.Needs { - if _, alreadySeen := seen[need]; alreadySeen { - continue - } - seen[need] = struct{}{} - mergedNeeds = append(mergedNeeds, need) +func normalizeBuiltinAugmentedNeeds(configuredJobName string, targetJobName string, augmentedNeeds []string, allJobs map[string]*Job) ([]string, error) { + normalizedNeeds := make([]string, 0, len(augmentedNeeds)) + for _, rawNeed := range augmentedNeeds { + need := normalizeBuiltinJobAlias(rawNeed) + if need == targetJobName { + return nil, fmt.Errorf("jobs.%s.needs: %q cannot depend on itself", configuredJobName, rawNeed) } - for _, need := range normalizedNeeds { - if _, alreadySeen := seen[need]; alreadySeen { - continue - } - seen[need] = struct{}{} - mergedNeeds = append(mergedNeeds, need) + if _, known := allJobs[need]; !known { + return nil, fmt.Errorf("jobs.%s.needs: unknown job %q", configuredJobName, rawNeed) } - targetJob.Needs = mergedNeeds - compilerJobsLog.Printf("Applied jobs.%s.needs augmentation to %q: %v", configuredJobName, targetJobName, normalizedNeeds) + normalizedNeeds = append(normalizedNeeds, need) } + return normalizedNeeds, nil +} - return nil +func mergeJobNeeds(existing []string, additional []string) []string { + seen := make(map[string]struct{}, len(existing)+len(additional)) + mergedNeeds := make([]string, 0, len(existing)+len(additional)) + for _, need := range append(append([]string{}, existing...), additional...) { + if _, alreadySeen := seen[need]; alreadySeen { + continue + } + seen[need] = struct{}{} + mergedNeeds = append(mergedNeeds, need) + } + return mergedNeeds } func validateRestrictedBuiltinSetupSteps(jobName string, hasSetupSteps bool) error { @@ -785,73 +783,62 @@ func insertPreStepsAtEarliestBoundary(steps []string, preSteps []string) []strin if len(preSteps) == 0 { return steps } + firstCheckoutIdx, firstTokenMintIdx, lastSetupIdx := findPreStepInsertionAnchors(steps) + insertIdx := determinePreStepInsertIdx(steps, firstCheckoutIdx, firstTokenMintIdx, lastSetupIdx) + if insertIdx > len(steps) { + insertIdx = len(steps) + } + result := make([]string, 0, safeAllocationCapacity(len(steps), len(preSteps))) + result = append(result, steps[:insertIdx]...) + result = append(result, preSteps...) + result = append(result, steps[insertIdx:]...) + return result +} +func findPreStepInsertionAnchors(steps []string) (int, int, int) { firstCheckoutIdx := -1 firstTokenMintIdx := -1 lastSetupIdx := -1 for i, step := range steps { if firstCheckoutIdx == -1 && strings.Contains(step, "uses: actions/checkout@") { - firstCheckoutIdx = i - // Walk backward to the checkout step's list-item boundary ("- "). - // If no boundary is found, keep the current index so insertion still - // occurs before the checkout uses-line. - for j := i; j >= 0; j-- { - trimmed := strings.TrimLeft(steps[j], " ") - if strings.HasPrefix(trimmed, "- ") { - firstCheckoutIdx = j - break - } - } + firstCheckoutIdx = findStepListBoundary(steps, i) } if firstTokenMintIdx == -1 && strings.Contains(step, "uses: actions/create-github-app-token@") { - firstTokenMintIdx = i - // Walk backward to the token-mint step's list-item boundary ("- "). - // If no boundary is found, keep the current index so insertion still - // occurs before the token-mint uses-line. - for j := i; j >= 0; j-- { - trimmed := strings.TrimLeft(steps[j], " ") - if strings.HasPrefix(trimmed, "- ") { - firstTokenMintIdx = j - break - } - } + firstTokenMintIdx = findStepListBoundary(steps, i) } if exactSetupStepIDPattern.MatchString(step) { lastSetupIdx = i } } + return firstCheckoutIdx, firstTokenMintIdx, lastSetupIdx +} + +func findStepListBoundary(steps []string, idx int) int { + for j := idx; j >= 0; j-- { + if strings.HasPrefix(strings.TrimLeft(steps[j], " "), "- ") { + return j + } + } + return idx +} - insertIdx := len(steps) +func determinePreStepInsertIdx(steps []string, firstCheckoutIdx int, firstTokenMintIdx int, lastSetupIdx int) int { if lastSetupIdx >= 0 { for i := lastSetupIdx + 1; i < len(steps); i++ { - trimmed := strings.TrimLeft(steps[i], " ") - if strings.HasPrefix(trimmed, "- ") { - insertIdx = i - break - } - } - if insertIdx == len(steps) { - compilerJobsLog.Print("No step boundary found after setup step; appending pre-steps at end") - } - } else if firstTokenMintIdx >= 0 { - insertIdx = firstTokenMintIdx - if firstCheckoutIdx >= 0 { - if firstCheckoutIdx < insertIdx { - insertIdx = firstCheckoutIdx + if strings.HasPrefix(strings.TrimLeft(steps[i], " "), "- ") { + return i } } - } else if firstCheckoutIdx >= 0 { - insertIdx = firstCheckoutIdx + compilerJobsLog.Print("No step boundary found after setup step; appending pre-steps at end") + return len(steps) } - if insertIdx > len(steps) { - insertIdx = len(steps) + if firstTokenMintIdx >= 0 && (firstCheckoutIdx == -1 || firstTokenMintIdx < firstCheckoutIdx) { + return firstTokenMintIdx } - - result := make([]string, 0, safeAllocationCapacity(len(steps), len(preSteps))) - result = append(result, steps[:insertIdx]...) - result = append(result, preSteps...) - result = append(result, steps[insertIdx:]...) - return result + if firstCheckoutIdx >= 0 { + return firstCheckoutIdx + } + return len(steps) } func (c *Compiler) extractPinnedJobSteps(fieldName string, jobName string, configMap map[string]any, data *WorkflowData) ([]string, error) { diff --git a/pkg/workflow/compiler_difc_proxy.go b/pkg/workflow/compiler_difc_proxy.go index ceb42ba3480..ad6675fee96 100644 --- a/pkg/workflow/compiler_difc_proxy.go +++ b/pkg/workflow/compiler_difc_proxy.go @@ -363,24 +363,7 @@ func injectProxyEnvIntoCustomSteps(customSteps string) string { if customSteps == "" { return customSteps } - - // Extract version comments from uses lines before unmarshaling. - // YAML treats "# comment" as a comment and strips it during Unmarshal, so we - // must capture them here and re-apply after processing to preserve annotations - // like "uses: actions/upload-artifact@sha # v7" in the compiled lock file. - // Without this, gh-aw-manifest falls back to recording the SHA as the version. - versionComments := make(map[string]string) // key: action@sha, value: " # vX" - for line := range strings.SplitSeq(customSteps, "\n") { - trimmed := strings.TrimSpace(line) - if strings.HasPrefix(trimmed, "uses:") && strings.Contains(trimmed, " # ") { - parts := strings.SplitN(trimmed, " # ", 2) - if len(parts) == 2 { - usesValue := strings.TrimSpace(strings.TrimPrefix(parts[0], "uses:")) - versionComments[usesValue] = " # " + parts[1] - } - } - } - + versionComments := captureProxyStepVersionComments(customSteps) var parsed struct { Steps []map[string]any `yaml:"steps"` } @@ -388,34 +371,7 @@ func injectProxyEnvIntoCustomSteps(customSteps string) string { difcProxyLog.Printf("injectProxyEnvIntoCustomSteps: could not parse custom steps, returning as-is: %v", err) return customSteps } - - proxyEnv := proxyEnvVars() - - // Convert each step to an ordered MapSlice with priority fields first so that - // name/uses stay ahead of env for stable diffs, then merge proxy env vars. - orderedSteps := make([]yaml.MapSlice, len(parsed.Steps)) - for i, step := range parsed.Steps { - envMap, ok := step["env"].(map[string]any) - if !ok { - envMap = make(map[string]any) - } - for k, v := range proxyEnv { - envMap[k] = v - } - step["env"] = envMap - - // Re-apply version comment to uses value so the comment survives re-serialization. - if usesVal, hasUses := step["uses"]; hasUses { - if usesStr, ok := usesVal.(string); ok { - if comment, hasComment := versionComments[usesStr]; hasComment { - step["uses"] = usesStr + comment - } - } - } - - orderedSteps[i] = OrderMapFields(step, constants.PriorityStepFields) - } - + orderedSteps := buildProxyInjectedOrderedSteps(parsed.Steps, versionComments) resultBytes, err := yaml.MarshalWithOptions( map[string]any{"steps": orderedSteps}, yaml.Indent(2), @@ -431,6 +387,52 @@ func injectProxyEnvIntoCustomSteps(customSteps string) string { return unquoteUsesWithComments(strings.TrimRight(string(resultBytes), "\n")) } +func captureProxyStepVersionComments(customSteps string) map[string]string { + versionComments := make(map[string]string) + for line := range strings.SplitSeq(customSteps, "\n") { + trimmed := strings.TrimSpace(line) + if strings.HasPrefix(trimmed, "uses:") && strings.Contains(trimmed, " # ") { + parts := strings.SplitN(trimmed, " # ", 2) + if len(parts) == 2 { + versionComments[strings.TrimSpace(strings.TrimPrefix(parts[0], "uses:"))] = " # " + parts[1] + } + } + } + return versionComments +} + +func buildProxyInjectedOrderedSteps(steps []map[string]any, versionComments map[string]string) []yaml.MapSlice { + proxyEnv := proxyEnvVars() + orderedSteps := make([]yaml.MapSlice, len(steps)) + for i, step := range steps { + mergeProxyEnvIntoStep(step, proxyEnv) + reapplyProxyVersionComment(step, versionComments) + orderedSteps[i] = OrderMapFields(step, constants.PriorityStepFields) + } + return orderedSteps +} + +func mergeProxyEnvIntoStep(step map[string]any, proxyEnv map[string]string) { + envMap, ok := step["env"].(map[string]any) + if !ok { + envMap = make(map[string]any) + } + for k, v := range proxyEnv { + envMap[k] = v + } + step["env"] = envMap +} + +func reapplyProxyVersionComment(step map[string]any, versionComments map[string]string) { + usesStr, ok := step["uses"].(string) + if !ok { + return + } + if comment, hasComment := versionComments[usesStr]; hasComment { + step["uses"] = usesStr + comment + } +} + // generateStopDIFCProxyStep generates a step that stops the DIFC proxy container // before the MCP gateway starts. The proxy must be stopped first to avoid // double-filtering: the gateway uses the same guard policy for the agent phase. diff --git a/pkg/workflow/compiler_experiments.go b/pkg/workflow/compiler_experiments.go index 7da48e94e99..cdb683858b4 100644 --- a/pkg/workflow/compiler_experiments.go +++ b/pkg/workflow/compiler_experiments.go @@ -133,99 +133,119 @@ func WorkflowStateBranchName(prefix, workflowID string) string { func extractOneExperimentConfig(name string, val any) *ExperimentConfig { switch v := val.(type) { case []string: - if len(v) >= 2 { - return &ExperimentConfig{Variants: v} - } + return newExperimentConfigFromVariants(v) case []any: - var variants []string - for _, item := range v { - if s, ok := item.(string); ok { - variants = append(variants, s) - } - } - if len(variants) >= 2 { - return &ExperimentConfig{Variants: variants} - } + return newExperimentConfigFromVariants(extractExperimentVariants(v)) case map[string]any: - // New object form: extract variants and optional metadata fields. - cfg := &ExperimentConfig{} - varRaw, ok := v["variants"] - if !ok { - experimentsLog.Printf("Skipping experiment %q: object form requires 'variants' field", name) - return nil - } - switch vv := varRaw.(type) { - case []string: - cfg.Variants = vv - case []any: - for _, item := range vv { - if s, ok := item.(string); ok { - cfg.Variants = append(cfg.Variants, s) - } - } - } - if len(cfg.Variants) < 2 { - experimentsLog.Printf("Skipping experiment %q: must have at least 2 variants", name) - return nil - } - if d, ok := v["description"].(string); ok { - cfg.Description = d - } - if m, ok := v["metric"].(string); ok { - cfg.Metric = m - } - if sd, ok := v["start_date"].(string); ok { - cfg.StartDate = sd - } - if ed, ok := v["end_date"].(string); ok { - cfg.EndDate = ed - } - if n, ok := extractIntField(v["issue"]); ok { - cfg.Issue = n - } - if weightRaw, ok := v["weight"]; ok { - cfg.Weight = extractIntSlice(weightRaw) - } - if h, ok := v["hypothesis"].(string); ok { - cfg.Hypothesis = h - } - if smRaw, ok := v["secondary_metrics"]; ok { - cfg.SecondaryMetrics = parseStringSliceAny(smRaw, nil) - } - if gmRaw, ok := v["guardrail_metrics"]; ok { - cfg.GuardrailMetrics = extractGuardrailMetrics(gmRaw) - } - if n, ok := extractIntField(v["min_samples"]); ok { - cfg.MinSamples = n - } - if at, ok := v["analysis_type"].(string); ok { - cfg.AnalysisType = at - } - if tagsRaw, ok := v["tags"]; ok { - cfg.Tags = parseStringSliceAny(tagsRaw, nil) - } - if notifyRaw, ok := v["notify"]; ok { - if notifyMap, ok := notifyRaw.(map[string]any); ok { - notify := &ExperimentNotify{} - hasNotify := false - if n, ok := extractIntField(notifyMap["discussion"]); ok { - notify.Discussion = n - hasNotify = true - } - if n, ok := extractIntField(notifyMap["issue"]); ok { - notify.Issue = n - hasNotify = true - } - if hasNotify { - cfg.Notify = notify - } - } - } - return cfg + return extractObjectExperimentConfig(name, v) } return nil } +func newExperimentConfigFromVariants(variants []string) *ExperimentConfig { + if len(variants) < 2 { + return nil + } + return &ExperimentConfig{Variants: variants} +} + +func extractExperimentVariants(items []any) []string { + var variants []string + for _, item := range items { + if s, ok := item.(string); ok { + variants = append(variants, s) + } + } + return variants +} + +func extractObjectExperimentConfig(name string, config map[string]any) *ExperimentConfig { + varRaw, ok := config["variants"] + if !ok { + experimentsLog.Printf("Skipping experiment %q: object form requires 'variants' field", name) + return nil + } + cfg := &ExperimentConfig{Variants: extractExperimentVariantValues(varRaw)} + if len(cfg.Variants) < 2 { + experimentsLog.Printf("Skipping experiment %q: must have at least 2 variants", name) + return nil + } + populateExperimentConfigMetadata(cfg, config) + cfg.Notify = extractExperimentNotify(config["notify"]) + return cfg +} + +func extractExperimentVariantValues(varRaw any) []string { + switch vv := varRaw.(type) { + case []string: + return vv + case []any: + return extractExperimentVariants(vv) + default: + return nil + } +} + +func populateExperimentConfigMetadata(cfg *ExperimentConfig, config map[string]any) { + if d, ok := config["description"].(string); ok { + cfg.Description = d + } + if m, ok := config["metric"].(string); ok { + cfg.Metric = m + } + if sd, ok := config["start_date"].(string); ok { + cfg.StartDate = sd + } + if ed, ok := config["end_date"].(string); ok { + cfg.EndDate = ed + } + if n, ok := extractIntField(config["issue"]); ok { + cfg.Issue = n + } + if weightRaw, ok := config["weight"]; ok { + cfg.Weight = extractIntSlice(weightRaw) + } + if h, ok := config["hypothesis"].(string); ok { + cfg.Hypothesis = h + } + if smRaw, ok := config["secondary_metrics"]; ok { + cfg.SecondaryMetrics = parseStringSliceAny(smRaw, nil) + } + if gmRaw, ok := config["guardrail_metrics"]; ok { + cfg.GuardrailMetrics = extractGuardrailMetrics(gmRaw) + } + if n, ok := extractIntField(config["min_samples"]); ok { + cfg.MinSamples = n + } + if at, ok := config["analysis_type"].(string); ok { + cfg.AnalysisType = at + } + if tagsRaw, ok := config["tags"]; ok { + cfg.Tags = parseStringSliceAny(tagsRaw, nil) + } +} + +func extractExperimentNotify(notifyRaw any) *ExperimentNotify { + notifyMap, ok := notifyRaw.(map[string]any) + if !ok { + return nil + } + notify := &ExperimentNotify{} + hasNotify := false + if n, ok := extractIntField(notifyMap["discussion"]); ok { + notify.Discussion = n + hasNotify = true + } + if n, ok := extractIntField(notifyMap["issue"]); ok { + notify.Issue = n + hasNotify = true + } + if !hasNotify { + return nil + } + return notify +} + // extractIntField converts a numeric any value to int. // Returns (int(value), true) on success; (0, false) when val is nil, not a supported // numeric type, negative, or out of int range. @@ -651,47 +671,64 @@ func (c *Compiler) buildPushExperimentsStateJob(data *WorkflowData) (*Job, error if len(data.Experiments) == 0 || data.ExperimentsStorage != ExperimentsStorageRepo { return nil, nil } - experimentsLog.Printf("Building push_experiments_state job (branch=%s)", experimentsBranchName(data.WorkflowID)) + steps := c.buildPushExperimentsStateJobSteps(data) + job := &Job{ + Name: pushExperimentsStateJobName, + RunsOn: c.formatFrameworkJobRunsOn(data), + If: buildPushExperimentsStateJobCondition(), + Permissions: "permissions:\n contents: write", + Needs: []string{string(constants.ActivationJobName)}, + Steps: steps, + } + return job, nil +} - var steps []string +func (c *Compiler) buildPushExperimentsStateJobSteps(data *WorkflowData) []string { + steps := c.buildPushExperimentsSetupSteps(data) + steps = append(steps, buildPushExperimentsCheckoutStep()) + steps = append(steps, c.generateGitConfigurationSteps()...) + steps = append(steps, c.buildPushExperimentsArtifactDownloadStep(data)) + steps = append(steps, buildPushExperimentsPushStep(data)) + if c.actionMode.IsDev() { + steps = append(steps, c.generateRestoreActionsSetupStep()) + } + return steps +} - // Setup step so the push_experiment_state.cjs script is available. +func (c *Compiler) buildPushExperimentsSetupSteps(data *WorkflowData) []string { setupActionRef := c.resolveActionReference("./actions/setup", data) - if setupActionRef != "" || c.actionMode.IsScript() { - steps = append(steps, c.generateCheckoutActionsFolder(data)...) - traceID := fmt.Sprintf("${{ needs.%s.outputs.setup-trace-id }}", constants.ActivationJobName) - parentSpanID := setupParentSpanNeedsExpr(constants.ActivationJobName) - steps = append(steps, c.generateSetupStep(data, setupActionRef, SetupActionDestination, false, traceID, parentSpanID)...) + if setupActionRef == "" && !c.actionMode.IsScript() { + return nil } + steps := append([]string{}, c.generateCheckoutActionsFolder(data)...) + traceID := fmt.Sprintf("${{ needs.%s.outputs.setup-trace-id }}", constants.ActivationJobName) + parentSpanID := setupParentSpanNeedsExpr(constants.ActivationJobName) + return append(steps, c.generateSetupStep(data, setupActionRef, SetupActionDestination, false, traceID, parentSpanID)...) +} - // Checkout step – configure git credentials without downloading workspace files. +func buildPushExperimentsCheckoutStep() string { var checkoutStep strings.Builder checkoutStep.WriteString(" - name: Checkout repository\n") fmt.Fprintf(&checkoutStep, " uses: %s\n", getActionPin("actions/checkout")) checkoutStep.WriteString(" with:\n") checkoutStep.WriteString(" persist-credentials: false\n") checkoutStep.WriteString(" sparse-checkout: .\n") - steps = append(steps, checkoutStep.String()) - - // Git configuration (author, email). - steps = append(steps, c.generateGitConfigurationSteps()...) + return checkoutStep.String() +} - // Download the experiment artifact uploaded by the activation job. - artifactName := experimentArtifactDownloadName(data) +func (c *Compiler) buildPushExperimentsArtifactDownloadStep(data *WorkflowData) string { var downloadStep strings.Builder downloadStep.WriteString(" - name: Download experiment artifact\n") fmt.Fprintf(&downloadStep, " uses: %s\n", c.getActionPin("actions/download-artifact")) downloadStep.WriteString(" continue-on-error: true\n") downloadStep.WriteString(" with:\n") - fmt.Fprintf(&downloadStep, " name: %s\n", artifactName) + fmt.Fprintf(&downloadStep, " name: %s\n", experimentArtifactDownloadName(data)) fmt.Fprintf(&downloadStep, " path: %s\n", experimentsCacheDir) - steps = append(steps, downloadStep.String()) - - // Push experiment state to the git branch via push_experiment_state.cjs. - // This helper uses pushSignedCommits to create verified (signed) commits. - branchName := experimentsBranchName(data.WorkflowID) + return downloadStep.String() +} +func buildPushExperimentsPushStep(data *WorkflowData) string { var pushStep strings.Builder pushStep.WriteString(" - name: Push experiment state to git\n") pushStep.WriteString(" id: push_experiments_state\n") @@ -702,37 +739,21 @@ func (c *Compiler) buildPushExperimentsStateJob(data *WorkflowData) (*Job, error pushStep.WriteString(" GITHUB_RUN_ID: ${{ github.run_id }}\n") pushStep.WriteString(" GITHUB_SERVER_URL: ${{ github.server_url }}\n") fmt.Fprintf(&pushStep, " GH_AW_EXPERIMENT_STATE_DIR: %s\n", experimentsCacheDir) - fmt.Fprintf(&pushStep, " GH_AW_EXPERIMENT_BRANCH: %s\n", branchName) + fmt.Fprintf(&pushStep, " GH_AW_EXPERIMENT_BRANCH: %s\n", experimentsBranchName(data.WorkflowID)) pushStep.WriteString(" with:\n") pushStep.WriteString(" script: |\n") pushStep.WriteString(" const { setupGlobals } = require('" + SetupActionDestination + "/setup_globals.cjs');\n") pushStep.WriteString(" setupGlobals(core, github, context, exec, io, getOctokit);\n") pushStep.WriteString(" const { main } = require('" + SetupActionDestination + "/push_experiment_state.cjs');\n") pushStep.WriteString(" await main();\n") - steps = append(steps, pushStep.String()) - - // Restore the checkout in dev mode (same reason as push_repo_memory). - if c.actionMode.IsDev() { - steps = append(steps, c.generateRestoreActionsSetupStep()) - } + return pushStep.String() +} - // The push_experiments_state job runs after the activation job succeeds. - // It does not depend on the agent job because experiment state was fully resolved in activation. +func buildPushExperimentsStateJobCondition() string { activationSucceeded := BuildEquals( BuildPropertyAccess(fmt.Sprintf("needs.%s.result", constants.ActivationJobName)), BuildStringLiteral("success"), ) notCancelled := &NotNode{Child: BuildFunctionCall("cancelled")} - jobCondition := RenderCondition(BuildAnd(BuildAnd(BuildFunctionCall("always"), notCancelled), activationSucceeded)) - - job := &Job{ - Name: pushExperimentsStateJobName, - RunsOn: c.formatFrameworkJobRunsOn(data), - If: jobCondition, - Permissions: "permissions:\n contents: write", - Needs: []string{string(constants.ActivationJobName)}, - Steps: steps, - } - - return job, nil + return RenderCondition(BuildAnd(BuildAnd(BuildFunctionCall("always"), notCancelled), activationSucceeded)) } diff --git a/pkg/workflow/compiler_github_actions_steps.go b/pkg/workflow/compiler_github_actions_steps.go index 05cf47f3c26..de0736eef7b 100644 --- a/pkg/workflow/compiler_github_actions_steps.go +++ b/pkg/workflow/compiler_github_actions_steps.go @@ -61,56 +61,61 @@ func generatePlaceholderSubstitutionStep(yaml *strings.Builder, expressionMappin if len(expressionMappings) == 0 { return } - compilerGitHubActionsStepsLog.Printf("Generating placeholder substitution step with %d mappings", len(expressionMappings)) + writePlaceholderSubstitutionStepHeader(yaml, indent, data) + writePlaceholderSubstitutionEnv(yaml, expressionMappings, indent) + writePlaceholderSubstitutionScript(yaml, expressionMappings, indent) +} - // Use actions/github-script to perform the substitutions +func writePlaceholderSubstitutionStepHeader(yaml *strings.Builder, indent string, data *WorkflowData) { yaml.WriteString(indent + "- name: Substitute placeholders\n") fmt.Fprintf(yaml, indent+" uses: %s\n", getCachedActionPin("actions/github-script", data)) yaml.WriteString(indent + " env:\n") yaml.WriteString(indent + " GH_AW_PROMPT: /tmp/gh-aw/aw-prompts/prompt.txt\n") +} - // Add all environment variables - // For static values (wrapped in quotes), output them directly without ${{ }} - // For GitHub expressions, wrap them in ${{ }} +func writePlaceholderSubstitutionEnv(yaml *strings.Builder, expressionMappings []*ExpressionMapping, indent string) { for _, mapping := range expressionMappings { - content := mapping.Content - // Check if this is a static quoted value (starts and ends with quotes) - if (strings.HasPrefix(content, "'") && strings.HasSuffix(content, "'")) || - (strings.HasPrefix(content, "\"") && strings.HasSuffix(content, "\"")) { - // Static value - output directly without ${{ }} wrapper - // Check if inner value is multi-line; if so use a YAML double-quoted scalar - // with escaped newlines to avoid invalid YAML. - innerValue := content[1 : len(content)-1] - if strings.Contains(innerValue, "\n") { - escaped := strings.ReplaceAll(innerValue, `\`, `\\`) - escaped = strings.ReplaceAll(escaped, `"`, `\"`) - escaped = strings.ReplaceAll(escaped, "\n", `\n`) - fmt.Fprintf(yaml, indent+" %s: \"%s\"\n", mapping.EnvVar, escaped) - } else { - fmt.Fprintf(yaml, indent+" %s: %s\n", mapping.EnvVar, content) - } - } else { - // GitHub expression - wrap in ${{ }} - fmt.Fprintf(yaml, indent+" %s: ${{ %s }}\n", mapping.EnvVar, content) - } + writePlaceholderEnvVar(yaml, indent, mapping) + } +} + +func writePlaceholderEnvVar(yaml *strings.Builder, indent string, mapping *ExpressionMapping) { + content := mapping.Content + if !isQuotedPlaceholderValue(content) { + fmt.Fprintf(yaml, indent+" %s: ${{ %s }}\n", mapping.EnvVar, content) + return + } + innerValue := content[1 : len(content)-1] + if strings.Contains(innerValue, "\n") { + fmt.Fprintf(yaml, indent+" %s: \"%s\"\n", mapping.EnvVar, escapeMultilinePlaceholderValue(innerValue)) + return } + fmt.Fprintf(yaml, indent+" %s: %s\n", mapping.EnvVar, content) +} +func isQuotedPlaceholderValue(content string) bool { + return (strings.HasPrefix(content, "'") && strings.HasSuffix(content, "'")) || + (strings.HasPrefix(content, "\"") && strings.HasSuffix(content, "\"")) +} + +func escapeMultilinePlaceholderValue(value string) string { + escaped := strings.ReplaceAll(value, `\`, `\\`) + escaped = strings.ReplaceAll(escaped, `"`, `\"`) + return strings.ReplaceAll(escaped, "\n", `\n`) +} + +func writePlaceholderSubstitutionScript(yaml *strings.Builder, expressionMappings []*ExpressionMapping, indent string) { yaml.WriteString(indent + " with:\n") yaml.WriteString(indent + " script: |\n") - - // Use setup_globals helper to make GitHub Actions objects available globally yaml.WriteString(indent + " const { setupGlobals } = require('" + SetupActionDestination + "/setup_globals.cjs');\n") yaml.WriteString(indent + " setupGlobals(core, github, context, exec, io, getOctokit);\n") yaml.WriteString(indent + " \n") - // Use require() to load script from copied files yaml.WriteString(indent + " const substitutePlaceholders = require('" + SetupActionDestination + "/substitute_placeholders.cjs');\n") yaml.WriteString(indent + " \n") - yaml.WriteString(indent + " // Call the substitution function\n") yaml.WriteString(indent + " return await substitutePlaceholders({\n") yaml.WriteString(indent + " file: process.env.GH_AW_PROMPT,\n") yaml.WriteString(indent + " substitutions: {\n") - for i, mapping := range expressionMappings { comma := "," if i == len(expressionMappings)-1 { @@ -118,7 +123,6 @@ func generatePlaceholderSubstitutionStep(yaml *strings.Builder, expressionMappin } fmt.Fprintf(yaml, indent+" %s: process.env.%s%s\n", mapping.EnvVar, mapping.EnvVar, comma) } - yaml.WriteString(indent + " }\n") yaml.WriteString(indent + " });\n") } diff --git a/pkg/workflow/compiler_github_mcp_steps.go b/pkg/workflow/compiler_github_mcp_steps.go index 605f10d9326..7857a257952 100644 --- a/pkg/workflow/compiler_github_mcp_steps.go +++ b/pkg/workflow/compiler_github_mcp_steps.go @@ -21,76 +21,74 @@ import ( // This applies regardless of whether a GitHub App token is configured, because repo-scoping // is not a substitute for author-integrity filtering inside a repository. func (c *Compiler) generateGitHubMCPLockdownDetectionStep(yaml *strings.Builder, data *WorkflowData) { - // Check if GitHub tool is present githubTool, hasGitHub := data.Tools["github"] if !hasGitHub || githubTool == false { githubConfigLog.Print("Skipping GitHub MCP lockdown detection step: GitHub tool not enabled") return } - - // NOTE: Do NOT skip this step when guard policies are explicitly configured. - // Even when min-integrity/repos are hardcoded, the step must still run to output - // the repository visibility via steps.determine-automatic-lockdown.outputs.visibility, - // which is referenced as sink-visibility in safe-outputs and other MCP server guard - // policies. Removing the step while leaving those references in place breaks workflows - // at runtime with undefined step output errors. githubConfigLog.Print("Generating automatic guard policy determination step for GitHub MCP server") + pinnedAction := resolveGitHubMCPLockdownPinnedAction(data) + lockdownConfig := extractGitHubMCPLockdownConfig(githubTool) + yaml.WriteString(" - name: Determine automatic lockdown mode for GitHub MCP Server\n") + yaml.WriteString(" id: determine-automatic-lockdown\n") + fmt.Fprintf(yaml, " uses: %s\n", pinnedAction) + writeGitHubMCPLockdownEnv(yaml, lockdownConfig) + yaml.WriteString(" with:\n") + yaml.WriteString(" script: |\n") + yaml.WriteString(" const determineAutomaticLockdown = require('${{ runner.temp }}/gh-aw/actions/determine_automatic_lockdown.cjs');\n") + yaml.WriteString(" await determineAutomaticLockdown(github, context, core);\n") +} - // Resolve the latest version of actions/github-script +type githubMCPLockdownConfig struct { + configuredMinIntegrity string + configuredRepos string + privateToPublicFlowsAllow bool +} + +func resolveGitHubMCPLockdownPinnedAction(data *WorkflowData) string { actionRepo := "actions/github-script" actionVersion := string(constants.DefaultGitHubScriptVersion) pinnedAction, err := getActionPinWithData(actionRepo, actionVersion, data) if err != nil { githubConfigLog.Printf("Failed to resolve %s@%s: %v", actionRepo, actionVersion, err) - // In strict mode, this error would have been returned by getActionPinWithData - // In normal mode, we fall back to using the version tag without pinning - pinnedAction = fmt.Sprintf("%s@%s", actionRepo, actionVersion) + return fmt.Sprintf("%s@%s", actionRepo, actionVersion) } + return pinnedAction +} - // Extract current guard policy configuration to pass as env vars so the step can - // detect whether each field is already configured and avoid overriding it. - configuredMinIntegrity := "" - configuredRepos := "" - privateToPublicFlowsAllow := false - if toolConfig, ok := githubTool.(map[string]any); ok { - if v, exists := toolConfig["min-integrity"]; exists { - configuredMinIntegrity = serializeEnvStringValue(v) - } - // Support both 'allowed-repos' (preferred) and deprecated 'repos' - if v, exists := toolConfig["allowed-repos"]; exists { - configuredRepos = serializeEnvStringValue(v) - } else if v, exists := toolConfig["repos"]; exists { - configuredRepos = serializeEnvStringValue(v) - } - // Detect private-to-public-flows: allow to inform the default repos value. - // When set to "allow", the user has explicitly opted in to cross-visibility data - // flows, so the repos default should be "all" rather than "public" even for - // public repositories. - if ptpFlows, _ := toolConfig["private-to-public-flows"].(string); ptpFlows == "allow" { - privateToPublicFlowsAllow = true - } +func extractGitHubMCPLockdownConfig(githubTool any) githubMCPLockdownConfig { + var cfg githubMCPLockdownConfig + toolConfig, ok := githubTool.(map[string]any) + if !ok { + return cfg + } + if v, exists := toolConfig["min-integrity"]; exists { + cfg.configuredMinIntegrity = serializeEnvStringValue(v) + } + if v, exists := toolConfig["allowed-repos"]; exists { + cfg.configuredRepos = serializeEnvStringValue(v) + } else if v, exists := toolConfig["repos"]; exists { + cfg.configuredRepos = serializeEnvStringValue(v) + } + if ptpFlows, _ := toolConfig["private-to-public-flows"].(string); ptpFlows == "allow" { + cfg.privateToPublicFlowsAllow = true } + return cfg +} - // Generate the step using the determine_automatic_lockdown.cjs action - yaml.WriteString(" - name: Determine automatic lockdown mode for GitHub MCP Server\n") - yaml.WriteString(" id: determine-automatic-lockdown\n") - fmt.Fprintf(yaml, " uses: %s\n", pinnedAction) +func writeGitHubMCPLockdownEnv(yaml *strings.Builder, cfg githubMCPLockdownConfig) { yaml.WriteString(" env:\n") yaml.WriteString(" GH_AW_GITHUB_TOKEN: ${{ secrets.GH_AW_GITHUB_TOKEN }}\n") yaml.WriteString(" GH_AW_GITHUB_MCP_SERVER_TOKEN: ${{ secrets.GH_AW_GITHUB_MCP_SERVER_TOKEN }}\n") - if configuredMinIntegrity != "" { - fmt.Fprintf(yaml, " GH_AW_GITHUB_MIN_INTEGRITY: %s\n", quoteYAMLEnvValue(configuredMinIntegrity)) + if cfg.configuredMinIntegrity != "" { + fmt.Fprintf(yaml, " GH_AW_GITHUB_MIN_INTEGRITY: %s\n", quoteYAMLEnvValue(cfg.configuredMinIntegrity)) } - if configuredRepos != "" { - fmt.Fprintf(yaml, " GH_AW_GITHUB_REPOS: %s\n", quoteYAMLEnvValue(configuredRepos)) + if cfg.configuredRepos != "" { + fmt.Fprintf(yaml, " GH_AW_GITHUB_REPOS: %s\n", quoteYAMLEnvValue(cfg.configuredRepos)) } - if privateToPublicFlowsAllow { + if cfg.privateToPublicFlowsAllow { yaml.WriteString(" GH_AW_PRIVATE_TO_PUBLIC_FLOWS: " + quoteYAMLEnvValue("allow") + "\n") } - yaml.WriteString(" with:\n") - yaml.WriteString(" script: |\n") - yaml.WriteString(" const determineAutomaticLockdown = require('${{ runner.temp }}/gh-aw/actions/determine_automatic_lockdown.cjs');\n") - yaml.WriteString(" await determineAutomaticLockdown(github, context, core);\n") } // serializeEnvStringValue converts a workflow config value to a string suitable for a @@ -211,55 +209,39 @@ func (c *Compiler) generateParseGuardVarsStep(yaml *strings.Builder, data *Workf githubConfigLog.Print("Skipping parse-guard-vars step: no explicit guard policies configured") return } - githubConfigLog.Print("Generating parse-guard-vars step for blocked-users, trusted-users and approval-labels") - - // Determine the compile-time static values (or user expression) for each field. - // These come from the parsed tools config so we don't lose data from the raw map. - var blockedUsersExtra, trustedUsersExtra, approvalLabelsExtra string - - if data.ParsedTools != nil && data.ParsedTools.GitHub != nil { - gh := data.ParsedTools.GitHub - switch { - case len(gh.BlockedUsers) > 0: - // Static list from frontmatter — join as comma-separated for the env var. - blockedUsersExtra = strings.Join(gh.BlockedUsers, ",") - case gh.BlockedUsersExpr != "": - // User-provided GitHub Actions expression — passed verbatim; GHA evaluates it. - blockedUsersExtra = gh.BlockedUsersExpr - } - switch { - case len(gh.TrustedUsers) > 0: - trustedUsersExtra = strings.Join(gh.TrustedUsers, ",") - case gh.TrustedUsersExpr != "": - trustedUsersExtra = gh.TrustedUsersExpr - } - switch { - case len(gh.ApprovalLabels) > 0: - approvalLabelsExtra = strings.Join(gh.ApprovalLabels, ",") - case gh.ApprovalLabelsExpr != "": - approvalLabelsExtra = gh.ApprovalLabelsExpr - } - } - + blockedUsersExtra, trustedUsersExtra, approvalLabelsExtra := extractGitHubGuardVarExtras(data) yaml.WriteString(" - name: Parse integrity filter lists\n") yaml.WriteString(" id: parse-guard-vars\n") yaml.WriteString(" env:\n") - - if blockedUsersExtra != "" { - fmt.Fprintf(yaml, " GH_AW_BLOCKED_USERS_EXTRA: %s\n", blockedUsersExtra) - } + writeOptionalGuardVarExtra(yaml, "GH_AW_BLOCKED_USERS_EXTRA", blockedUsersExtra) fmt.Fprintf(yaml, " GH_AW_BLOCKED_USERS_VAR: ${{ vars.%s || '' }}\n", constants.EnvVarGitHubBlockedUsers) + writeOptionalGuardVarExtra(yaml, "GH_AW_TRUSTED_USERS_EXTRA", trustedUsersExtra) + fmt.Fprintf(yaml, " GH_AW_TRUSTED_USERS_VAR: ${{ vars.%s || '' }}\n", constants.EnvVarGitHubTrustedUsers) + writeOptionalGuardVarExtra(yaml, "GH_AW_APPROVAL_LABELS_EXTRA", approvalLabelsExtra) + fmt.Fprintf(yaml, " GH_AW_APPROVAL_LABELS_VAR: ${{ vars.%s || '' }}\n", constants.EnvVarGitHubApprovalLabels) + yaml.WriteString(" run: bash \"${RUNNER_TEMP}/gh-aw/actions/parse_guard_list.sh\"\n") +} - if trustedUsersExtra != "" { - fmt.Fprintf(yaml, " GH_AW_TRUSTED_USERS_EXTRA: %s\n", trustedUsersExtra) +func extractGitHubGuardVarExtras(data *WorkflowData) (string, string, string) { + if data.ParsedTools == nil || data.ParsedTools.GitHub == nil { + return "", "", "" } - fmt.Fprintf(yaml, " GH_AW_TRUSTED_USERS_VAR: ${{ vars.%s || '' }}\n", constants.EnvVarGitHubTrustedUsers) + gh := data.ParsedTools.GitHub + return pickGuardVarExtra(gh.BlockedUsers, gh.BlockedUsersExpr), + pickGuardVarExtra(gh.TrustedUsers, gh.TrustedUsersExpr), + pickGuardVarExtra(gh.ApprovalLabels, gh.ApprovalLabelsExpr) +} - if approvalLabelsExtra != "" { - fmt.Fprintf(yaml, " GH_AW_APPROVAL_LABELS_EXTRA: %s\n", approvalLabelsExtra) +func pickGuardVarExtra(values []string, expr string) string { + if len(values) > 0 { + return strings.Join(values, ",") } - fmt.Fprintf(yaml, " GH_AW_APPROVAL_LABELS_VAR: ${{ vars.%s || '' }}\n", constants.EnvVarGitHubApprovalLabels) + return expr +} - yaml.WriteString(" run: bash \"${RUNNER_TEMP}/gh-aw/actions/parse_guard_list.sh\"\n") +func writeOptionalGuardVarExtra(yaml *strings.Builder, envName string, value string) { + if value != "" { + fmt.Fprintf(yaml, " %s: %s\n", envName, value) + } } diff --git a/pkg/workflow/compiler_jobs.go b/pkg/workflow/compiler_jobs.go index a844b10bb2d..8aa4b1ef3e0 100644 --- a/pkg/workflow/compiler_jobs.go +++ b/pkg/workflow/compiler_jobs.go @@ -176,85 +176,70 @@ func (c *Compiler) getCustomJobsReferencedInPromptWithNoActivationDep(data *Work // This function orchestrates the building of all job types by delegating to focused helper functions. func (c *Compiler) buildJobs(data *WorkflowData, markdownPath string) error { compilerJobsLog.Printf("Building jobs for workflow: %s", markdownPath) - - // Use the already-parsed frontmatter from WorkflowData (populated by ParseWorkflowFile / - // ParseWorkflowString) instead of re-reading and re-parsing the file on every compilation. - // Note: RawFrontmatter has already been through preprocessScheduleFields, so shorthand - // triggers (e.g. "on: daily") are already expanded into their structured form. - // The consumers (needsRoleCheck, hasWorkflowRunTrigger) only inspect event keys in the - // "on" field, which is exactly what we need here. frontmatter := data.RawFrontmatter - - // Extract lock filename for timestamp check lockFilename := filepath.Base(stringutil.MarkdownToLockFile(markdownPath)) - - // Resolve custom safe-output actions early so that tool schemas (derived from action.yml) - // are available when buildMainJobWrapper → generateMCPSetup → generateToolsMetaJSON → - // generateDynamicTools runs. Without this early resolution the dynamic_tools entry for - // each action tool would have an empty schema because Inputs/ActionDescription are nil. - if data.SafeOutputs != nil && len(data.SafeOutputs.Actions) > 0 { - c.resolveAllActions(data, markdownPath) - } - - // Build pre-activation and activation jobs + c.resolveSafeOutputActionsForJobs(data, markdownPath) _, activationJobCreated, err := c.buildPreActivationAndActivationJobs(data, frontmatter, lockFilename) if err != nil { return err } + if err := c.buildCoreWorkflowJobs(data, markdownPath, activationJobCreated); err != nil { + return err + } + if err := c.finalizeWorkflowJobs(data); err != nil { + return err + } + compilerJobsLog.Print("Successfully built all jobs for workflow") + return nil +} + +func (c *Compiler) resolveSafeOutputActionsForJobs(data *WorkflowData, markdownPath string) { + if data.SafeOutputs != nil && len(data.SafeOutputs.Actions) > 0 { + c.resolveAllActions(data, markdownPath) + } +} - // Build main workflow job +func (c *Compiler) buildCoreWorkflowJobs(data *WorkflowData, markdownPath string, activationJobCreated bool) error { if err := c.buildMainJobWrapper(data, activationJobCreated); err != nil { return err } - - // Build safe outputs jobs if configured if err := c.buildSafeOutputsJobs(data, string(constants.AgentJobName), markdownPath); err != nil { return fmt.Errorf("failed to build safe outputs jobs: %w", err) } - - // Build BinEval evals job if evals are declared in frontmatter. - if evalsJob, err := c.buildEvalsJob(data); err != nil { - return fmt.Errorf("failed to build evals job: %w", err) - } else if evalsJob != nil { - if err := c.jobManager.AddJob(evalsJob); err != nil { - return fmt.Errorf("failed to add evals job: %w", err) - } + if err := c.addEvalsJob(data); err != nil { + return err } - - // Apply jobs..pre-steps customizations to already-created built-in jobs - // before processing non-built-in custom jobs. if err := c.applyBuiltinJobPreSteps(data); err != nil { return fmt.Errorf("failed to apply built-in job pre-steps: %w", err) } - - // Build additional custom jobs from frontmatter jobs section if len(data.Jobs) > 0 { compilerJobsLog.Printf("Building %d custom jobs from frontmatter", len(data.Jobs)) } if err := c.buildCustomJobs(data, activationJobCreated); err != nil { return fmt.Errorf("failed to build custom jobs: %w", err) } + return c.buildMemoryManagementJobs(data) +} - // Build memory management jobs (repo-memory and cache-memory) - if err := c.buildMemoryManagementJobs(data); err != nil { - return err +func (c *Compiler) addEvalsJob(data *WorkflowData) error { + evalsJob, err := c.buildEvalsJob(data) + if err != nil { + return fmt.Errorf("failed to build evals job: %w", err) + } + if evalsJob == nil { + return nil + } + if err := c.jobManager.AddJob(evalsJob); err != nil { + return fmt.Errorf("failed to add evals job: %w", err) } + return nil +} - // Apply additive jobs..needs augmentations once all jobs are created, - // so referenced custom/imported jobs can be validated against the final job set. +func (c *Compiler) finalizeWorkflowJobs(data *WorkflowData) error { if err := c.applyBuiltinJobNeedsAugmentations(data); err != nil { return fmt.Errorf("failed to apply built-in job needs augmentations: %w", err) } - - // Final pass: ensure conclusion job depends on ALL remaining workflow jobs. - // This guarantees conclusion always runs last, even for custom user-defined jobs - // (e.g. post-issue, super_linter) that were not explicitly added to its needs. - if err := c.ensureConclusionIsLastJob(); err != nil { - return err - } - - compilerJobsLog.Print("Successfully built all jobs for workflow") - return nil + return c.ensureConclusionIsLastJob() } // buildPreActivationAndActivationJobs builds the pre-activation and activation jobs if needed. diff --git a/pkg/workflow/compiler_main_job.go b/pkg/workflow/compiler_main_job.go index 714976a6337..00c30250657 100644 --- a/pkg/workflow/compiler_main_job.go +++ b/pkg/workflow/compiler_main_job.go @@ -21,26 +21,8 @@ func isBuiltinJobName(jobName string) bool { func (c *Compiler) buildMainJob(data *WorkflowData, activationJobCreated bool) (*Job, error) { workflowLog.Printf("Building main job for workflow: %s", data.Name) var steps []string - - setupActionRef := c.resolveActionReference("./actions/setup", data) - if setupActionRef != "" || c.actionMode.IsScript() { - compilerMainJobLog.Printf("Adding actions-folder checkout and setup steps (ref=%q, scriptMode=%v)", setupActionRef, c.actionMode.IsScript()) - steps = append(steps, c.generateCheckoutActionsFolder(data)...) - agentTraceID := fmt.Sprintf("${{ needs.%s.outputs.setup-trace-id }}", constants.ActivationJobName) - agentParentSpanID := setupParentSpanNeedsExpr(constants.ActivationJobName) - steps = append(steps, c.generateSetupStep(data, setupActionRef, SetupActionDestination, false, agentTraceID, agentParentSpanID)...) - } - // Set runtime paths that depend on RUNNER_TEMP via $GITHUB_ENV. - // These cannot be set in job-level env: because the runner context is not - // available there (only in step-level env: and run: blocks). - if data.SafeOutputs != nil { - compilerMainJobLog.Print("Adding runtime-paths step for safe-outputs") - steps = append(steps, c.generateSetRuntimePathsStep()...) - } - + steps = append(steps, c.buildMainJobSetupAndRuntimeSteps(data)...) jobCondition := c.buildMainJobCondition(data, activationJobCreated) - - // Build agent step content (checkout app tokens minted here to avoid masked-value drops). var stepBuilder strings.Builder if err := c.generateMainJobSteps(&stepBuilder, data); err != nil { return nil, fmt.Errorf("failed to generate main job steps: %w", err) @@ -48,10 +30,8 @@ func (c *Compiler) buildMainJob(data *WorkflowData, activationJobCreated bool) ( if stepsContent := stepBuilder.String(); stepsContent != "" { steps = append(steps, stepsContent) } - depends, engineEnvContent := c.buildMainJobDependencies(data, activationJobCreated) c.warnBuiltinJobEnvReferences(depends, engineEnvContent) - outputs := c.buildMainJobOutputs(data) env := c.buildMainJobEnv(data) agentConcurrency := GenerateJobConcurrencyConfig(data) @@ -59,13 +39,37 @@ func (c *Compiler) buildMainJob(data *WorkflowData, activationJobCreated bool) ( if err != nil { return nil, err } + steps = c.appendMainJobScriptCleanup(steps) + compilerMainJobLog.Printf("Built main job: steps=%d, needs=%v, outputs=%d", len(steps), depends, len(outputs)) + return c.finalizeMainJob(data, jobCondition, permissions, agentConcurrency, env, steps, depends, outputs), nil +} +func (c *Compiler) buildMainJobSetupAndRuntimeSteps(data *WorkflowData) []string { + var steps []string + setupActionRef := c.resolveActionReference("./actions/setup", data) + if setupActionRef != "" || c.actionMode.IsScript() { + compilerMainJobLog.Printf("Adding actions-folder checkout and setup steps (ref=%q, scriptMode=%v)", setupActionRef, c.actionMode.IsScript()) + steps = append(steps, c.generateCheckoutActionsFolder(data)...) + agentTraceID := fmt.Sprintf("${{ needs.%s.outputs.setup-trace-id }}", constants.ActivationJobName) + agentParentSpanID := setupParentSpanNeedsExpr(constants.ActivationJobName) + steps = append(steps, c.generateSetupStep(data, setupActionRef, SetupActionDestination, false, agentTraceID, agentParentSpanID)...) + } + if data.SafeOutputs != nil { + compilerMainJobLog.Print("Adding runtime-paths step for safe-outputs") + steps = append(steps, c.generateSetRuntimePathsStep()...) + } + return steps +} + +func (c *Compiler) appendMainJobScriptCleanup(steps []string) []string { if c.actionMode.IsScript() { compilerMainJobLog.Print("Adding script-mode cleanup step") - steps = append(steps, c.generateScriptModeCleanupStep()) + return append(steps, c.generateScriptModeCleanupStep()) } + return steps +} - compilerMainJobLog.Printf("Built main job: steps=%d, needs=%v, outputs=%d", len(steps), depends, len(outputs)) +func (c *Compiler) finalizeMainJob(data *WorkflowData, jobCondition string, permissions string, agentConcurrency string, env map[string]string, steps []string, depends []string, outputs map[string]string) *Job { return &Job{ Name: string(constants.AgentJobName), If: jobCondition, @@ -79,5 +83,5 @@ func (c *Compiler) buildMainJob(data *WorkflowData, activationJobCreated bool) ( Steps: steps, Needs: depends, Outputs: outputs, - }, nil + } } diff --git a/pkg/workflow/compiler_main_job_helpers.go b/pkg/workflow/compiler_main_job_helpers.go index b412a53fec2..f013190c7bf 100644 --- a/pkg/workflow/compiler_main_job_helpers.go +++ b/pkg/workflow/compiler_main_job_helpers.go @@ -256,69 +256,68 @@ func (c *Compiler) buildMainJobOutputs(data *WorkflowData) map[string]string { // buildMainJobEnv builds the job-level environment variable map for the main agent job. func (c *Compiler) buildMainJobEnv(data *WorkflowData) map[string]string { var env map[string]string + env = applyPlaywrightMainJobEnv(env, data) + env = applySafeOutputsMainJobEnv(env, data) + env = applyWorkflowIDMainJobEnv(env, data) + env = applyProjectUTCMainJobEnv(env, c.getCompiledProjectUTCOffset()) + return env +} - // Disable the Chromium process sandbox for playwright CLI mode. - // GitHub Actions runners are containerised environments where kernel namespace - // sandboxing is unavailable, which causes playwright-cli to abort with - // "Playwright can't run in this sandbox environment". - if isPlaywrightCLIMode(data.Tools) { - if env == nil { - env = make(map[string]string) - } - env["PLAYWRIGHT_MCP_SANDBOX"] = "false" +func applyPlaywrightMainJobEnv(env map[string]string, data *WorkflowData) map[string]string { + if !isPlaywrightCLIMode(data.Tools) { + return env } + env = ensureStringMap(env) + env["PLAYWRIGHT_MCP_SANDBOX"] = "false" + return env +} - if data.SafeOutputs != nil { - compilerMainJobLog.Printf("Configuring safe-outputs job env for main job (uploadAssets=%v)", data.SafeOutputs.UploadAssets != nil) - if env == nil { - env = make(map[string]string) - } - - // Set GH_AW_MCP_LOG_DIR for safe outputs MCP server logging - // Store in mcp-logs directory so it's included in mcp-logs artifact - env["GH_AW_MCP_LOG_DIR"] = constants.TmpMcpLogsSafeOutputsDir - - // Note: GH_AW_SAFE_OUTPUTS, GH_AW_SAFE_OUTPUTS_CONFIG_PATH, and - // GH_AW_SAFE_OUTPUTS_TOOLS_PATH are set via a run step (see generateSetRuntimePathsStep) - // because the runner context is not available in job-level env: blocks. - - // Add asset-related environment variables - // These must always be set (even to empty) because awmg v0.0.12+ validates ${VAR} references - if data.SafeOutputs.UploadAssets != nil { - env["GH_AW_ASSETS_BRANCH"] = fmt.Sprintf("%q", data.SafeOutputs.UploadAssets.BranchName) - env["GH_AW_ASSETS_MAX_SIZE_KB"] = strconv.Itoa(data.SafeOutputs.UploadAssets.MaxSizeKB) - env["GH_AW_ASSETS_ALLOWED_EXTS"] = fmt.Sprintf("%q", strings.Join(data.SafeOutputs.UploadAssets.AllowedExts, ",")) - } else { - // Set empty defaults when upload-assets is not configured - env["GH_AW_ASSETS_BRANCH"] = `""` - env["GH_AW_ASSETS_MAX_SIZE_KB"] = "0" - env["GH_AW_ASSETS_ALLOWED_EXTS"] = `""` - } +func applySafeOutputsMainJobEnv(env map[string]string, data *WorkflowData) map[string]string { + if data.SafeOutputs == nil { + return env + } + compilerMainJobLog.Printf("Configuring safe-outputs job env for main job (uploadAssets=%v)", data.SafeOutputs.UploadAssets != nil) + env = ensureStringMap(env) + env["GH_AW_MCP_LOG_DIR"] = constants.TmpMcpLogsSafeOutputsDir + applySafeOutputsAssetEnv(env, data.SafeOutputs) + env["DEFAULT_BRANCH"] = "${{ github.event.repository.default_branch }}" + return env +} - // DEFAULT_BRANCH is used by safeoutputs MCP server - // Use repository default branch from GitHub context - env["DEFAULT_BRANCH"] = "${{ github.event.repository.default_branch }}" +func applySafeOutputsAssetEnv(env map[string]string, safeOutputs *SafeOutputsConfig) { + if safeOutputs.UploadAssets == nil { + env["GH_AW_ASSETS_BRANCH"] = `""` + env["GH_AW_ASSETS_MAX_SIZE_KB"] = "0" + env["GH_AW_ASSETS_ALLOWED_EXTS"] = `""` + return } + env["GH_AW_ASSETS_BRANCH"] = fmt.Sprintf("%q", safeOutputs.UploadAssets.BranchName) + env["GH_AW_ASSETS_MAX_SIZE_KB"] = strconv.Itoa(safeOutputs.UploadAssets.MaxSizeKB) + env["GH_AW_ASSETS_ALLOWED_EXTS"] = fmt.Sprintf("%q", strings.Join(safeOutputs.UploadAssets.AllowedExts, ",")) +} - // Set GH_AW_WORKFLOW_ID_SANITIZED for cache-memory keys - // This contains the workflow ID with all hyphens removed and lowercased - // Used in cache keys to avoid spaces and special characters - if data.WorkflowID != "" { - if env == nil { - env = make(map[string]string) - } - env["GH_AW_WORKFLOW_ID_SANITIZED"] = SanitizeWorkflowIDForCacheKey(data.WorkflowID) +func applyWorkflowIDMainJobEnv(env map[string]string, data *WorkflowData) map[string]string { + if data.WorkflowID == "" { + return env } + env = ensureStringMap(env) + env["GH_AW_WORKFLOW_ID_SANITIZED"] = SanitizeWorkflowIDForCacheKey(data.WorkflowID) + return env +} - // Bake the repository project UTC offset (from aw.json) into job env so runtime - // JavaScript helpers do not need to read aw.json on the runner. - if utcOffset := c.getCompiledProjectUTCOffset(); utcOffset != "" { - if env == nil { - env = make(map[string]string) - } - env["GH_AW_PROJECT_UTC"] = fmt.Sprintf("%q", utcOffset) +func applyProjectUTCMainJobEnv(env map[string]string, utcOffset string) map[string]string { + if utcOffset == "" { + return env } + env = ensureStringMap(env) + env["GH_AW_PROJECT_UTC"] = fmt.Sprintf("%q", utcOffset) + return env +} +func ensureStringMap(env map[string]string) map[string]string { + if env == nil { + return make(map[string]string) + } return env } diff --git a/pkg/workflow/compiler_orchestrator_engine.go b/pkg/workflow/compiler_orchestrator_engine.go index 7d3e1884327..9f83e0baa51 100644 --- a/pkg/workflow/compiler_orchestrator_engine.go +++ b/pkg/workflow/compiler_orchestrator_engine.go @@ -258,67 +258,102 @@ func (c *Compiler) resolveEngineFromIncludesAndImports( engineConfig *EngineConfig, model string, ) (string, *EngineConfig, string, error) { + allEngines, err := c.collectResolvedEngines(result.Markdown, markdownDir, importsResult) + if err != nil { + return "", nil, "", err + } + engineSetting, err = c.validateAndRegisterResolvedEngines(engineSetting, allEngines) + if err != nil { + return "", nil, "", err + } + engineConfig, model, err = c.applyImportedEngineConfig(allEngines, engineConfig, model) + if err != nil { + return "", nil, "", err + } + engineSetting, engineConfig = c.finalizeResolvedEngineConfig(engineSetting, engineConfig) + return engineSetting, engineConfig, model, nil +} + +func (c *Compiler) collectResolvedEngines(markdown string, markdownDir string, importsResult *parser.ImportsResult) ([]string, error) { orchestratorEngineLog.Printf("Expanding includes for engine configurations") - includedEngines, err := parser.ExpandIncludesForEngines(result.Markdown, markdownDir) + includedEngines, err := parser.ExpandIncludesForEngines(markdown, markdownDir) if err != nil { orchestratorEngineLog.Printf("Failed to expand includes for engines: %v", err) - return "", nil, "", fmt.Errorf("failed to expand includes for engines: %w", err) + return nil, fmt.Errorf("failed to expand includes for engines: %w", err) } - allEngines := append(importsResult.MergedEngines, includedEngines...) + return append(importsResult.MergedEngines, includedEngines...), nil +} + +func (c *Compiler) validateAndRegisterResolvedEngines(engineSetting string, allEngines []string) (string, error) { orchestratorEngineLog.Printf("Validating single engine specification") finalEngineSetting, err := c.validateSingleEngineSpecification(engineSetting, allEngines) if err != nil { orchestratorEngineLog.Printf("Engine specification validation failed: %v", err) - return "", nil, "", err + return "", err } if finalEngineSetting != "" { engineSetting = finalEngineSetting } for _, engineJSON := range allEngines { if err := c.registerNamedEngineDefinitionFromJSON(engineJSON); err != nil { - return "", nil, "", fmt.Errorf("failed to register engine definition from included file: %w", err) + return "", fmt.Errorf("failed to register engine definition from included file: %w", err) } } - if engineConfig == nil && len(allEngines) > 0 { - orchestratorEngineLog.Printf("Extracting engine config from included file") - var extractedModel string - engineConfig, extractedModel, err = c.extractEngineConfigFromJSON(allEngines[0]) - if err != nil { - orchestratorEngineLog.Printf("Failed to extract engine config: %v", err) - return "", nil, "", fmt.Errorf("failed to extract engine config from included file: %w", err) - } - // Preserve the model from the main workflow frontmatter if already set; - // only fall back to the imported/shared workflow's model when the main - // workflow does not specify one (main workflow model takes precedence). - if model == "" { - model = extractedModel - } - if err := c.validateAndRegisterInlineEngineConfig(engineConfig); err != nil { - return "", nil, "", err - } - } else if model == "" && len(allEngines) > 0 { - // engineConfig is non-nil (e.g. from top-level max-ai-credits or other - // budget fields) but model has not been set by the main workflow. Extract - // just the model from the imported engine config so that an engine.model - // pin in an imported file is not silently dropped. - _, extractedModel, extractErr := c.extractEngineConfigFromJSON(allEngines[0]) - if extractErr == nil && extractedModel != "" { - model = extractedModel - orchestratorEngineLog.Printf("Applied model from imported engine config: %s", model) - } + return engineSetting, nil +} + +func (c *Compiler) applyImportedEngineConfig(allEngines []string, engineConfig *EngineConfig, model string) (*EngineConfig, string, error) { + if len(allEngines) == 0 { + return engineConfig, model, nil + } + if engineConfig == nil { + return c.extractPrimaryImportedEngineConfig(allEngines[0], model) + } + if model != "" { + return engineConfig, model, nil + } + return engineConfig, c.extractImportedEngineModel(allEngines[0]), nil +} + +func (c *Compiler) extractPrimaryImportedEngineConfig(engineJSON string, model string) (*EngineConfig, string, error) { + orchestratorEngineLog.Printf("Extracting engine config from included file") + engineConfig, extractedModel, err := c.extractEngineConfigFromJSON(engineJSON) + if err != nil { + orchestratorEngineLog.Printf("Failed to extract engine config: %v", err) + return nil, "", fmt.Errorf("failed to extract engine config from included file: %w", err) + } + if model == "" { + model = extractedModel + } + if err := c.validateAndRegisterInlineEngineConfig(engineConfig); err != nil { + return nil, "", err } + return engineConfig, model, nil +} + +func (c *Compiler) extractImportedEngineModel(engineJSON string) string { + _, extractedModel, err := c.extractEngineConfigFromJSON(engineJSON) + if err == nil && extractedModel != "" { + orchestratorEngineLog.Printf("Applied model from imported engine config: %s", extractedModel) + return extractedModel + } + return "" +} + +func (c *Compiler) finalizeResolvedEngineConfig(engineSetting string, engineConfig *EngineConfig) (string, *EngineConfig) { if engineSetting == "" { defaultEngine := c.engineRegistry.GetDefaultEngine() engineSetting = defaultEngine.GetID() workflowLog.Printf("No 'engine:' setting found, defaulting to: %s", engineSetting) } if engineConfig == nil { - engineConfig = &EngineConfig{ID: engineSetting} - } else if engineConfig.ID == "" && engineSetting != "" { + return engineSetting, &EngineConfig{ID: engineSetting} + } + if engineConfig.ID == "" && engineSetting != "" { engineConfig.ID = engineSetting orchestratorEngineLog.Printf("Normalized engineConfig.ID from engineSetting: %s", engineSetting) } - return engineSetting, engineConfig, model, nil + return engineSetting, engineConfig } // applyEngineImportDefaults merges import-derived engine defaults into engineConfig. @@ -335,9 +370,25 @@ func (c *Compiler) applyEngineImportDefaults( preservedMaxRuns int, preservedMaxTurnCacheMisses int, ) (*EngineConfig, string) { - if engineConfig == nil { - engineConfig = &EngineConfig{ID: engineSetting} + engineConfig = ensureEngineImportConfig(engineConfig, engineSetting) + applyPreservedEngineBudgetLimits(engineConfig, preservedMaxTurns, preservedMaxAICredits, preservedMaxRuns, preservedMaxTurnCacheMisses) + applyImportedEngineExecutionDefaults(engineConfig, importsResult) + applyImportedEngineTimeoutDefaults(engineConfig, importsResult) + if model == "" && importsResult.MergedEngineModel != "" { + model = importsResult.MergedEngineModel + orchestratorEngineLog.Printf("Applied model preference from import: %s", model) + } + return engineConfig, model +} + +func ensureEngineImportConfig(engineConfig *EngineConfig, engineSetting string) *EngineConfig { + if engineConfig != nil { + return engineConfig } + return &EngineConfig{ID: engineSetting} +} + +func applyPreservedEngineBudgetLimits(engineConfig *EngineConfig, preservedMaxTurns string, preservedMaxAICredits int64, preservedMaxRuns int, preservedMaxTurnCacheMisses int) { if preservedMaxTurns != "" { engineConfig.MaxTurns = preservedMaxTurns } @@ -350,51 +401,82 @@ func (c *Compiler) applyEngineImportDefaults( if preservedMaxTurnCacheMisses > 0 { engineConfig.MaxTurnCacheMisses = preservedMaxTurnCacheMisses } - if engineConfig.MaxTurns == "" && importsResult.MergedMaxTurns != "" { - var importedMaxTurns any - if err := json.Unmarshal([]byte(importsResult.MergedMaxTurns), &importedMaxTurns); err == nil { - if parsed := parseMaxTurnsValue(importedMaxTurns); parsed != "" { - engineConfig.MaxTurns = parsed - orchestratorEngineLog.Printf("Applied max-turns from import") - } +} + +func applyImportedEngineExecutionDefaults(engineConfig *EngineConfig, importsResult *parser.ImportsResult) { + applyImportedEngineMaxTurns(engineConfig, importsResult.MergedMaxTurns) + applyImportedEngineMaxToolDenials(engineConfig, importsResult.MergedMaxToolDenials) + applyImportedEngineMaxRuns(engineConfig, importsResult.MergedMaxRuns) + applyImportedEngineMaxAICredits(engineConfig, importsResult.MergedMaxAICredits) + applyImportedEngineMaxTurnCacheMisses(engineConfig, importsResult.MergedMaxTurnCacheMisses) +} + +func applyImportedEngineMaxTurns(engineConfig *EngineConfig, raw string) { + if engineConfig.MaxTurns != "" || raw == "" { + return + } + var imported any + if err := json.Unmarshal([]byte(raw), &imported); err == nil { + if parsed := parseMaxTurnsValue(imported); parsed != "" { + engineConfig.MaxTurns = parsed + orchestratorEngineLog.Printf("Applied max-turns from import") } } - if engineConfig.MaxToolDenials == "" && importsResult.MergedMaxToolDenials != "" { - var importedMaxToolDenials any - if err := json.Unmarshal([]byte(importsResult.MergedMaxToolDenials), &importedMaxToolDenials); err == nil { - if parsed := parseMaxToolDenialsValue(importedMaxToolDenials); parsed != "" { - engineConfig.MaxToolDenials = parsed - orchestratorEngineLog.Printf("Applied max-tool-denials from import") - } +} + +func applyImportedEngineMaxToolDenials(engineConfig *EngineConfig, raw string) { + if engineConfig.MaxToolDenials != "" || raw == "" { + return + } + var imported any + if err := json.Unmarshal([]byte(raw), &imported); err == nil { + if parsed := parseMaxToolDenialsValue(imported); parsed != "" { + engineConfig.MaxToolDenials = parsed + orchestratorEngineLog.Printf("Applied max-tool-denials from import") } } - if engineConfig.MaxRuns <= 0 && importsResult.MergedMaxRuns != "" { - var importedMaxRuns any - if err := json.Unmarshal([]byte(importsResult.MergedMaxRuns), &importedMaxRuns); err == nil { - if parsed := parseMaxRunsValue(importedMaxRuns); parsed > 0 { - engineConfig.MaxRuns = parsed - orchestratorEngineLog.Printf("Applied max-runs from import") - } +} + +func applyImportedEngineMaxRuns(engineConfig *EngineConfig, raw string) { + if engineConfig.MaxRuns > 0 || raw == "" { + return + } + var imported any + if err := json.Unmarshal([]byte(raw), &imported); err == nil { + if parsed := parseMaxRunsValue(imported); parsed > 0 { + engineConfig.MaxRuns = parsed + orchestratorEngineLog.Printf("Applied max-runs from import") } } - if engineConfig.MaxAICredits == 0 && importsResult.MergedMaxAICredits != "" { - var importedMaxAICredits any - if err := json.Unmarshal([]byte(importsResult.MergedMaxAICredits), &importedMaxAICredits); err == nil { - if parsed := parseMaxAICreditsValue(importedMaxAICredits); parsed != 0 { - engineConfig.MaxAICredits = parsed - orchestratorEngineLog.Printf("Applied max-ai-credits from import") - } +} + +func applyImportedEngineMaxAICredits(engineConfig *EngineConfig, raw string) { + if engineConfig.MaxAICredits != 0 || raw == "" { + return + } + var imported any + if err := json.Unmarshal([]byte(raw), &imported); err == nil { + if parsed := parseMaxAICreditsValue(imported); parsed != 0 { + engineConfig.MaxAICredits = parsed + orchestratorEngineLog.Printf("Applied max-ai-credits from import") } } - if engineConfig.MaxTurnCacheMisses <= 0 && importsResult.MergedMaxTurnCacheMisses != "" { - var importedMaxTurnCacheMisses any - if err := json.Unmarshal([]byte(importsResult.MergedMaxTurnCacheMisses), &importedMaxTurnCacheMisses); err == nil { - if parsed := parseMaxTurnCacheMissesValue(importedMaxTurnCacheMisses); parsed > 0 { - engineConfig.MaxTurnCacheMisses = parsed - orchestratorEngineLog.Printf("Applied max-turn-cache-misses from import") - } +} + +func applyImportedEngineMaxTurnCacheMisses(engineConfig *EngineConfig, raw string) { + if engineConfig.MaxTurnCacheMisses > 0 || raw == "" { + return + } + var imported any + if err := json.Unmarshal([]byte(raw), &imported); err == nil { + if parsed := parseMaxTurnCacheMissesValue(imported); parsed > 0 { + engineConfig.MaxTurnCacheMisses = parsed + orchestratorEngineLog.Printf("Applied max-turn-cache-misses from import") } } +} + +func applyImportedEngineTimeoutDefaults(engineConfig *EngineConfig, importsResult *parser.ImportsResult) { if engineConfig.MCPToolTimeout == "" && importsResult.MergedEngineMCPToolTimeout != "" { engineConfig.MCPToolTimeout = importsResult.MergedEngineMCPToolTimeout orchestratorEngineLog.Printf("Applied engine.mcp.tool-timeout from import: %s", engineConfig.MCPToolTimeout) @@ -403,11 +485,6 @@ func (c *Compiler) applyEngineImportDefaults( engineConfig.MCPSessionTimeout = importsResult.MergedEngineMCPSessionTimeout orchestratorEngineLog.Printf("Applied engine.mcp.session-timeout from import: %s", engineConfig.MCPSessionTimeout) } - if model == "" && importsResult.MergedEngineModel != "" { - model = importsResult.MergedEngineModel - orchestratorEngineLog.Printf("Applied model preference from import: %s", model) - } - return engineConfig, model } func (c *Compiler) resolveEngineRuntimeConfig(engineSetting string, engineConfig *EngineConfig) (CodingAgentEngine, []map[string]any, error) { diff --git a/pkg/workflow/compiler_orchestrator_frontmatter.go b/pkg/workflow/compiler_orchestrator_frontmatter.go index 7e1f60aa520..95c3aee9728 100644 --- a/pkg/workflow/compiler_orchestrator_frontmatter.go +++ b/pkg/workflow/compiler_orchestrator_frontmatter.go @@ -83,195 +83,179 @@ func (c *Compiler) validateEngineBeforeSchema( func (c *Compiler) parseFrontmatterSection(markdownPath string) (*frontmatterParseResult, error) { orchestratorFrontmatterLog.Printf("Starting frontmatter parsing: %s", markdownPath) workflowLog.Printf("Reading file: %s", markdownPath) - - // Clean the path to prevent path traversal issues (gosec G304) - // filepath.Clean removes ".." and other problematic path elements cleanPath := filepath.Clean(markdownPath) - - // Read the file - content, err := os.ReadFile(cleanPath) + content, contentString, err := c.readFrontmatterSource(cleanPath) if err != nil { - orchestratorFrontmatterLog.Printf("Failed to read file: %s, error: %v", cleanPath, err) - // Keep the user-facing message while avoiding exposure of os.PathError internals. - return nil, fmt.Errorf("failed to read file: %w", frontmatterReadError{message: err.Error()}) + return nil, err } - contentString := string(content) - workflowLog.Printf("File size: %d bytes", len(content)) - - // Parse frontmatter and markdown - orchestratorFrontmatterLog.Printf("Parsing frontmatter from file: %s", cleanPath) - result, err := parser.ExtractFrontmatterFromContent(contentString) + result, err := c.extractWorkflowFrontmatter(cleanPath, contentString) if err != nil { - orchestratorFrontmatterLog.Printf("Frontmatter extraction failed: %v", err) - // Use FrontmatterStart from result if available, otherwise default to line 2 (after opening ---) - frontmatterStart := 2 - if result != nil && result.FrontmatterStart > 0 { - frontmatterStart = result.FrontmatterStart - } - return nil, c.createFrontmatterError(cleanPath, contentString, err, frontmatterStart) + return nil, err } - if len(result.Frontmatter) == 0 { orchestratorFrontmatterLog.Print("No frontmatter found in file") return nil, errors.New("no frontmatter found") } - - // Preprocess schedule fields to convert human-friendly format to cron expressions if err := c.preprocessScheduleFields(result.Frontmatter, cleanPath, contentString); err != nil { orchestratorFrontmatterLog.Printf("Schedule preprocessing failed: %v", err) return nil, err } - - // Create a copy of frontmatter without internal markers for schema validation - // Keep the original frontmatter with markers for YAML generation frontmatterForValidation := c.copyFrontmatterWithoutInternalMarkers(result.Frontmatter) - - // Check if user accidentally used "triggers:" instead of the correct "on:" keyword if _, hasTriggers := frontmatterForValidation["triggers"]; hasTriggers { return nil, fmt.Errorf("%s: invalid frontmatter key 'triggers:' — use 'on:' to define workflow triggers", cleanPath) } + if sharedResult, err := c.handleFrontmatterWithoutOn(cleanPath, content, result, frontmatterForValidation); sharedResult != nil || err != nil { + return sharedResult, err + } + if result.Markdown == "" { + orchestratorFrontmatterLog.Print("No markdown content found for main workflow") + return nil, errors.New("no markdown content found") + } + if err := c.validateMainWorkflowFrontmatter(cleanPath, content, result, frontmatterForValidation); err != nil { + return nil, err + } + c.emitMainWorkflowMarkdownWarnings(cleanPath, result.Markdown) + workflowLog.Printf("Frontmatter: %d chars, Markdown: %d chars", len(result.Frontmatter), len(result.Markdown)) + return newFrontmatterParseResult(cleanPath, content, result, frontmatterForValidation), nil +} - // Check if "on" field is missing - if so, treat as a shared/imported workflow - _, hasOnField := frontmatterForValidation["on"] - if !hasOnField { - // Check if this is a redirect-only placeholder (has a redirect field but no 'on' trigger). - // Redirect-only files are distinct from regular shared workflows: they are placeholders - // that point to a workflow's new canonical location and are not intended to be imported. - // They occur when `gh aw add` downloads a workflow that has been moved but the redirect - // was not resolved to the full content during download. - if redirectVal, hasRedirect := frontmatterForValidation["redirect"]; hasRedirect { - if redirectStr, ok := redirectVal.(string); ok { - if redirectTarget := strings.TrimSpace(redirectStr); redirectTarget != "" { - detectionLog.Printf("Redirect-only workflow detected: redirect=%s", redirectTarget) - return &frontmatterParseResult{ - cleanPath: cleanPath, - content: content, - frontmatterResult: result, - frontmatterForValidation: frontmatterForValidation, - markdownDir: filepath.Dir(cleanPath), - isRedirectOnly: true, - redirectTarget: redirectTarget, - }, nil - } - } - } - - detectionLog.Printf("No 'on' field detected - treating as shared agentic workflow") +func (c *Compiler) readFrontmatterSource(cleanPath string) ([]byte, string, error) { + content, err := os.ReadFile(cleanPath) + if err != nil { + orchestratorFrontmatterLog.Printf("Failed to read file: %s, error: %v", cleanPath, err) + return nil, "", fmt.Errorf("failed to read file: %w", frontmatterReadError{message: err.Error()}) + } + return content, string(content), nil +} - // Validate as an included/shared workflow (uses main_workflow_schema with forbidden field checks) - if err := parser.ValidateIncludedFileFrontmatterWithSchemaAndLocation(frontmatterForValidation, cleanPath); err != nil { - orchestratorFrontmatterLog.Printf("Shared workflow validation failed: %v", err) - return nil, err - } +func (c *Compiler) extractWorkflowFrontmatter(cleanPath string, contentString string) (*parser.FrontmatterResult, error) { + orchestratorFrontmatterLog.Printf("Parsing frontmatter from file: %s", cleanPath) + result, err := parser.ExtractFrontmatterFromContent(contentString) + if err == nil { + return result, nil + } + orchestratorFrontmatterLog.Printf("Frontmatter extraction failed: %v", err) + frontmatterStart := 2 + if result != nil && result.FrontmatterStart > 0 { + frontmatterStart = result.FrontmatterStart + } + return nil, c.createFrontmatterError(cleanPath, contentString, err, frontmatterStart) +} - return &frontmatterParseResult{ - cleanPath: cleanPath, - content: content, - frontmatterResult: result, - frontmatterForValidation: frontmatterForValidation, - markdownDir: filepath.Dir(cleanPath), - isSharedWorkflow: true, - }, nil +func (c *Compiler) handleFrontmatterWithoutOn(cleanPath string, content []byte, result *parser.FrontmatterResult, frontmatterForValidation map[string]any) (*frontmatterParseResult, error) { + if _, hasOnField := frontmatterForValidation["on"]; hasOnField { + return nil, nil + } + if redirectResult := buildRedirectOnlyFrontmatterResult(cleanPath, content, result, frontmatterForValidation); redirectResult != nil { + return redirectResult, nil + } + detectionLog.Printf("No 'on' field detected - treating as shared agentic workflow") + if err := parser.ValidateIncludedFileFrontmatterWithSchemaAndLocation(frontmatterForValidation, cleanPath); err != nil { + orchestratorFrontmatterLog.Printf("Shared workflow validation failed: %v", err) + return nil, err } + sharedResult := newFrontmatterParseResult(cleanPath, content, result, frontmatterForValidation) + sharedResult.isSharedWorkflow = true + return sharedResult, nil +} - // For main workflows (with 'on' field), markdown content is required - if result.Markdown == "" { - orchestratorFrontmatterLog.Print("No markdown content found for main workflow") - return nil, errors.New("no markdown content found") +func buildRedirectOnlyFrontmatterResult(cleanPath string, content []byte, result *parser.FrontmatterResult, frontmatterForValidation map[string]any) *frontmatterParseResult { + redirectVal, hasRedirect := frontmatterForValidation["redirect"] + if !hasRedirect { + return nil } + redirectStr, ok := redirectVal.(string) + if !ok { + return nil + } + redirectTarget := strings.TrimSpace(redirectStr) + if redirectTarget == "" { + return nil + } + detectionLog.Printf("Redirect-only workflow detected: redirect=%s", redirectTarget) + redirectResult := newFrontmatterParseResult(cleanPath, content, result, frontmatterForValidation) + redirectResult.isRedirectOnly = true + redirectResult.redirectTarget = redirectTarget + return redirectResult +} +func (c *Compiler) validateMainWorkflowFrontmatter(cleanPath string, content []byte, result *parser.FrontmatterResult, frontmatterForValidation map[string]any) error { if err := c.validateEngineBeforeSchema(cleanPath, content, result, frontmatterForValidation); err != nil { orchestratorFrontmatterLog.Printf("String engine pre-validation failed: %v", err) - return nil, err + return err + } + if err := c.runMainWorkflowFrontmatterValidators(cleanPath, frontmatterForValidation); err != nil { + return err } + return c.validateMainWorkflowMarkdown(result.Markdown) +} - // Validate main workflow frontmatter contains only expected entries +func (c *Compiler) runMainWorkflowFrontmatterValidators(cleanPath string, frontmatterForValidation map[string]any) error { orchestratorFrontmatterLog.Printf("Validating main workflow frontmatter schema") - if err := parser.ValidateMainWorkflowFrontmatterWithSchemaAndLocation(frontmatterForValidation, cleanPath); err != nil { - orchestratorFrontmatterLog.Printf("Main workflow frontmatter validation failed: %v", err) - return nil, err - } - if err := validateFrontmatterSkills(frontmatterForValidation); err != nil { - orchestratorFrontmatterLog.Printf("Skills frontmatter validation failed: %v", err) - return nil, err + checks := []func() error{ + func() error { + return parser.ValidateMainWorkflowFrontmatterWithSchemaAndLocation(frontmatterForValidation, cleanPath) + }, + func() error { return validateFrontmatterSkills(frontmatterForValidation) }, + func() error { return ValidateEventFilters(frontmatterForValidation) }, + func() error { return c.validatePushBranchScopeFrontmatter(frontmatterForValidation) }, + func() error { return ValidateEventTypes(frontmatterForValidation) }, + func() error { return ValidateGlobPatterns(frontmatterForValidation) }, + func() error { return validateRunsOn(frontmatterForValidation, cleanPath) }, } - - // Validate event filter mutual exclusivity (branches/branches-ignore, paths/paths-ignore) - if err := ValidateEventFilters(frontmatterForValidation); err != nil { - orchestratorFrontmatterLog.Printf("Event filter validation failed: %v", err) - return nil, err + for _, check := range checks { + if err := check(); err != nil { + return err + } } + return nil +} - // Validate that push triggers are scoped to specific branches or tags to prevent fan-out. - // In strict mode this is an error; in non-strict mode it is downgraded to a warning. +func (c *Compiler) validatePushBranchScopeFrontmatter(frontmatterForValidation map[string]any) error { if err := ValidatePushBranchScope(frontmatterForValidation); err != nil { if c.effectiveStrictMode(frontmatterForValidation) { orchestratorFrontmatterLog.Printf("Push branch/tag scope validation failed: %v", err) - return nil, err + return err } orchestratorFrontmatterLog.Printf("Push branch/tag scope warning (non-strict mode): %v", err) fmt.Fprintln(os.Stderr, console.FormatWarningMessage(err.Error())) c.IncrementWarningCount() } + return nil +} - // Validate event type names in the 'on:' section for potential typos - if err := ValidateEventTypes(frontmatterForValidation); err != nil { - orchestratorFrontmatterLog.Printf("Event type validation failed: %v", err) - return nil, err - } - - // Validate glob pattern syntax in event filters (branches, tags, paths, etc.) - if err := ValidateGlobPatterns(frontmatterForValidation); err != nil { - orchestratorFrontmatterLog.Printf("Glob pattern validation failed: %v", err) - return nil, err - } - - // Validate that the runs-on field does not specify unsupported runner types (e.g. macOS) - if err := validateRunsOn(frontmatterForValidation, cleanPath); err != nil { - orchestratorFrontmatterLog.Printf("runs-on validation failed: %v", err) - return nil, err - } - - // Validate that @include/@import directives are not used inside template regions - if err := validateNoIncludesInTemplateRegions(result.Markdown); err != nil { +func (c *Compiler) validateMainWorkflowMarkdown(markdown string) error { + if err := validateNoIncludesInTemplateRegions(markdown); err != nil { orchestratorFrontmatterLog.Printf("Template region validation failed: %v", err) - return nil, fmt.Errorf("template region validation failed: %w", err) + return fmt.Errorf("template region validation failed: %w", err) } - - // Validate that pre-expanded __GH_AW_EXPERIMENTS_*__ placeholders are not used in template conditions - if err := validateNoPreExpandedExperimentPlaceholders(result.Markdown); err != nil { + if err := validateNoPreExpandedExperimentPlaceholders(markdown); err != nil { orchestratorFrontmatterLog.Printf("Pre-expanded experiment placeholder validation failed: %v", err) - return nil, fmt.Errorf("template condition validation failed: %w", err) + return fmt.Errorf("template condition validation failed: %w", err) } + return nil +} - // Warn when experiment comparison expressions use double-quoted string literals. - // GitHub Actions expression syntax only supports single-quoted string literals, so - // the compiler converts double quotes to single quotes automatically — but authors - // should fix the source to use single quotes to keep it consistent with the output. - for _, w := range detectDoubleQuotedExperimentComparisons(result.Markdown) { +func (c *Compiler) emitMainWorkflowMarkdownWarnings(cleanPath string, markdown string) { + for _, w := range detectDoubleQuotedExperimentComparisons(markdown) { fmt.Fprintln(os.Stderr, formatCompilerMessage(cleanPath, "warning", w)) c.IncrementWarningCount() } - - // Warn when template separators are embedded in the middle of a line. - // Keeping separators on their own lines improves compatibility with the - // template renderer and avoids brittle inline condition blocks. - for _, w := range detectMidlineTemplateSeparators(result.Markdown) { + for _, w := range detectMidlineTemplateSeparators(markdown) { fmt.Fprintln(os.Stderr, formatCompilerMessage(cleanPath, "warning", w)) c.IncrementWarningCount() } +} - workflowLog.Printf("Frontmatter: %d chars, Markdown: %d chars", len(result.Frontmatter), len(result.Markdown)) - +func newFrontmatterParseResult(cleanPath string, content []byte, result *parser.FrontmatterResult, frontmatterForValidation map[string]any) *frontmatterParseResult { return &frontmatterParseResult{ cleanPath: cleanPath, content: content, frontmatterResult: result, frontmatterForValidation: frontmatterForValidation, markdownDir: filepath.Dir(cleanPath), - isSharedWorkflow: false, - }, nil + } } // copyFrontmatterWithoutInternalMarkers creates a copy of frontmatter without internal marker fields. diff --git a/pkg/workflow/compiler_orchestrator_workflow.go b/pkg/workflow/compiler_orchestrator_workflow.go index 3f206ec5f12..8be882a8cf9 100644 --- a/pkg/workflow/compiler_orchestrator_workflow.go +++ b/pkg/workflow/compiler_orchestrator_workflow.go @@ -401,35 +401,42 @@ func (c *Compiler) extractAdditionalConfigurations( safeOutputs *SafeOutputsConfig, ) error { orchestratorWorkflowLog.Print("Extracting additional configurations") + toolsConfig, err := c.populateWorkflowMemoryConfigs(workflowData, tools) + if err != nil { + return err + } + c.populateWorkflowTriggerConfigs(frontmatter, workflowData, importsResult) + if err := c.populateWorkflowSafeOutputConfigs(frontmatter, markdownDir, workflowData, importsResult, markdown, safeOutputs, toolsConfig); err != nil { + return err + } + return c.populateWorkflowExperimentsAndEvals(frontmatter, workflowData) +} - // Extract cache-memory config and check for errors +func (c *Compiler) populateWorkflowMemoryConfigs(workflowData *WorkflowData, tools map[string]any) (*ToolsConfig, error) { cacheMemoryConfig, err := c.extractCacheMemoryConfigFromMap(tools) if err != nil { - return err + return nil, err } workflowData.CacheMemoryConfig = cacheMemoryConfig - - // Extract repo-memory config and check for errors toolsConfig, err := ParseToolsConfig(tools) if err != nil { - return err + return nil, err } repoMemoryConfig, err := c.extractRepoMemoryConfig(toolsConfig, workflowData.WorkflowID) if err != nil { - return err + return nil, err } workflowData.RepoMemoryConfig = repoMemoryConfig + return toolsConfig, nil +} - // Extract and process mcp-scripts and safe-outputs +func (c *Compiler) populateWorkflowTriggerConfigs(frontmatter map[string]any, workflowData *WorkflowData, importsResult *parser.ImportsResult) { workflowData.Command, workflowData.CommandEvents, workflowData.CommandCentralized, workflowData.CommandPlaceholder = c.extractCommandConfig(frontmatter) workflowData.LabelCommand, workflowData.LabelCommandEvents, workflowData.LabelCommandDecentralized, workflowData.LabelCommandRemoveLabel = c.extractLabelCommandConfig(frontmatter) workflowData.Jobs = c.extractJobsFromFrontmatter(frontmatter) - - // Merge jobs from imported YAML workflows if importsResult.MergedJobs != "" && importsResult.MergedJobs != "{}" { workflowData.Jobs = c.mergeJobsFromYAMLImports(workflowData.Jobs, importsResult.MergedJobs) } - workflowData.Roles = c.extractRoles(frontmatter) workflowData.Bots = expandBotNames(mergeBots(c.extractBots(frontmatter), importsResult.MergedBots)) workflowData.LabelNames = c.extractLabelNames(frontmatter) @@ -441,111 +448,95 @@ func (c *Compiler) extractAdditionalConfigurations( workflowData.ActivationGitHubToken = c.resolveActivationGitHubToken(frontmatter, importsResult) workflowData.ActivationGitHubApp = c.resolveActivationGitHubApp(frontmatter, importsResult) workflowData.TopLevelGitHubApp = resolveTopLevelGitHubApp(frontmatter, importsResult) +} - // Use the already extracted output configuration +func (c *Compiler) populateWorkflowSafeOutputConfigs(frontmatter map[string]any, markdownDir string, workflowData *WorkflowData, importsResult *parser.ImportsResult, markdown string, safeOutputs *SafeOutputsConfig, toolsConfig *ToolsConfig) error { workflowData.SafeOutputs = safeOutputs - - // Extract comment-memory from tools and attach to safe-outputs configuration. - // comment-memory now belongs under tools: next to cache-memory and repo-memory. - commentMemoryConfig := c.extractCommentMemoryConfig(toolsConfig) - if commentMemoryConfig != nil { - if workflowData.SafeOutputs == nil { - workflowData.SafeOutputs = &SafeOutputsConfig{} - } - workflowData.SafeOutputs.CommentMemory = commentMemoryConfig - } - - // Extract mcp-scripts configuration + attachCommentMemoryConfig(workflowData, c.extractCommentMemoryConfig(toolsConfig)) workflowData.MCPScripts = c.extractMCPScriptsConfig(frontmatter) - - // Merge mcp-scripts from imports if len(importsResult.MergedMCPScripts) > 0 { workflowData.MCPScripts = c.mergeMCPScripts(workflowData.MCPScripts, importsResult.MergedMCPScripts) } - - // Extract safe-jobs from safe-outputs.jobs location - topSafeJobs := extractSafeJobsFromFrontmatter(frontmatter) - - // Process @include directives to extract additional safe-outputs configurations - includedSafeOutputsConfigs, err := parser.ExpandIncludesForSafeOutputs(markdown, markdownDir) + allSafeOutputsConfigs, err := c.collectSafeOutputsConfigs(markdownDir, importsResult, markdown) if err != nil { - return fmt.Errorf("failed to expand includes for safe-outputs: %w", err) + return err + } + if err := c.mergeWorkflowSafeOutputs(frontmatter, workflowData, safeOutputs, allSafeOutputsConfigs); err != nil { + return err } + applyDefaultCreateIssue(workflowData) + applyTopLevelGitHubAppFallbacks(workflowData) + return nil +} - // Combine imported safe-outputs with included safe-outputs - var allSafeOutputsConfigs []string - if len(importsResult.MergedSafeOutputs) > 0 { - allSafeOutputsConfigs = append(allSafeOutputsConfigs, importsResult.MergedSafeOutputs...) +func attachCommentMemoryConfig(workflowData *WorkflowData, commentMemoryConfig *CommentMemoryConfig) { + if commentMemoryConfig == nil { + return } - if len(includedSafeOutputsConfigs) > 0 { - allSafeOutputsConfigs = append(allSafeOutputsConfigs, includedSafeOutputsConfigs...) + if workflowData.SafeOutputs == nil { + workflowData.SafeOutputs = &SafeOutputsConfig{} } + workflowData.SafeOutputs.CommentMemory = commentMemoryConfig +} - // Merge safe-jobs from all safe-outputs configurations (imported and included) +func (c *Compiler) collectSafeOutputsConfigs(markdownDir string, importsResult *parser.ImportsResult, markdown string) ([]string, error) { + includedSafeOutputsConfigs, err := parser.ExpandIncludesForSafeOutputs(markdown, markdownDir) + if err != nil { + return nil, fmt.Errorf("failed to expand includes for safe-outputs: %w", err) + } + allSafeOutputsConfigs := append([]string{}, importsResult.MergedSafeOutputs...) + allSafeOutputsConfigs = append(allSafeOutputsConfigs, includedSafeOutputsConfigs...) + return allSafeOutputsConfigs, nil +} + +func (c *Compiler) mergeWorkflowSafeOutputs(frontmatter map[string]any, workflowData *WorkflowData, safeOutputs *SafeOutputsConfig, allSafeOutputsConfigs []string) error { + topSafeJobs := extractSafeJobsFromFrontmatter(frontmatter) includedSafeJobs, err := c.mergeSafeJobsFromIncludedConfigs(topSafeJobs, allSafeOutputsConfigs) if err != nil { return fmt.Errorf("failed to merge safe-jobs from includes: %w", err) } - - // Merge app configuration from included safe-outputs configurations includedApp, err := c.mergeAppFromIncludedConfigs(workflowData.SafeOutputs, allSafeOutputsConfigs) if err != nil { return fmt.Errorf("failed to merge app from includes: %w", err) } + ensureWorkflowSafeOutputsTargets(workflowData, includedSafeJobs, includedApp) + rawSafeOutputsMap, _ := frontmatter["safe-outputs"].(map[string]any) + mergedSafeOutputs, err := c.MergeSafeOutputs(workflowData.SafeOutputs, allSafeOutputsConfigs, rawSafeOutputsMap) + if err != nil { + return fmt.Errorf("failed to merge safe-outputs from imports: %w", err) + } + workflowData.SafeOutputs = mergedSafeOutputs + applyImportedSafeOutputThreatDetectionDefault(workflowData, safeOutputs, allSafeOutputsConfigs) + return nil +} - // Ensure SafeOutputs exists and populate the Jobs field with merged jobs +func ensureWorkflowSafeOutputsTargets(workflowData *WorkflowData, includedSafeJobs map[string]*SafeJobConfig, includedApp *GitHubAppConfig) { if workflowData.SafeOutputs == nil && len(includedSafeJobs) > 0 { workflowData.SafeOutputs = &SafeOutputsConfig{} } - // Always use the merged includedSafeJobs as it contains both main and imported jobs if workflowData.SafeOutputs != nil && len(includedSafeJobs) > 0 { workflowData.SafeOutputs.Jobs = includedSafeJobs } - - // Populate the App field if it's not set in the top-level workflow but is in an included config if workflowData.SafeOutputs != nil && workflowData.SafeOutputs.GitHubApp == nil && includedApp != nil { workflowData.SafeOutputs.GitHubApp = includedApp } +} - // Merge safe-outputs types from imports. - // Pass the raw safe-outputs map from frontmatter so MergeSafeOutputs can distinguish - // between types the user explicitly configured and types that were auto-defaulted by - // extractSafeOutputsConfig. Without this, auto-defaults (e.g. threat-detection) would - // prevent imported configurations for those types from being merged. - rawSafeOutputsMap, _ := frontmatter["safe-outputs"].(map[string]any) - mergedSafeOutputs, err := c.MergeSafeOutputs(workflowData.SafeOutputs, allSafeOutputsConfigs, rawSafeOutputsMap) - if err != nil { - return fmt.Errorf("failed to merge safe-outputs from imports: %w", err) +func applyImportedSafeOutputThreatDetectionDefault(workflowData *WorkflowData, safeOutputs *SafeOutputsConfig, allSafeOutputsConfigs []string) { + if safeOutputs != nil || workflowData.SafeOutputs == nil || workflowData.SafeOutputs.ThreatDetection != nil { + return } - workflowData.SafeOutputs = mergedSafeOutputs - - // Apply default threat detection when safe-outputs came entirely from imports/includes - // (i.e. the main frontmatter has no safe-outputs: section). In this case the merge - // produces a non-nil SafeOutputs but leaves ThreatDetection nil, which would suppress - // the detection gate on the safe_outputs job. Mirroring the behaviour of - // extractSafeOutputsConfig for direct frontmatter declarations, we enable detection by - // default unless any imported config explicitly sets threat-detection: false. - if safeOutputs == nil && workflowData.SafeOutputs != nil && workflowData.SafeOutputs.ThreatDetection == nil { - if !isThreatDetectionExplicitlyDisabledInConfigs(allSafeOutputsConfigs) { - orchestratorWorkflowLog.Print("Applying default threat-detection for safe-outputs assembled from imports/includes") - workflowData.SafeOutputs.ThreatDetection = &ThreatDetectionConfig{} - } + if isThreatDetectionExplicitlyDisabledInConfigs(allSafeOutputsConfigs) { + return } + orchestratorWorkflowLog.Print("Applying default threat-detection for safe-outputs assembled from imports/includes") + workflowData.SafeOutputs.ThreatDetection = &ThreatDetectionConfig{} +} - // Auto-inject create-issues if safe-outputs is configured but has no non-builtin outputs. - // This ensures every workflow with safe-outputs has at least one meaningful action handler. - applyDefaultCreateIssue(workflowData) - - // Apply the top-level github-app as a fallback for all nested github-app token minting operations. - // This runs last so that all section-specific configurations have been resolved first. - applyTopLevelGitHubAppFallbacks(workflowData) - - // Extract experiments configuration once; derive the simple variants map from the configs. +func (c *Compiler) populateWorkflowExperimentsAndEvals(frontmatter map[string]any, workflowData *WorkflowData) error { workflowData.ExperimentConfigs = extractExperimentConfigsFromFrontmatter(frontmatter) workflowData.Experiments = experimentVariantsFromConfigs(workflowData.ExperimentConfigs) workflowData.ExperimentsStorage = extractExperimentsStorageFromFrontmatter(frontmatter) - - // Extract BinEval evals configuration. evalsConfig, err := c.parseEvalsFromFrontmatter(frontmatter) if err != nil { return fmt.Errorf("invalid evals configuration: %w", err) @@ -554,7 +545,6 @@ func (c *Compiler) extractAdditionalConfigurations( if err := validateExperimentMetricReferences(workflowData.ExperimentConfigs, workflowData.Evals); err != nil { return fmt.Errorf("invalid experiments configuration: %w", err) } - return nil } @@ -625,98 +615,81 @@ func (c *Compiler) processOnSectionAndFilters( cleanPath string, ) error { orchestratorWorkflowLog.Print("Processing on section and filters") - - // Process stop-after configuration from the on: section - if err := c.processStopAfterConfiguration(frontmatter, workflowData, cleanPath); err != nil { - return err - } - - // Process skip-if-match configuration from the on: section - if err := c.processSkipIfMatchConfiguration(frontmatter, workflowData); err != nil { - return err - } - - // Process skip-if-no-match configuration from the on: section - if err := c.processSkipIfNoMatchConfiguration(frontmatter, workflowData); err != nil { - return err - } - - // Process skip-if-check-failing configuration from the on: section - if err := c.processSkipIfCheckFailingConfiguration(frontmatter, workflowData); err != nil { - return err - } - - // Process manual-approval configuration from the on: section - if err := c.processManualApprovalConfiguration(frontmatter, workflowData); err != nil { + if err := c.processOnSectionCore(frontmatter, workflowData, cleanPath); err != nil { return err } + c.applyOnSectionFilters(workflowData, frontmatter) + return c.populateOnSectionDerivedFields(frontmatter, workflowData) +} - // Parse the "on" section for command triggers, reactions, and other events - if err := c.parseOnSection(frontmatter, workflowData, cleanPath); err != nil { - return err +func (c *Compiler) processOnSectionCore(frontmatter map[string]any, workflowData *WorkflowData, cleanPath string) error { + checks := []func() error{ + func() error { return c.processStopAfterConfiguration(frontmatter, workflowData, cleanPath) }, + func() error { return c.processSkipIfMatchConfiguration(frontmatter, workflowData) }, + func() error { return c.processSkipIfNoMatchConfiguration(frontmatter, workflowData) }, + func() error { return c.processSkipIfCheckFailingConfiguration(frontmatter, workflowData) }, + func() error { return c.processManualApprovalConfiguration(frontmatter, workflowData) }, + func() error { return c.parseOnSection(frontmatter, workflowData, cleanPath) }, + func() error { return c.applyDefaults(workflowData, cleanPath) }, } - - // Apply defaults - if err := c.applyDefaults(workflowData, cleanPath); err != nil { - return err + for _, check := range checks { + if err := check(); err != nil { + return err + } } + return nil +} - // Apply pull request draft filter if specified +func (c *Compiler) applyOnSectionFilters(workflowData *WorkflowData, frontmatter map[string]any) { c.applyPullRequestDraftFilter(workflowData, frontmatter) - - // Apply pull request fork filter if specified c.applyPullRequestForkFilter(workflowData, frontmatter) - - // Apply pull request stack filter (default: latest stacked PR only) c.applyPullRequestStackFilter(workflowData, frontmatter) - - // Apply label filter if specified c.applyLabelFilter(workflowData, frontmatter) +} - // Extract on.steps for pre-activation step injection +func (c *Compiler) populateOnSectionDerivedFields(frontmatter map[string]any, workflowData *WorkflowData) error { onSteps, err := extractOnSteps(frontmatter) if err != nil { return err } - - // Apply action pinning to on.steps - if len(onSteps) > 0 { - anySteps := make([]any, len(onSteps)) - for i, s := range onSteps { - anySteps[i] = s - } - typedSteps, convErr := SliceToSteps(anySteps) - if convErr == nil { - typedSteps, convErr = applyActionPinsToTypedSteps(typedSteps, workflowData) - if convErr != nil { - return fmt.Errorf("on.steps: %w", convErr) - } - for i, s := range typedSteps { - onSteps[i] = s.ToMap() - } - } else { - orchestratorWorkflowLog.Printf("Failed to convert on.steps to typed steps for action pinning: %v", convErr) - } + onSteps, err = applyOnSectionStepPins(onSteps, workflowData) + if err != nil { + return err } - workflowData.OnSteps = onSteps - - // Extract on.permissions for pre-activation job permissions workflowData.OnPermissions = extractOnPermissions(frontmatter) - - // Extract on.needs for pre-activation/activation job dependencies onNeeds, err := extractOnNeeds(frontmatter) if err != nil { return err } workflowData.OnNeeds = onNeeds - - // Extract on.restore-memory to opt in to pre-activation memory restore for on.steps. onRestoreMemory, err := extractOnRestoreMemory(frontmatter) if err != nil { return err } workflowData.OnRestoreMemory = onRestoreMemory - return nil } + +func applyOnSectionStepPins(onSteps []map[string]any, workflowData *WorkflowData) ([]map[string]any, error) { + if len(onSteps) == 0 { + return onSteps, nil + } + anySteps := make([]any, len(onSteps)) + for i, s := range onSteps { + anySteps[i] = s + } + typedSteps, err := SliceToSteps(anySteps) + if err != nil { + orchestratorWorkflowLog.Printf("Failed to convert on.steps to typed steps for action pinning: %v", err) + return onSteps, nil + } + typedSteps, err = applyActionPinsToTypedSteps(typedSteps, workflowData) + if err != nil { + return nil, fmt.Errorf("on.steps: %w", err) + } + for i, s := range typedSteps { + onSteps[i] = s.ToMap() + } + return onSteps, nil +} diff --git a/pkg/workflow/compiler_safe_output_jobs.go b/pkg/workflow/compiler_safe_output_jobs.go index bf71421a987..850912c48b4 100644 --- a/pkg/workflow/compiler_safe_output_jobs.go +++ b/pkg/workflow/compiler_safe_output_jobs.go @@ -19,35 +19,47 @@ func (c *Compiler) buildSafeOutputsJobs(data *WorkflowData, jobName, markdownPat return nil } compilerSafeOutputJobsLog.Print("Building safe outputs jobs") + state := &safeOutputsJobBuildState{threatDetectionEnabled: IsDetectionJobEnabled(data.SafeOutputs)} + if err := c.addDetectionJobIfNeeded(data, state); err != nil { + return err + } + if err := c.addPrimarySafeOutputJobs(data, jobName, markdownPath, state); err != nil { + return err + } + if err := c.addOptionalSafeOutputJobs(data, jobName, markdownPath, state); err != nil { + return err + } + unlockJob, err := c.addUnlockJobIfNeeded(data, state.threatDetectionEnabled) + if err != nil { + return err + } + return c.addConclusionSafeOutputJob(data, jobName, state, unlockJob) +} - // Detection is always enabled for safe-outputs workflows unless threat-detection is explicitly - // disabled (threat-detection: false) or the engine is disabled with no custom steps - // (threat-detection: { engine: false } with no steps). ThreatDetection is nil only when - // explicitly disabled. When engine is false with no custom steps, the detection job has - // nothing to run so it is skipped entirely. - threatDetectionEnabled := IsDetectionJobEnabled(data.SafeOutputs) +type safeOutputsJobBuildState struct { + threatDetectionEnabled bool + safeOutputJobNames []string +} - // Build the separate detection job. Detection runs by default for all safe-outputs workflows - // and is only skipped when ThreatDetection is nil (i.e. threat-detection: false was set). - // The detection job runs after the agent job, downloads the agent artifact, - // and outputs detection_success and detection_conclusion for downstream jobs. - if threatDetectionEnabled { - detectionJob, err := c.buildDetectionJob(data) - if err != nil { - return fmt.Errorf("failed to build detection job: %w", err) - } - if detectionJob != nil { - if err := c.jobManager.AddJob(detectionJob); err != nil { - return fmt.Errorf("failed to add detection job: %w", err) - } - compilerSafeOutputJobsLog.Print("Added separate detection job") - } +func (c *Compiler) addDetectionJobIfNeeded(data *WorkflowData, state *safeOutputsJobBuildState) error { + if !state.threatDetectionEnabled { + return nil } + detectionJob, err := c.buildDetectionJob(data) + if err != nil { + return fmt.Errorf("failed to build detection job: %w", err) + } + if detectionJob == nil { + return nil + } + if err := c.jobManager.AddJob(detectionJob); err != nil { + return fmt.Errorf("failed to add detection job: %w", err) + } + compilerSafeOutputJobsLog.Print("Added separate detection job") + return nil +} - // Track safe output job names to establish dependencies for conclusion job - var safeOutputJobNames []string - - // Build consolidated safe outputs job containing all safe output operations as steps +func (c *Compiler) addPrimarySafeOutputJobs(data *WorkflowData, jobName, markdownPath string, state *safeOutputsJobBuildState) error { consolidatedJob, consolidatedStepNames, err := c.buildConsolidatedSafeOutputsJob(data, jobName, markdownPath) if err != nil { return fmt.Errorf("failed to build consolidated safe outputs job: %w", err) @@ -56,105 +68,103 @@ func (c *Compiler) buildSafeOutputsJobs(data *WorkflowData, jobName, markdownPat if err := c.jobManager.AddJob(consolidatedJob); err != nil { return fmt.Errorf("failed to add consolidated safe outputs job: %w", err) } - safeOutputJobNames = append(safeOutputJobNames, consolidatedJob.Name) + state.safeOutputJobNames = append(state.safeOutputJobNames, consolidatedJob.Name) compilerSafeOutputJobsLog.Printf("Added consolidated safe outputs job with %d steps: %v", len(consolidatedStepNames), consolidatedStepNames) } - - // Build safe-jobs if configured - // Safe-jobs should depend on agent job (always) AND detection job (if threat detection is enabled) - // These custom safe-jobs should also be included in the conclusion job's dependencies - safeJobNames, err := c.buildSafeJobs(data, threatDetectionEnabled) + safeJobNames, err := c.buildSafeJobs(data, state.threatDetectionEnabled) if err != nil { return fmt.Errorf("failed to build safe-jobs: %w", err) } - // Add custom safe-job names to the list of safe output jobs - safeOutputJobNames = append(safeOutputJobNames, safeJobNames...) + state.safeOutputJobNames = append(state.safeOutputJobNames, safeJobNames...) compilerSafeOutputJobsLog.Printf("Added %d custom safe-job names to conclusion dependencies", len(safeJobNames)) + return nil +} - // Build upload_assets job as a separate job if configured - // This needs to be separate from the consolidated safe_outputs job because it requires: - // 1. Git configuration for pushing to orphaned branches - // 2. Checkout with proper credentials - // 3. Different permissions (contents: write) - if data.SafeOutputs != nil && data.SafeOutputs.UploadAssets != nil { - compilerSafeOutputJobsLog.Print("Building separate upload_assets job") - uploadAssetsJob, err := c.buildUploadAssetsJob(data, jobName, threatDetectionEnabled) - if err != nil { - return fmt.Errorf("failed to build upload_assets job: %w", err) - } - if err := c.jobManager.AddJob(uploadAssetsJob); err != nil { - return fmt.Errorf("failed to add upload_assets job: %w", err) - } - safeOutputJobNames = append(safeOutputJobNames, uploadAssetsJob.Name) - compilerSafeOutputJobsLog.Printf("Added separate upload_assets job") +func (c *Compiler) addOptionalSafeOutputJobs(data *WorkflowData, jobName, markdownPath string, state *safeOutputsJobBuildState) error { + if err := c.addUploadAssetsJobIfNeeded(data, jobName, state); err != nil { + return err } - - // Build upload_code_scanning_sarif job as a separate job if create-code-scanning-alert is configured. - // This job runs after safe_outputs and only when the safe_outputs job exported a SARIF file. - // It is separate to avoid the checkout step (needed to restore HEAD to github.sha) from - // interfering with other safe-output operations in the consolidated safe_outputs job. - if data.SafeOutputs != nil && data.SafeOutputs.CreateCodeScanningAlerts != nil && - !isHandlerStaged(templatableBoolIsTrue(data.SafeOutputs.Staged), data.SafeOutputs.CreateCodeScanningAlerts.Staged) { - compilerSafeOutputJobsLog.Print("Building separate upload_code_scanning_sarif job") - codeScanningJob, err := c.buildCodeScanningUploadJob(data) - if err != nil { - return fmt.Errorf("failed to build upload_code_scanning_sarif job: %w", err) - } - if err := c.jobManager.AddJob(codeScanningJob); err != nil { - return fmt.Errorf("failed to add upload_code_scanning_sarif job: %w", err) - } - safeOutputJobNames = append(safeOutputJobNames, codeScanningJob.Name) - compilerSafeOutputJobsLog.Printf("Added separate upload_code_scanning_sarif job") + if err := c.addCodeScanningJobIfNeeded(data, state); err != nil { + return err } - - // Build conditional call-workflow fan-out jobs if configured. - // Each allowed worker gets its own `uses:` job with an `if:` condition that - // checks whether safe_outputs selected it. Only one runs per execution. callWorkflowJobNames, err := c.buildCallWorkflowJobs(data, markdownPath) if err != nil { return fmt.Errorf("failed to build call-workflow fan-out jobs: %w", err) } - safeOutputJobNames = append(safeOutputJobNames, callWorkflowJobNames...) + state.safeOutputJobNames = append(state.safeOutputJobNames, callWorkflowJobNames...) compilerSafeOutputJobsLog.Printf("Added %d call-workflow fan-out jobs", len(callWorkflowJobNames)) + return nil +} + +func (c *Compiler) addUploadAssetsJobIfNeeded(data *WorkflowData, jobName string, state *safeOutputsJobBuildState) error { + if data.SafeOutputs == nil || data.SafeOutputs.UploadAssets == nil { + return nil + } + compilerSafeOutputJobsLog.Print("Building separate upload_assets job") + uploadAssetsJob, err := c.buildUploadAssetsJob(data, jobName, state.threatDetectionEnabled) + if err != nil { + return fmt.Errorf("failed to build upload_assets job: %w", err) + } + if err := c.jobManager.AddJob(uploadAssetsJob); err != nil { + return fmt.Errorf("failed to add upload_assets job: %w", err) + } + state.safeOutputJobNames = append(state.safeOutputJobNames, uploadAssetsJob.Name) + compilerSafeOutputJobsLog.Print("Added separate upload_assets job") + return nil +} + +func (c *Compiler) addCodeScanningJobIfNeeded(data *WorkflowData, state *safeOutputsJobBuildState) error { + if data.SafeOutputs == nil || data.SafeOutputs.CreateCodeScanningAlerts == nil || + isHandlerStaged(templatableBoolIsTrue(data.SafeOutputs.Staged), data.SafeOutputs.CreateCodeScanningAlerts.Staged) { + return nil + } + compilerSafeOutputJobsLog.Print("Building separate upload_code_scanning_sarif job") + codeScanningJob, err := c.buildCodeScanningUploadJob(data) + if err != nil { + return fmt.Errorf("failed to build upload_code_scanning_sarif job: %w", err) + } + if err := c.jobManager.AddJob(codeScanningJob); err != nil { + return fmt.Errorf("failed to add upload_code_scanning_sarif job: %w", err) + } + state.safeOutputJobNames = append(state.safeOutputJobNames, codeScanningJob.Name) + compilerSafeOutputJobsLog.Print("Added separate upload_code_scanning_sarif job") + return nil +} - // Build dedicated unlock job if lock-for-agent is enabled - // This job is separate from conclusion to ensure it always runs, even if other jobs fail - // It depends on agent and detection (if enabled) to run after workflow execution completes +func (c *Compiler) addUnlockJobIfNeeded(data *WorkflowData, threatDetectionEnabled bool) (*Job, error) { unlockJob, err := c.buildUnlockJob(data, threatDetectionEnabled) if err != nil { - return fmt.Errorf("failed to build unlock job: %w", err) + return nil, fmt.Errorf("failed to build unlock job: %w", err) } - if unlockJob != nil { - if err := c.jobManager.AddJob(unlockJob); err != nil { - return fmt.Errorf("failed to add unlock job: %w", err) - } - compilerSafeOutputJobsLog.Print("Added dedicated unlock job") + if unlockJob == nil { + return nil, nil } + if err := c.jobManager.AddJob(unlockJob); err != nil { + return nil, fmt.Errorf("failed to add unlock job: %w", err) + } + compilerSafeOutputJobsLog.Print("Added dedicated unlock job") + return unlockJob, nil +} - // Build conclusion job if add-comment is configured OR if command trigger is configured with reactions - // This job runs last, after all safe output jobs (and push_repo_memory if configured), to update the activation comment on failure - // The buildConclusionJob function itself will decide whether to create the job based on the configuration - conclusionJob, err := c.buildConclusionJob(data, jobName, safeOutputJobNames) +func (c *Compiler) addConclusionSafeOutputJob(data *WorkflowData, jobName string, state *safeOutputsJobBuildState, unlockJob *Job) error { + conclusionJob, err := c.buildConclusionJob(data, jobName, state.safeOutputJobNames) if err != nil { return fmt.Errorf("failed to build conclusion job: %w", err) } - if conclusionJob != nil { - // If unlock job exists, conclusion should depend on it to run after unlock completes - if unlockJob != nil { - conclusionJob.Needs = append(conclusionJob.Needs, "unlock") - compilerSafeOutputJobsLog.Printf("Added unlock job dependency to conclusion job") - } - // If push_repo_memory job exists, conclusion should depend on it - // Check if the job was already created (it's created in buildJobs) - if _, exists := c.jobManager.GetJob("push_repo_memory"); exists { - conclusionJob.Needs = append(conclusionJob.Needs, "push_repo_memory") - compilerSafeOutputJobsLog.Printf("Added push_repo_memory dependency to conclusion job") - } - if err := c.jobManager.AddJob(conclusionJob); err != nil { - return fmt.Errorf("failed to add conclusion job: %w", err) - } + if conclusionJob == nil { + return nil + } + if unlockJob != nil { + conclusionJob.Needs = append(conclusionJob.Needs, "unlock") + compilerSafeOutputJobsLog.Printf("Added unlock job dependency to conclusion job") + } + if _, exists := c.jobManager.GetJob("push_repo_memory"); exists { + conclusionJob.Needs = append(conclusionJob.Needs, "push_repo_memory") + compilerSafeOutputJobsLog.Printf("Added push_repo_memory dependency to conclusion job") + } + if err := c.jobManager.AddJob(conclusionJob); err != nil { + return fmt.Errorf("failed to add conclusion job: %w", err) } - return nil } @@ -186,167 +196,142 @@ func (c *Compiler) buildCallWorkflowJobs(data *WorkflowData, markdownPath string } compilerSafeOutputJobsLog.Printf("Building %d call-workflow fan-out jobs", len(config.Workflows)) - - var jobNames []string - + jobNames := make([]string, 0, len(config.Workflows)) for _, workflowName := range config.Workflows { - // Build the job name: "call-{sanitized-workflow-name}" - // sanitizeJobName normalizes underscores to hyphens (NormalizeSafeOutputIdentifier + dash conversion) - sanitizedName := sanitizeJobName(workflowName) - jobName := "call-" + sanitizedName - - // Determine the relative path to the worker workflow file - workflowPath, ok := config.WorkflowFiles[workflowName] - if !ok || workflowPath == "" { - // Fallback: construct path from name - workflowPath = fmt.Sprintf("./.github/workflows/%s.lock.yml", workflowName) + callJob, jobName, err := c.buildSingleCallWorkflowJob(data, markdownPath, workflowName, config) + if err != nil { + return nil, err } - - // Build the with: block. Forward one entry per declared workflow_call input - // on the worker, derived from the payload, so that worker steps can reference - // inputs. directly without parsing JSON. The canonical `payload` - // envelope is only forwarded when the worker explicitly declares a `payload` - // input; GitHub Actions rejects a `uses:` step that passes an input the - // called workflow does not declare, so it must not be added unconditionally. - jobNeeds := []string{"safe_outputs"} - with := map[string]any{} - - if markdownPath != "" { - fileResult, findErr := findWorkflowFile(workflowName, markdownPath) - if findErr != nil { - compilerSafeOutputJobsLog.Printf("Warning: could not find worker workflow file for '%s': %v. "+ - "Typed inputs will not be forwarded in the with: block.", workflowName, findErr) - } else { - var workflowInputs map[string]any - var inputErr error - switch { - case fileResult.lockExists: - workflowInputs, inputErr = extractWorkflowCallInputs(fileResult.lockPath) - case fileResult.ymlExists: - workflowInputs, inputErr = extractWorkflowCallInputs(fileResult.ymlPath) - case fileResult.mdExists: - workflowInputs, inputErr = extractMDWorkflowCallInputs(fileResult.mdPath) - default: - compilerSafeOutputJobsLog.Printf("Warning: no worker file found for '%s'; "+ - "typed inputs will not be forwarded in the with: block.", workflowName) - } - if inputErr != nil { - compilerSafeOutputJobsLog.Printf("Warning: could not extract workflow_call inputs for '%s': %v. "+ - "Typed inputs will not be forwarded in the with: block.", workflowName, inputErr) - } else if workflowInputs != nil { - typedInputCount := 0 - for inputName := range workflowInputs { - if inputName == "payload" { - // The worker explicitly declares the canonical payload - // envelope input; forward the raw transport rather than a - // fromJSON expression. - with["payload"] = "${{ needs.safe_outputs.outputs.call_workflow_payload }}" - continue - } - with[inputName] = buildCallWorkflowInputExpression(inputName) - typedInputCount++ - } - compilerSafeOutputJobsLog.Printf("Forwarding %d typed inputs for call-workflow job '%s'", typedInputCount, jobName) - } - - } + if err := c.jobManager.AddJob(callJob); err != nil { + return nil, fmt.Errorf("failed to add call-workflow job '%s': %w", jobName, err) } + jobNames = append(jobNames, jobName) + compilerSafeOutputJobsLog.Printf("Added call-workflow job: %s (uses: %s)", jobName, callJob.Uses) + } + return jobNames, nil +} - callJob := &Job{ - Name: jobName, - Needs: jobNeeds, - If: fmt.Sprintf("needs.safe_outputs.outputs.call_workflow_name == '%s'", workflowName), - Uses: workflowPath, - With: with, - } +func (c *Compiler) buildSingleCallWorkflowJob(data *WorkflowData, markdownPath, workflowName string, config *CallWorkflowConfig) (*Job, string, error) { + jobName := "call-" + sanitizeJobName(workflowName) + callJob := &Job{ + Name: jobName, + Needs: []string{"safe_outputs"}, + If: fmt.Sprintf("needs.safe_outputs.outputs.call_workflow_name == '%s'", workflowName), + Uses: resolveCallWorkflowPath(config, workflowName), + With: buildCallWorkflowInputs(workflowName, markdownPath, jobName), + } + c.configureCallWorkflowSecrets(callJob, workflowName, markdownPath, jobName) + c.applyCallWorkflowPermissions(callJob, data, workflowName, markdownPath, jobName) + return callJob, jobName, nil +} - // Infer the minimal set of secrets required by the worker workflow so we can - // pass them explicitly instead of using secrets: inherit. This requires the - // worker to have been compiled with on.workflow_call.secrets declarations. - // If the worker has not yet been compiled (no .lock.yml/.yml), or declares no - // secrets, fall back to secrets: inherit for backward compatibility. - if markdownPath != "" { - workerSecrets, secretsErr := extractCallWorkflowSecrets(workflowName, markdownPath) - if secretsErr != nil { - compilerSafeOutputJobsLog.Printf("Warning: could not extract secrets for call-workflow job '%s': %v. "+ - "Falling back to secrets: inherit.", jobName, secretsErr) - callJob.SecretsInherit = true - } else if len(workerSecrets) == 0 { - // No secrets were extracted from the worker. This can mean either the - // worker declares no workflow_call secrets or its compiled file was not - // found yet. Fall back to secrets: inherit for backward compatibility. - compilerSafeOutputJobsLog.Printf("No workflow_call secrets could be extracted for worker '%s' "+ - "(worker may declare none or its compiled file may not exist yet); using secrets: inherit", workflowName) - callJob.SecretsInherit = true - } else { - // Map each declared secret explicitly. - callJob.Secrets = make(map[string]string, len(workerSecrets)) - for _, s := range workerSecrets { - callJob.Secrets[s] = fmt.Sprintf("${{ secrets.%s }}", s) - } - compilerSafeOutputJobsLog.Printf("Mapped %d explicit secrets for call-workflow job '%s'", len(workerSecrets), jobName) - } - } else { - callJob.SecretsInherit = true - } +func resolveCallWorkflowPath(config *CallWorkflowConfig, workflowName string) string { + if workflowPath := config.WorkflowFiles[workflowName]; workflowPath != "" { + return workflowPath + } + return fmt.Sprintf("./.github/workflows/%s.lock.yml", workflowName) +} - // Compute the call- job's permission envelope as the union of: - // 1. The caller's own declared permissions (the base scope the caller controls). - // 2. The worker's job-level permissions (the minimum the worker needs to run). - // GitHub validates reusable workflow calls against the caller job's declared - // permissions and rejects the run at startup when the caller grants less than - // the worker requires. Taking the union ensures the call job always holds a - // sufficient grant without requiring the caller's markdown to enumerate every - // permission the worker needs. - callerPerms := data.CachedPermissions - if callerPerms == nil { - callerPerms = NewPermissionsParser(data.Permissions).ToPermissions() +func buildCallWorkflowInputs(workflowName, markdownPath, jobName string) map[string]any { + if markdownPath == "" { + return map[string]any{} + } + fileResult, findErr := findWorkflowFile(workflowName, markdownPath) + if findErr != nil { + compilerSafeOutputJobsLog.Printf("Warning: could not find worker workflow file for '%s': %v. Typed inputs will not be forwarded in the with: block.", workflowName, findErr) + return map[string]any{} + } + workflowInputs, err := loadCallWorkflowInputs(fileResult, workflowName) + if err != nil || workflowInputs == nil { + if err != nil { + compilerSafeOutputJobsLog.Printf("Warning: could not extract workflow_call inputs for '%s': %v. Typed inputs will not be forwarded in the with: block.", workflowName, err) } - - effectivePerms := callerPerms - var importedPerms *callWorkflowPermissionImport - var permErr error - if markdownPath != "" { - importedPerms, permErr = extractCallWorkflowPermissionImport(workflowName, markdownPath) - if permErr != nil { - // Non-fatal: log and continue. The worker file may not exist yet (it may be - // compiled in the same batch), in which case we fall back to the caller's - // own declared permissions. - compilerSafeOutputJobsLog.Printf("Could not extract worker permissions for call-workflow job '%s' (falling back to caller-only permissions): %v", jobName, permErr) - } else if importedPerms != nil && importedPerms.permissions != nil { - // Compute the union by merging caller and worker permissions into a - // fresh map-based Permissions. Starting from a blank slate (rather - // than a clone of callerPerms) ensures shorthand values like - // "read-all" are correctly expanded before the worker's explicit - // scopes are merged on top — cloning a shorthand Permissions and then - // merging a map into it would clear the shorthand field without first - // expanding it, silently dropping the caller's baseline grant. - merged := NewPermissions() - merged.Merge(callerPerms) - merged.Merge(importedPerms.permissions) - effectivePerms = merged - compilerSafeOutputJobsLog.Printf("Merged caller and worker permissions for call-workflow job '%s'", jobName) - } + return map[string]any{} + } + with := map[string]any{} + typedInputCount := 0 + for inputName := range workflowInputs { + if inputName == "payload" { + with["payload"] = "${{ needs.safe_outputs.outputs.call_workflow_payload }}" + continue } + with[inputName] = buildCallWorkflowInputExpression(inputName) + typedInputCount++ + } + compilerSafeOutputJobsLog.Printf("Forwarding %d typed inputs for call-workflow job '%s'", typedInputCount, jobName) + return with +} - if effectivePerms != nil { - rendered := effectivePerms.RenderToYAML() - if rendered != "" { - callJob.PermissionsComment = buildCallWorkflowPermissionsComment(workflowName, importedPerms) - callJob.Permissions = rendered - compilerSafeOutputJobsLog.Printf("Set permissions on call-workflow job '%s': %s", jobName, rendered) - } - } +func loadCallWorkflowInputs(fileResult *findWorkflowFileResult, workflowName string) (map[string]any, error) { + switch { + case fileResult.lockExists: + return extractWorkflowCallInputs(fileResult.lockPath) + case fileResult.ymlExists: + return extractWorkflowCallInputs(fileResult.ymlPath) + case fileResult.mdExists: + return extractMDWorkflowCallInputs(fileResult.mdPath) + default: + compilerSafeOutputJobsLog.Printf("Warning: no worker file found for '%s'; typed inputs will not be forwarded in the with: block.", workflowName) + return nil, nil + } +} - if err := c.jobManager.AddJob(callJob); err != nil { - return nil, fmt.Errorf("failed to add call-workflow job '%s': %w", jobName, err) - } +func (c *Compiler) configureCallWorkflowSecrets(callJob *Job, workflowName, markdownPath, jobName string) { + if markdownPath == "" { + callJob.SecretsInherit = true + return + } + workerSecrets, secretsErr := extractCallWorkflowSecrets(workflowName, markdownPath) + if secretsErr != nil { + compilerSafeOutputJobsLog.Printf("Warning: could not extract secrets for call-workflow job '%s': %v. Falling back to secrets: inherit.", jobName, secretsErr) + callJob.SecretsInherit = true + return + } + if len(workerSecrets) == 0 { + compilerSafeOutputJobsLog.Printf("No workflow_call secrets could be extracted for worker '%s' (worker may declare none or its compiled file may not exist yet); using secrets: inherit", workflowName) + callJob.SecretsInherit = true + return + } + callJob.Secrets = make(map[string]string, len(workerSecrets)) + for _, secret := range workerSecrets { + callJob.Secrets[secret] = fmt.Sprintf("${{ secrets.%s }}", secret) + } + compilerSafeOutputJobsLog.Printf("Mapped %d explicit secrets for call-workflow job '%s'", len(workerSecrets), jobName) +} - jobNames = append(jobNames, jobName) - compilerSafeOutputJobsLog.Printf("Added call-workflow job: %s (uses: %s)", jobName, workflowPath) +func (c *Compiler) applyCallWorkflowPermissions(callJob *Job, data *WorkflowData, workflowName, markdownPath, jobName string) { + effectivePerms, importedPerms := computeCallWorkflowPermissions(data, workflowName, markdownPath, jobName) + if effectivePerms == nil { + return } + if rendered := effectivePerms.RenderToYAML(); rendered != "" { + callJob.PermissionsComment = buildCallWorkflowPermissionsComment(workflowName, importedPerms) + callJob.Permissions = rendered + compilerSafeOutputJobsLog.Printf("Set permissions on call-workflow job '%s': %s", jobName, rendered) + } +} - return jobNames, nil +func computeCallWorkflowPermissions(data *WorkflowData, workflowName, markdownPath, jobName string) (*Permissions, *callWorkflowPermissionImport) { + callerPerms := data.CachedPermissions + if callerPerms == nil { + callerPerms = NewPermissionsParser(data.Permissions).ToPermissions() + } + if markdownPath == "" { + return callerPerms, nil + } + importedPerms, permErr := extractCallWorkflowPermissionImport(workflowName, markdownPath) + if permErr != nil { + compilerSafeOutputJobsLog.Printf("Could not extract worker permissions for call-workflow job '%s' (falling back to caller-only permissions): %v", jobName, permErr) + return callerPerms, nil + } + if importedPerms == nil || importedPerms.permissions == nil { + return callerPerms, importedPerms + } + merged := NewPermissions() + merged.Merge(callerPerms) + merged.Merge(importedPerms.permissions) + compilerSafeOutputJobsLog.Printf("Merged caller and worker permissions for call-workflow job '%s'", jobName) + return merged, importedPerms } func buildCallWorkflowInputExpression(inputName string) string { diff --git a/pkg/workflow/compiler_safe_outputs_steps.go b/pkg/workflow/compiler_safe_outputs_steps.go index 7ed712cf284..d2e776fba29 100644 --- a/pkg/workflow/compiler_safe_outputs_steps.go +++ b/pkg/workflow/compiler_safe_outputs_steps.go @@ -89,252 +89,192 @@ func (c *Compiler) buildSharedPRCheckoutSteps(data *WorkflowData) []string { // with a single dispatcher step that processes all safe output types. func (c *Compiler) buildHandlerManagerStep(data *WorkflowData) ([]string, error) { consolidatedSafeOutputsStepsLog.Print("Building handler manager step") + steps := c.buildHandlerManagerTokenMintingSteps(data) + steps = append(steps, + " - name: Process Safe Outputs\n", + " id: process_safe_outputs\n", + fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/github-script", data)), + " env:\n", + " GH_AW_AGENT_OUTPUT: ${{ steps.setup-agent-output-env.outputs.GH_AW_AGENT_OUTPUT }}\n", + " GH_AW_COMMENT_ID: ${{ needs.activation.outputs.comment_id }}\n", + ) + var err error + steps, err = c.appendHandlerManagerEnvSection(steps, data) + if err != nil { + return nil, err + } + steps = append(steps, " with:\n") + configToken := "" + if data.SafeOutputs != nil && data.SafeOutputs.GitHubToken != "" { + configToken = data.SafeOutputs.GitHubToken + } + c.addSafeOutputGitHubTokenForConfig(&steps, data, configToken) + steps = append(steps, " script: |\n", generateGitHubScriptWithRequire("process_safe_outputs.cjs")) + return steps, nil +} - var steps []string +func (c *Compiler) buildHandlerManagerTokenMintingSteps(data *WorkflowData) []string { + if data.SafeOutputs == nil { + return nil + } + steps := c.buildPerHandlerAppTokenSteps(data.SafeOutputs) + return append(steps, c.buildDispatchRepositoryTokenSteps(data.SafeOutputs)...) +} - // Add per-handler GitHub App token minting steps before the handler manager step. - // For each registered handler that has a per-handler github-app configured, mint a - // dedicated token step whose permissions are scoped to only that handler's needs. - // This implements the principle of least privilege: a workflow can assign separate - // apps to different outputs so each token only carries the permissions it requires. - if data.SafeOutputs != nil { - for _, handler := range safeOutputHandlers { - if handler.StructField == "" { - continue - } - handlerApp := getHandlerGitHubApp(data.SafeOutputs, handler.StructField) - if handlerApp == nil || handler.PermissionBuilder == nil { - continue - } - handlerPermissions := handler.PermissionBuilder(data.SafeOutputs) - if handlerPermissions == nil { - continue - } - stepID := handler.Key + "-app-token" - consolidatedSafeOutputsStepsLog.Printf("Adding per-handler GitHub App token minting step for %s", handler.Key) - steps = append(steps, c.buildGitHubAppTokenMintStepWithMeta( - handlerApp, - handlerPermissions, - "", - "", - fmt.Sprintf("Generate GitHub App token (%s)", handler.Key), - stepID, - )...) +func (c *Compiler) buildPerHandlerAppTokenSteps(safeOutputs *SafeOutputsConfig) []string { + var steps []string + for _, handler := range safeOutputHandlers { + if handler.StructField == "" { + continue } - - if data.SafeOutputs.DispatchRepository != nil && len(data.SafeOutputs.DispatchRepository.Tools) > 0 { - toolKeys := make([]string, 0, len(data.SafeOutputs.DispatchRepository.Tools)) - for toolKey := range data.SafeOutputs.DispatchRepository.Tools { - toolKeys = append(toolKeys, toolKey) - } - sort.Strings(toolKeys) - globalStaged := templatableBoolIsTrue(data.SafeOutputs.Staged) - for _, toolKey := range toolKeys { - tool := data.SafeOutputs.DispatchRepository.Tools[toolKey] - if tool == nil || tool.GitHubApp == nil || isHandlerStaged(globalStaged, tool.Staged) { - continue - } - stepID := dispatchRepositoryToolAppTokenStepID(toolKey) - consolidatedSafeOutputsStepsLog.Printf("Adding dispatch-repository GitHub App token minting step for %s", toolKey) - steps = append(steps, c.buildGitHubAppTokenMintStepWithMeta( - tool.GitHubApp, - NewPermissionsContentsWrite(), - "", - "", - fmt.Sprintf("Generate GitHub App token (dispatch-repository %s)", toolKey), - stepID, - )...) - } + handlerApp := getHandlerGitHubApp(safeOutputs, handler.StructField) + if handlerApp == nil || handler.PermissionBuilder == nil { + continue } + handlerPermissions := handler.PermissionBuilder(safeOutputs) + if handlerPermissions == nil { + continue + } + stepID := handler.Key + "-app-token" + consolidatedSafeOutputsStepsLog.Printf("Adding per-handler GitHub App token minting step for %s", handler.Key) + steps = append(steps, c.buildGitHubAppTokenMintStepWithMeta(handlerApp, handlerPermissions, "", "", fmt.Sprintf("Generate GitHub App token (%s)", handler.Key), stepID)...) } + return steps +} - // Step name and metadata - steps = append(steps, " - name: Process Safe Outputs\n") - steps = append(steps, " id: process_safe_outputs\n") - steps = append(steps, fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/github-script", data))) +func (c *Compiler) buildDispatchRepositoryTokenSteps(safeOutputs *SafeOutputsConfig) []string { + if safeOutputs.DispatchRepository == nil || len(safeOutputs.DispatchRepository.Tools) == 0 { + return nil + } + toolKeys := make([]string, 0, len(safeOutputs.DispatchRepository.Tools)) + for toolKey := range safeOutputs.DispatchRepository.Tools { + toolKeys = append(toolKeys, toolKey) + } + sort.Strings(toolKeys) + globalStaged := templatableBoolIsTrue(safeOutputs.Staged) + var steps []string + for _, toolKey := range toolKeys { + tool := safeOutputs.DispatchRepository.Tools[toolKey] + if tool == nil || tool.GitHubApp == nil || isHandlerStaged(globalStaged, tool.Staged) { + continue + } + stepID := dispatchRepositoryToolAppTokenStepID(toolKey) + consolidatedSafeOutputsStepsLog.Printf("Adding dispatch-repository GitHub App token minting step for %s", toolKey) + steps = append(steps, c.buildGitHubAppTokenMintStepWithMeta(tool.GitHubApp, NewPermissionsContentsWrite(), "", "", fmt.Sprintf("Generate GitHub App token (dispatch-repository %s)", toolKey), stepID)...) + } + return steps +} - // Environment variables - steps = append(steps, " env:\n") - steps = append(steps, " GH_AW_AGENT_OUTPUT: ${{ steps.setup-agent-output-env.outputs.GH_AW_AGENT_OUTPUT }}\n") - steps = append(steps, " GH_AW_COMMENT_ID: ${{ needs.activation.outputs.comment_id }}\n") +func (c *Compiler) appendHandlerManagerEnvSection(steps []string, data *WorkflowData) ([]string, error) { + domainsStr, err := c.resolveHandlerManagerAllowedDomains(data) + if err != nil { + return nil, err + } + steps = appendHandlerManagerBaseEnv(steps, data, domainsStr) + c.addCustomSafeOutputEnvVars(&steps, data) + c.addHandlerManagerConfigEnvVar(&steps, data) + c.addAllSafeOutputConfigEnvVars(&steps, data) + return c.appendHandlerManagerTokenEnv(steps, data), nil +} - // Add allowed domains configuration for URL sanitization in safe output handlers. - // Without this, sanitizeContent() in safe_output_handler_manager.cjs only allows - // default GitHub domains, causing user-configured allowed domains to be redacted. - var domainsStr string +func (c *Compiler) resolveHandlerManagerAllowedDomains(data *WorkflowData) (string, error) { if data.SafeOutputs != nil && len(data.SafeOutputs.AllowedDomains) > 0 { - // allowed-domains: additional domains unioned with engine/network base set; supports ecosystem identifiers - expanded, err := c.computeExpandedAllowedDomainsForSanitization(data) - if err != nil { - return nil, err - } - domainsStr = expanded - } else { - computed, err := c.computeAllowedDomainsForSanitization(data) - if err != nil { - return nil, err - } - domainsStr = computed + return c.computeExpandedAllowedDomainsForSanitization(data) } + return c.computeAllowedDomainsForSanitization(data) +} + +func appendHandlerManagerBaseEnv(steps []string, data *WorkflowData, domainsStr string) []string { if domainsStr != "" { steps = append(steps, fmt.Sprintf(" GH_AW_ALLOWED_DOMAINS: %q\n", domainsStr)) } if data.SafeOutputs != nil && data.SafeOutputs.URLs != "" { steps = append(steps, fmt.Sprintf(" GH_AW_SAFE_OUTPUTS_URLS: %q\n", data.SafeOutputs.URLs)) } - // Pass GitHub server/API URLs so buildAllowedDomains() can add GHES domains dynamically - steps = append(steps, " GITHUB_SERVER_URL: ${{ github.server_url }}\n") - steps = append(steps, " GITHUB_API_URL: ${{ github.api_url }}\n") - - // Note: The project handler manager has been removed. - // All project-related operations are now handled by the unified handler. - - // Add GH_AW_SAFE_OUTPUT_JOBS so the handler manager knows which message types are - // handled by custom safe-output job steps and should be silently skipped rather than - // reported as "No handler loaded for message type '...'". - if customJobsJSON := buildCustomSafeOutputJobsJSON(data); customJobsJSON != "" { - steps = append(steps, fmt.Sprintf(" GH_AW_SAFE_OUTPUT_JOBS: %q\n", customJobsJSON)) - consolidatedSafeOutputsStepsLog.Print("Added GH_AW_SAFE_OUTPUT_JOBS env var for custom safe job types") + steps = append(steps, " GITHUB_SERVER_URL: ${{ github.server_url }}\n", " GITHUB_API_URL: ${{ github.api_url }}\n") + for _, entry := range []struct{ key, value, log string }{ + {"GH_AW_SAFE_OUTPUT_JOBS", buildCustomSafeOutputJobsJSON(data), "Added GH_AW_SAFE_OUTPUT_JOBS env var for custom safe job types"}, + {"GH_AW_SAFE_OUTPUT_SCRIPTS", buildCustomSafeOutputScriptsJSON(data), "Added GH_AW_SAFE_OUTPUT_SCRIPTS env var for custom script handlers"}, + {"GH_AW_SAFE_OUTPUT_ACTIONS", buildCustomSafeOutputActionsJSON(data), "Added GH_AW_SAFE_OUTPUT_ACTIONS env var for custom action handlers"}, + } { + if entry.value == "" { + continue + } + steps = append(steps, fmt.Sprintf(" %s: %q\n", entry.key, entry.value)) + consolidatedSafeOutputsStepsLog.Print(entry.log) } + return steps +} - // Add GH_AW_SAFE_OUTPUT_SCRIPTS so the handler manager can load inline script handlers. - // The env var maps normalized script names to their .cjs filenames in the actions folder. - if customScriptsJSON := buildCustomSafeOutputScriptsJSON(data); customScriptsJSON != "" { - steps = append(steps, fmt.Sprintf(" GH_AW_SAFE_OUTPUT_SCRIPTS: %q\n", customScriptsJSON)) - consolidatedSafeOutputsStepsLog.Print("Added GH_AW_SAFE_OUTPUT_SCRIPTS env var for custom script handlers") - } +func (c *Compiler) appendHandlerManagerTokenEnv(steps []string, data *WorkflowData) []string { + steps = appendHandlerManagerCITriggerEnv(steps, data) + steps = appendHandlerManagerProjectEnv(steps, data.SafeOutputs) + steps = appendHandlerManagerAssignmentEnv(steps, data.SafeOutputs) + steps = appendHandlerManagerAgentSessionEnv(steps, data.SafeOutputs) + return appendHandlerManagerGitTokenEnv(steps, data) +} - // Add GH_AW_SAFE_OUTPUT_ACTIONS so the handler manager can load custom action handlers. - // The env var maps normalized action names to themselves (reserved for future extensibility). - if customActionsJSON := buildCustomSafeOutputActionsJSON(data); customActionsJSON != "" { - steps = append(steps, fmt.Sprintf(" GH_AW_SAFE_OUTPUT_ACTIONS: %q\n", customActionsJSON)) - consolidatedSafeOutputsStepsLog.Print("Added GH_AW_SAFE_OUTPUT_ACTIONS env var for custom action handlers") +func appendHandlerManagerCITriggerEnv(steps []string, data *WorkflowData) []string { + if !usesPatchesAndCheckouts(data.SafeOutputs) { + return steps } - - // Add custom safe output env vars - c.addCustomSafeOutputEnvVars(&steps, data) - - // Add handler manager config as JSON - c.addHandlerManagerConfigEnvVar(&steps, data) - - // Add all safe output configuration env vars (still needed by individual handlers) - c.addAllSafeOutputConfigEnvVars(&steps, data) - - // Add extra empty commit token if create-pull-request or push-to-pull-request-branch is configured. - // This token is used to push an empty commit after code changes to trigger CI events, - // working around the GITHUB_TOKEN limitation where events don't trigger other workflows. - // Only emit this env var when one of these safe outputs is actually configured. - if usesPatchesAndCheckouts(data.SafeOutputs) { - var ciTriggerToken string - if data.SafeOutputs.CreatePullRequests != nil && data.SafeOutputs.CreatePullRequests.GithubTokenForExtraEmptyCommit != "" { - ciTriggerToken = data.SafeOutputs.CreatePullRequests.GithubTokenForExtraEmptyCommit - } else if data.SafeOutputs.PushToPullRequestBranch != nil && data.SafeOutputs.PushToPullRequestBranch.GithubTokenForExtraEmptyCommit != "" { - ciTriggerToken = data.SafeOutputs.PushToPullRequestBranch.GithubTokenForExtraEmptyCommit - } - - switch ciTriggerToken { - case "app": - steps = append(steps, " GH_AW_CI_TRIGGER_TOKEN: ${{ steps.safe-outputs-app-token.outputs.token || '' }}\n") - consolidatedSafeOutputsStepsLog.Print("Extra empty commit using GitHub App token") - default: - // Use the magic GH_AW_CI_TRIGGER_TOKEN secret (default behavior when not explicitly configured) - steps = append(steps, fmt.Sprintf(" GH_AW_CI_TRIGGER_TOKEN: %s\n", getEffectiveCITriggerGitHubToken(ciTriggerToken))) - consolidatedSafeOutputsStepsLog.Print("Extra empty commit using GH_AW_CI_TRIGGER_TOKEN") - } + ciTriggerToken := "" + if data.SafeOutputs.CreatePullRequests != nil && data.SafeOutputs.CreatePullRequests.GithubTokenForExtraEmptyCommit != "" { + ciTriggerToken = data.SafeOutputs.CreatePullRequests.GithubTokenForExtraEmptyCommit + } else if data.SafeOutputs.PushToPullRequestBranch != nil && data.SafeOutputs.PushToPullRequestBranch.GithubTokenForExtraEmptyCommit != "" { + ciTriggerToken = data.SafeOutputs.PushToPullRequestBranch.GithubTokenForExtraEmptyCommit } + if ciTriggerToken == "app" { + consolidatedSafeOutputsStepsLog.Print("Extra empty commit using GitHub App token") + return append(steps, " GH_AW_CI_TRIGGER_TOKEN: ${{ steps.safe-outputs-app-token.outputs.token || '' }}\n") + } + consolidatedSafeOutputsStepsLog.Print("Extra empty commit using GH_AW_CI_TRIGGER_TOKEN") + return append(steps, fmt.Sprintf(" GH_AW_CI_TRIGGER_TOKEN: %s\n", getEffectiveCITriggerGitHubToken(ciTriggerToken))) +} - // Add GH_AW_PROJECT_URL and GH_AW_PROJECT_GITHUB_TOKEN environment variables for project operations - // These are set from the project URL and token configured in any project-related safe-output: - // - update-project - // - create-project-status-update - // - create-project - // - // The project field is REQUIRED in update-project and create-project-status-update (enforced by schema validation) - // Agents can optionally override this per-message by including a project field in their output - // - // Note: If multiple project configs are present, we prefer update-project > create-project-status-update > create-project - // This is only relevant for the environment variables - each configuration must explicitly specify its own settings - projectURL, projectToken := resolveProjectURLAndToken(data.SafeOutputs) - +func appendHandlerManagerProjectEnv(steps []string, safeOutputs *SafeOutputsConfig) []string { + projectURL, projectToken := resolveProjectURLAndToken(safeOutputs) if projectURL != "" { steps = append(steps, fmt.Sprintf(" GH_AW_PROJECT_URL: %q\n", projectURL)) } - if projectToken != "" { steps = append(steps, fmt.Sprintf(" GH_AW_PROJECT_GITHUB_TOKEN: %s\n", projectToken)) } + return steps +} - // Add GH_AW_ASSIGN_TO_AGENT_TOKEN when assign-to-agent is configured OR when create-issue - // or create-pull-request is configured with copilot in assignees. All handlers create a - // dedicated Octokit using this token (agent token preference chain), which is required - // because the Copilot assignment API only accepts PATs (not GitHub App tokens). This env - // var is evaluated as a GitHub Actions expression, so it resolves to the actual token value - // before the step runs. - if data.SafeOutputs != nil && data.SafeOutputs.AssignToAgent != nil { - agentTokenStr := getEffectiveCopilotCodingAgentGitHubToken(data.SafeOutputs.AssignToAgent.GitHubToken) - //nolint:gosec // G101: False positive - this is a GitHub Actions expression template, not a hardcoded credential - steps = append(steps, fmt.Sprintf(" GH_AW_ASSIGN_TO_AGENT_TOKEN: %s\n", agentTokenStr)) +func appendHandlerManagerAssignmentEnv(steps []string, safeOutputs *SafeOutputsConfig) []string { + switch { + case safeOutputs != nil && safeOutputs.AssignToAgent != nil: + steps = append(steps, fmt.Sprintf(" GH_AW_ASSIGN_TO_AGENT_TOKEN: %s\n", getEffectiveCopilotCodingAgentGitHubToken(safeOutputs.AssignToAgent.GitHubToken))) consolidatedSafeOutputsStepsLog.Print("Added GH_AW_ASSIGN_TO_AGENT_TOKEN env var for assign-to-agent handler") - } else if data.SafeOutputs != nil && data.SafeOutputs.CreateIssues != nil && hasCopilotAssignee(data.SafeOutputs.CreateIssues.Assignees) { - agentTokenStr := getEffectiveCopilotCodingAgentGitHubToken(data.SafeOutputs.CreateIssues.GitHubToken) - //nolint:gosec // G101: False positive - this is a GitHub Actions expression template, not a hardcoded credential - steps = append(steps, fmt.Sprintf(" GH_AW_ASSIGN_TO_AGENT_TOKEN: %s\n", agentTokenStr)) + case safeOutputs != nil && safeOutputs.CreateIssues != nil && hasCopilotAssignee(safeOutputs.CreateIssues.Assignees): + steps = append(steps, fmt.Sprintf(" GH_AW_ASSIGN_TO_AGENT_TOKEN: %s\n", getEffectiveCopilotCodingAgentGitHubToken(safeOutputs.CreateIssues.GitHubToken))) consolidatedSafeOutputsStepsLog.Print("Added GH_AW_ASSIGN_TO_AGENT_TOKEN env var for create-issue copilot assignment handler") - } else if data.SafeOutputs != nil && data.SafeOutputs.CreatePullRequests != nil && hasCopilotAssignee(data.SafeOutputs.CreatePullRequests.Assignees) { - agentTokenStr := getEffectiveCopilotCodingAgentGitHubToken(data.SafeOutputs.CreatePullRequests.GitHubToken) - //nolint:gosec // G101: False positive - this is a GitHub Actions expression template, not a hardcoded credential - steps = append(steps, fmt.Sprintf(" GH_AW_ASSIGN_TO_AGENT_TOKEN: %s\n", agentTokenStr)) + case safeOutputs != nil && safeOutputs.CreatePullRequests != nil && hasCopilotAssignee(safeOutputs.CreatePullRequests.Assignees): + steps = append(steps, fmt.Sprintf(" GH_AW_ASSIGN_TO_AGENT_TOKEN: %s\n", getEffectiveCopilotCodingAgentGitHubToken(safeOutputs.CreatePullRequests.GitHubToken))) consolidatedSafeOutputsStepsLog.Print("Added GH_AW_ASSIGN_TO_AGENT_TOKEN env var for create-pull-request copilot assignment handler") } + return steps +} - // Add GH_AW_AGENT_SESSION_TOKEN when create-agent-session is configured. - // The create_agent_session handler passes this token as GH_TOKEN to the gh CLI - // (agent token preference chain), which is required because the default GITHUB_TOKEN - // does not have permission to create agent sessions via gh agent-task create. - if data.SafeOutputs != nil && data.SafeOutputs.CreateAgentSessions != nil { - agentSessionTokenStr := getEffectiveCopilotCodingAgentGitHubToken(data.SafeOutputs.CreateAgentSessions.GitHubToken) - //nolint:gosec // G101: False positive - this is a GitHub Actions expression template, not a hardcoded credential - steps = append(steps, fmt.Sprintf(" GH_AW_AGENT_SESSION_TOKEN: %s\n", agentSessionTokenStr)) - consolidatedSafeOutputsStepsLog.Print("Added GH_AW_AGENT_SESSION_TOKEN env var for create-agent-session handler") +func appendHandlerManagerAgentSessionEnv(steps []string, safeOutputs *SafeOutputsConfig) []string { + if safeOutputs == nil || safeOutputs.CreateAgentSessions == nil { + return steps } + steps = append(steps, fmt.Sprintf(" GH_AW_AGENT_SESSION_TOKEN: %s\n", getEffectiveCopilotCodingAgentGitHubToken(safeOutputs.CreateAgentSessions.GitHubToken))) + consolidatedSafeOutputsStepsLog.Print("Added GH_AW_AGENT_SESSION_TOKEN env var for create-agent-session handler") + return steps +} - // When create-pull-request or push-to-pull-request-branch is configured with a custom token - // (including GitHub App), expose that token as GITHUB_TOKEN so that git CLI operations in - // the JavaScript handlers can authenticate. The create_pull_request.cjs handler reads - // process.env.GITHUB_TOKEN to enable dynamic repo checkout for multi-repo/cross-repo - // scenarios (allowed-repos). Without this, the handler falls back to the default - // repo-scoped token which lacks access to other repos. - if usesPatchesAndCheckouts(data.SafeOutputs) { - gitToken, isCustom := resolvePRCheckoutToken(data.SafeOutputs, NewCheckoutManager(data.CheckoutConfigs)) - // Only override GITHUB_TOKEN when a custom token (app or PAT) is explicitly configured. - // When no custom token is set, the default repo-scoped GITHUB_TOKEN from GitHub Actions - // is already in the environment and overriding it with the same default is unnecessary. - if isCustom { - //nolint:gosec // G101: False positive - this is a GitHub Actions expression template, not a hardcoded credential - steps = append(steps, fmt.Sprintf(" GITHUB_TOKEN: %s\n", gitToken)) - consolidatedSafeOutputsStepsLog.Printf("Adding GITHUB_TOKEN env var for cross-repo git CLI operations") - } +func appendHandlerManagerGitTokenEnv(steps []string, data *WorkflowData) []string { + if !usesPatchesAndCheckouts(data.SafeOutputs) { + return steps } - - // With section for github-token - // Use the standard safe-outputs token for the shared github-script client. - // Project operations use GH_AW_PROJECT_GITHUB_TOKEN from env with dedicated handler logic. - steps = append(steps, " with:\n") - // Token precedence for the handler manager step: - // 1. Safe-outputs level token (so.GitHubToken) - // 2. Magic secret fallback via getEffectiveSafeOutputGitHubToken() - // - // Note: We do NOT fall back to per-output tokens (add-comment, create-issue, etc.) - // because those are specific to their operations. The handler manager needs a - // general-purpose token for the github-script client. - configToken := "" - if data.SafeOutputs != nil && data.SafeOutputs.GitHubToken != "" { - configToken = data.SafeOutputs.GitHubToken + gitToken, isCustom := resolvePRCheckoutToken(data.SafeOutputs, NewCheckoutManager(data.CheckoutConfigs)) + if !isCustom { + return steps } - c.addSafeOutputGitHubTokenForConfig(&steps, data, configToken) - - steps = append(steps, " script: |\n") - steps = append(steps, generateGitHubScriptWithRequire("process_safe_outputs.cjs")) - - return steps, nil + consolidatedSafeOutputsStepsLog.Printf("Adding GITHUB_TOKEN env var for cross-repo git CLI operations") + return append(steps, fmt.Sprintf(" GITHUB_TOKEN: %s\n", gitToken)) } diff --git a/pkg/workflow/compiler_string_api.go b/pkg/workflow/compiler_string_api.go index a78a8fd97d0..13eb43c55af 100644 --- a/pkg/workflow/compiler_string_api.go +++ b/pkg/workflow/compiler_string_api.go @@ -56,19 +56,35 @@ func (c *Compiler) CompileToYAML(workflowData *WorkflowData, markdownPath string // The virtualPath is used for error messages and lock file naming (e.g., "workflow.md"). func (c *Compiler) ParseWorkflowString(content string, virtualPath string) (*WorkflowData, error) { workflowLog.Printf("ParseWorkflowString: parsing %d bytes with virtual path %s", len(content), virtualPath) - cleanPath := filepath.Clean(virtualPath) contentBytes := []byte(content) - - // Store content so downstream code can use it instead of reading from disk. - // Cleared in CompileToYAML after compilation completes. c.contentOverride = content - - // Enable inline prompt mode for string-based compilation (Wasm/browser) - // since runtime-import macros cannot resolve without filesystem access c.inlinePrompt = true + parseResult, err := c.parseWorkflowStringFrontmatter(content, cleanPath, contentBytes) + if err != nil { + return nil, err + } + engineSetup, err := c.setupEngineAndImports(parseResult.frontmatterResult, parseResult.cleanPath, parseResult.content, parseResult.markdownDir) + if err != nil { + return nil, err + } + toolsResult, err := c.processToolsAndMarkdown(parseResult.frontmatterResult, parseResult.cleanPath, parseResult.markdownDir, engineSetup.agenticEngine, engineSetup.engineSetting, engineSetup.importsResult) + if err != nil { + return nil, err + } + workflowData := c.buildInitialWorkflowData(parseResult.frontmatterResult, toolsResult, engineSetup, engineSetup.importsResult) + workflowData.WorkflowID = GetWorkflowIDFromPath(cleanPath) + if err := c.validateParsedWorkflowString(workflowData, cleanPath); err != nil { + return nil, err + } + c.populateParsedWorkflowRuntime(workflowData) + if err := c.finalizeParsedWorkflowString(workflowData, parseResult, engineSetup.importsResult, cleanPath); err != nil { + return nil, err + } + return workflowData, nil +} - // Parse frontmatter directly from content string +func (c *Compiler) parseWorkflowStringFrontmatter(content, cleanPath string, contentBytes []byte) (*frontmatterParseResult, error) { result, err := parser.ExtractFrontmatterFromContent(content) if err != nil { frontmatterStart := 2 @@ -77,118 +93,85 @@ func (c *Compiler) ParseWorkflowString(content string, virtualPath string) (*Wor } return nil, c.createFrontmatterError(cleanPath, content, err, frontmatterStart) } - if len(result.Frontmatter) == 0 { return nil, errors.New("no frontmatter found") } - compilerStringAPILog.Printf("ParseWorkflowString: extracted frontmatter with %d fields", len(result.Frontmatter)) - - // Preprocess schedule fields if err := c.preprocessScheduleFields(result.Frontmatter, cleanPath, content); err != nil { return nil, err } - frontmatterForValidation := c.copyFrontmatterWithoutInternalMarkers(result.Frontmatter) - - // Check if "on" field is missing - distinguish redirect-only placeholders from shared workflows - _, hasOnField := frontmatterForValidation["on"] - if !hasOnField { - // Check if this is a redirect-only placeholder (has redirect field but no 'on' trigger). - // Redirect-only files are distinct from regular shared workflows: they are placeholders - // pointing to a workflow's new canonical location and should not be treated as importable components. - if redirectVal, hasRedirect := frontmatterForValidation["redirect"]; hasRedirect { - if redirectStr, ok := redirectVal.(string); ok { - if redirectTarget := strings.TrimSpace(redirectStr); redirectTarget != "" { - compilerStringAPILog.Printf("ParseWorkflowString: redirect-only workflow detected: redirect=%s", redirectTarget) - return nil, &RedirectOnlyWorkflowError{Path: cleanPath, Target: redirectTarget} - } - } - } - compilerStringAPILog.Printf("ParseWorkflowString: no 'on' field, treating as shared workflow: %s", cleanPath) - return nil, &SharedWorkflowError{Path: cleanPath} - } - - if err := c.validateEngineBeforeSchema(cleanPath, contentBytes, result, frontmatterForValidation); err != nil { - compilerStringAPILog.Printf("ParseWorkflowString: string engine pre-validation failed for %s", cleanPath) - return nil, err - } - - // Validate frontmatter against schema - if err := parser.ValidateMainWorkflowFrontmatterWithSchemaAndLocation(frontmatterForValidation, cleanPath); err != nil { - compilerStringAPILog.Printf("ParseWorkflowString: schema validation failed for %s", cleanPath) + if err := c.validateWorkflowStringFrontmatter(cleanPath, contentBytes, result, frontmatterForValidation); err != nil { return nil, err } - - compilerStringAPILog.Printf("ParseWorkflowString: frontmatter validated, frontmatter_fields=%d", len(frontmatterForValidation)) - - // Build parse result to reuse the rest of the orchestrator pipeline - parseResult := &frontmatterParseResult{ + return &frontmatterParseResult{ cleanPath: cleanPath, content: contentBytes, frontmatterResult: result, frontmatterForValidation: frontmatterForValidation, markdownDir: filepath.Dir(cleanPath), isSharedWorkflow: false, - } - - // Setup engine and process imports - engineSetup, err := c.setupEngineAndImports(parseResult.frontmatterResult, parseResult.cleanPath, parseResult.content, parseResult.markdownDir) - if err != nil { - return nil, err - } - - // Process tools and markdown - toolsResult, err := c.processToolsAndMarkdown(parseResult.frontmatterResult, parseResult.cleanPath, parseResult.markdownDir, engineSetup.agenticEngine, engineSetup.engineSetting, engineSetup.importsResult) - if err != nil { - return nil, err - } - - // Build initial workflow data structure - workflowData := c.buildInitialWorkflowData(parseResult.frontmatterResult, toolsResult, engineSetup, engineSetup.importsResult) - workflowData.WorkflowID = GetWorkflowIDFromPath(cleanPath) + }, nil +} - // Validate bash tool configuration - if err := validateBashToolConfig(workflowData.ParsedTools, workflowData.Name); err != nil { - return nil, fmt.Errorf("%s: %w", cleanPath, err) +func (c *Compiler) validateWorkflowStringFrontmatter(cleanPath string, contentBytes []byte, result *parser.FrontmatterResult, frontmatterForValidation map[string]any) error { + if err := validateWorkflowStringOnField(cleanPath, frontmatterForValidation); err != nil { + return err } - - // Validate optional engine.mcp.session-timeout configuration. - if err := c.validateEngineMCPSessionTimeout(workflowData); err != nil { - return nil, fmt.Errorf("%s: %w", cleanPath, err) + if err := c.validateEngineBeforeSchema(cleanPath, contentBytes, result, frontmatterForValidation); err != nil { + compilerStringAPILog.Printf("ParseWorkflowString: string engine pre-validation failed for %s", cleanPath) + return err } - - // Validate optional engine.mcp.tool-timeout configuration. - if err := c.validateEngineMCPToolTimeout(workflowData); err != nil { - return nil, fmt.Errorf("%s: %w", cleanPath, err) + if err := parser.ValidateMainWorkflowFrontmatterWithSchemaAndLocation(frontmatterForValidation, cleanPath); err != nil { + compilerStringAPILog.Printf("ParseWorkflowString: schema validation failed for %s", cleanPath) + return err } + compilerStringAPILog.Printf("ParseWorkflowString: frontmatter validated, frontmatter_fields=%d", len(frontmatterForValidation)) + return nil +} - // Validate GitHub tool configuration - if err := validateGitHubToolConfig(workflowData.ParsedTools, workflowData.Name); err != nil { - return nil, fmt.Errorf("%s: %w", cleanPath, err) +func validateWorkflowStringOnField(cleanPath string, frontmatterForValidation map[string]any) error { + if _, hasOnField := frontmatterForValidation["on"]; hasOnField { + return nil } - - // Validate GitHub tool read-only configuration - if err := validateGitHubReadOnly(workflowData.ParsedTools, workflowData.Name); err != nil { - return nil, fmt.Errorf("%s: %w", cleanPath, err) + if redirectVal, hasRedirect := frontmatterForValidation["redirect"]; hasRedirect { + if redirectStr, ok := redirectVal.(string); ok { + if redirectTarget := strings.TrimSpace(redirectStr); redirectTarget != "" { + compilerStringAPILog.Printf("ParseWorkflowString: redirect-only workflow detected: redirect=%s", redirectTarget) + return &RedirectOnlyWorkflowError{Path: cleanPath, Target: redirectTarget} + } + } } + compilerStringAPILog.Printf("ParseWorkflowString: no 'on' field, treating as shared workflow: %s", cleanPath) + return &SharedWorkflowError{Path: cleanPath} +} - // Validate GitHub guard policy configuration - if err := validateGitHubGuardPolicy(workflowData.ParsedTools, workflowData.Name); err != nil { - return nil, fmt.Errorf("%s: %w", cleanPath, err) +func (c *Compiler) validateParsedWorkflowString(workflowData *WorkflowData, cleanPath string) error { + validators := []func() error{ + func() error { return validateBashToolConfig(workflowData.ParsedTools, workflowData.Name) }, + func() error { return c.validateEngineMCPSessionTimeout(workflowData) }, + func() error { return c.validateEngineMCPToolTimeout(workflowData) }, + func() error { return validateGitHubToolConfig(workflowData.ParsedTools, workflowData.Name) }, + func() error { return validateGitHubReadOnly(workflowData.ParsedTools, workflowData.Name) }, + func() error { return validateGitHubGuardPolicy(workflowData.ParsedTools, workflowData.Name) }, + } + for _, validate := range validators { + if err := validate(); err != nil { + return fmt.Errorf("%s: %w", cleanPath, err) + } } emitGitHubLockdownGuardPolicyWarning(c, workflowData.ParsedTools, cleanPath) - - // Validate integrity-reactions feature configuration var gatewayConfig *MCPGatewayRuntimeConfig if workflowData.SandboxConfig != nil { gatewayConfig = workflowData.SandboxConfig.MCP } if err := validateIntegrityReactions(workflowData.ParsedTools, workflowData.Name, workflowData, gatewayConfig); err != nil { - return nil, fmt.Errorf("%s: %w", cleanPath, err) + return fmt.Errorf("%s: %w", cleanPath, err) } + return nil +} - // Setup action cache and resolver +func (c *Compiler) populateParsedWorkflowRuntime(workflowData *WorkflowData) { actionCache, actionResolver := c.getSharedActionResolver() workflowData.Ctx = c.ctx workflowData.ActionCache = actionCache @@ -196,31 +179,22 @@ func (c *Compiler) ParseWorkflowString(content string, virtualPath string) (*Wor workflowData.ActionPinWarnings = c.actionPinWarnings workflowData.ActionPinMappings = c.getActionPinMappings() workflowData.ContainerPinMappings = c.getContainerPinMappings() +} - // Extract YAML configuration sections +func (c *Compiler) finalizeParsedWorkflowString(workflowData *WorkflowData, parseResult *frontmatterParseResult, importsResult *parser.ImportsResult, cleanPath string) error { if err := c.extractYAMLSections(parseResult.frontmatterResult.Frontmatter, workflowData); err != nil { - return nil, fmt.Errorf("failed to extract YAML sections: %w", err) + return fmt.Errorf("failed to extract YAML sections: %w", err) } - - // Merge features from imports - if len(engineSetup.importsResult.MergedFeatures) > 0 { - compilerStringAPILog.Printf("ParseWorkflowString: merging %d features from imports", len(engineSetup.importsResult.MergedFeatures)) - mergedFeatures, err := c.MergeFeatures(workflowData.Features, engineSetup.importsResult.MergedFeatures) + if len(importsResult.MergedFeatures) > 0 { + compilerStringAPILog.Printf("ParseWorkflowString: merging %d features from imports", len(importsResult.MergedFeatures)) + mergedFeatures, err := c.MergeFeatures(workflowData.Features, importsResult.MergedFeatures) if err != nil { - return nil, fmt.Errorf("failed to merge features from imports: %w", err) + return fmt.Errorf("failed to merge features from imports: %w", err) } workflowData.Features = mergedFeatures } - - // Process and merge custom steps - if err := c.processAndMergeSteps(parseResult.frontmatterResult.Frontmatter, workflowData, engineSetup.importsResult); err != nil { - return nil, err - } - - // Apply defaults - if err := c.applyDefaults(workflowData, cleanPath); err != nil { - return nil, err + if err := c.processAndMergeSteps(parseResult.frontmatterResult.Frontmatter, workflowData, importsResult); err != nil { + return err } - - return workflowData, nil + return c.applyDefaults(workflowData, cleanPath) } diff --git a/pkg/workflow/compiler_unlock_job.go b/pkg/workflow/compiler_unlock_job.go index 6c91dbffe72..a4f093b6547 100644 --- a/pkg/workflow/compiler_unlock_job.go +++ b/pkg/workflow/compiler_unlock_job.go @@ -19,105 +19,89 @@ var compilerUnlockJobLog = logger.New("workflow:compiler_unlock_job") // The job depends on agent and detection (if enabled) to ensure unlock happens after workflow execution. func (c *Compiler) buildUnlockJob(data *WorkflowData, threatDetectionEnabled bool) (*Job, error) { compilerUnlockJobLog.Print("Building dedicated unlock job") - if !data.LockForAgent { compilerUnlockJobLog.Print("Lock-for-agent not enabled, skipping unlock job") return nil, nil } + steps, err := c.buildUnlockJobSteps(data) + if err != nil { + return nil, err + } + if c.actionMode.IsScript() { + steps = append(steps, c.generateScriptModeCleanupStep()) + } + needs := buildUnlockJobNeeds(threatDetectionEnabled) + permissions := c.buildUnlockJobPermissions(data) + jobCondition := buildUnlockJobCondition() + compilerUnlockJobLog.Printf("Job built successfully: dependencies=%v", needs) + + job := &Job{ + Name: "unlock", + Needs: needs, + If: RenderCondition(jobCondition), + RunsOn: c.formatFrameworkJobRunsOn(data), + Permissions: permissions, + Steps: steps, + TimeoutMinutes: 5, // Short timeout - unlock is a quick operation + } - var steps []string + return job, nil +} - // Add setup step to copy scripts +func (c *Compiler) buildUnlockJobSteps(data *WorkflowData) ([]string, error) { setupActionRef := c.resolveActionReference("./actions/setup", data) if setupActionRef == "" && !c.actionMode.IsScript() { return nil, errors.New("setup action reference is required but could not be resolved") } - - // For dev mode (local action path), checkout the actions folder first - steps = append(steps, c.generateCheckoutActionsFolder(data)...) - - // Unlock job doesn't need project support - // Unlock job depends on activation, reuse its trace ID + steps := append([]string{}, c.generateCheckoutActionsFolder(data)...) unlockTraceID := fmt.Sprintf("${{ needs.%s.outputs.setup-trace-id }}", constants.ActivationJobName) unlockParentSpanID := setupParentSpanNeedsExpr(constants.ActivationJobName) steps = append(steps, c.generateSetupStep(data, setupActionRef, SetupActionDestination, false, unlockTraceID, unlockParentSpanID)...) + steps = append(steps, buildUnlockIssueStep(data)...) + return steps, nil +} - // Add unlock step - // Build condition: only unlock if issue was locked by activation job - // Must match lock condition: event type is 'issues' or 'issue_comment' - eventTypeCheck := BuildOr( - BuildEventTypeEquals("issues"), - BuildEventTypeEquals("issue_comment"), +func buildUnlockIssueStep(data *WorkflowData) []string { + unlockCondition := BuildAnd( + BuildOr(BuildEventTypeEquals("issues"), BuildEventTypeEquals("issue_comment")), + BuildEquals(BuildPropertyAccess(fmt.Sprintf("needs.%s.outputs.issue_locked", constants.ActivationJobName)), BuildStringLiteral("true")), ) - lockedOutputCheck := BuildEquals( - BuildPropertyAccess(fmt.Sprintf("needs.%s.outputs.issue_locked", constants.ActivationJobName)), - BuildStringLiteral("true"), - ) - - unlockCondition := BuildAnd(eventTypeCheck, lockedOutputCheck) - - steps = append(steps, " - name: Unlock issue after agentic workflow\n") - steps = append(steps, " id: unlock-issue\n") - steps = append(steps, fmt.Sprintf(" if: %s\n", RenderCondition(unlockCondition))) - steps = append(steps, fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/github-script", data))) - steps = append(steps, " with:\n") - steps = append(steps, " script: |\n") - steps = append(steps, generateGitHubScriptWithRequire("unlock-issue.cjs")) - compilerUnlockJobLog.Print("Added unlock issue step to dedicated unlock job") + return []string{ + " - name: Unlock issue after agentic workflow\n", + " id: unlock-issue\n", + fmt.Sprintf(" if: %s\n", RenderCondition(unlockCondition)), + fmt.Sprintf(" uses: %s\n", getCachedActionPin("actions/github-script", data)), + " with:\n", + " script: |\n", + generateGitHubScriptWithRequire("unlock-issue.cjs"), + } +} - // Build the condition for this job: - // 1. always() - run even if agent or other jobs fail - // 2. activation was not skipped - skip unlock when activation was never triggered - // (e.g. the event did not match any trigger condition, so no locking occurred) - // 3. issue was locked (checked at step level for clarity in workflow YAML) - alwaysFunc := BuildFunctionCall("always") - activationNotSkipped := BuildNotEquals( - BuildPropertyAccess(fmt.Sprintf("needs.%s.result", constants.ActivationJobName)), - BuildStringLiteral("skipped"), +func buildUnlockJobCondition() ConditionNode { + return BuildAnd( + BuildFunctionCall("always"), + BuildNotEquals(BuildPropertyAccess(fmt.Sprintf("needs.%s.result", constants.ActivationJobName)), BuildStringLiteral("skipped")), ) - jobCondition := BuildAnd(alwaysFunc, activationNotSkipped) +} - // Create the unlock job - // This job depends on activation (for issue_locked output) and agent (to run after workflow) - // When threat detection is enabled, it also depends on the detection job +func buildUnlockJobNeeds(threatDetectionEnabled bool) []string { needs := []string{string(constants.ActivationJobName), string(constants.AgentJobName)} if threatDetectionEnabled { needs = append(needs, string(constants.DetectionJobName)) compilerUnlockJobLog.Print("Added detection job dependency to unlock job") } + return needs +} - // Determine permissions - need contents: read for dev mode checkout, issues: write for unlocking - var permissions string +func (c *Compiler) buildUnlockJobPermissions(data *WorkflowData) string { needsContentsRead := (c.actionMode.IsDev() || c.actionMode.IsScript()) && len(c.generateCheckoutActionsFolder(data)) > 0 if needsContentsRead { perms := NewPermissionsContentsRead() - // Add issues write permission for unlocking - perms.Set(PermissionIssues, PermissionWrite) - permissions = perms.RenderToYAML() - } else { - // Only need issues write permission - perms := NewPermissions() perms.Set(PermissionIssues, PermissionWrite) - permissions = perms.RenderToYAML() - } - - compilerUnlockJobLog.Printf("Job built successfully: dependencies=%v", needs) - - // In script mode, explicitly add a cleanup step (mirrors post.js in dev/release/action mode). - if c.actionMode.IsScript() { - steps = append(steps, c.generateScriptModeCleanupStep()) + return perms.RenderToYAML() } - - job := &Job{ - Name: "unlock", - Needs: needs, - If: RenderCondition(jobCondition), - RunsOn: c.formatFrameworkJobRunsOn(data), - Permissions: permissions, - Steps: steps, - TimeoutMinutes: 5, // Short timeout - unlock is a quick operation - } - - return job, nil + perms := NewPermissions() + perms.Set(PermissionIssues, PermissionWrite) + return perms.RenderToYAML() } diff --git a/pkg/workflow/compiler_workflow_call.go b/pkg/workflow/compiler_workflow_call.go index 25bb0ac32ed..1fc306372de 100644 --- a/pkg/workflow/compiler_workflow_call.go +++ b/pkg/workflow/compiler_workflow_call.go @@ -97,10 +97,7 @@ func (c *Compiler) injectWorkflowCallOutputs(onSection string, safeOutputs *Safe if safeOutputs == nil || !strings.Contains(onSection, "workflow_call") { return onSection } - workflowCallLog.Print("Injecting workflow_call outputs for safe-output results") - - // Build the auto-generated outputs map based on configured safe output types generatedOutputs := buildWorkflowCallOutputsMap(safeOutputs) if len(generatedOutputs) == 0 { workflowCallLog.Print("No workflow_call outputs to inject (no safe-output types configured)") @@ -108,69 +105,12 @@ func (c *Compiler) injectWorkflowCallOutputs(onSection string, safeOutputs *Safe } workflowCallLog.Printf("Generated %d workflow_call outputs to inject", len(generatedOutputs)) - - // Parse the on section YAML - var onData map[string]any - if err := yaml.Unmarshal([]byte(onSection), &onData); err != nil { - workflowCallLog.Printf("Warning: failed to parse on section for workflow_call outputs injection: %v", err) - return onSection - } - - // Get the 'on' map - onMap, ok := onData["on"].(map[string]any) + _, onMap, workflowCallMap, ok := parseWorkflowCallSection(onSection, false) if !ok { return onSection } - - // Get the workflow_call entry - workflowCallVal, hasWorkflowCall := onMap["workflow_call"] - if !hasWorkflowCall { - return onSection - } - - // Convert workflow_call to a map (it may be nil if declared without options) - var workflowCallMap map[string]any - if workflowCallVal == nil { - workflowCallMap = make(map[string]any) - } else if m, ok := workflowCallVal.(map[string]any); ok { - workflowCallMap = m - } else { - workflowCallMap = make(map[string]any) - } - - // Merge auto-generated outputs with any existing user-defined outputs. - // User-defined outputs take precedence (their keys overwrite generated ones). - mergedOutputs := make(map[string]workflowCallOutputEntry) - maps.Copy(mergedOutputs, generatedOutputs) - if existingOutputs, hasOutputs := workflowCallMap["outputs"].(map[string]any); hasOutputs { - for k, v := range existingOutputs { - // User-defined entries may be maps with description+value or plain strings - if outputMap, ok := v.(map[string]any); ok { - entry := workflowCallOutputEntry{} - if desc, ok := outputMap["description"].(string); ok { - entry.Description = desc - } - if val, ok := outputMap["value"].(string); ok { - entry.Value = val - } - mergedOutputs[k] = entry - } - } - } - - workflowCallLog.Printf("Merged workflow_call outputs: total=%d", len(mergedOutputs)) - workflowCallMap["outputs"] = mergedOutputs - onMap["workflow_call"] = workflowCallMap - - // Re-marshal to YAML - newOnData := map[string]any{"on": onMap} - newYAML, err := yaml.Marshal(newOnData) - if err != nil { - workflowCallLog.Printf("Warning: failed to marshal on section with workflow_call outputs: %v", err) - return onSection - } - - return strings.TrimSuffix(string(newYAML), "\n") + workflowCallMap["outputs"] = mergeWorkflowCallOutputs(workflowCallMap, generatedOutputs) + return marshalWorkflowCallSection(onSection, onMap, workflowCallMap, "outputs") } // buildWorkflowCallOutputsMap constructs the outputs map for on.workflow_call.outputs @@ -263,91 +203,131 @@ func injectWorkflowCallSecretsSection(onSection string, secrets []string) string sort.Strings(secretsToInject) workflowCallLog.Printf("Injecting %d workflow_call secrets declarations", len(secretsToInject)) + _, onMap, workflowCallMap, ok := parseWorkflowCallSection(onSection, true) + if !ok { + return onSection + } + workflowCallMap["secrets"] = buildWorkflowCallSecretsMap(workflowCallMap, secretsToInject) + return marshalWorkflowCallSection(onSection, onMap, workflowCallMap, "secrets") +} - // Parse the on section YAML. +func parseWorkflowCallSection(onSection string, normalizeOn bool) (map[string]any, map[string]any, map[string]any, bool) { var onData map[string]any if err := yaml.Unmarshal([]byte(onSection), &onData); err != nil { - workflowCallLog.Printf("Warning: failed to parse on section for workflow_call secrets injection: %v", err) - return onSection + workflowCallLog.Printf("Warning: failed to parse on section for workflow_call injection: %v", err) + return nil, nil, nil, false + } + onMap, ok := extractWorkflowCallOnMap(onData, normalizeOn) + if !ok { + return nil, nil, nil, false } + workflowCallMap, ok := normalizeWorkflowCallMap(onMap["workflow_call"]) + return onData, onMap, workflowCallMap, ok +} - // Normalize onData["on"] to map[string]any, handling string and slice shorthand forms. +func extractWorkflowCallOnMap(onData map[string]any, normalize bool) (map[string]any, bool) { rawOn, hasOn := onData["on"] if !hasOn { - return onSection + return nil, false } - var onMap map[string]any + if !normalize { + onMap, ok := rawOn.(map[string]any) + return onMap, ok + } + onMap, ok := normalizeOnSectionValue(rawOn) + if ok { + onData["on"] = onMap + } + return onMap, ok +} + +func normalizeOnSectionValue(rawOn any) (map[string]any, bool) { switch v := rawOn.(type) { case map[string]any: - onMap = v + return v, true case string: - onMap = map[string]any{v: nil} + return map[string]any{v: nil}, true case []any: - onMap = make(map[string]any, len(v)) + onMap := make(map[string]any, len(v)) for _, event := range v { if eventName, ok := event.(string); ok { onMap[eventName] = nil } } + return onMap, true case []string: - onMap = make(map[string]any, len(v)) + onMap := make(map[string]any, len(v)) for _, eventName := range v { onMap[eventName] = nil } + return onMap, true default: - return onSection + return nil, false } - onData["on"] = onMap +} - workflowCallVal, hasWorkflowCall := onMap["workflow_call"] - if !hasWorkflowCall { - return onSection +func normalizeWorkflowCallMap(workflowCallVal any) (map[string]any, bool) { + if workflowCallVal == nil { + return make(map[string]any), true } + if workflowCallMap, ok := workflowCallVal.(map[string]any); ok { + return workflowCallMap, true + } + return make(map[string]any), true +} - // Convert workflow_call to a map (it may be nil if declared without options). - var workflowCallMap map[string]any - if workflowCallVal == nil { - workflowCallMap = make(map[string]any) - } else if m, ok := workflowCallVal.(map[string]any); ok { - workflowCallMap = m - } else { - workflowCallMap = make(map[string]any) +func mergeWorkflowCallOutputs(workflowCallMap map[string]any, generatedOutputs map[string]workflowCallOutputEntry) map[string]workflowCallOutputEntry { + mergedOutputs := make(map[string]workflowCallOutputEntry) + maps.Copy(mergedOutputs, generatedOutputs) + if existingOutputs, hasOutputs := workflowCallMap["outputs"].(map[string]any); hasOutputs { + for k, v := range existingOutputs { + if outputMap, ok := v.(map[string]any); ok { + mergedOutputs[k] = workflowCallOutputEntry{ + Description: stringValueFromMap(outputMap, "description"), + Value: stringValueFromMap(outputMap, "value"), + } + } + } } + workflowCallLog.Printf("Merged workflow_call outputs: total=%d", len(mergedOutputs)) + return mergedOutputs +} - // Build the auto-generated secrets map (required: false for all entries). +func buildWorkflowCallSecretsMap(workflowCallMap map[string]any, secretsToInject []string) map[string]any { generatedSecrets := make(map[string]workflowCallSecretEntry, len(secretsToInject)) for _, name := range secretsToInject { generatedSecrets[name] = workflowCallSecretEntry{Required: false} } - - // Merge: user-defined entries take precedence over generated ones. if existingSecrets, hasSecrets := workflowCallMap["secrets"].(map[string]any); hasSecrets { for k, v := range existingSecrets { if entryMap, ok := v.(map[string]any); ok { - entry := workflowCallSecretEntry{} - if req, ok := entryMap["required"].(bool); ok { - entry.Required = req - } - generatedSecrets[k] = entry + generatedSecrets[k] = workflowCallSecretEntry{Required: boolValueFromMap(entryMap, "required")} } } } - - // Convert to a plain map for marshaling. secretsOut := make(map[string]any, len(generatedSecrets)) for k, v := range generatedSecrets { secretsOut[k] = map[string]any{"required": v.Required} } - workflowCallMap["secrets"] = secretsOut - onMap["workflow_call"] = workflowCallMap + return secretsOut +} - // Re-marshal to YAML. - newOnData := map[string]any{"on": onMap} - newYAML, err := yaml.Marshal(newOnData) +func marshalWorkflowCallSection(onSection string, onMap map[string]any, workflowCallMap map[string]any, kind string) string { + onMap["workflow_call"] = workflowCallMap + newYAML, err := yaml.Marshal(map[string]any{"on": onMap}) if err != nil { - workflowCallLog.Printf("Warning: failed to marshal on section with workflow_call secrets: %v", err) + workflowCallLog.Printf("Warning: failed to marshal on section with workflow_call %s: %v", kind, err) return onSection } - return strings.TrimSuffix(string(newYAML), "\n") } + +func stringValueFromMap(values map[string]any, key string) string { + value, _ := values[key].(string) + return value +} + +func boolValueFromMap(values map[string]any, key string) bool { + value, _ := values[key].(bool) + return value +} diff --git a/pkg/workflow/compiler_yaml.go b/pkg/workflow/compiler_yaml.go index 73ec659c3c3..ff6c448e825 100644 --- a/pkg/workflow/compiler_yaml.go +++ b/pkg/workflow/compiler_yaml.go @@ -91,131 +91,91 @@ func (c *Compiler) generateWorkflowBody(yaml *strings.Builder, data *WorkflowDat func (c *Compiler) generateYAML(data *WorkflowData, markdownPath string) (string, []string, []string, error) { compilerYamlLog.Printf("Generating YAML for workflow: %s", data.Name) - - // Compute frontmatter hash BEFORE building jobs so that the stable hash is - // available to heredoc-delimiter generation throughout job construction. - // Using the hex-encoded SHA-256 frontmatter hash string as an HMAC key keeps - // the compiled lock file identical across repeated compilations of the same workflow. - var frontmatterHash string - var bodyHash string - if markdownPath != "" { - baseDir := filepath.Dir(markdownPath) - cache := parser.NewImportCache(baseDir) - - // computeWorkflowHash calls the parsed-content path when RawMarkdown is - // available (fast path), falling back to a disk read otherwise. - computeWorkflowHash := func( - fromParsed func() (string, error), - fromFile func() (string, error), - ) (string, error) { - if data.RawMarkdown != "" { - return fromParsed() - } - compilerYamlLog.Printf("RawMarkdown not set; falling back to reading file from disk: %s", markdownPath) - return fromFile() - } - - hash, err := computeWorkflowHash( - func() (string, error) { - return parser.ComputeFrontmatterHashFromParsedContent(data.FrontmatterYAML, data.RawMarkdown, data.RawFrontmatter, baseDir, cache, parser.DefaultFileReader) - }, - func() (string, error) { - return parser.ComputeFrontmatterHashFromFileWithParsedFrontmatter(markdownPath, data.RawFrontmatter, cache, parser.DefaultFileReader) - }, - ) - if err != nil { - return "", nil, nil, fmt.Errorf("failed to generate workflow YAML: could not compute stable frontmatter hash for %q: %w", markdownPath, err) - } - frontmatterHash = hash - compilerYamlLog.Printf("Computed frontmatter hash: %s", hash) - - // Compute body hash to cover changes to the markdown body that are not captured - // by the frontmatter hash. This enables stale-check: full detection. - bHash, bErr := computeWorkflowHash( - func() (string, error) { - return parser.ComputeBodyHashFromParsedContent(data.RawMarkdown, data.FrontmatterYAML, baseDir, parser.DefaultFileReader) - }, - func() (string, error) { - return parser.ComputeBodyHashFromFile(markdownPath) - }, - ) - if bErr != nil { - compilerYamlLog.Printf("Warning: could not compute body hash for %q: %v", markdownPath, bErr) - // Non-fatal: continue without body hash - } else { - bodyHash = bHash - compilerYamlLog.Printf("Computed body hash: %s", bodyHash) - } + frontmatterHash, bodyHash, err := c.computeWorkflowHashes(data, markdownPath) + if err != nil { + return "", nil, nil, err } - // Store hash on WorkflowData so job-building helpers (MCP renderers, prompt - // step generators, etc.) can derive stable heredoc delimiters from it. data.FrontmatterHash = frontmatterHash - - // Build all jobs and validate dependencies if err := c.buildJobsAndValidate(data, markdownPath); err != nil { return "", nil, nil, fmt.Errorf("failed to build and validate jobs: %w", err) } + yamlContent, secrets, actions := c.buildFinalWorkflowYAML(data, frontmatterHash, bodyHash) + yamlContent = c.finalizeGeneratedYAML(yamlContent, data) + + compilerYamlLog.Printf("Successfully generated YAML for workflow: %s (%d bytes)", data.Name, len(yamlContent)) + return yamlContent, secrets, actions, nil +} - // Pre-allocate builder capacity based on estimated workflow size. - // Copilot/Claude workflows with safe-outputs typically compile to ~70–90 KB. - // 96 KB avoids the first reallocation for the common case. The performance - // benefit of this function comes from eliminating the intermediate copies - // that RenderToYAML + WriteString used to incur, not from capacity reduction. +func (c *Compiler) computeWorkflowHashes(data *WorkflowData, markdownPath string) (string, string, error) { + if markdownPath == "" { + return "", "", nil + } + baseDir := filepath.Dir(markdownPath) + cache := parser.NewImportCache(baseDir) + hash, err := c.computeWorkflowHashValue(data, markdownPath, func() (string, error) { + return parser.ComputeFrontmatterHashFromParsedContent(data.FrontmatterYAML, data.RawMarkdown, data.RawFrontmatter, baseDir, cache, parser.DefaultFileReader) + }, func() (string, error) { + return parser.ComputeFrontmatterHashFromFileWithParsedFrontmatter(markdownPath, data.RawFrontmatter, cache, parser.DefaultFileReader) + }) + if err != nil { + return "", "", fmt.Errorf("failed to generate workflow YAML: could not compute stable frontmatter hash for %q: %w", markdownPath, err) + } + bodyHash, bodyErr := c.computeWorkflowHashValue(data, markdownPath, func() (string, error) { + return parser.ComputeBodyHashFromParsedContent(data.RawMarkdown, data.FrontmatterYAML, baseDir, parser.DefaultFileReader) + }, func() (string, error) { + return parser.ComputeBodyHashFromFile(markdownPath) + }) + if bodyErr != nil { + compilerYamlLog.Printf("Warning: could not compute body hash for %q: %v", markdownPath, bodyErr) + bodyHash = "" + } + return hash, bodyHash, nil +} + +func (c *Compiler) computeWorkflowHashValue(data *WorkflowData, markdownPath string, fromParsed func() (string, error), fromFile func() (string, error)) (string, error) { + if data.RawMarkdown != "" { + return fromParsed() + } + compilerYamlLog.Printf("RawMarkdown not set; falling back to reading file from disk: %s", markdownPath) + return fromFile() +} + +func (c *Compiler) buildFinalWorkflowYAML(data *WorkflowData, frontmatterHash, bodyHash string) (string, []string, []string) { const initialBuilderCapacity = 96 * 1024 + bodyContent, secrets, actions := c.renderWorkflowBodyWithMetadata(data, initialBuilderCapacity) var yaml strings.Builder yaml.Grow(initialBuilderCapacity) + c.generateWorkflowHeader(&yaml, data, frontmatterHash, bodyHash, secrets, actions) + yaml.WriteString(bodyContent) + return yaml.String(), secrets, actions +} - // Generate workflow body first so we can collect secrets and custom actions - // for inclusion in the header comment. - var body strings.Builder - body.Grow(initialBuilderCapacity) - c.generateWorkflowBody(&body, data) - bodyContent := body.String() - - // Collect secrets and external action references from the generated body. - // These are returned to the caller so they can be used for safe update enforcement - // without requiring a second scan of the full YAML content. +func (c *Compiler) renderWorkflowBodyWithMetadata(data *WorkflowData, capacity int) (string, []string, []string) { + bodyContent := c.renderWorkflowBody(data, capacity) secrets := CollectSecretReferences(bodyContent) actions := CollectActionReferences(bodyContent) - - // If this workflow has a workflow_call trigger, inject on.workflow_call.secrets: - // declarations so callers can map secrets explicitly instead of using secrets: inherit. - // We update data.On and regenerate the body so the compiled output includes the - // declarations. The set of secrets does not change between the two passes (the - // injected declarations do not add new ${{ secrets.* }} references). if hasWorkflowCallTrigger(data.On) && len(secrets) > 0 { updatedOn := injectWorkflowCallSecretsSection(data.On, secrets) if updatedOn != data.On { data.On = updatedOn - body.Reset() - body.Grow(initialBuilderCapacity) - c.generateWorkflowBody(&body, data) - bodyContent = body.String() + bodyContent = c.renderWorkflowBody(data, capacity) compilerYamlLog.Printf("Regenerated workflow body with on.workflow_call.secrets declarations") } } + return bodyContent, secrets, actions +} - // Generate workflow header comments (including metadata as first line, plus secrets/actions lists) - c.generateWorkflowHeader(&yaml, data, frontmatterHash, bodyHash, secrets, actions) - - // Append the workflow body - yaml.WriteString(bodyContent) - - yamlContent := yaml.String() +func (c *Compiler) renderWorkflowBody(data *WorkflowData, capacity int) string { + var body strings.Builder + body.Grow(capacity) + c.generateWorkflowBody(&body, data) + return body.String() +} - // If we're in non-cloning trial mode and this workflow has issue triggers, - // replace github.event.issue.number with inputs.issue_number +func (c *Compiler) finalizeGeneratedYAML(yamlContent string, data *WorkflowData) string { if c.trialMode && c.hasIssueTrigger(data.On) { compilerYamlLog.Print("Trial mode enabled, replacing issue number references") yamlContent = c.replaceIssueNumberReferences(yamlContent) } - - // Normalize assembled YAML whitespace. This clears indentation-only blank lines - // everywhere, trims trailing whitespace on structural YAML lines, preserves - // block-scalar payload content, caps over-long structural blank runs, and - // ensures the file ends with exactly one trailing newline. - yamlContent = normalizeBlankLines(yamlContent) - - compilerYamlLog.Printf("Successfully generated YAML for workflow: %s (%d bytes)", data.Name, len(yamlContent)) - return yamlContent, secrets, actions, nil + return normalizeBlankLines(yamlContent) } diff --git a/pkg/workflow/compiler_yaml_normalize.go b/pkg/workflow/compiler_yaml_normalize.go index 5f80a60daba..b97a29fcfc5 100644 --- a/pkg/workflow/compiler_yaml_normalize.go +++ b/pkg/workflow/compiler_yaml_normalize.go @@ -29,91 +29,11 @@ const maxConsecutiveBlankLines = 2 // the input byte-by-byte and builds the result with a single pre-allocated strings.Builder. func normalizeBlankLines(yamlContent string) string { compilerYAMLNormalizeLog.Printf("Normalizing blank lines in %d bytes of YAML", len(yamlContent)) - var b strings.Builder - b.Grow(len(yamlContent)) - - // lastNonBlankEnd tracks the builder length immediately after writing the last - // non-blank line (including its trailing newline). It starts at 0 and is only - // advanced when a substantive line is written, so it stays 0 when all lines - // are whitespace-only or the input is empty. Every line — blank or not — still - // gets a '\n' written to b, so b.Len() and lastNonBlankEnd may diverge when - // there are trailing blank lines. - lastNonBlankEnd := 0 - // blankRun counts consecutive blank lines emitted since the last non-blank - // structural line, so runs longer than maxConsecutiveBlankLines can be - // collapsed outside block scalars. - blankRun := 0 - inBlockScalar := false - pendingBlockScalar := false - blockScalarHeaderIndent := 0 - blockScalarIndent := 0 + state := newYAMLNormalizationState(len(yamlContent)) pos := 0 for pos < len(yamlContent) { - // Find the end of the current line. end := strings.IndexByte(yamlContent[pos:], '\n') - var line string - if end == -1 { - line = yamlContent[pos:] - } else { - line = yamlContent[pos : pos+end] - } - - processStructuralLine := true - trimmed := strings.TrimRight(line, " \t") - if pendingBlockScalar || inBlockScalar { - if trimmed == "" { - // Whitespace-only lines inside block scalars are still semantically - // blank, so emit them as empty lines but never cap the run. - b.WriteByte('\n') - processStructuralLine = false - } else { - lineIndent := countLeadingSpaces(line) - if pendingBlockScalar { - if lineIndent <= blockScalarHeaderIndent { - pendingBlockScalar = false - } else { - blockScalarIndent = lineIndent - inBlockScalar = true - pendingBlockScalar = false - } - } - if inBlockScalar { - if lineIndent < blockScalarIndent { - inBlockScalar = false - } else { - b.WriteString(line) - b.WriteByte('\n') - lastNonBlankEnd = b.Len() - processStructuralLine = false - } - } - } - } - - if processStructuralLine { - if trimmed == "" { - // Blank structural line: emit at most maxConsecutiveBlankLines in a - // row so yamllint's empty-lines rule is never exceeded. lastNonBlankEnd - // is NOT updated here so that trailing blank lines (including a blank - // final "line" produced by a file that ends with "\n\n" or by - // whitespace-only text after the last real line) are excluded from the - // returned slice. - if blankRun < maxConsecutiveBlankLines { - b.WriteByte('\n') - blankRun++ - } - } else { - b.WriteString(trimmed) - b.WriteByte('\n') - lastNonBlankEnd = b.Len() - blankRun = 0 - if headerIndent, ok := blockScalarHeaderIndentForLine(trimmed); ok { - pendingBlockScalar = true - blockScalarHeaderIndent = headerIndent - } - } - } - + state.processLine(extractYAMLLine(yamlContent, pos, end)) if end == -1 { break } @@ -124,15 +44,91 @@ func normalizeBlankLines(yamlContent string) string { // (empty input or all-whitespace). Return a single newline, which matches // the original strings.TrimRight(…, "\n") + "\n" behaviour for that case. // NOTE: b.String()[:0] must NOT be used here; the early return is intentional. - if lastNonBlankEnd == 0 { + if state.lastNonBlankEnd == 0 { compilerYAMLNormalizeLog.Print("Input contained no non-blank lines, returning single newline") return "\n" } - // Slice the builder string to drop trailing blank lines. b.String() copies - // the builder's internal buffer into a new string once; the slice avoids a - // second copy that a separate strings.Builder trim would incur. - compilerYAMLNormalizeLog.Printf("Normalized YAML to %d bytes", lastNonBlankEnd) - return b.String()[:lastNonBlankEnd] + compilerYAMLNormalizeLog.Printf("Normalized YAML to %d bytes", state.lastNonBlankEnd) + return state.result() +} + +type yamlNormalizationState struct { + builder strings.Builder + lastNonBlankEnd int + blankRun int + inBlockScalar bool + pendingBlockScalar bool + blockScalarHeaderIndent int + blockScalarIndent int +} + +func newYAMLNormalizationState(capacity int) *yamlNormalizationState { + var builder strings.Builder + builder.Grow(capacity) + return &yamlNormalizationState{builder: builder} +} + +func extractYAMLLine(content string, pos, end int) string { + if end == -1 { + return content[pos:] + } + return content[pos : pos+end] +} + +func (s *yamlNormalizationState) processLine(line string) { + trimmed := strings.TrimRight(line, " \t") + if s.processBlockScalarLine(line, trimmed) { + return + } + s.processStructuralLine(trimmed) +} + +func (s *yamlNormalizationState) processBlockScalarLine(line, trimmed string) bool { + if !s.pendingBlockScalar && !s.inBlockScalar { + return false + } + if trimmed == "" { + s.builder.WriteByte('\n') + return true + } + lineIndent := countLeadingSpaces(line) + if s.pendingBlockScalar { + s.pendingBlockScalar = false + if lineIndent > s.blockScalarHeaderIndent { + s.blockScalarIndent = lineIndent + s.inBlockScalar = true + } + } + if !s.inBlockScalar || lineIndent < s.blockScalarIndent { + s.inBlockScalar = false + return false + } + s.builder.WriteString(line) + s.builder.WriteByte('\n') + s.lastNonBlankEnd = s.builder.Len() + return true +} + +func (s *yamlNormalizationState) processStructuralLine(trimmed string) { + if trimmed == "" { + if s.blankRun < maxConsecutiveBlankLines { + s.builder.WriteByte('\n') + s.blankRun++ + } + return + } + s.builder.WriteString(trimmed) + s.builder.WriteByte('\n') + s.lastNonBlankEnd = s.builder.Len() + s.blankRun = 0 + if headerIndent, ok := blockScalarHeaderIndentForLine(trimmed); ok { + s.pendingBlockScalar = true + s.blockScalarHeaderIndent = headerIndent + } +} + +func (s *yamlNormalizationState) result() string { + return s.builder.String()[:s.lastNonBlankEnd] } func countLeadingSpaces(line string) int { diff --git a/pkg/workflow/copilot_engine_tools.go b/pkg/workflow/copilot_engine_tools.go index 44a8cf78764..fba8777f7d1 100644 --- a/pkg/workflow/copilot_engine_tools.go +++ b/pkg/workflow/copilot_engine_tools.go @@ -35,6 +35,13 @@ import ( var copilotEngineToolsLog = logger.New("workflow:copilot_engine_tools") +var copilotBuiltInTools = map[string]struct{}{ + "bash": {}, + "edit": {}, + "web-search": {}, + "playwright": {}, +} + // sanitizeCopilotShellCommand truncates a bash tool command at the first single // quote to produce a safe prefix for the Copilot CLI --allow-tool shell() argument. // @@ -59,214 +66,208 @@ func sanitizeCopilotShellCommand(cmdStr string) (string, bool) { // Returns a sorted list of arguments ready to be passed to the Copilot CLI. func (e *CopilotEngine) computeCopilotToolArguments(tools map[string]any, safeOutputs *SafeOutputsConfig, mcpScripts *MCPScriptsConfig, workflowData *WorkflowData) []string { copilotEngineToolsLog.Printf("Computing tool arguments: tools=%d", len(tools)) - if tools == nil { - tools = make(map[string]any) + tools = ensureToolMap(tools) + if hasCopilotBashWildcard(tools) { + copilotEngineToolsLog.Print("Bash wildcard detected, using --allow-all-tools") + return []string{"--allow-all-tools"} } - var args []string - hasRestrictedBashAllowlist := false + args, hasRestrictedBashAllowlist := collectCopilotBashToolArguments(tools) + args = e.addRestrictedBashMountArguments(args, hasRestrictedBashAllowlist, tools, safeOutputs, mcpScripts, workflowData) + args = appendCopilotBuiltinToolArguments(args, tools, safeOutputs, mcpScripts) + args = appendCopilotMCPServerArguments(args, tools) + args = sortAndDeduplicateAllowToolArgs(args) - // Check if bash has wildcard - if so, use --allow-all-tools instead - if bashConfig, hasBash := tools["bash"]; hasBash { - if bashCommands, ok := bashConfig.([]any); ok { - // Check for :* or * wildcard - if present, allow all tools - for _, cmd := range bashCommands { - if cmdStr, ok := cmd.(string); ok { - if cmdStr == ":*" || cmdStr == "*" { - // Use --allow-all-tools flag instead of individual tool permissions - copilotEngineToolsLog.Print("Bash wildcard detected, using --allow-all-tools") - return []string{"--allow-all-tools"} - } - } - } - } + copilotEngineToolsLog.Printf("Computed %d tool arguments", len(args)/2) + return args +} + +func ensureToolMap(tools map[string]any) map[string]any { + if tools != nil { + return tools } + return make(map[string]any) +} - // Handle bash/shell tools (when no wildcard) - if bashConfig, hasBash := tools["bash"]; hasBash { - if bashCommands, ok := bashConfig.([]any); ok { - hasRestrictedBashAllowlist = true - // Add specific shell commands - for _, cmd := range bashCommands { - if cmdStr, ok := cmd.(string); ok { - // Normalize trailing " *" wildcard (e.g. "jq *" → "jq") so that - // all engines emit the canonical prefix form (shell(jq)) regardless - // of whether the command was written with or without the wildcard. - cmdStr, _ = normalizeBashCommand(cmdStr) - // For stem commands (like dotnet, npm, cargo), Copilot CLI uses - // subcommand matching. When the user specifies just the base command - // (e.g., "dotnet"), append :* so "dotnet build", "dotnet test", etc. - // are all permitted. Skip if the command already has a colon (explicit - // matching) or a space (user already specified the subcommand). - if !strings.Contains(cmdStr, ":") && !strings.Contains(cmdStr, " ") && constants.CopilotStemCommands[cmdStr] { - args = append(args, "--allow-tool", fmt.Sprintf("shell(%s:*)", cmdStr)) - } else { - sanitized, wasSanitized := sanitizeCopilotShellCommand(cmdStr) - if wasSanitized { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage( - fmt.Sprintf("bash tool %q contains single quotes that crash Copilot CLI; "+ - "truncated to safe prefix %q for shell() prefix-matching. "+ - "Use %q in your workflow to silence this warning.", - cmdStr, sanitized, sanitized))) - } - args = append(args, "--allow-tool", fmt.Sprintf("shell(%s)", sanitized)) - } - } - } - } else { - // Bash with no specific commands or null value - allow all shell - args = append(args, "--allow-tool", "shell") +func hasCopilotBashWildcard(tools map[string]any) bool { + bashCommands, ok := tools["bash"].([]any) + if !ok { + return false + } + for _, cmd := range bashCommands { + cmdStr, ok := cmd.(string) + if ok && (cmdStr == ":*" || cmdStr == "*") { + return true } } + return false +} - // When MCP tools are mounted as CLI commands and bash uses a restricted allowlist, - // ensure mounted MCP CLI commands are executable via shell(:*). - // This avoids Copilot CLI permission blocks for mounted commands such as safeoutputs. - if hasRestrictedBashAllowlist { - effectiveWorkflowData := buildCLIWorkflowDataForMounts(workflowData, tools, safeOutputs, mcpScripts) +func collectCopilotBashToolArguments(tools map[string]any) ([]string, bool) { + bashConfig, hasBash := tools["bash"] + if !hasBash { + return nil, false + } + bashCommands, ok := bashConfig.([]any) + if !ok { + return []string{"--allow-tool", "shell"}, false + } - for _, serverName := range getMountedCLIServerNamesIfBashRestricted(effectiveWorkflowData, tools, safeOutputs, mcpScripts) { - args = append(args, "--allow-tool", fmt.Sprintf("shell(%s:*)", serverName)) - } - // When playwright is configured in CLI mode, playwright-cli must be executable. - // Automatically add shell(playwright-cli:*) to the restricted bash allowlist. - if workflowData != nil && isPlaywrightCLIMode(workflowData.Tools) { - args = append(args, "--allow-tool", "shell(playwright-cli:*)") - } - // When GitHub CLI mode is enabled (tools.github.mode: gh-proxy), GitHub access - // goes through the gh CLI, so allow shell(gh:*). - if isGitHubCLIModeEnabled(effectiveWorkflowData) { - args = append(args, "--allow-tool", "shell(gh:*)") + args := make([]string, 0, len(bashCommands)*2) + for _, cmd := range bashCommands { + cmdStr, ok := cmd.(string) + if !ok { + continue } + args = append(args, "--allow-tool", formatCopilotShellAllowTool(cmdStr)) } + return args, true +} - // Handle edit tools requirement for file write access - // Note: safe-outputs do not need write permission as they use MCP +func formatCopilotShellAllowTool(cmdStr string) string { + cmdStr, _ = normalizeBashCommand(cmdStr) + if !strings.Contains(cmdStr, ":") && !strings.Contains(cmdStr, " ") && constants.CopilotStemCommands[cmdStr] { + return fmt.Sprintf("shell(%s:*)", cmdStr) + } + + sanitized, wasSanitized := sanitizeCopilotShellCommand(cmdStr) + if wasSanitized { + fmt.Fprintln(os.Stderr, console.FormatWarningMessage( + fmt.Sprintf("bash tool %q contains single quotes that crash Copilot CLI; truncated to safe prefix %q for shell() prefix-matching. Use %q in your workflow to silence this warning.", + cmdStr, sanitized, sanitized))) + } + return fmt.Sprintf("shell(%s)", sanitized) +} + +func (e *CopilotEngine) addRestrictedBashMountArguments(args []string, hasRestrictedBashAllowlist bool, tools map[string]any, safeOutputs *SafeOutputsConfig, mcpScripts *MCPScriptsConfig, workflowData *WorkflowData) []string { + if !hasRestrictedBashAllowlist { + return args + } + + effectiveWorkflowData := buildCLIWorkflowDataForMounts(workflowData, tools, safeOutputs, mcpScripts) + for _, serverName := range getMountedCLIServerNamesIfBashRestricted(effectiveWorkflowData, tools, safeOutputs, mcpScripts) { + args = append(args, "--allow-tool", fmt.Sprintf("shell(%s:*)", serverName)) + } + if workflowData != nil && isPlaywrightCLIMode(workflowData.Tools) { + args = append(args, "--allow-tool", "shell(playwright-cli:*)") + } + if isGitHubCLIModeEnabled(effectiveWorkflowData) { + args = append(args, "--allow-tool", "shell(gh:*)") + } + return args +} + +func appendCopilotBuiltinToolArguments(args []string, tools map[string]any, safeOutputs *SafeOutputsConfig, mcpScripts *MCPScriptsConfig) []string { if _, hasEdit := tools["edit"]; hasEdit { copilotEngineToolsLog.Print("Edit tool enabled, adding write permission") args = append(args, "--allow-tool", "write") } - - // Handle safe_outputs MCP server - allow all tools if safe outputs are enabled - // This includes both safeOutputs config and safeOutputs.Jobs if HasSafeOutputsEnabled(safeOutputs) { copilotEngineToolsLog.Print("Safe-outputs enabled, adding MCP server permission") args = append(args, "--allow-tool", constants.SafeOutputsMCPServerID.String()) } - - // Handle mcp_scripts MCP server - allow the server if mcp-scripts are configured and feature flag is enabled if IsMCPScriptsEnabled(mcpScripts) { args = append(args, "--allow-tool", constants.MCPScriptsMCPServerID.String()) } - - // Handle web-fetch builtin tool (Copilot CLI uses web_fetch with underscore) if _, hasWebFetch := tools["web-fetch"]; hasWebFetch { copilotEngineToolsLog.Print("Web-fetch tool enabled, adding web_fetch permission") - // web-fetch -> web_fetch args = append(args, "--allow-tool", "web_fetch") } + return args +} - // Built-in tool names that should be skipped when processing MCP servers - // Note: GitHub is NOT included here because it needs MCP configuration in CLI mode - // Note: web-fetch is NOT included here because it needs explicit --allow-tool argument - builtInTools := map[string]struct { - }{ - "bash": {}, - "edit": {}, - "web-search": {}, - "playwright": {}, - } - - // Handle MCP server tools +func appendCopilotMCPServerArguments(args []string, tools map[string]any) []string { for toolName, toolConfig := range tools { - // Skip built-in tools we've already handled - if setutil.Contains(builtInTools, toolName) { + if setutil.Contains(copilotBuiltInTools, toolName) { continue } - - // GitHub is a special case - it's an MCP server but doesn't have explicit MCP config in the workflow - // It gets MCP configuration through the parser's processBuiltinMCPTool if toolName == "github" { - if toolConfigMap, ok := toolConfig.(map[string]any); ok { - if allowed, hasAllowed := toolConfigMap["allowed"]; hasAllowed { - if allowedList, ok := allowed.([]any); ok { - // Process allowed list in a single pass - hasWildcard := false - for _, allowedTool := range allowedList { - if toolStr, ok := allowedTool.(string); ok { - if toolStr == "*" { - // Wildcard means allow entire GitHub MCP server - hasWildcard = true - } else { - // Add individual tool permission - args = append(args, "--allow-tool", fmt.Sprintf("github(%s)", toolStr)) - } - } - } - - // Add server-level permission only if wildcard was present - if hasWildcard { - args = append(args, "--allow-tool", "github") - } - } - } else { - // No allowed field specified - allow entire GitHub MCP server - args = append(args, "--allow-tool", "github") - } - } else { - // GitHub tool exists but is not a map (e.g., github: null) - allow entire server - args = append(args, "--allow-tool", "github") - } + args = appendGitHubToolArguments(args, toolConfig) continue } + args = appendCustomMCPToolArguments(args, toolName, toolConfig) + } + return args +} - // Check if this is an MCP server configuration - if toolConfigMap, ok := toolConfig.(map[string]any); ok { - if hasMcp, _ := hasMCPConfig(toolConfigMap); hasMcp { - copilotEngineToolsLog.Printf("Adding custom MCP server permission: %s", toolName) - // Allow the entire MCP server - args = append(args, "--allow-tool", toolName) +func appendGitHubToolArguments(args []string, toolConfig any) []string { + toolConfigMap, ok := toolConfig.(map[string]any) + if !ok { + return append(args, "--allow-tool", "github") + } + allowed, hasAllowed := toolConfigMap["allowed"] + if !hasAllowed { + return append(args, "--allow-tool", "github") + } + allowedList, ok := allowed.([]any) + if !ok { + return args + } - // If it has specific allowed tools, add them individually - if allowed, hasAllowed := toolConfigMap["allowed"]; hasAllowed { - if allowedList, ok := allowed.([]any); ok { - for _, allowedTool := range allowedList { - if toolStr, ok := allowedTool.(string); ok { - args = append(args, "--allow-tool", fmt.Sprintf("%s(%s)", toolName, toolStr)) - } - } - } - } - } + hasWildcard := false + for _, allowedTool := range allowedList { + toolStr, ok := allowedTool.(string) + if !ok { + continue } + if toolStr == "*" { + hasWildcard = true + continue + } + args = append(args, "--allow-tool", fmt.Sprintf("github(%s)", toolStr)) + } + if hasWildcard { + args = append(args, "--allow-tool", "github") + } + return args +} + +func appendCustomMCPToolArguments(args []string, toolName string, toolConfig any) []string { + toolConfigMap, ok := toolConfig.(map[string]any) + if !ok { + return args + } + hasMcp, _ := hasMCPConfig(toolConfigMap) + if !hasMcp { + return args } - // Sort and deduplicate values, then rebuild args. - // Deduplication is needed because sanitizeCopilotShellCommand can truncate - // multiple different commands to the same safe prefix (e.g. several jq filters - // all become "jq"), producing duplicate --allow-tool shell(jq) entries. - if len(args) > 0 { - var values []string - for i := 1; i < len(args); i += 2 { - values = append(values, args[i]) + copilotEngineToolsLog.Printf("Adding custom MCP server permission: %s", toolName) + args = append(args, "--allow-tool", toolName) + allowedList, ok := toolConfigMap["allowed"].([]any) + if !ok { + return args + } + for _, allowedTool := range allowedList { + toolStr, ok := allowedTool.(string) + if ok { + args = append(args, "--allow-tool", fmt.Sprintf("%s(%s)", toolName, toolStr)) } - sort.Strings(values) + } + return args +} - // Rebuild args with sorted, deduplicated values - newArgs := make([]string, 0, len(args)) - prev := "" - for _, value := range values { - if value == prev { - continue - } - newArgs = append(newArgs, "--allow-tool", value) - prev = value - } - args = newArgs +func sortAndDeduplicateAllowToolArgs(args []string) []string { + if len(args) == 0 { + return args } - copilotEngineToolsLog.Printf("Computed %d tool arguments", len(args)/2) - return args + values := make([]string, 0, len(args)/2) + for i := 1; i < len(args); i += 2 { + values = append(values, args[i]) + } + sort.Strings(values) + + newArgs := make([]string, 0, len(args)) + prev := "" + for _, value := range values { + if value == prev { + continue + } + newArgs = append(newArgs, "--allow-tool", value) + prev = value + } + return newArgs } // generateCopilotToolArgumentsComment generates a multi-line comment showing each tool argument. diff --git a/pkg/workflow/gemini_engine.go b/pkg/workflow/gemini_engine.go index 1bb7cdaab3a..f7525951934 100644 --- a/pkg/workflow/gemini_engine.go +++ b/pkg/workflow/gemini_engine.go @@ -187,234 +187,147 @@ func (e *GeminiEngine) GetPreBundleSteps(workflowData *WorkflowData) []GitHubAct // GetExecutionSteps returns the GitHub Actions steps for executing Gemini func (e *GeminiEngine) GetExecutionSteps(workflowData *WorkflowData, logFile string) []GitHubActionStep { geminiLog.Printf("Generating execution steps for Gemini engine: workflow=%s, firewall=%v", workflowData.Name, isFirewallEnabled(workflowData)) - - var steps []GitHubActionStep - - // Write .gemini/settings.json with context.includeDirectories and tools.core. - // This step runs after the MCP gateway setup (which may have written mcpServers config) - // and merges the context/tools settings into any existing settings.json. settingsStep := e.generateGeminiSettingsStep(workflowData) - steps = append(steps, settingsStep) - - // Build gemini CLI arguments based on configuration - var geminiArgs []string - - // Model is passed via the native GEMINI_MODEL environment variable only when explicitly - // configured. When not configured, the Gemini CLI uses its built-in default model. - // This avoids embedding the value directly in the shell command (which fails template injection - // validation for GitHub Actions expressions like ${{ inputs.model }}). modelConfigured := workflowData.Model != "" + firewallEnabled := isFirewallEnabled(workflowData) + vertexWIF := isGeminiVertexWIF(workflowData) + geminiCommand := e.buildGeminiCLICommand(workflowData) + command := e.buildGeminiExecutionCommand(workflowData, logFile, geminiCommand, firewallEnabled) + env := e.buildGeminiExecutionEnv(workflowData, firewallEnabled, vertexWIF, modelConfigured) + step := e.buildGeminiExecutionStep(workflowData, command, env) + return []GitHubActionStep{settingsStep, step} +} - // Gemini CLI reads MCP config from .gemini/settings.json (project-level) - // The conversion script (convert_gateway_config_gemini.sh) writes settings.json - // during the MCP setup step, so no --mcp-config flag is needed here. - - // Auto-approve all tool executions (equivalent to Codex's --dangerously-bypass-approvals-and-sandbox) - // Without this, Gemini CLI's default approval mode rejects tool calls with "Tool execution denied by policy" - geminiArgs = append(geminiArgs, "--yolo") - - // Skip the workspace trust check so --yolo is not overridden to "default" approval mode. - // Gemini CLI v1.x checks whether the working directory is trusted and overrides --yolo - // with "default" approval mode (exit code 55) when the folder is untrusted. - // GEMINI_CLI_TRUST_WORKSPACE=true (also set in the step env) handles the same case via - // environment variable, but --skip-trust is more reliable when AWF's sandbox does not - // forward all host environment variables into the container. - geminiArgs = append(geminiArgs, "--skip-trust") - - // Add streaming JSON output (JSONL format, compatible with the log parser) - geminiArgs = append(geminiArgs, "--output-format", "stream-json") - - // Note: the --prompt argument is appended raw after shellJoinArgs below because it contains - // a shell command substitution ("$(cat ...)") that must NOT go through shellEscapeArg — - // single-quoting it would prevent shell expansion at runtime. - - // Build the command +func (e *GeminiEngine) buildGeminiCLICommand(workflowData *WorkflowData) string { + geminiArgs := []string{"--yolo", "--skip-trust", "--output-format", "stream-json"} commandName := "gemini" if workflowData.EngineConfig != nil && workflowData.EngineConfig.Command != "" { commandName = workflowData.EngineConfig.Command } - - // Append the prompt arg raw (not through shellJoinArgs) to preserve shell expansion geminiCommand := fmt.Sprintf(`%s %s --prompt "$(cat /tmp/gh-aw/aw-prompts/prompt.txt)"`, commandName, shellJoinArgs(geminiArgs)) - geminiCommand = getWorkspaceCommandPrefixFor(workflowData.EngineConfig) + geminiCommand + return getWorkspaceCommandPrefixFor(workflowData.EngineConfig) + geminiCommand +} - // Build the full command with AWF wrapping if enabled - var command string - firewallEnabled := isFirewallEnabled(workflowData) - if firewallEnabled { - // Get allowed domains: prefer the pre-warmed cache on WorkflowData to avoid - // re-running the expensive map+sort operation. - var allowedDomains string - if workflowData.CachedAllowedDomainsComputed { - allowedDomains = workflowData.CachedAllowedDomainsStr - } else { - allowedDomains = GetAllowedDomainsForEngine(constants.GeminiEngine, - workflowData.NetworkPermissions, - workflowData.Tools, - workflowData.Runtimes, - ) - } - // Add GHES/custom API target domains to the firewall allow-list when engine.api-target is set - if workflowData.EngineConfig != nil && workflowData.EngineConfig.APITarget != "" { - allowedDomains = mergeAPITargetDomains(allowedDomains, workflowData.EngineConfig.APITarget) - } - - npmPathSetup := GetNpmBinPathSetup() - geminiCommandWithPath := fmt.Sprintf("%s && %s", npmPathSetup, geminiCommand) - // Add MCP CLI bin directory to PATH when cli-proxy is enabled - if mcpCLIPath := GetMCPCLIPathSetup(workflowData); mcpCLIPath != "" { - geminiCommandWithPath = fmt.Sprintf("%s && %s", mcpCLIPath, geminiCommandWithPath) - } - - command = BuildAWFCommand(AWFCommandConfig{ - EngineName: "gemini", - EngineCommand: geminiCommandWithPath, - LogFile: logFile, - WorkflowData: workflowData, - UsesTTY: false, - AllowedDomains: allowedDomains, - // Create the agent step summary file before AWF starts so it is accessible - // inside the sandbox. The agent writes its step summary content here, and the - // file is appended to $GITHUB_STEP_SUMMARY after secret redaction. - PathSetup: "touch " + AgentStepSummaryPath, - // Exclude every env var whose step-env value is a secret so the agent - // cannot read raw token values via bash tools (env / printenv). - ExcludeEnvVarNames: ComputeAWFExcludeEnvVarNames(workflowData, e.GetRequiredSecretNames(workflowData)), - }) - } else { - command = fmt.Sprintf(`set -o pipefail +func (e *GeminiEngine) buildGeminiExecutionCommand(workflowData *WorkflowData, logFile, geminiCommand string, firewallEnabled bool) string { + if !firewallEnabled { + return fmt.Sprintf(`set -o pipefail printf '%%s' "$(date +%%s%%3N)" > %s touch %s (umask 177 && touch %s) %s 2>&1 | tee -a %s`, AgentCLIStartMsPath, AgentStepSummaryPath, logFile, geminiCommand, logFile) } - // Build environment variables - vertexWIF := isGeminiVertexWIF(workflowData) + geminiCommandWithPath := fmt.Sprintf("%s && %s", GetNpmBinPathSetup(), geminiCommand) + if mcpCLIPath := GetMCPCLIPathSetup(workflowData); mcpCLIPath != "" { + geminiCommandWithPath = fmt.Sprintf("%s && %s", mcpCLIPath, geminiCommandWithPath) + } + return BuildAWFCommand(AWFCommandConfig{ + EngineName: "gemini", + EngineCommand: geminiCommandWithPath, + LogFile: logFile, + WorkflowData: workflowData, + UsesTTY: false, + AllowedDomains: e.geminiAllowedDomains(workflowData), + PathSetup: "touch " + AgentStepSummaryPath, + ExcludeEnvVarNames: ComputeAWFExcludeEnvVarNames(workflowData, e.GetRequiredSecretNames(workflowData)), + }) +} + +func (e *GeminiEngine) geminiAllowedDomains(workflowData *WorkflowData) string { + allowedDomains := workflowData.CachedAllowedDomainsStr + if !workflowData.CachedAllowedDomainsComputed { + allowedDomains = GetAllowedDomainsForEngine(constants.GeminiEngine, workflowData.NetworkPermissions, workflowData.Tools, workflowData.Runtimes) + } + if workflowData.EngineConfig != nil && workflowData.EngineConfig.APITarget != "" { + allowedDomains = mergeAPITargetDomains(allowedDomains, workflowData.EngineConfig.APITarget) + } + return allowedDomains +} + +func (e *GeminiEngine) buildGeminiExecutionEnv(workflowData *WorkflowData, firewallEnabled, vertexWIF, modelConfigured bool) map[string]string { + env := e.baseGeminiExecutionEnv(workflowData, vertexWIF) + e.applyGeminiRuntimeEnv(env, workflowData, firewallEnabled, modelConfigured) + applyEngineCwdEnv(env, workflowData) + if workflowData.EngineConfig != nil && len(workflowData.EngineConfig.Env) > 0 { + maps.Copy(env, workflowData.EngineConfig.Env) + } + e.applyGeminiAgentEnv(env, workflowData) + if vertexWIF { + applyGeminiVertexWIFEnv(env, workflowData.EngineConfig.Auth) + } + return env +} + +func (e *GeminiEngine) baseGeminiExecutionEnv(workflowData *WorkflowData, vertexWIF bool) map[string]string { env := map[string]string{ - "GH_AW_PROMPT": constants.AwPromptsFile, - // Tag the step as a GitHub AW agentic execution for discoverability by agents - "GITHUB_AW": "true", - "GITHUB_WORKSPACE": "${{ github.workspace }}", - "RUNNER_TEMP": "${{ runner.temp }}", - // Override GITHUB_STEP_SUMMARY with a path that exists inside the sandbox. - // The runner's original path is unreachable within the AWF isolated filesystem; - // we create this file before the agent starts and append it to the real - // $GITHUB_STEP_SUMMARY after secret redaction. - "GITHUB_STEP_SUMMARY": AgentStepSummaryPath, - // Enable verbose debug logging from Gemini CLI for better diagnostics. - // Gemini CLI uses the npm 'debug' package, and 'gemini-cli:*' enables all - // internal Gemini CLI debug channels (see: https://gemini-cli-docs.pages.dev/cli/configuration). - // Non-JSON debug lines are gracefully skipped by ParseLogMetrics. - "DEBUG": "gemini-cli:*", - // Trust the workspace to prevent Gemini CLI v1.x from overriding --yolo to default - // approval mode when the workspace is untrusted, which causes exit code 55. + "DEBUG": "gemini-cli:*", "GEMINI_CLI_TRUST_WORKSPACE": "true", + "GH_AW_PROMPT": constants.AwPromptsFile, + "GITHUB_AW": "true", + "GITHUB_STEP_SUMMARY": AgentStepSummaryPath, + "GITHUB_WORKSPACE": "${{ github.workspace }}", + "RUNNER_TEMP": "${{ runner.temp }}", } if !vertexWIF { - // Set static API key when WIF is not configured. - // When WIF is active, authentication is handled by the AWF api-proxy sidecar - // via the AWF_AUTH_GCP_* env vars set through engine.auth. env["GEMINI_API_KEY"] = "${{ secrets.GEMINI_API_KEY }}" } injectWorkflowCallNetworkAllowedEnv(env, workflowData) - // Indicate the phase: "agent" for the main run, "detection" for threat detection, - // and "evals" for the eval harness execution. - // Include the compiler version so agents can identify which gh-aw version generated the workflow + return env +} + +func (e *GeminiEngine) applyGeminiRuntimeEnv(env map[string]string, workflowData *WorkflowData, firewallEnabled, modelConfigured bool) { env["GH_AW_PHASE"] = workflowRunPhase(workflowData) if IsRelease() { env["GH_AW_VERSION"] = GetVersion() } else { env["GH_AW_VERSION"] = "dev" } - - // Add MCP config env var if needed (points to .gemini/settings.json for Gemini) if HasMCPServers(workflowData) { env["GH_AW_MCP_CONFIG"] = "${{ github.workspace }}/.gemini/settings.json" } - - // When the firewall (AWF) is enabled with --enable-api-proxy, point Gemini CLI at the - // LLM gateway sidecar instead of the real googleapis.com endpoint. if firewallEnabled { env["GEMINI_API_BASE_URL"] = fmt.Sprintf("http://host.docker.internal:%d", constants.GeminiLLMGatewayPort) - - // Set git identity environment variables so the first git commit succeeds inside the - // container. AWF's --env-all forwards these to the container, ensuring git does not - // rely on the host-side ~/.gitconfig which is not visible in the sandbox. maps.Copy(env, getGitIdentityEnvVars()) } - - // Add safe outputs env applySafeOutputEnvToMap(env, workflowData) - - // Propagate W3C trace context so engine spans nest under the gh-aw.agent.setup span. applyTraceContextEnvToMap(env) - if workflowData.EngineConfig != nil && workflowData.EngineConfig.MaxTurns != "" { env["GH_AW_MAX_TURNS"] = workflowData.EngineConfig.MaxTurns } else { env["GH_AW_MAX_TURNS"] = compilerenv.BuildDefaultMaxTurnsExpression() } - - // Set the model environment variable only when explicitly configured. - // When model is configured, use the native GEMINI_MODEL env var - the Gemini CLI reads it - // directly, avoiding the need to embed the value in the shell command (which would fail - // template injection validation for GitHub Actions expressions like ${{ inputs.model }}). - // When model is not configured, let the Gemini CLI use its built-in default model. if modelConfigured { geminiLog.Printf("Setting %s env var for model: %s", constants.GeminiCLIModelEnvVar, workflowData.Model) env[constants.GeminiCLIModelEnvVar] = workflowData.Model } +} - // Add custom environment variables from engine config. - // This allows users to override the default engine token expression (e.g. - // GEMINI_API_KEY: ${{ secrets.MY_ORG_GEMINI_KEY }}) via engine.env. - applyEngineCwdEnv(env, workflowData) - if workflowData.EngineConfig != nil && len(workflowData.EngineConfig.Env) > 0 { - maps.Copy(env, workflowData.EngineConfig.Env) - } - - // Add custom environment variables from agent config +func (e *GeminiEngine) applyGeminiAgentEnv(env map[string]string, workflowData *WorkflowData) { agentConfig := getAgentConfig(workflowData) - if agentConfig != nil && len(agentConfig.Env) > 0 { - maps.Copy(env, agentConfig.Env) - geminiLog.Printf("Added %d custom env vars from agent config", len(agentConfig.Env)) + if agentConfig == nil || len(agentConfig.Env) == 0 { + return } + maps.Copy(env, agentConfig.Env) + geminiLog.Printf("Added %d custom env vars from agent config", len(agentConfig.Env)) +} - // Apply Vertex AI WIF env vars AFTER engine.env and agent.env merges to ensure - // they cannot be overridden by user-provided engine.env values. - if vertexWIF { - auth := workflowData.EngineConfig.Auth - // Gemini CLI v0.39+ selects Vertex AI backend when this is set to "true". - env["GOOGLE_GENAI_USE_VERTEXAI"] = "true" - env["GOOGLE_CLOUD_PROJECT"] = auth.GoogleProject - location := auth.GoogleLocation - if location == "" { - location = "us-central1" - } - env["GOOGLE_CLOUD_LOCATION"] = location +func applyGeminiVertexWIFEnv(env map[string]string, auth *EngineAuthConfig) { + location := auth.GoogleLocation + if location == "" { + location = "us-central1" } + env["GOOGLE_CLOUD_LOCATION"] = location + env["GOOGLE_CLOUD_PROJECT"] = auth.GoogleProject + env["GOOGLE_GENAI_USE_VERTEXAI"] = "true" +} - // Generate the execution step +func (e *GeminiEngine) buildGeminiExecutionStep(workflowData *WorkflowData, command string, env map[string]string) GitHubActionStep { stepLines := []string{ " - name: Execute Gemini CLI", " id: agentic_execution", + " timeout-minutes: " + resolveStepTimeoutValue(workflowData), } - - // Add timeout at step level (GitHub Actions standard) - stepLines = append(stepLines, " timeout-minutes: "+resolveStepTimeoutValue(workflowData)) - - // Filter environment variables for security - allowedSecrets := e.GetRequiredSecretNames(workflowData) - filteredEnv := FilterEnvForSecrets(env, allowedSecrets) - - // Inject GH_TOKEN for CLI proxy (added after filtering since it uses a special - // fallback expression that is always allowed when cli-proxy is enabled) + filteredEnv := FilterEnvForSecrets(env, e.GetRequiredSecretNames(workflowData)) addCliProxyGHTokenToEnv(filteredEnv, workflowData) - - // Format step with command and env - stepLines = FormatStepWithCommandAndEnv(stepLines, command, filteredEnv) - - steps = append(steps, GitHubActionStep(stepLines)) - return steps + return GitHubActionStep(FormatStepWithCommandAndEnv(stepLines, command, filteredEnv)) } diff --git a/pkg/workflow/permissions_compiler_validator.go b/pkg/workflow/permissions_compiler_validator.go index 16b7bfec537..6eb2f31ecfa 100644 --- a/pkg/workflow/permissions_compiler_validator.go +++ b/pkg/workflow/permissions_compiler_validator.go @@ -52,123 +52,117 @@ var permissionsCompilerLog = logger.New("workflow:permissions_compiler_validator // branch security, GitHub MCP toolset permissions, and the id-token write warning. // It returns the parsed *Permissions for reuse in subsequent validation steps. func (c *Compiler) validatePermissions(workflowData *WorkflowData, markdownPath string) (*Permissions, error) { - // Use the cached *Permissions object when available to avoid repeated YAML parsing. - // CachedPermissions is populated by applyDefaults after all permission mutations are applied. - // Fall back to parsing from the raw string for code paths that bypass applyDefaults - // (e.g., tests that construct WorkflowData directly). - var workflowPermissions *Permissions + workflowPermissions := getWorkflowPermissions(workflowData) + if err := validateCachedPermissionScopeNames(workflowData, markdownPath); err != nil { + return nil, err + } + if err := c.validateCorePermissionRules(workflowData, workflowPermissions, markdownPath); err != nil { + return nil, err + } + if err := c.validateGitHubToolPermissions(workflowData, workflowPermissions, markdownPath); err != nil { + return nil, err + } + if err := validateOIDCPermissions(workflowData, workflowPermissions); err != nil { + return nil, formatCompilerError(markdownPath, "error", err.Error(), err) + } + c.emitPermissionWarnings(workflowData, workflowPermissions, markdownPath) + return workflowPermissions, nil +} + +func getWorkflowPermissions(workflowData *WorkflowData) *Permissions { if workflowData.CachedPermissions != nil { - workflowPermissions = workflowData.CachedPermissions - } else { - workflowPermissions = NewPermissionsParser(workflowData.Permissions).ToPermissions() + return workflowData.CachedPermissions } + return NewPermissionsParser(workflowData.Permissions).ToPermissions() +} - // Validate permission scope names for typos (e.g. "contnts" → "contents") +func validateCachedPermissionScopeNames(workflowData *WorkflowData, markdownPath string) error { workflowLog.Printf("Validating permission scope names") - var scopeValidationErr error - if workflowData.CachedPermissionScopeNamesSet { - scopeValidationErr = workflowData.CachedPermissionScopeNamesErr - } else { + scopeValidationErr := workflowData.CachedPermissionScopeNamesErr + if !workflowData.CachedPermissionScopeNamesSet { scopeValidationErr = ValidatePermissionScopeNames(workflowData.Permissions) } - if scopeValidationErr != nil { - return nil, formatCompilerError(markdownPath, "error", scopeValidationErr.Error(), scopeValidationErr) + if scopeValidationErr == nil { + return nil } + return formatCompilerError(markdownPath, "error", scopeValidationErr.Error(), scopeValidationErr) +} - // Validate dangerous permissions - workflowLog.Printf("Validating dangerous permissions") - if err := validateDangerousPermissions(workflowData, workflowPermissions); err != nil { - return nil, formatCompilerError(markdownPath, "error", err.Error(), err) +func (c *Compiler) validateCorePermissionRules(workflowData *WorkflowData, workflowPermissions *Permissions, markdownPath string) error { + if err := validateFormattedPermissionCheck("dangerous permissions", markdownPath, func() error { + return validateDangerousPermissions(workflowData, workflowPermissions) + }); err != nil { + return err } - - // Validate GitHub App-only permissions require a GitHub App to be configured - workflowLog.Printf("Validating GitHub App-only permissions") - if err := validateGitHubAppOnlyPermissions(workflowData, workflowPermissions); err != nil { - return nil, formatCompilerError(markdownPath, "error", err.Error(), err) + if err := validateFormattedPermissionCheck("GitHub App-only permissions", markdownPath, func() error { + return validateGitHubAppOnlyPermissions(workflowData, workflowPermissions) + }); err != nil { + return err } - - // Validate tools.github.github-app.permissions does not use "write" - workflowLog.Printf("Validating GitHub MCP app permissions (no write)") - if err := validateGitHubMCPAppPermissionsNoWrite(workflowData); err != nil { - return nil, formatCompilerError(markdownPath, "error", err.Error(), err) + if err := validateFormattedPermissionCheck("GitHub MCP app permissions (no write)", markdownPath, func() error { + return validateGitHubMCPAppPermissionsNoWrite(workflowData) + }); err != nil { + return err } - - // Warn when github-app.permissions is set in contexts that don't support it - warnGitHubAppPermissionsUnsupportedContexts(workflowData) - - // Validate workflow_run triggers have branch restrictions workflowLog.Printf("Validating workflow_run triggers for branch restrictions") if err := c.validateWorkflowRunBranches(workflowData, markdownPath); err != nil { - return nil, err + return err } - - // Validate pull_request_target trigger security workflowLog.Printf("Validating pull_request_target trigger security") if err := c.validatePullRequestTargetTrigger(workflowData, markdownPath); err != nil { - return nil, err + return err } + warnGitHubAppPermissionsUnsupportedContexts(workflowData) + return nil +} - // Validate permissions against GitHub MCP toolsets - workflowLog.Printf("Validating permissions for GitHub MCP toolsets") - if workflowData.ParsedTools != nil && workflowData.ParsedTools.GitHub != nil { - // Check if GitHub tool was explicitly configured in frontmatter - // If permissions exist but tools.github was NOT explicitly configured, - // skip validation and let the GitHub MCP server handle permission issues - hasPermissions := workflowData.Permissions != "" - - workflowLog.Printf("Permission validation check: hasExplicitGitHubTool=%v, hasPermissions=%v", - workflowData.HasExplicitGitHubTool, hasPermissions) - - // Skip validation if permissions exist but GitHub tool was auto-added (not explicit) - if hasPermissions && !workflowData.HasExplicitGitHubTool { - workflowLog.Printf("Skipping permission validation: permissions exist but tools.github not explicitly configured") - } else { - // Validate permissions using the typed GitHub tool configuration. - // Pass the cached parsed toolsets from applyDefaults to avoid a redundant - // ParseGitHubToolsets call inside ValidatePermissions. - validationResult := ValidatePermissions(workflowPermissions, workflowData.ParsedTools.GitHub, workflowData.CachedParsedToolsets) - - if validationResult.HasValidationIssues { - // Format the validation message - message := FormatValidationMessage(validationResult, c.strictMode) - - if len(validationResult.MissingPermissions) > 0 { - downgradeToWarning := c.strictMode && shouldDowngradeDefaultToolsetPermissionError(workflowData.ParsedTools.GitHub) - if c.strictMode && !downgradeToWarning { - // In strict mode, missing permissions are errors - return nil, formatCompilerError(markdownPath, "error", message, nil) - } - - if downgradeToWarning { - message += "\n\n" + missingPermissionsDefaultToolsetWarning - } +func validateFormattedPermissionCheck(name string, markdownPath string, fn func() error) error { + workflowLog.Printf("Validating %s", name) + if err := fn(); err != nil { + return formatCompilerError(markdownPath, "error", err.Error(), err) + } + return nil +} - // Emit to stderr once per markdown path + warning fingerprint. - // Prefer frontmatter hash when available; otherwise use the formatted - // message as a fallback fingerprint for code paths/tests where the hash - // is not set. - warningFingerprint := workflowData.FrontmatterHash - if warningFingerprint == "" { - warningFingerprint = message - } - if c.permissionWarningShown[markdownPath] != warningFingerprint { - // In non-strict mode, missing permissions are warnings. - // In strict mode with default-only toolsets, this is intentionally downgraded to warning. - fmt.Fprintln(os.Stderr, formatCompilerMessage(markdownPath, "warning", message)) - c.permissionWarningShown[markdownPath] = warningFingerprint - } - c.IncrementWarningCount() - } - } - } +func (c *Compiler) validateGitHubToolPermissions(workflowData *WorkflowData, workflowPermissions *Permissions, markdownPath string) error { + workflowLog.Printf("Validating permissions for GitHub MCP toolsets") + if workflowData.ParsedTools == nil || workflowData.ParsedTools.GitHub == nil { + return nil + } + hasPermissions := workflowData.Permissions != "" + workflowLog.Printf("Permission validation check: hasExplicitGitHubTool=%v, hasPermissions=%v", workflowData.HasExplicitGitHubTool, hasPermissions) + if hasPermissions && !workflowData.HasExplicitGitHubTool { + workflowLog.Printf("Skipping permission validation: permissions exist but tools.github not explicitly configured") + return nil + } + validationResult := ValidatePermissions(workflowPermissions, workflowData.ParsedTools.GitHub, workflowData.CachedParsedToolsets) + if !validationResult.HasValidationIssues || len(validationResult.MissingPermissions) == 0 { + return nil + } + message := FormatValidationMessage(validationResult, c.strictMode) + if c.strictMode && !shouldDowngradeDefaultToolsetPermissionError(workflowData.ParsedTools.GitHub) { + return formatCompilerError(markdownPath, "error", message, nil) } + if shouldDowngradeDefaultToolsetPermissionError(workflowData.ParsedTools.GitHub) { + message += "\n\n" + missingPermissionsDefaultToolsetWarning + } + emitPermissionValidationWarning(c, workflowData, markdownPath, message) + return nil +} - // Enforce required id-token: write permission for OIDC auth users. - if err := validateOIDCPermissions(workflowData, workflowPermissions); err != nil { - return nil, formatCompilerError(markdownPath, "error", err.Error(), err) +func emitPermissionValidationWarning(c *Compiler, workflowData *WorkflowData, markdownPath string, message string) { + warningFingerprint := workflowData.FrontmatterHash + if warningFingerprint == "" { + warningFingerprint = message + } + if c.permissionWarningShown[markdownPath] != warningFingerprint { + fmt.Fprintln(os.Stderr, formatCompilerMessage(markdownPath, "warning", message)) + c.permissionWarningShown[markdownPath] = warningFingerprint } + c.IncrementWarningCount() +} - // Emit warning if id-token: write permission is detected +func (c *Compiler) emitPermissionWarnings(workflowData *WorkflowData, workflowPermissions *Permissions, markdownPath string) { workflowLog.Printf("Checking for id-token: write permission") if level, exists := workflowPermissions.Get(PermissionIdToken); exists && level == PermissionWrite { warningMsg := `This workflow grants id-token: write permission @@ -178,18 +172,21 @@ Ensure proper audience validation and trust policies are configured.` c.IncrementWarningCount() } if shouldEmitCopilotRequestsEnableTip(workflowData, workflowPermissions) && !c.repositoryOwnerIsIndividualUser() { - if !c.copilotRequestsTipShown[markdownPath] { - if c.batchMode { - c.copilotTipNeeded = true - } else { - tipMsg := `Tip: set permissions.copilot-requests: write to use GitHub Actions token-based inference with the Copilot engine instead of a personal access token (COPILOT_GITHUB_TOKEN). This option requires that your organization has centralized Copilot billing enabled and may not be available in all organizations — see https://github.github.com/gh-aw/reference/billing/ for details.` - fmt.Fprintln(os.Stderr, formatCompilerMessage(markdownPath, "info", tipMsg)) - } - c.copilotRequestsTipShown[markdownPath] = true - } + c.emitCopilotRequestsPermissionTip(markdownPath) } +} - return workflowPermissions, nil +func (c *Compiler) emitCopilotRequestsPermissionTip(markdownPath string) { + if c.copilotRequestsTipShown[markdownPath] { + return + } + if c.batchMode { + c.copilotTipNeeded = true + } else { + tipMsg := `Tip: set permissions.copilot-requests: write to use GitHub Actions token-based inference with the Copilot engine instead of a personal access token (COPILOT_GITHUB_TOKEN). This option requires that your organization has centralized Copilot billing enabled and may not be available in all organizations — see https://github.github.com/gh-aw/reference/billing/ for details.` + fmt.Fprintln(os.Stderr, formatCompilerMessage(markdownPath, "info", tipMsg)) + } + c.copilotRequestsTipShown[markdownPath] = true } // repositoryOwnerIsIndividualUser reports whether the repository owner is confirmed diff --git a/pkg/workflow/universal_llm_consumer_engine.go b/pkg/workflow/universal_llm_consumer_engine.go index 3818da7d5c3..7ae63819c67 100644 --- a/pkg/workflow/universal_llm_consumer_engine.go +++ b/pkg/workflow/universal_llm_consumer_engine.go @@ -245,130 +245,109 @@ func (e *UniversalLLMConsumerEngine) BuildCLIEngineExecutionSteps( ) []GitHubActionStep { universalLLMConsumerLog.Printf("Generating execution steps for %s engine: workflow=%s, firewall=%v", cfg.DefaultCommandName, workflowData.Name, isFirewallEnabled(workflowData)) + modelConfigured := workflowData.Model != "" + firewallEnabled := isFirewallEnabled(workflowData) + engineCommand := e.buildUniversalCLICommand(workflowData, cfg) + command := e.buildUniversalCLIExecutionCommand(workflowData, logFile, cfg, engineCommand, firewallEnabled, modelConfigured) + env := e.buildUniversalCLIExecutionEnv(workflowData, cfg, firewallEnabled, modelConfigured) + step := e.buildUniversalCLIExecutionStep(workflowData, cfg, command, env) + return prependUniversalCLIConfigStep(cfg.ConfigStep, step) +} - var steps []GitHubActionStep - - // Prepend the config step (writes permissions JSON to workspace). - if len(cfg.ConfigStep) > 0 { - steps = append(steps, cfg.ConfigStep) +func prependUniversalCLIConfigStep(configStep, step GitHubActionStep) []GitHubActionStep { + if len(configStep) == 0 { + return []GitHubActionStep{step} } + return []GitHubActionStep{configStep, step} +} - modelConfigured := workflowData.Model != "" - - // Build CLI command: run "". +func (e *UniversalLLMConsumerEngine) buildUniversalCLICommand(workflowData *WorkflowData, cfg UniversalCLIEngineExecutionConfig) string { cliArgs := append([]string{}, cfg.ExtraCLIArgs...) - promptArg := fmt.Sprintf("\"$(cat %s)\"", constants.AwPromptsFile) commandName := cfg.DefaultCommandName if workflowData.EngineConfig != nil && workflowData.EngineConfig.Command != "" { commandName = workflowData.EngineConfig.Command } - engineCommand := fmt.Sprintf("%s run %s %s", commandName, shellJoinArgs(cliArgs), promptArg) - engineCommand = getWorkspaceCommandPrefixFor(workflowData.EngineConfig) + engineCommand + engineCommand := fmt.Sprintf("%s run %s \"$(cat %s)\"", commandName, shellJoinArgs(cliArgs), constants.AwPromptsFile) + return getWorkspaceCommandPrefixFor(workflowData.EngineConfig) + engineCommand +} - firewallEnabled := isFirewallEnabled(workflowData) - var command string +func (e *UniversalLLMConsumerEngine) buildUniversalCLIExecutionCommand(workflowData *WorkflowData, logFile string, cfg UniversalCLIEngineExecutionConfig, engineCommand string, firewallEnabled, modelConfigured bool) string { if firewallEnabled { - model := "" - if modelConfigured { - model = workflowData.Model - } - // Get allowed domains: prefer the pre-warmed cache on WorkflowData to avoid - // re-running the expensive map+sort operation. - var allowedDomains string - if workflowData.CachedAllowedDomainsComputed { - allowedDomains = workflowData.CachedAllowedDomainsStr - } else { - // The model was validated before reaching here, so a malformed model - // (e.g. leading slash) must never occur. Panic is the correct response - // to an internal invariant violation. - allowedDomains = mustGetAllowedDomainsForEngineWithModel( - cfg.EngineConstant, - model, - workflowData.NetworkPermissions, - workflowData.Tools, - workflowData.Runtimes, - ) - } - - npmPathSetup := GetNpmBinPathSetup() - // Propagate no_proxy inside the AWF container. --env-all forwards NO_PROXY - // from the YAML env block, but Bun (and other runtimes) also check the - // lowercase variant, so we export it explicitly from the uppercase value. - engineCommandWithPath := fmt.Sprintf("export no_proxy=\"${NO_PROXY:-}\" && %s && %s", npmPathSetup, engineCommand) + engineCommandWithPath := fmt.Sprintf("export no_proxy=\"${NO_PROXY:-}\" && %s && %s", GetNpmBinPathSetup(), engineCommand) if mcpCLIPath := GetMCPCLIPathSetup(workflowData); mcpCLIPath != "" { engineCommandWithPath = fmt.Sprintf("%s && %s", mcpCLIPath, engineCommandWithPath) } - - command = BuildAWFCommand(AWFCommandConfig{ + return BuildAWFCommand(AWFCommandConfig{ EngineName: cfg.DefaultCommandName, EngineCommand: engineCommandWithPath, LogFile: logFile, WorkflowData: workflowData, UsesTTY: false, - AllowedDomains: allowedDomains, + AllowedDomains: resolveUniversalCLIAllowedDomains(workflowData, cfg, modelConfigured), }) - } else if cfg.WriteTimestamp { - command = fmt.Sprintf("set -o pipefail\nexport no_proxy=\"${NO_PROXY:-}\"\nprintf '%%s' \"$(date +%%s%%3N)\" > %s\n%s 2>&1 | tee -a %s", + } + if cfg.WriteTimestamp { + return fmt.Sprintf("set -o pipefail\nexport no_proxy=\"${NO_PROXY:-}\"\nprintf '%%s' \"$(date +%%s%%3N)\" > %s\n%s 2>&1 | tee -a %s", AgentCLIStartMsPath, engineCommand, logFile) - } else { - command = fmt.Sprintf("set -o pipefail\nexport no_proxy=\"${NO_PROXY:-}\"\n%s 2>&1 | tee -a %s", engineCommand, logFile) } + return fmt.Sprintf("set -o pipefail\nexport no_proxy=\"${NO_PROXY:-}\"\n%s 2>&1 | tee -a %s", engineCommand, logFile) +} +func resolveUniversalCLIAllowedDomains(workflowData *WorkflowData, cfg UniversalCLIEngineExecutionConfig, modelConfigured bool) string { + if workflowData.CachedAllowedDomainsComputed { + return workflowData.CachedAllowedDomainsStr + } + model := "" + if modelConfigured { + model = workflowData.Model + } + return mustGetAllowedDomainsForEngineWithModel( + cfg.EngineConstant, + model, + workflowData.NetworkPermissions, + workflowData.Tools, + workflowData.Runtimes, + ) +} + +func (e *UniversalLLMConsumerEngine) buildUniversalCLIExecutionEnv(workflowData *WorkflowData, cfg UniversalCLIEngineExecutionConfig, firewallEnabled, modelConfigured bool) map[string]string { env := map[string]string{ "GH_AW_PROMPT": constants.AwPromptsFile, "GITHUB_WORKSPACE": "${{ github.workspace }}", + "NO_PROXY": constants.AWFNoProxyHosts, "RUNNER_TEMP": "${{ runner.temp }}", - // Set NO_PROXY so that the AWF agent's HTTP client skips the squid proxy - // for local endpoints. The lowercase no_proxy variant is exported inside - // the run script rather than as a YAML env key because GitHub's workflow - // parser rejects case-insensitive duplicate env keys (NO_PROXY/no_proxy), - // which causes workflow_dispatch to fail with "failed to parse workflow". - "NO_PROXY": constants.AWFNoProxyHosts, } injectWorkflowCallNetworkAllowedEnv(env, workflowData) e.ApplyUniversalProviderEnv(env, workflowData, firewallEnabled) - if HasMCPServers(workflowData) { env["GH_AW_MCP_CONFIG"] = "${{ github.workspace }}/" + cfg.MCPConfigFile } - applySafeOutputEnvToMap(env, workflowData) - - // Propagate W3C trace context so engine spans nest under the gh-aw.agent.setup span. applyTraceContextEnvToMap(env) - if workflowData.EngineConfig != nil && workflowData.EngineConfig.MaxTurns != "" { env["GH_AW_MAX_TURNS"] = workflowData.EngineConfig.MaxTurns } else { env["GH_AW_MAX_TURNS"] = compilerenv.BuildDefaultMaxTurnsExpression() } - - // Model env var (only when explicitly configured and the engine supports it). if modelConfigured && cfg.ModelEnvVarName != "" { universalLLMConsumerLog.Printf("Setting %s env var for model: %s", cfg.ModelEnvVarName, workflowData.Model) env[cfg.ModelEnvVarName] = workflowData.Model } - - // Custom env from engine config (allows provider key override). applyEngineCwdEnv(env, workflowData) if workflowData.EngineConfig != nil && len(workflowData.EngineConfig.Env) > 0 { maps.Copy(env, workflowData.EngineConfig.Env) } - - // Agent config env. - agentConfig := getAgentConfig(workflowData) - if agentConfig != nil && len(agentConfig.Env) > 0 { + if agentConfig := getAgentConfig(workflowData); agentConfig != nil && len(agentConfig.Env) > 0 { maps.Copy(env, agentConfig.Env) } + return env +} +func (e *UniversalLLMConsumerEngine) buildUniversalCLIExecutionStep(workflowData *WorkflowData, cfg UniversalCLIEngineExecutionConfig, command string, env map[string]string) GitHubActionStep { stepLines := []string{ " - name: " + cfg.StepName, " id: agentic_execution", } - allowedSecrets := e.GetUniversalRequiredSecretNames(workflowData) - filteredEnv := FilterEnvForSecrets(env, allowedSecrets) - stepLines = FormatStepWithCommandAndEnv(stepLines, command, filteredEnv) - - steps = append(steps, GitHubActionStep(stepLines)) - return steps + filteredEnv := FilterEnvForSecrets(env, e.GetUniversalRequiredSecretNames(workflowData)) + return GitHubActionStep(FormatStepWithCommandAndEnv(stepLines, command, filteredEnv)) } From 8cce734f9a275184ce428646bfb32d35b3e4b4fd Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Sun, 2 Aug 2026 18:53:25 +0000 Subject: [PATCH 3/4] refactor: reduce function lengths in pkg/workflow and pkg/cli (slices 1-3) Co-authored-by: pelikhan <4175913+pelikhan@users.noreply.github.com> --- pkg/cli/access_log.go | 115 +++-- pkg/cli/actions_build_command.go | 80 ++-- pkg/cli/add_interactive_engine.go | 332 +++++++-------- pkg/cli/add_interactive_git.go | 403 ++++++++++-------- pkg/cli/add_interactive_orchestrator.go | 161 ++++--- pkg/cli/add_interactive_schedule.go | 336 +++++++-------- pkg/cli/add_interactive_workflow.go | 252 ++++++----- pkg/cli/add_package_manifest.go | 366 +++++++++------- pkg/cli/add_skill_rewrite.go | 149 +++---- pkg/cli/add_wizard_command.go | 134 +++--- pkg/cli/add_workflow_pr.go | 151 +++---- pkg/cli/add_workflow_resolution.go | 339 ++++++++------- pkg/cli/audit_agentic_analysis.go | 258 +++++------ pkg/cli/audit_comparison.go | 137 +++--- pkg/cli/audit_expanded.go | 236 +++++----- pkg/cli/audit_report_experiments.go | 108 ++--- pkg/workflow/safe_outputs_actions.go | 12 + pkg/workflow/safe_outputs_app_config.go | 12 + pkg/workflow/safe_outputs_config_base.go | 4 + .../safe_outputs_config_extraction.go | 4 + .../safe_outputs_config_generation.go | 8 + pkg/workflow/safe_outputs_config_global.go | 4 + pkg/workflow/safe_outputs_config_runtime.go | 4 + pkg/workflow/safe_outputs_data_schema.go | 4 + pkg/workflow/safe_outputs_handler_registry.go | 169 ++++---- pkg/workflow/safe_outputs_jobs.go | 4 + pkg/workflow/safe_outputs_max_validation.go | 4 + pkg/workflow/safe_outputs_messages_config.go | 4 + pkg/workflow/safe_outputs_permissions.go | 4 + ...utputs_steps_shell_expansion_validation.go | 4 + .../safe_outputs_tools_computation.go | 5 + pkg/workflow/safe_outputs_tools_generation.go | 8 + .../safe_outputs_tools_repo_params.go | 4 + pkg/workflow/safe_outputs_validation.go | 4 + .../safe_outputs_validation_config.go | 4 + pkg/workflow/tools_parser.go | 12 + 36 files changed, 1994 insertions(+), 1841 deletions(-) diff --git a/pkg/cli/access_log.go b/pkg/cli/access_log.go index 2e66f38d9d0..a903423c68d 100644 --- a/pkg/cli/access_log.go +++ b/pkg/cli/access_log.go @@ -86,85 +86,76 @@ func (d *DomainAnalysis) AddMetrics(other LogAnalysis) { // parseSquidAccessLog parses a squid access log file and extracts domain information func parseSquidAccessLog(logPath string, verbose bool) (*DomainAnalysis, error) { accessLogLog.Printf("Parsing squid access log: %s", logPath) - file, err := os.Open(logPath) if err != nil { accessLogLog.Printf("Failed to open access log %s: %v", logPath, err) return nil, fmt.Errorf("failed to open access log: %w", err) } defer file.Close() - analysis := &DomainAnalysis{} + allowedDomainsSet := make(map[string]struct{}) + blockedDomainsSet := make(map[string]struct{}) + if err := scanSquidAccessLog(file, analysis, allowedDomainsSet, blockedDomainsSet, verbose); err != nil { + return nil, err + } + sort.Strings(analysis.AllowedDomains) + sort.Strings(analysis.BlockedDomains) + accessLogLog.Printf("Parsed access log: total_requests=%d, allowed=%d, blocked=%d, unique_allowed_domains=%d, unique_blocked_domains=%d", analysis.TotalRequests, analysis.AllowedRequests, analysis.BlockedRequests, len(analysis.AllowedDomains), len(analysis.BlockedDomains)) + return analysis, nil +} - allowedDomainsSet := make(map[string]struct { - }) - blockedDomainsSet := make(map[string]struct { - }) - +func scanSquidAccessLog(file *os.File, analysis *DomainAnalysis, allowedDomainsSet, blockedDomainsSet map[string]struct{}, verbose bool) error { scanner := bufio.NewScanner(file) for scanner.Scan() { - line := strings.TrimSpace(scanner.Text()) - if line == "" || strings.HasPrefix(line, "#") { - continue - } - - entry, err := parseSquidLogLine(line) - if err != nil { - if verbose { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Failed to parse log line: %v", err))) - } - continue - } - - analysis.TotalRequests++ - - // Extract domain from URL - domain := stringutil.ExtractDomainFromURL(entry.URL) - if domain == "" { - continue - } - - // Determine if request was allowed or blocked based on status code - // Squid typically returns: - // - 200, 206, 304: Allowed/successful - // - 403: Forbidden (blocked by ACL) - // - 407: Proxy authentication required - // - 502, 503: Connection/upstream errors - statusCode := entry.Status - isAllowed := statusCode == "TCP_HIT/200" || statusCode == "TCP_MISS/200" || - statusCode == "TCP_REFRESH_MODIFIED/200" || statusCode == "TCP_IMS_HIT/304" || - strings.Contains(statusCode, "/200") || strings.Contains(statusCode, "/206") || - strings.Contains(statusCode, "/304") - - if isAllowed { - analysis.AllowedRequests++ - if !setutil.Contains(allowedDomainsSet, domain) { - allowedDomainsSet[domain] = struct { - }{} - analysis.AllowedDomains = append(analysis.AllowedDomains, domain) - } - } else { - analysis.BlockedRequests++ - if !setutil.Contains(blockedDomainsSet, domain) { - blockedDomainsSet[domain] = struct { - }{} - analysis.BlockedDomains = append(analysis.BlockedDomains, domain) - } + if err := processSquidLogLine(scanner.Text(), analysis, allowedDomainsSet, blockedDomainsSet, verbose); err != nil { + return err } } - if err := scanner.Err(); err != nil { - return nil, fmt.Errorf("error reading access log: %w", err) + return fmt.Errorf("error reading access log: %w", err) } + return nil +} - // Sort domains for consistent output - sort.Strings(analysis.AllowedDomains) - sort.Strings(analysis.BlockedDomains) +func processSquidLogLine(rawLine string, analysis *DomainAnalysis, allowedDomainsSet, blockedDomainsSet map[string]struct{}, verbose bool) error { + line := strings.TrimSpace(rawLine) + if line == "" || strings.HasPrefix(line, "#") { + return nil + } + entry, err := parseSquidLogLine(line) + if err != nil { + if verbose { + fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Failed to parse log line: %v", err))) + } + return nil + } + analysis.TotalRequests++ + domain := stringutil.ExtractDomainFromURL(entry.URL) + if domain == "" { + return nil + } + if isAllowedSquidStatus(entry.Status) { + analysis.AllowedRequests++ + appendUniqueDomain(&analysis.AllowedDomains, allowedDomainsSet, domain) + return nil + } + analysis.BlockedRequests++ + appendUniqueDomain(&analysis.BlockedDomains, blockedDomainsSet, domain) + return nil +} - accessLogLog.Printf("Parsed access log: total_requests=%d, allowed=%d, blocked=%d, unique_allowed_domains=%d, unique_blocked_domains=%d", - analysis.TotalRequests, analysis.AllowedRequests, analysis.BlockedRequests, len(analysis.AllowedDomains), len(analysis.BlockedDomains)) +func isAllowedSquidStatus(statusCode string) bool { + return statusCode == "TCP_HIT/200" || statusCode == "TCP_MISS/200" || + statusCode == "TCP_REFRESH_MODIFIED/200" || statusCode == "TCP_IMS_HIT/304" || + strings.Contains(statusCode, "/200") || strings.Contains(statusCode, "/206") || strings.Contains(statusCode, "/304") +} - return analysis, nil +func appendUniqueDomain(domains *[]string, seen map[string]struct{}, domain string) { + if setutil.Contains(seen, domain) { + return + } + seen[domain] = struct{}{} + *domains = append(*domains, domain) } // parseSquidLogLine parses a single squid access log line diff --git a/pkg/cli/actions_build_command.go b/pkg/cli/actions_build_command.go index 4c8f193eec7..1c12f29448f 100644 --- a/pkg/cli/actions_build_command.go +++ b/pkg/cli/actions_build_command.go @@ -201,55 +201,66 @@ func validateActionYml(actionPath string) error { // buildAction builds a single action by bundling its dependencies func buildAction(actionsDir, actionName string) error { actionsBuildLog.Printf("Building action: %s", actionName) - fmt.Fprintln(os.Stderr, console.FormatInfoMessage("\n📦 Building action: "+actionName)) - actionPath := filepath.Join(actionsDir, actionName) - - // Validate action.yml - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(" ✓ Validating action.yml")) - if err := validateActionYml(actionPath); err != nil { + if err := validateBuiltAction(actionPath); err != nil { return err } - - // Special handling for setup: build shell script with embedded files if actionName == "setup" { return buildSetupAction(actionsDir, actionName) } - - // Check if this is a composite action (doesn't need JavaScript bundling) isComposite, err := isCompositeAction(actionPath) if err != nil { return fmt.Errorf("failed to check action type: %w", err) } - if isComposite { fmt.Fprintln(os.Stderr, console.FormatInfoMessage(" ✓ Composite action - no JavaScript bundling needed")) return nil } + return bundleJavaScriptAction(actionPath, actionName) +} + +func validateBuiltAction(actionPath string) error { + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(" ✓ Validating action.yml")) + return validateActionYml(actionPath) +} + +func bundleJavaScriptAction(actionPath, actionName string) error { + outputPath, sourceContent, err := readActionSourceFile(actionPath) + if err != nil { + return err + } + dependencies := getActionDependencies(actionName) + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf(" ✓ Found %d dependencies", len(dependencies)))) + files := collectActionDependencyFiles(dependencies) + outputContent, err := buildBundledActionSource(sourceContent, files) + if err != nil { + return err + } + if err := os.WriteFile(outputPath, []byte(outputContent), constants.FilePermSensitive); err != nil { + return fmt.Errorf("failed to write output file: %w", err) + } + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(" ✓ Built "+outputPath)) + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf(" ✓ Embedded %d files", len(files)))) + return nil +} +func readActionSourceFile(actionPath string) (string, []byte, error) { srcPath := filepath.Join(actionPath, "src", "index.js") outputPath := filepath.Join(actionPath, "index.js") - - // Check if source file exists if _, err := os.Stat(srcPath); os.IsNotExist(err) { - return fmt.Errorf("source file not found: %s", srcPath) + return "", nil, fmt.Errorf("source file not found: %s", srcPath) } - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(" ✓ Reading source file")) sourceContent, err := os.ReadFile(srcPath) if err != nil { - return fmt.Errorf("failed to read source file: %w", err) + return "", nil, fmt.Errorf("failed to read source file: %w", err) } + return outputPath, sourceContent, nil +} - // Get dependencies for this action - dependencies := getActionDependencies(actionName) - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf(" ✓ Found %d dependencies", len(dependencies)))) - - // Get all JavaScript sources +func collectActionDependencyFiles(dependencies []string) map[string]string { sources := workflow.GetJavaScriptSources() - - // Read dependency files files := make(map[string]string) for _, dep := range dependencies { if content, ok := sources[dep]; ok { @@ -259,30 +270,17 @@ func buildAction(actionsDir, actionName string) error { fmt.Fprintln(os.Stderr, console.FormatWarningMessage(" ⚠ Warning: Could not find "+dep)) } } + return files +} - // Generate FILES object with embedded content +func buildBundledActionSource(sourceContent []byte, files map[string]string) (string, error) { filesJSON, err := json.MarshalIndent(files, "", " ") if err != nil { - return fmt.Errorf("failed to marshal files: %w", err) + return "", fmt.Errorf("failed to marshal files: %w", err) } - - // Indent the JSON for proper embedding indentedJSON := strings.ReplaceAll(string(filesJSON), "\n", "\n ") indentedJSON = " " + strings.TrimPrefix(indentedJSON, " ") - - // Replace the FILES placeholder in source - // Match: const FILES = { ... }; - outputContent := filesConstPattern.ReplaceAllString(string(sourceContent), fmt.Sprintf("const FILES = %s;", strings.TrimSpace(indentedJSON))) - - // Write output file with restrictive permissions (0600 for security) - if err := os.WriteFile(outputPath, []byte(outputContent), constants.FilePermSensitive); err != nil { - return fmt.Errorf("failed to write output file: %w", err) - } - - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(" ✓ Built "+outputPath)) - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf(" ✓ Embedded %d files", len(files)))) - - return nil + return filesConstPattern.ReplaceAllString(string(sourceContent), fmt.Sprintf("const FILES = %s;", strings.TrimSpace(indentedJSON))), nil } // isCompositeAction checks if an action uses the 'composite' runtime diff --git a/pkg/cli/add_interactive_engine.go b/pkg/cli/add_interactive_engine.go index 194f5873555..90e06020272 100644 --- a/pkg/cli/add_interactive_engine.go +++ b/pkg/cli/add_interactive_engine.go @@ -1,6 +1,7 @@ package cli import ( + "context" "errors" "fmt" "os" @@ -18,193 +19,186 @@ import ( // selectAIEngineAndKey prompts the user to select an AI engine and provide API key func (c *AddInteractiveConfig) selectAIEngineAndKey() error { addInteractiveLog.Print("Starting coding agent selection") - - // First, check which secrets already exist in the repository if err := c.checkExistingSecrets(); err != nil { return err } + workflowSpecifiedEngine := c.workflowSpecifiedEngine() + defaultEngine := c.defaultInteractiveEngine(workflowSpecifiedEngine) + if c.EngineOverride != "" { + fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Using coding agent: "+c.EngineOverride)) + return c.configureEngineAPISecret(c.EngineOverride) + } + c.printWorkflowSpecifiedEngine(workflowSpecifiedEngine) + selectedEngine, err := c.promptForInteractiveEngine(defaultEngine, workflowSpecifiedEngine) + if err != nil { + return err + } + c.EngineOverride = selectedEngine + fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Selected engine: "+selectedEngine)) + return c.configureEngineAPISecret(selectedEngine) +} - // Determine default engine based on existing secrets, workflow preference, then environment - // Priority order: flag override > existing secrets > workflow frontmatter > environment > default - defaultEngine := string(constants.DefaultEngine) - workflowSpecifiedEngine := "" - - // Check if workflow specifies a preferred engine in frontmatter - if c.resolvedWorkflows != nil && len(c.resolvedWorkflows.Workflows) > 0 { - for _, wf := range c.resolvedWorkflows.Workflows { - if wf.Engine != "" { - workflowSpecifiedEngine = wf.Engine - addInteractiveLog.Printf("Workflow specifies engine in frontmatter: %s", wf.Engine) - break - } +func (c *AddInteractiveConfig) workflowSpecifiedEngine() string { + if c.resolvedWorkflows == nil || len(c.resolvedWorkflows.Workflows) == 0 { + return "" + } + for _, wf := range c.resolvedWorkflows.Workflows { + if wf.Engine != "" { + addInteractiveLog.Printf("Workflow specifies engine in frontmatter: %s", wf.Engine) + return wf.Engine } } + return "" +} - // If engine is explicitly overridden via flag, use that +func (c *AddInteractiveConfig) defaultInteractiveEngine(workflowSpecifiedEngine string) string { if c.EngineOverride != "" { - defaultEngine = c.EngineOverride - } else { - // Priority 1: Check existing repository secrets using EngineOptions - // This takes precedence over workflow preference since users should use what's already available - for _, opt := range constants.EngineOptions { - if setutil.Contains(c.existingSecrets, opt.SecretName) { - defaultEngine = opt.Value - addInteractiveLog.Printf("Found existing secret %s, recommending engine: %s", opt.SecretName, opt.Value) - break - } - } - - // Priority 2: If no existing secret found, use workflow frontmatter preference - if defaultEngine == string(constants.DefaultEngine) && workflowSpecifiedEngine != "" { - defaultEngine = workflowSpecifiedEngine - } + return c.EngineOverride + } + if engine := c.defaultEngineFromSecrets(); engine != "" { + return engine + } + if workflowSpecifiedEngine != "" { + return workflowSpecifiedEngine + } + if engine := defaultEngineFromEnvironment(); engine != "" { + return engine + } + return string(constants.DefaultEngine) +} - // Priority 3: Check environment variables if no existing secret or workflow preference found - if defaultEngine == string(constants.DefaultEngine) && workflowSpecifiedEngine == "" { - for _, opt := range constants.EngineOptions { - envVar := opt.SecretName - if opt.EnvVarName != "" { - envVar = opt.EnvVarName - } - if lookupEnv(envVar) != "" { - defaultEngine = opt.Value - addInteractiveLog.Printf("Found env var %s, recommending engine: %s", envVar, opt.Value) - break - } - } +func (c *AddInteractiveConfig) defaultEngineFromSecrets() string { + for _, opt := range constants.EngineOptions { + if setutil.Contains(c.existingSecrets, opt.SecretName) { + addInteractiveLog.Printf("Found existing secret %s, recommending engine: %s", opt.SecretName, opt.Value) + return opt.Value } } + return "" +} - // If engine is already overridden, skip selection - if c.EngineOverride != "" { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Using coding agent: "+c.EngineOverride)) - return c.configureEngineAPISecret(c.EngineOverride) +func defaultEngineFromEnvironment() string { + for _, opt := range constants.EngineOptions { + envVar := opt.SecretName + if opt.EnvVarName != "" { + envVar = opt.EnvVarName + } + if lookupEnv(envVar) != "" { + addInteractiveLog.Printf("Found env var %s, recommending engine: %s", envVar, opt.Value) + return opt.Value + } } + return "" +} - // Inform user if workflow specifies an engine +func (c *AddInteractiveConfig) printWorkflowSpecifiedEngine(workflowSpecifiedEngine string) { if workflowSpecifiedEngine != "" { fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Workflow specifies engine: "+workflowSpecifiedEngine)) } +} + +func (c *AddInteractiveConfig) promptForInteractiveEngine(defaultEngine, workflowSpecifiedEngine string) (string, error) { + engineOptions := reorderEngineOptions(buildInteractiveEngineOptions(c.existingSecrets, workflowSpecifiedEngine), defaultEngine) + var selectedEngine string + fmt.Fprintln(os.Stderr, "") + form := console.NewSelectForm(huh.NewSelect[string]().Title("Which coding agent would you like to use?").Description("This determines which coding agent processes your workflows").Options(engineOptions...).Value(&selectedEngine)) + if err := form.RunWithContext(c.Ctx); err != nil { + return "", fmt.Errorf("failed to select coding agent: %w", err) + } + return selectedEngine, nil +} - // Build engine options with notes about existing secrets and workflow specification. - // The list of engines is derived from the catalog to ensure all registered engines appear. +func buildInteractiveEngineOptions(existingSecrets map[string]struct{}, workflowSpecifiedEngine string) []huh.Option[string] { catalog := workflow.NewEngineCatalog(workflow.NewEngineRegistry()) - engineOptions := sliceutil.Map(catalog.All(), func(def *workflow.EngineDefinition) huh.Option[string] { - opt := constants.GetEngineOption(def.ID) + return sliceutil.Map(catalog.All(), func(def *workflow.EngineDefinition) huh.Option[string] { label := fmt.Sprintf("%s - %s", def.DisplayName, def.Description) - // Add markers for secret availability and workflow specification. - // opt may be nil for catalog engines not yet represented in EngineOptions; - // in that case we conservatively show '[no secret]'. - if opt != nil && setutil.Contains(c.existingSecrets, opt.SecretName) { - label += " [secret exists]" - } else { - label += " [no secret]" - } + label += interactiveEngineSecretMarker(existingSecrets, def.ID) if def.ID == workflowSpecifiedEngine { label += " [specified in workflow]" } return huh.NewOption(label, def.ID) }) +} - var selectedEngine string +func interactiveEngineSecretMarker(existingSecrets map[string]struct{}, engineID string) string { + opt := constants.GetEngineOption(engineID) + if opt != nil && setutil.Contains(existingSecrets, opt.SecretName) { + return " [secret exists]" + } + return " [no secret]" +} - // Set the default selection by moving it to front +func reorderEngineOptions(engineOptions []huh.Option[string], defaultEngine string) []huh.Option[string] { for i, opt := range engineOptions { - if opt.Value == defaultEngine { - if i > 0 { - engineOptions[0], engineOptions[i] = engineOptions[i], engineOptions[0] - } + if opt.Value == defaultEngine && i > 0 { + engineOptions[0], engineOptions[i] = engineOptions[i], engineOptions[0] break } } - - fmt.Fprintln(os.Stderr, "") - form := console.NewSelectForm( - huh.NewSelect[string](). - Title("Which coding agent would you like to use?"). - Description("This determines which coding agent processes your workflows"). - Options(engineOptions...). - Value(&selectedEngine), - ) - - if err := form.RunWithContext(c.Ctx); err != nil { - return fmt.Errorf("failed to select coding agent: %w", err) - } - - c.EngineOverride = selectedEngine - fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Selected engine: "+selectedEngine)) - - return c.configureEngineAPISecret(selectedEngine) + return engineOptions } +// configureEngineAPISecret collects the API key for the selected engine using the unified engine secrets functions // configureEngineAPISecret collects the API key for the selected engine using the unified engine secrets functions func (c *AddInteractiveConfig) configureEngineAPISecret(engine string) error { addInteractiveLog.Printf("Collecting API key for engine: %s", engine) - - // If --no-secret flag is set, skip secrets configuration entirely. - // Note: for Copilot workflows, --no-secret implies the PAT path; users who want - // copilot-requests (org billing) should not pass --no-secret. if c.SkipSecret { - opt := constants.GetEngineOption(engine) - if opt != nil { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Skipping %s secret setup (--no-secret flag set).", opt.SecretName))) - } else { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Skipping secret setup (--no-secret flag set).")) - } + printSkippedEngineSecretSetup(engine) return nil } - - // For Copilot, ask the user whether to use copilot-requests (org billing) or a PAT. - // Only prompt when an interactive context is available (wizard path); default to PAT otherwise. - if engine == string(constants.CopilotEngine) && c.Ctx != nil { - if err := c.selectCopilotAuthMethod(); err != nil { - return err - } - if c.UseCopilotRequests { - return nil - } + if err := c.maybeConfigureCopilotAuth(engine); err != nil || c.UseCopilotRequests { + return err } - - // If user doesn't have write access, skip secrets configuration. - // Users without write access cannot configure repository secrets. if !c.hasWriteAccess { - opt := constants.GetEngineOption(engine) - if opt != nil { - fmt.Fprintln(os.Stderr, "") - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Skipping %s secret setup — write access is required to configure repository secrets.", opt.SecretName))) - fmt.Fprintln(os.Stderr, "") - fmt.Fprintln(os.Stderr, "Once you have write access or an admin configures the repository, set the secret with:") - fmt.Fprintln(os.Stderr, console.FormatCommandMessage(fmt.Sprintf(" gh aw secrets set %s --repo %s", opt.SecretName, c.RepoOverride))) - } + printEngineSecretWriteAccessMessage(engine, c.RepoOverride) return nil } + return c.ensureEngineSecretConfigured(engine) +} - // Use the unified checkAndEnsureEngineSecrets function - config := EngineSecretConfig{ - Ctx: c.Ctx, - RepoSlug: c.RepoOverride, - Engine: engine, - Verbose: c.Verbose, - ExistingSecrets: c.existingSecrets, - IncludeSystemSecrets: false, // Don't include system secrets in add-wizard - IncludeOptional: false, +func printSkippedEngineSecretSetup(engine string) { + if opt := constants.GetEngineOption(engine); opt != nil { + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Skipping %s secret setup (--no-secret flag set).", opt.SecretName))) + return } + fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Skipping secret setup (--no-secret flag set).")) +} - if err := checkAndEnsureEngineSecretsForEngine(config); err != nil { +func (c *AddInteractiveConfig) maybeConfigureCopilotAuth(engine string) error { + if engine != string(constants.CopilotEngine) || c.Ctx == nil { + return nil + } + if err := c.selectCopilotAuthMethod(); err != nil { return err } + return nil +} - // Update existingSecrets to reflect that the secret was uploaded - // This prevents duplicate secret uploads in createWorkflowPRAndConfigureSecret later +func printEngineSecretWriteAccessMessage(engine, repoOverride string) { opt := constants.GetEngineOption(engine) - if opt != nil { + if opt == nil { + return + } + fmt.Fprintln(os.Stderr, "") + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(fmt.Sprintf("Skipping %s secret setup — write access is required to configure repository secrets.", opt.SecretName))) + fmt.Fprintln(os.Stderr, "") + fmt.Fprintln(os.Stderr, "Once you have write access or an admin configures the repository, set the secret with:") + fmt.Fprintln(os.Stderr, console.FormatCommandMessage(fmt.Sprintf(" gh aw secrets set %s --repo %s", opt.SecretName, repoOverride))) +} + +func (c *AddInteractiveConfig) ensureEngineSecretConfigured(engine string) error { + config := EngineSecretConfig{Ctx: c.Ctx, RepoSlug: c.RepoOverride, Engine: engine, Verbose: c.Verbose, ExistingSecrets: c.existingSecrets, IncludeSystemSecrets: false, IncludeOptional: false} + if err := checkAndEnsureEngineSecretsForEngine(config); err != nil { + return err + } + if opt := constants.GetEngineOption(engine); opt != nil { c.existingSecrets[opt.SecretName] = struct{}{} addInteractiveLog.Printf("Updated existingSecrets to include %s after upload", opt.SecretName) } - return nil } +// authMethodCopilotRequests is the wizard option value for Copilot org-billing authentication // authMethodCopilotRequests is the wizard option value for Copilot org-billing authentication // (permissions.copilot-requests: write). Extracted as a package-level constant so both the // form definition and applyCopilotAuthMethodChoice reference the same sentinel. @@ -215,56 +209,45 @@ const authMethodCopilotRequests = "copilot-requests" // Sets c.UseCopilotRequests when org billing is chosen. func (c *AddInteractiveConfig) selectCopilotAuthMethod() error { addInteractiveLog.Print("Prompting user for Copilot authentication method") + probe, copilotRequestsLabel := c.probeCopilotAuthMethod() + options := buildCopilotAuthOptions(probe, copilotRequestsLabel) + if probe.InfoNote != "" { + fmt.Fprintln(os.Stderr, console.FormatInfoMessage(probe.InfoNote)) + } + fmt.Fprintln(os.Stderr, "") + authMethod, err := runCopilotAuthMethodForm(c.Ctx, probe, options) + if err != nil { + return err + } + c.applyCopilotAuthMethodChoice(authMethod) + return nil +} - const authMethodPAT = "pat" - - // Detect org Copilot CLI billing status before building the form. - // c.RepoOverride is in "owner/repo" format; we need just the org login. - // When no org login is available the result is inconclusive (same as a - // non-200 response) so the user still sees the info note. +func (c *AddInteractiveConfig) probeCopilotAuthMethod() (orgCopilotBillingProbeResult, string) { copilotRequestsLabel := "Use copilot-requests (org's Copilot billing, no PAT)" - - var probe orgCopilotBillingProbeResult + probe := orgCopilotBillingProbeResult{InfoNote: copilotBillingInconclusiveNote} if orgLogin, _, found := strings.Cut(c.RepoOverride, "/"); found && orgLogin != "" { probe = probeCopilotBillingForOrg(c.Ctx, orgLogin) - } else { - probe = orgCopilotBillingProbeResult{ - InfoNote: copilotBillingInconclusiveNote, - } } c.copilotCLIBillingStatus = probe.BillingStatus - copilotRequestsLabel += probe.LabelSuffix - if probe.InfoNote != "" { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage(probe.InfoNote)) - } - - fmt.Fprintln(os.Stderr, "") + return probe, copilotRequestsLabel + probe.LabelSuffix +} - // Build select options. - // When billing is confirmed enabled, copilot-requests is listed first (pre-selected). - // When billing is disabled or inconclusive, PAT is listed first (default selection). - // The copilot-requests option is always shown; when disabled a validation guard - // prevents it from being submitted. +func buildCopilotAuthOptions(probe orgCopilotBillingProbeResult, copilotRequestsLabel string) []huh.Option[string] { + const authMethodPAT = "pat" patOpt := huh.NewOption("Use a Personal Access Token (PAT) as COPILOT_GITHUB_TOKEN", authMethodPAT) copilotRequestsOpt := huh.NewOption(copilotRequestsLabel, authMethodCopilotRequests) - - var options []huh.Option[string] - switch probe.BillingStatus { - case "enabled": - // copilot-requests pre-selected - options = []huh.Option[string]{copilotRequestsOpt.Selected(true), patOpt} - default: - // PAT is default (first) for disabled or inconclusive - options = []huh.Option[string]{patOpt.Selected(true), copilotRequestsOpt} + if probe.BillingStatus == "enabled" { + return []huh.Option[string]{copilotRequestsOpt.Selected(true), patOpt} } + return []huh.Option[string]{patOpt.Selected(true), copilotRequestsOpt} +} +func runCopilotAuthMethodForm(ctx context.Context, probe orgCopilotBillingProbeResult, options []huh.Option[string]) (string, error) { var authMethod string - selectField := huh.NewSelect[string](). - Title("How would you like Copilot workflows to authenticate?"). - Description("copilot-requests uses the org's Copilot billing seat — no PAT required.\nPAT uses a fine-grained personal access token stored as COPILOT_GITHUB_TOKEN (requires repo write access to configure)."). - Options(options...). - Value(&authMethod) - + description := "copilot-requests uses the org's Copilot billing seat — no PAT required.\n" + + "PAT uses a fine-grained personal access token stored as COPILOT_GITHUB_TOKEN (requires repo write access to configure)." + selectField := huh.NewSelect[string]().Title("How would you like Copilot workflows to authenticate?").Description(description).Options(options...).Value(&authMethod) if probe.Disabled { selectField = selectField.Validate(func(v string) error { if v == authMethodCopilotRequests { @@ -273,15 +256,10 @@ func (c *AddInteractiveConfig) selectCopilotAuthMethod() error { return nil }) } - - form := console.NewSelectForm(selectField) - - if err := form.RunWithContext(c.Ctx); err != nil { - return fmt.Errorf("failed to select Copilot authentication method: %w", err) + if err := console.NewSelectForm(selectField).RunWithContext(ctx); err != nil { + return "", fmt.Errorf("failed to select Copilot authentication method: %w", err) } - - c.applyCopilotAuthMethodChoice(authMethod) - return nil + return authMethod, nil } // applyCopilotAuthMethodChoice records the user's Copilot auth method selection and prints diff --git a/pkg/cli/add_interactive_git.go b/pkg/cli/add_interactive_git.go index bb4ae9a995e..041d706a21b 100644 --- a/pkg/cli/add_interactive_git.go +++ b/pkg/cli/add_interactive_git.go @@ -25,175 +25,191 @@ func isAlreadyMergedGHError(err error) bool { // createWorkflowPRAndConfigureSecret creates the PR, merges it, and adds the secret func (c *AddInteractiveConfig) createWorkflowPRAndConfigureSecret(ctx context.Context, workflowFiles, initFiles []string, secretName, secretValue string) error { addInteractiveLog.Print("Applying changes") - fmt.Fprintln(os.Stderr, "") - - // Add the workflow using existing implementation with --create-pull-request - // Pass the resolved workflows to avoid re-fetching them - // Pass Quiet=true to suppress detailed output (already shown earlier in interactive mode) - // This returns the result including PR number and HasWorkflowDispatch - opts := AddOptions{ - Verbose: c.Verbose, - Quiet: true, - EngineOverride: c.EngineOverride, - Name: "", - Force: false, - AppendText: c.AppendText, - CreatePR: true, - NoGitattributes: c.NoGitattributes, - WorkflowDir: c.WorkflowDir, - NoStopAfter: c.NoStopAfter, - StopAfter: c.StopAfter, - DisableSecurityScanner: c.DisableSecurityScanner, - AddCopilotRequestsPermission: c.UseCopilotRequests, + result, err := c.addResolvedWorkflowsForPR(ctx) + if err != nil { + return err + } + if err := c.handleWorkflowPullRequest(result); err != nil { + return err } + return c.configureRepositorySecret(secretName, secretValue) +} + +func (c *AddInteractiveConfig) addResolvedWorkflowsForPR(ctx context.Context) (*AddWorkflowsResult, error) { + opts := AddOptions{Verbose: c.Verbose, Quiet: true, EngineOverride: c.EngineOverride, AppendText: c.AppendText, CreatePR: true, NoGitattributes: c.NoGitattributes, WorkflowDir: c.WorkflowDir, NoStopAfter: c.NoStopAfter, StopAfter: c.StopAfter, DisableSecurityScanner: c.DisableSecurityScanner, AddCopilotRequestsPermission: c.UseCopilotRequests} result, err := AddResolvedWorkflows(ctx, c.WorkflowSpecs, c.resolvedWorkflows, opts) if err != nil { - return fmt.Errorf("failed to add workflow: %w", err) + return nil, fmt.Errorf("failed to add workflow: %w", err) } c.addResult = result + return result, nil +} - // Step 8b: Optionally merge the PR – loop until merged, confirmed-merged, or user exits +func (c *AddInteractiveConfig) handleWorkflowPullRequest(result *AddWorkflowsResult) error { if result.PRNumber == 0 { - if result.PRURL == "" { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Requested workflow files already exist locally; no pull request was created.")) - return nil + return c.handleMissingPullRequest(result) + } + fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Pull request created: "+result.PRURL)) + fmt.Fprintln(os.Stderr, "") + return c.runPullRequestMergeLoop(result) +} + +func (c *AddInteractiveConfig) handleMissingPullRequest(result *AddWorkflowsResult) error { + if result.PRURL == "" { + fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Requested workflow files already exist locally; no pull request was created.")) + return nil + } + fmt.Fprintln(os.Stderr, console.FormatWarningMessage("Could not determine PR number")) + fmt.Fprintln(os.Stderr, "Please merge the PR manually from the GitHub web interface.") + return nil +} + +type mergeLoopState struct { + mergeDone bool + mergeFailed bool + userReviewing bool +} + +type mergeAction string + +const ( + mergeActionAttempt mergeAction = "attempt" + mergeActionEditTitle mergeAction = "editTitle" + mergeActionReview mergeAction = "review" + mergeActionConfirmed mergeAction = "confirmed" + mergeActionExit mergeAction = "exit" +) + +func (c *AddInteractiveConfig) runPullRequestMergeLoop(result *AddWorkflowsResult) error { + state := &mergeLoopState{} + for !state.mergeDone { + action, err := promptMergeAction(result.PRURL, state) + if err != nil { + return fmt.Errorf("failed to get user input: %w", err) } - fmt.Fprintln(os.Stderr, console.FormatWarningMessage("Could not determine PR number")) - fmt.Fprintln(os.Stderr, "Please merge the PR manually from the GitHub web interface.") - } else { - fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Pull request created: "+result.PRURL)) + if err := c.handleMergeAction(result, state, action); err != nil { + return err + } + } + return nil +} + +func promptMergeAction(prURL string, state *mergeLoopState) (mergeAction, error) { + var chosen mergeAction + selectForm := console.NewSelectForm(huh.NewSelect[mergeAction]().Title("What would you like to do with pull request " + prURL + "?").Options(buildMergeActionOptions(state)...).Value(&chosen)) + if err := selectForm.Run(); err != nil { + return "", err + } + return chosen, nil +} + +func buildMergeActionOptions(state *mergeLoopState) []huh.Option[mergeAction] { + options := []huh.Option[mergeAction]{huh.NewOption("Attempt to merge", mergeActionAttempt)} + if state.mergeFailed { + options = append(options, huh.NewOption("Edit PR title and retry", mergeActionEditTitle)) + } + if state.userReviewing { + options = append(options, huh.NewOption("PR has been manually merged", mergeActionConfirmed), huh.NewOption("Exit, I'm done here", mergeActionExit)) + return options + } + return append(options, huh.NewOption("I'll review/merge myself", mergeActionReview), huh.NewOption("Exit", mergeActionExit)) +} + +func (c *AddInteractiveConfig) handleMergeAction(result *AddWorkflowsResult, state *mergeLoopState, action mergeAction) error { + switch action { + case mergeActionAttempt: + return c.attemptMergePullRequest(result, state) + case mergeActionEditTitle: + return c.editPullRequestTitle(result.PRNumber, state) + case mergeActionReview: + state.userReviewing = true + fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Please review and merge the pull request: "+result.PRURL)) fmt.Fprintln(os.Stderr, "") + case mergeActionConfirmed: + fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Great – continuing with the merged pull request")) + state.mergeDone = true + case mergeActionExit: + fmt.Fprintln(os.Stderr, "") + fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Exiting. You can merge the pull request later: "+result.PRURL)) + return errors.New("user exited before PR was merged") + } + return nil +} - // mergeAction values used in the select loop - type mergeAction string - const ( - mergeActionAttempt mergeAction = "attempt" - mergeActionEditTitle mergeAction = "editTitle" - mergeActionReview mergeAction = "review" - mergeActionConfirmed mergeAction = "confirmed" - mergeActionExit mergeAction = "exit" - ) - - mergeDone := false // true when the PR is merged (or confirmed merged) - mergeFailed := false // true after an unsuccessful merge attempt - userReviewing := false // true after the user chose "I'll review myself" - - for !mergeDone { - // Build option list based on current state - var options []huh.Option[mergeAction] - - options = append(options, huh.NewOption("Attempt to merge", mergeActionAttempt)) - - if mergeFailed { - options = append(options, huh.NewOption("Edit PR title and retry", mergeActionEditTitle)) - } - - if userReviewing { - options = append(options, huh.NewOption("PR has been manually merged", mergeActionConfirmed)) - } else { - options = append(options, huh.NewOption("I'll review/merge myself", mergeActionReview)) - } - - if userReviewing { - options = append(options, huh.NewOption("Exit, I'm done here", mergeActionExit)) - } else { - options = append(options, huh.NewOption("Exit", mergeActionExit)) - } - - var chosen mergeAction - selectForm := console.NewSelectForm( - huh.NewSelect[mergeAction](). - Title("What would you like to do with pull request " + result.PRURL + "?"). - Options(options...). - Value(&chosen), - ) - - if selectErr := selectForm.Run(); selectErr != nil { - return fmt.Errorf("failed to get user input: %w", selectErr) - } - - switch chosen { - case mergeActionAttempt: - if mergeErr := c.mergePullRequest(result.PRNumber); mergeErr != nil { - if isAlreadyMergedGHError(mergeErr) { - fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Merged pull request "+result.PRURL)) - mergeDone = true - } else { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Failed to merge PR: %v", mergeErr))) - if mergeFailed { - fmt.Fprintln(os.Stderr, "Please merge the PR manually: "+result.PRURL) - } - mergeFailed = true - } - } else { - fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Merged pull request "+result.PRURL)) - mergeDone = true - } - - case mergeActionEditTitle: - var newTitle string - titleForm := console.NewInputForm( - huh.NewInput(). - Title("Enter new PR title"). - Description("Add a prefix if required, for example: feat: or fix:"). - Value(&newTitle), - ) - if titleErr := titleForm.Run(); titleErr != nil { - return fmt.Errorf("failed to get user input: %w", titleErr) - } - newTitle = strings.TrimSpace(newTitle) - if newTitle == "" { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage("PR title cannot be empty, keeping current title")) - } else if editErr := editPRTitle(result.PRNumber, newTitle, c.RepoOverride); editErr != nil { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Failed to update PR title: %v", editErr))) - } else { - fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("PR title updated to: "+newTitle)) - mergeFailed = false - } - - case mergeActionReview: - userReviewing = true - fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Please review and merge the pull request: "+result.PRURL)) - fmt.Fprintln(os.Stderr, "") - - case mergeActionConfirmed: - fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Great – continuing with the merged pull request")) - mergeDone = true - - case mergeActionExit: - fmt.Fprintln(os.Stderr, "") - fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Exiting. You can merge the pull request later: "+result.PRURL)) - return errors.New("user exited before PR was merged") - } +func (c *AddInteractiveConfig) attemptMergePullRequest(result *AddWorkflowsResult, state *mergeLoopState) error { + if err := c.mergePullRequest(result.PRNumber); err != nil { + if isAlreadyMergedGHError(err) { + fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Merged pull request "+result.PRURL)) + state.mergeDone = true + return nil + } + fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Failed to merge PR: %v", err))) + if state.mergeFailed { + fmt.Fprintln(os.Stderr, "Please merge the PR manually: "+result.PRURL) } + state.mergeFailed = true + return nil } + fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Merged pull request "+result.PRURL)) + state.mergeDone = true + return nil +} - // Step 8c: Add the secret (skip if no secret configured or already exists in repository) - if secretName == "" { - // No secret to configure (e.g., user doesn't have write access to the repository) - } else if secretValue == "" { - // Secret already exists in repo, nothing to do +func (c *AddInteractiveConfig) editPullRequestTitle(prNumber int, state *mergeLoopState) error { + newTitle, err := promptForPullRequestTitle() + if err != nil { + return fmt.Errorf("failed to get user input: %w", err) + } + if newTitle == "" { + fmt.Fprintln(os.Stderr, console.FormatWarningMessage("PR title cannot be empty, keeping current title")) + return nil + } + if err := editPRTitle(prNumber, newTitle, c.RepoOverride); err != nil { + fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Failed to update PR title: %v", err))) + return nil + } + fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("PR title updated to: "+newTitle)) + state.mergeFailed = false + return nil +} + +func promptForPullRequestTitle() (string, error) { + var newTitle string + titleForm := console.NewInputForm(huh.NewInput().Title("Enter new PR title").Description("Add a prefix if required, for example: feat: or fix:").Value(&newTitle)) + if err := titleForm.Run(); err != nil { + return "", err + } + return strings.TrimSpace(newTitle), nil +} + +func (c *AddInteractiveConfig) configureRepositorySecret(secretName, secretValue string) error { + switch { + case secretName == "": + return nil + case secretValue == "": if c.Verbose { fmt.Fprintln(os.Stderr, "") fmt.Fprintln(os.Stderr, console.FormatSuccessMessage(fmt.Sprintf("Secret '%s' already configured", secretName))) } - } else { - fmt.Fprintln(os.Stderr, "") - fmt.Fprintln(os.Stderr, console.FormatProgressMessage(fmt.Sprintf("Adding secret '%s' to repository...", secretName))) - - if err := c.addRepositorySecret(secretName, secretValue); err != nil { - fmt.Fprintln(os.Stderr, console.FormatErrorMessage(fmt.Sprintf("Failed to add secret: %v", err))) - fmt.Fprintln(os.Stderr, "") - fmt.Fprintln(os.Stderr, "Please add the secret manually:") - fmt.Fprintln(os.Stderr, " 1. Go to your repository Settings → Secrets and variables → Actions") - fmt.Fprintf(os.Stderr, " 2. Click 'New repository secret' and add '%s'\n", secretName) - return fmt.Errorf("failed to add secret: %w", err) - } - - fmt.Fprintln(os.Stderr, console.FormatSuccessMessage(fmt.Sprintf("Secret '%s' added", secretName))) + return nil + default: + return c.addConfiguredRepositorySecret(secretName, secretValue) } +} +func (c *AddInteractiveConfig) addConfiguredRepositorySecret(secretName, secretValue string) error { + fmt.Fprintln(os.Stderr, "") + fmt.Fprintln(os.Stderr, console.FormatProgressMessage(fmt.Sprintf("Adding secret '%s' to repository...", secretName))) + if err := c.addRepositorySecret(secretName, secretValue); err != nil { + fmt.Fprintln(os.Stderr, console.FormatErrorMessage(fmt.Sprintf("Failed to add secret: %v", err))) + fmt.Fprintln(os.Stderr, "") + fmt.Fprintln(os.Stderr, "Please add the secret manually:") + fmt.Fprintln(os.Stderr, " 1. Go to your repository Settings → Secrets and variables → Actions") + fmt.Fprintf(os.Stderr, " 2. Click 'New repository secret' and add '%s'\n", secretName) + return fmt.Errorf("failed to add secret: %w", err) + } + fmt.Fprintln(os.Stderr, console.FormatSuccessMessage(fmt.Sprintf("Secret '%s' added", secretName))) return nil } @@ -202,70 +218,81 @@ func (c *AddInteractiveConfig) createWorkflowPRAndConfigureSecret(ctx context.Co // the merged workflow files, which are required when offering to run the workflow. func (c *AddInteractiveConfig) updateLocalBranch() error { addInteractiveLog.Print("Updating local branch with merged changes") + defaultBranch := c.detectDefaultBranch() + addInteractiveLog.Printf("Default branch: %s", defaultBranch) + if err := c.fetchDefaultBranch(defaultBranch); err != nil { + return err + } + if err := c.switchToDefaultBranch(defaultBranch); err != nil { + return err + } + if err := c.pullDefaultBranch(defaultBranch); err != nil { + return err + } + if c.Verbose { + fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Local branch updated with merged changes")) + } + return nil +} - // Get the default branch name using gh +func (c *AddInteractiveConfig) detectDefaultBranch() string { output, err := workflow.RunGHCombined("Getting default branch...", "repo", "view", "--repo", c.RepoOverride, "--json", "defaultBranchRef", "--jq", ".defaultBranchRef.name") - defaultBranch := "" if err == nil { - defaultBranch = strings.TrimSpace(string(output)) - } - - // Fallback: query the local origin remote directly (works even when gh repo - // view fails, e.g. forks without a default remote set). - if defaultBranch == "" { - addInteractiveLog.Print("gh repo view failed, trying git ls-remote to detect default branch") - lsCmd := exec.Command("git", "ls-remote", "--symref", "origin", "HEAD") - lsOutput, lsErr := lsCmd.CombinedOutput() - if lsErr == nil { - defaultBranch = parseDefaultBranchFromLsRemote(string(lsOutput)) + if defaultBranch := strings.TrimSpace(string(output)); defaultBranch != "" { + return defaultBranch } } + return fallbackDefaultBranch() +} - if defaultBranch == "" { - defaultBranch = "main" +func fallbackDefaultBranch() string { + addInteractiveLog.Print("gh repo view failed, trying git ls-remote to detect default branch") + cmd := exec.Command("git", "ls-remote", "--symref", "origin", "HEAD") + output, err := cmd.CombinedOutput() + if err == nil { + if defaultBranch := parseDefaultBranchFromLsRemote(string(output)); defaultBranch != "" { + return defaultBranch + } } - addInteractiveLog.Printf("Default branch: %s", defaultBranch) + return "main" +} - // Fetch the latest changes from origin +func (c *AddInteractiveConfig) fetchDefaultBranch(defaultBranch string) error { if c.Verbose { fmt.Fprintln(os.Stderr, console.FormatProgressMessage("Fetching latest changes from GitHub...")) } - - fetchCmd := exec.Command("git", "fetch", "origin", defaultBranch) - fetchOutput, err := fetchCmd.CombinedOutput() - if err != nil { - return fmt.Errorf("git fetch failed: %w (output: %s)", err, string(fetchOutput)) + cmd := exec.Command("git", "fetch", "origin", defaultBranch) + if output, err := cmd.CombinedOutput(); err != nil { + return fmt.Errorf("git fetch failed: %w (output: %s)", err, string(output)) } + return nil +} - // Switch to the default branch so the working tree contains the merged workflow - // files. Without this, users on a feature branch won't have the files locally and - // the subsequent "run workflow" step will fail with "workflow file not found". +func (c *AddInteractiveConfig) switchToDefaultBranch(defaultBranch string) error { currentBranch, err := getCurrentBranch() if err != nil { addInteractiveLog.Printf("Could not determine current branch: %v", err) currentBranch = "" } - - if currentBranch != defaultBranch { - addInteractiveLog.Printf("Switching from %q to default branch %q", currentBranch, defaultBranch) - if err := switchBranch(defaultBranch, c.Verbose); err != nil { - return fmt.Errorf("failed to switch to default branch %s: %w", defaultBranch, err) - } + if currentBranch == defaultBranch { + return nil } - - pullCmd := exec.Command("git", "pull", "origin", defaultBranch) - pullOutput, err := pullCmd.CombinedOutput() - if err != nil { - return fmt.Errorf("git pull failed: %w (output: %s)", err, string(pullOutput)) + addInteractiveLog.Printf("Switching from %q to default branch %q", currentBranch, defaultBranch) + if err := switchBranch(defaultBranch, c.Verbose); err != nil { + return fmt.Errorf("failed to switch to default branch %s: %w", defaultBranch, err) } + return nil +} - if c.Verbose { - fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Local branch updated with merged changes")) +func (c *AddInteractiveConfig) pullDefaultBranch(defaultBranch string) error { + cmd := exec.Command("git", "pull", "origin", defaultBranch) + if output, err := cmd.CombinedOutput(); err != nil { + return fmt.Errorf("git pull failed: %w (output: %s)", err, string(output)) } - return nil } +// checkCleanWorkingDirectory verifies the working directory has no uncommitted changes. // checkCleanWorkingDirectory verifies the working directory has no uncommitted changes. // This is checked early in the interactive flow to avoid failing later during PR creation. func (c *AddInteractiveConfig) checkCleanWorkingDirectory() error { diff --git a/pkg/cli/add_interactive_orchestrator.go b/pkg/cli/add_interactive_orchestrator.go index 4617933eec6..ce7cf9136af 100644 --- a/pkg/cli/add_interactive_orchestrator.go +++ b/pkg/cli/add_interactive_orchestrator.go @@ -70,130 +70,117 @@ type AddInteractiveConfig struct { // as it will be overwritten by the provided ctx. func RunAddInteractive(ctx context.Context, config *AddInteractiveConfig) error { addInteractiveLog.Print("Starting interactive add workflow") + if err := validateInteractiveAddEnvironment(); err != nil { + return err + } + config.Ctx = ctx + configureInteractiveGitHubHost(config.Verbose) + console.ShowWelcomeBanner("This tool will walk you through adding an automated workflow to your repository.") + bootstrapProfile, err := runInteractiveAddPreparation(config) + if err != nil { + return err + } + filesToAdd, initFiles, err := runInteractiveAddPlanning(config) + if err != nil { + return err + } + secretName, secretValue, err := config.resolveInteractiveSecret() + if err != nil { + return err + } + if err := config.confirmChanges(filesToAdd, initFiles, secretName, secretValue); err != nil { + return err + } + if err := config.createWorkflowPRAndConfigureSecret(ctx, filesToAdd, initFiles, secretName, secretValue); err != nil { + return err + } + if err := config.applyInteractiveBootstrap(ctx, bootstrapProfile); err != nil { + return err + } + return config.checkStatusAndOfferRun(ctx) +} - // Assert this function is not running in automated unit tests or CI. - // GO_TEST_MODE intentionally uses GetBoolFromEnv so common boolean spellings - // are treated consistently across test and automation environments, while - // IsRunningInCI centralizes the broader CI environment detection logic. +func validateInteractiveAddEnvironment() error { if envutil.GetBoolFromEnv("GO_TEST_MODE", false, addInteractiveLog) || IsRunningInCI() { return errors.New("interactive add cannot be used in automated tests or CI environments") } + return nil +} - // Set context on the config - config.Ctx = ctx - - // Auto-detect GHES host from git remote if not already set - if os.Getenv("GH_HOST") == "" { //nolint:osgetenvlibrary - detectedHost := getHostFromOriginRemote() - if detectedHost != "github.com" { - addInteractiveLog.Printf("Auto-detected GHES host from git remote: %s", detectedHost) - workflow.SetDefaultGHHost(detectedHost) - if config.Verbose { - fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Auto-detected GitHub Enterprise host: "+detectedHost)) - } - } +func configureInteractiveGitHubHost(verbose bool) { + if os.Getenv("GH_HOST") != "" { //nolint:osgetenvlibrary + return } + detectedHost := getHostFromOriginRemote() + if detectedHost == "github.com" { + return + } + addInteractiveLog.Printf("Auto-detected GHES host from git remote: %s", detectedHost) + workflow.SetDefaultGHHost(detectedHost) + if verbose { + fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Auto-detected GitHub Enterprise host: "+detectedHost)) + } +} - // Step 1: Welcome message - console.ShowWelcomeBanner("This tool will walk you through adding an automated workflow to your repository.") - - // Step 1b: Resolve workflows early to get descriptions and validate specs +func runInteractiveAddPreparation(config *AddInteractiveConfig) (*resolvedBootstrapProfile, error) { if err := config.resolveWorkflows(); err != nil { - return err + return nil, err } - - // Step 1c: Show workflow descriptions if available config.showWorkflowDescriptions() - - // Step 2: Check gh auth status if err := config.checkGHAuthStatus(); err != nil { - return err + return nil, err } - - // Step 3: Check git repository and get org/repo if err := config.checkGitRepository(); err != nil { - return err + return nil, err } - - // Step 3b: Check working directory is clean (must be clean for PR creation later) if err := config.checkCleanWorkingDirectory(); err != nil { - return err + return nil, err } - - // Step 4: Check GitHub Actions is enabled if err := config.checkActionsEnabled(); err != nil { - return err + return nil, err } - - // Step 5: Check user permissions if err := config.checkUserPermissions(); err != nil { - return err + return nil, err } - - var bootstrapProfile *resolvedBootstrapProfile - if config.resolvedWorkflows != nil { - bootstrapProfile = config.resolvedWorkflows.BootstrapProfile + if config.resolvedWorkflows == nil { + return nil, nil } - // All config steps run post-install in the exact order they are declared in the - // manifest. We no longer split them into a pre-install and post-install phase so - // that the declared ordering is preserved. - remainingBootstrapProfile := bootstrapProfile + return config.resolvedWorkflows.BootstrapProfile, nil +} - // Step 6: Select coding agent and collect API key +func runInteractiveAddPlanning(config *AddInteractiveConfig) ([]string, []string, error) { if err := config.selectAIEngineAndKey(); err != nil { - return err + return nil, nil, err } - initFiles, err := ensureAddRepositoryInitializedWithDetails(config.EngineOverride, config.Verbose, config.NoGitattributes) if err != nil { - return err + return nil, nil, err } - - // Step 7: Determine files to add filesToAdd, _, err := config.determineFilesToAdd() if err != nil { - return err + return nil, nil, err } - - // Step 7b: Offer schedule frequency selection for scheduled workflows if err := config.selectScheduleFrequency(); err != nil { - return err - } - - // Step 8: Confirm with user - var secretName, secretValue string - if config.hasWriteAccess && !config.SkipSecret && !config.UseCopilotRequests { - secretName, secretValue, err = config.resolveEngineApiKeyCredential() - if err != nil { - return err - } - } - - if err := config.confirmChanges(filesToAdd, initFiles, secretName, secretValue); err != nil { - return err + return nil, nil, err } + return filesToAdd, initFiles, nil +} - // Step 9: Apply changes (create PR, merge, add secret) - if err := config.createWorkflowPRAndConfigureSecret(ctx, filesToAdd, initFiles, secretName, secretValue); err != nil { - return err +func (c *AddInteractiveConfig) resolveInteractiveSecret() (string, string, error) { + if !c.hasWriteAccess || c.SkipSecret || c.UseCopilotRequests { + return "", "", nil } + return c.resolveEngineApiKeyCredential() +} - // Step 9b: Apply bootstrap config steps interactively (if the package declares any) - if remainingBootstrapProfile != nil { - if config.hasWriteAccess { - if err := executeBootstrapConfigForAdd(ctx, config.RepoOverride, config.WorkflowSpecs, remainingBootstrapProfile, config.UseCopilotRequests, config.Verbose); err != nil { - return err - } - } else { - printBootstrapConfigTODO(os.Stderr, remainingBootstrapProfile) - } +func (c *AddInteractiveConfig) applyInteractiveBootstrap(ctx context.Context, bootstrapProfile *resolvedBootstrapProfile) error { + if bootstrapProfile == nil { + return nil } - - // Step 10: Check status and offer to run - if err := config.checkStatusAndOfferRun(ctx); err != nil { - return err + if c.hasWriteAccess { + return executeBootstrapConfigForAdd(ctx, c.RepoOverride, c.WorkflowSpecs, bootstrapProfile, c.UseCopilotRequests, c.Verbose) } - + printBootstrapConfigTODO(os.Stderr, bootstrapProfile) return nil } diff --git a/pkg/cli/add_interactive_schedule.go b/pkg/cli/add_interactive_schedule.go index 5e3961d6625..7b538e9fd82 100644 --- a/pkg/cli/add_interactive_schedule.go +++ b/pkg/cli/add_interactive_schedule.go @@ -51,148 +51,156 @@ func detectWorkflowScheduleInfo(content string) scheduleDetection { if err != nil || result.Frontmatter == nil { return scheduleDetection{} } - onValue, exists := result.Frontmatter["on"] if !exists { return scheduleDetection{} } - - // Case 1: on is a simple string (e.g., "on: daily" or "on: 0 * * * *") if onStr, ok := onValue.(string); ok { - _, _, parseErr := parser.ParseSchedule(onStr) - if parseErr == nil { - return scheduleDetection{ - RawExpr: onStr, - Frequency: classifyScheduleFrequency(onStr), - IsUpdatable: true, - IsOnMap: false, - } - } + return detectStringSchedule(onStr) + } + onMap, ok := onValue.(map[string]any) + if !ok { return scheduleDetection{} } + return detectMappedSchedule(onMap) +} - // Case 2: on is a map — extract schedule value if present - if onMap, ok := onValue.(map[string]any); ok { - schedValue, hasSchedule := onMap["schedule"] - if !hasSchedule { - return scheduleDetection{} - } - - // Determine if on: has triggers beyond schedule / workflow_dispatch - isMultiTrigger := false - for key := range onMap { - if key != "schedule" && key != "workflow_dispatch" { - isMultiTrigger = true - scheduleWizardLog.Printf("Multi-trigger on: map detected (trigger '%s')", key) - break - } - } +func detectStringSchedule(onStr string) scheduleDetection { + if _, _, err := parser.ParseSchedule(onStr); err != nil { + return scheduleDetection{} + } + return scheduleDetection{RawExpr: onStr, Frequency: classifyScheduleFrequency(onStr), IsUpdatable: true} +} - // Schedule as string shorthand (e.g., "schedule: daily") - if schedStr, ok := schedValue.(string); ok { - return scheduleDetection{ - RawExpr: schedStr, - Frequency: classifyScheduleFrequency(schedStr), - IsUpdatable: true, - IsMultiTrigger: isMultiTrigger, - IsOnMap: true, - } - } +func detectMappedSchedule(onMap map[string]any) scheduleDetection { + schedValue, hasSchedule := onMap["schedule"] + if !hasSchedule { + return scheduleDetection{} + } + isMultiTrigger := scheduleHasMultipleTriggers(onMap) + if schedStr, ok := schedValue.(string); ok { + return scheduleDetection{RawExpr: schedStr, Frequency: classifyScheduleFrequency(schedStr), IsUpdatable: true, IsMultiTrigger: isMultiTrigger, IsOnMap: true} + } + return detectArraySchedule(schedValue, isMultiTrigger) +} - // Schedule as array (e.g., "schedule:\n - cron: daily") - if schedArray, ok := schedValue.([]any); ok && len(schedArray) > 0 { - // Workflows with multiple cron entries cannot be safely rewritten to a single - // frequency, so mark them as not updatable. - if len(schedArray) > 1 { - scheduleWizardLog.Printf("Multiple cron entries (%d) detected — not updatable", len(schedArray)) - return scheduleDetection{} - } - if item, ok := schedArray[0].(map[string]any); ok { - if cronVal, ok := item["cron"].(string); ok { - return scheduleDetection{ - RawExpr: cronVal, - Frequency: classifyScheduleFrequency(cronVal), - IsUpdatable: true, - IsMultiTrigger: isMultiTrigger, - IsOnMap: true, - } - } - } +func scheduleHasMultipleTriggers(onMap map[string]any) bool { + for key := range onMap { + if key != "schedule" && key != "workflow_dispatch" { + scheduleWizardLog.Printf("Multi-trigger on: map detected (trigger '%s')", key) + return true } } + return false +} - return scheduleDetection{} +func detectArraySchedule(schedValue any, isMultiTrigger bool) scheduleDetection { + schedArray, ok := schedValue.([]any) + if !ok || len(schedArray) == 0 { + return scheduleDetection{} + } + if len(schedArray) > 1 { + scheduleWizardLog.Printf("Multiple cron entries (%d) detected — not updatable", len(schedArray)) + return scheduleDetection{} + } + item, ok := schedArray[0].(map[string]any) + if !ok { + return scheduleDetection{} + } + cronVal, ok := item["cron"].(string) + if !ok { + return scheduleDetection{} + } + return scheduleDetection{RawExpr: cronVal, Frequency: classifyScheduleFrequency(cronVal), IsUpdatable: true, IsMultiTrigger: isMultiTrigger, IsOnMap: true} } +// classifyScheduleFrequency determines which standard frequency a schedule expression represents. // classifyScheduleFrequency determines which standard frequency a schedule expression represents. // Returns one of: "hourly", "3-hourly", "daily", "weekly", "monthly", or "custom". func classifyScheduleFrequency(scheduleStr string) string { normalized := strings.ToLower(strings.TrimSpace(scheduleStr)) + if frequency, ok := classifyFriendlyScheduleFrequency(normalized); ok { + return frequency + } + if frequency, ok := classifyFuzzyScheduleFrequency(normalized); ok { + return frequency + } + if frequency, ok := classifyCronScheduleFrequency(scheduleStr); ok { + return frequency + } + return "custom" +} - // Direct friendly-format matches +func classifyFriendlyScheduleFrequency(normalized string) (string, bool) { switch normalized { case "hourly", "every 1h", "every 1 hour", "every 1 hours": - return "hourly" + return "hourly", true case "every 3h", "every 3 hours": - return "3-hourly" + return "3-hourly", true case "daily": - return "daily" + return "daily", true case "weekly": - return "weekly" + return "weekly", true + default: + return "", false } +} - // Fuzzy cron placeholder matches (produced by the compiler during preprocessing) - if strings.HasPrefix(normalized, "fuzzy:hourly/1 ") || normalized == "fuzzy:hourly/1" { //nolint:tolowerequalfold - return "hourly" - } - if strings.HasPrefix(normalized, "fuzzy:hourly/3 ") || normalized == "fuzzy:hourly/3" { //nolint:tolowerequalfold - return "3-hourly" - } - if strings.HasPrefix(normalized, "fuzzy:daily") { - return "daily" - } - if strings.HasPrefix(normalized, "fuzzy:weekly") { - return "weekly" +func classifyFuzzyScheduleFrequency(normalized string) (string, bool) { + switch { + case strings.HasPrefix(normalized, "fuzzy:hourly/1 ") || normalized == "fuzzy:hourly/1": + return "hourly", true + case strings.HasPrefix(normalized, "fuzzy:hourly/3 ") || normalized == "fuzzy:hourly/3": + return "3-hourly", true + case strings.HasPrefix(normalized, "fuzzy:daily"): + return "daily", true + case strings.HasPrefix(normalized, "fuzzy:weekly"): + return "weekly", true + default: + return "", false } +} - // Cron expression checks +func classifyCronScheduleFrequency(scheduleStr string) (string, bool) { if parser.IsHourlyCron(scheduleStr) { - fields := strings.Fields(scheduleStr) - if len(fields) == 5 { - interval := strings.TrimPrefix(fields[1], "*/") - switch interval { - case "1": - return "hourly" - case "3": - return "3-hourly" - } - } - return "custom" + return classifyHourlyCronFrequency(scheduleStr), true } - if parser.IsDailyCron(scheduleStr) { - return "daily" + return "daily", true } - if parser.IsWeeklyCron(scheduleStr) { - return "weekly" + return "weekly", true + } + if isMonthlyCron(scheduleStr) { + return "monthly", true } + return "", false +} - // Monthly cron: M H * * where is a specific numeric date (e.g. "1", "15"). - // fields: [0]=minute [1]=hour [2]=day-of-month [3]=month [4]=day-of-week - // Excludes interval expressions like "*/2" so that "0 0 */2 * *" (every-2-days) is - // correctly classified as "custom" rather than "monthly". +func classifyHourlyCronFrequency(scheduleStr string) string { fields := strings.Fields(scheduleStr) - if len(fields) == 5 && fields[3] == "*" && fields[4] == "*" { - day := fields[2] // day-of-month field - if day != "*" && !strings.ContainsAny(day, "*/-,") { - return "monthly" + if len(fields) == 5 { + interval := strings.TrimPrefix(fields[1], "*/") + if interval == "1" { + return "hourly" + } + if interval == "3" { + return "3-hourly" } } - return "custom" } +func isMonthlyCron(scheduleStr string) bool { + fields := strings.Fields(scheduleStr) + if len(fields) != 5 || fields[3] != "*" || fields[4] != "*" { + return false + } + day := fields[2] + return day != "*" && !strings.ContainsAny(day, "*/-,") +} + +// selectScheduleFrequency presents a schedule-frequency selection form to the user when the // selectScheduleFrequency presents a schedule-frequency selection form to the user when the // workflow being added has a schedule trigger. If the user picks a different frequency the // resolved workflow content is updated in memory so the change is reflected in the PR. @@ -200,84 +208,80 @@ func (c *AddInteractiveConfig) selectScheduleFrequency() error { if c.resolvedWorkflows == nil || len(c.resolvedWorkflows.Workflows) == 0 { return nil } - for _, wf := range c.resolvedWorkflows.Workflows { - content := string(wf.Content) - detection := detectWorkflowScheduleInfo(content) - if !detection.IsUpdatable { - continue - } - - rawExpr := detection.RawExpr - currentFreq := detection.Frequency - scheduleWizardLog.Printf("Detected schedule: expr=%q, freq=%s, multiTrigger=%v", rawExpr, currentFreq, detection.IsMultiTrigger) - - // Build the ordered option list - options := buildScheduleOptions(rawExpr, currentFreq) - - fmt.Fprintln(os.Stderr, "") - fmt.Fprintln(os.Stderr, console.FormatInfoMessage("This workflow runs on a schedule.")) - - var selected string - form := console.NewSelectForm( - huh.NewSelect[string](). - Title("How often should this workflow run?"). - Description("Current schedule: " + rawExpr). - Options(options...). - Value(&selected), - ) - - if err := form.RunWithContext(c.Ctx); err != nil { - return fmt.Errorf("failed to select schedule frequency: %w", err) + if err := c.selectWorkflowScheduleFrequency(wf); err != nil { + return err } + } + return nil +} - scheduleWizardLog.Printf("User selected frequency: %s", selected) - - // "custom" or same frequency means keep as-is - if selected == "custom" || selected == currentFreq { - scheduleWizardLog.Printf("Schedule unchanged: keeping %q", rawExpr) - continue +func (c *AddInteractiveConfig) selectWorkflowScheduleFrequency(wf *ResolvedWorkflow) error { + content := string(wf.Content) + detection := detectWorkflowScheduleInfo(content) + if !detection.IsUpdatable { + return nil + } + scheduleWizardLog.Printf("Detected schedule: expr=%q, freq=%s, multiTrigger=%v", detection.RawExpr, detection.Frequency, detection.IsMultiTrigger) + selected, err := c.promptForScheduleFrequency(detection) + if err != nil || selected == "custom" || selected == detection.Frequency { + if err == nil && (selected == "custom" || selected == detection.Frequency) { + scheduleWizardLog.Printf("Schedule unchanged: keeping %q", detection.RawExpr) } + return err + } + return updateResolvedWorkflowSchedule(wf, content, detection, selected) +} - // Look up the schedule expression for the chosen frequency - var newExpr string - for _, opt := range standardScheduleFrequencies { - if opt.Value == selected { - newExpr = opt.Expression - break - } - } - if newExpr == "" { - continue - } +func (c *AddInteractiveConfig) promptForScheduleFrequency(detection scheduleDetection) (string, error) { + options := buildScheduleOptions(detection.RawExpr, detection.Frequency) + var selected string + fmt.Fprintln(os.Stderr, "") + fmt.Fprintln(os.Stderr, console.FormatInfoMessage("This workflow runs on a schedule.")) + form := console.NewSelectForm(huh.NewSelect[string]().Title("How often should this workflow run?").Description("Current schedule: " + detection.RawExpr).Options(options...).Value(&selected)) + if err := form.RunWithContext(c.Ctx); err != nil { + return "", fmt.Errorf("failed to select schedule frequency: %w", err) + } + scheduleWizardLog.Printf("User selected frequency: %s", selected) + return selected, nil +} - // Update the workflow content in memory. - // When on: is a mapping, update only the schedule sub-key so other triggers - // (e.g., workflow_dispatch, push) are preserved. - // When on: is a scalar string, replace the on: field value directly. - var updatedContent string - var updateErr error - if detection.IsOnMap { - updatedContent, updateErr = UpdateScheduleInOnBlock(content, newExpr) - } else { - updatedContent, updateErr = UpdateFieldInFrontmatter(content, "on", newExpr) - } - if updateErr != nil { - scheduleWizardLog.Printf("Failed to update schedule (isOnMap=%v): %v", detection.IsOnMap, updateErr) - fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Could not update schedule: %v", updateErr))) - continue - } +func updateResolvedWorkflowSchedule(wf *ResolvedWorkflow, content string, detection scheduleDetection, selected string) error { + newExpr := scheduleExpressionForFrequency(selected) + if newExpr == "" { + return nil + } + updatedContent, err := rewriteWorkflowSchedule(content, detection, newExpr) + if err != nil { + scheduleWizardLog.Printf("Failed to update schedule (isOnMap=%v): %v", detection.IsOnMap, err) + fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Could not update schedule: %v", err))) + return nil + } + wf.Content = []byte(updatedContent) + if wf.SourceInfo != nil { + wf.SourceInfo.Content = []byte(updatedContent) + } + fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Schedule updated to: "+selected)) + return nil +} - wf.Content = []byte(updatedContent) - if wf.SourceInfo != nil { - wf.SourceInfo.Content = []byte(updatedContent) +func scheduleExpressionForFrequency(selected string) string { + for _, opt := range standardScheduleFrequencies { + if opt.Value == selected { + return opt.Expression } - fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Schedule updated to: "+selected)) } + return "" +} - return nil +func rewriteWorkflowSchedule(content string, detection scheduleDetection, newExpr string) (string, error) { + if detection.IsOnMap { + return UpdateScheduleInOnBlock(content, newExpr) + } + return UpdateFieldInFrontmatter(content, "on", newExpr) } +// buildScheduleOptions constructs the huh option list for the schedule frequency form. // buildScheduleOptions constructs the huh option list for the schedule frequency form. // The default option (matching the current frequency) is placed first. func buildScheduleOptions(rawExpr, currentFreq string) []huh.Option[string] { diff --git a/pkg/cli/add_interactive_workflow.go b/pkg/cli/add_interactive_workflow.go index cb011f1dc2a..833803c7d5b 100644 --- a/pkg/cli/add_interactive_workflow.go +++ b/pkg/cli/add_interactive_workflow.go @@ -16,89 +16,128 @@ import ( // checkStatusAndOfferRun checks if the workflow appears in status and offers to run it func (c *AddInteractiveConfig) checkStatusAndOfferRun(ctx context.Context) error { addInteractiveLog.Print("Checking workflow status and offering to run") - - // Wait a moment for GitHub to process the merge fmt.Fprintln(os.Stderr, "") - - // Use spinner only in non-verbose mode (spinner can't be restarted after stop) - var spinner *console.SpinnerWrapper - if !c.Verbose { - spinner = console.NewSpinner("Waiting for workflow to be available...") - spinner.Start() + workflowName := c.primaryWorkflowName() + workflowFound, err := c.waitForWorkflowAvailability(ctx, workflowName) + if err != nil { + return err } + if !workflowFound { + c.printWorkflowNotReadyMessage() + return nil + } + fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Workflow is ready")) + if !c.shouldOfferWorkflowRun() || c.handleCodespacesWorkflowRun() { + return nil + } + if !c.promptToRunWorkflowNow(ctx) { + c.showFinalInstructions() + return nil + } + c.runWorkflowNow(ctx, workflowName) + c.showFinalInstructions() + return nil +} - // Try a few times to see the workflow in status - var workflowFound bool +func (c *AddInteractiveConfig) waitForWorkflowAvailability(ctx context.Context, workflowName string) (bool, error) { + spinner := c.startWorkflowAvailabilitySpinner() + defer stopWorkflowAvailabilitySpinner(spinner) for i := range 5 { - // Wait 2 seconds before each check (including the first) - timer := time.NewTimer(2 * time.Second) - select { - case <-ctx.Done(): - timer.Stop() - if spinner != nil { - spinner.Stop() - } - return ctx.Err() - case <-timer.C: - // Continue with check + if err := waitForWorkflowAvailabilityAttempt(ctx); err != nil { + return false, err } - - workflowName := c.primaryWorkflowName() - if workflowName != "" { - if c.Verbose { - fmt.Fprintf(os.Stderr, "Checking workflow status (attempt %d/5) for: %s\n", i+1, workflowName) - } - // Check if workflow is in status - statuses, err := findWorkflowsByFilenamePattern(workflowName, c.RepoOverride, c.Verbose) - if err != nil { - if c.Verbose { - fmt.Fprintf(os.Stderr, "Status check error: %v\n", err) - } - } else if len(statuses) > 0 { - if c.Verbose { - fmt.Fprintf(os.Stderr, "Found %d workflow(s) matching pattern\n", len(statuses)) - } - workflowFound = true - break - } else if c.Verbose { - fmt.Fprintln(os.Stderr, "No workflows found matching pattern yet") - } + found, err := c.checkWorkflowAvailabilityAttempt(workflowName, i) + if found { + return true, nil + } + if err != nil { + continue } } + return false, nil +} + +func (c *AddInteractiveConfig) startWorkflowAvailabilitySpinner() *console.SpinnerWrapper { + if c.Verbose { + return nil + } + spinner := console.NewSpinner("Waiting for workflow to be available...") + spinner.Start() + return spinner +} +func stopWorkflowAvailabilitySpinner(spinner *console.SpinnerWrapper) { if spinner != nil { spinner.Stop() } +} - if !workflowFound { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage("Could not verify workflow status.")) - fmt.Fprintf(os.Stderr, "You can check status with: %s status\n", string(constants.CLIExtensionPrefix)) - c.showFinalInstructions() +func waitForWorkflowAvailabilityAttempt(ctx context.Context) error { + timer := time.NewTimer(2 * time.Second) + defer timer.Stop() + select { + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: return nil } +} - fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Workflow is ready")) - - // Only offer to run if workflow has workflow_dispatch trigger - if c.addResult == nil || !c.addResult.HasWorkflowDispatch { - addInteractiveLog.Print("Workflow does not have workflow_dispatch trigger, skipping run offer") - c.showFinalInstructions() - return nil +func (c *AddInteractiveConfig) checkWorkflowAvailabilityAttempt(workflowName string, attempt int) (bool, error) { + if workflowName == "" { + return false, nil + } + if c.Verbose { + fmt.Fprintf(os.Stderr, "Checking workflow status (attempt %d/5) for: %s\n", attempt+1, workflowName) + } + statuses, err := findWorkflowsByFilenamePattern(workflowName, c.RepoOverride, c.Verbose) + if err != nil { + if c.Verbose { + fmt.Fprintf(os.Stderr, "Status check error: %v\n", err) + } + return false, err + } + if len(statuses) > 0 { + if c.Verbose { + fmt.Fprintf(os.Stderr, "Found %d workflow(s) matching pattern\n", len(statuses)) + } + return true, nil } + if c.Verbose { + fmt.Fprintln(os.Stderr, "No workflows found matching pattern yet") + } + return false, nil +} - // In Codespaces, don't offer to trigger - provide link to Actions page instead - if isRunningInCodespace() { - addInteractiveLog.Print("Running in Codespaces, skipping run offer and showing Actions link") - fmt.Fprintln(os.Stderr, "") - fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Running in GitHub Codespaces - please trigger the workflow manually from the Actions page")) - fmt.Fprintf(os.Stderr, "🔗 https://github.com/%s/actions\n", c.RepoOverride) - c.showFinalInstructions() - return nil +func (c *AddInteractiveConfig) printWorkflowNotReadyMessage() { + fmt.Fprintln(os.Stderr, console.FormatWarningMessage("Could not verify workflow status.")) + fmt.Fprintf(os.Stderr, "You can check status with: %s status\n", string(constants.CLIExtensionPrefix)) + c.showFinalInstructions() +} + +func (c *AddInteractiveConfig) shouldOfferWorkflowRun() bool { + if c.addResult != nil && c.addResult.HasWorkflowDispatch { + return true } + addInteractiveLog.Print("Workflow does not have workflow_dispatch trigger, skipping run offer") + c.showFinalInstructions() + return false +} - // Ask if user wants to run the workflow +func (c *AddInteractiveConfig) handleCodespacesWorkflowRun() bool { + if !isRunningInCodespace() { + return false + } + addInteractiveLog.Print("Running in Codespaces, skipping run offer and showing Actions link") fmt.Fprintln(os.Stderr, "") - runNow := true // Default to yes + fmt.Fprintln(os.Stderr, console.FormatInfoMessage("Running in GitHub Codespaces - please trigger the workflow manually from the Actions page")) + fmt.Fprintf(os.Stderr, "🔗 https://github.com/%s/actions\n", c.RepoOverride) + c.showFinalInstructions() + return true +} + +func (c *AddInteractiveConfig) promptToRunWorkflowNow(ctx context.Context) bool { + runNow := true form := console.NewConfirmForm( huh.NewConfirm(). Title("Would you like to run the workflow once now?"). @@ -107,60 +146,53 @@ func (c *AddInteractiveConfig) checkStatusAndOfferRun(ctx context.Context) error Negative("No, I'll run later"). Value(&runNow), ) - if err := form.RunWithContext(ctx); err != nil { - return nil // Not critical, just skip + return false } + return runNow +} - if !runNow { - c.showFinalInstructions() - return nil +func (c *AddInteractiveConfig) runWorkflowNow(ctx context.Context, workflowName string) { + if workflowName == "" { + return } + fmt.Fprintln(os.Stderr, "") + c.updateWorkflowRunBranchState() + if err := RunSpecificWorkflowInteractively(ctx, RunWorkflowOptions{ + WorkflowName: workflowName, + Verbose: c.Verbose, + EngineOverride: c.EngineOverride, + RepoOverride: c.RepoOverride, + }); err != nil { + fmt.Fprintln(os.Stderr, console.FormatErrorMessage(fmt.Sprintf("Failed to run workflow: %v", err))) + return + } + c.printTriggeredWorkflowRunURL(workflowName) +} - // Run the workflow interactively (collects inputs if the workflow has them) - workflowName := c.primaryWorkflowName() - if workflowName != "" { - fmt.Fprintln(os.Stderr, "") - - // Pull the merged workflow files now that we know GitHub has processed the - // merge (workflowFound is true). Doing this here—rather than immediately - // after the PR merge—avoids a race where git fetch runs before GitHub's git - // objects have been updated, which caused "workflow file not found" errors. - if !c.Verbose { - fmt.Fprintln(os.Stderr, "Updating local branch (this may take a few seconds)...") - } - if err := c.updateLocalBranch(); err != nil { - addInteractiveLog.Printf("Failed to update local branch: %v", err) - fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Could not update local branch: %v", err))) - fmt.Fprintln(os.Stderr, "You may need to switch to your repository's default branch (for example 'main') and run 'git pull' manually before running the workflow.") - } - if !c.Verbose { - fmt.Fprintln(os.Stderr, "Finished updating local branch.") - } - - if err := RunSpecificWorkflowInteractively(ctx, RunWorkflowOptions{ - WorkflowName: workflowName, - Verbose: c.Verbose, - EngineOverride: c.EngineOverride, - RepoOverride: c.RepoOverride, - }); err != nil { - fmt.Fprintln(os.Stderr, console.FormatErrorMessage(fmt.Sprintf("Failed to run workflow: %v", err))) - c.showFinalInstructions() - return nil - } - - // Get the run URL for step 10 - runInfo, err := getLatestWorkflowRunWithRetry(workflowName+".lock.yml", c.RepoOverride, c.Verbose) - if err == nil && runInfo.URL != "" { - fmt.Fprintln(os.Stderr, "") - fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Workflow triggered successfully!")) - fmt.Fprintln(os.Stderr, "") - fmt.Fprintf(os.Stderr, "🔗 View workflow run: %s\n", runInfo.URL) - } +func (c *AddInteractiveConfig) updateWorkflowRunBranchState() { + if !c.Verbose { + fmt.Fprintln(os.Stderr, "Updating local branch (this may take a few seconds)...") } + if err := c.updateLocalBranch(); err != nil { + addInteractiveLog.Printf("Failed to update local branch: %v", err) + fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Could not update local branch: %v", err))) + fmt.Fprintln(os.Stderr, "You may need to switch to your repository's default branch (for example 'main') and run 'git pull' manually before running the workflow.") + } + if !c.Verbose { + fmt.Fprintln(os.Stderr, "Finished updating local branch.") + } +} - c.showFinalInstructions() - return nil +func (c *AddInteractiveConfig) printTriggeredWorkflowRunURL(workflowName string) { + runInfo, err := getLatestWorkflowRunWithRetry(workflowName+".lock.yml", c.RepoOverride, c.Verbose) + if err != nil || runInfo.URL == "" { + return + } + fmt.Fprintln(os.Stderr, "") + fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Workflow triggered successfully!")) + fmt.Fprintln(os.Stderr, "") + fmt.Fprintf(os.Stderr, "🔗 View workflow run: %s\n", runInfo.URL) } // findWorkflowsByFilenamePattern is a helper to find workflows registered in GitHub by filename pattern. diff --git a/pkg/cli/add_package_manifest.go b/pkg/cli/add_package_manifest.go index 82455e64a7f..313b55c1cc4 100644 --- a/pkg/cli/add_package_manifest.go +++ b/pkg/cli/add_package_manifest.go @@ -77,99 +77,90 @@ func (e packageRemoteNotFoundError) Unwrap() []error { } func resolveRepositoryPackage(ctx context.Context, repoSpec *RepoSpec, host string) (*resolvedRepositoryPackage, error) { - parts := strings.SplitN(repoSpec.RepoSlug, "/", 2) - if len(parts) != 2 { - return nil, fmt.Errorf("invalid repository slug: %s", repoSpec.RepoSlug) - } - - owner := parts[0] - repo := parts[1] - // At manifest-fetch time there is no resolved package metadata yet. - ref := repositoryPackageEffectiveRef(repoSpec, nil) - if ref == "" { - if isGhAwRepository(repoSpec.RepoSlug) { - if latestRelease, err := getRepositoryPackageLatestRelease(ctx, repoSpec.RepoSlug, host); err == nil { - ref = latestRelease - } else { - addPackageManifestLog.Printf("failed to resolve latest release for %s (host=%q): %v", repoSpec.RepoSlug, host, err) - } - } - if ref == "" { - ref = "main" - if defaultBranch, err := getRepositoryPackageDefaultBranch(ctx, repoSpec.RepoSlug, host); err == nil { - ref = defaultBranch - } else { - addPackageManifestLog.Printf("failed to resolve default branch for %s (host=%q), falling back to %q: %v", repoSpec.RepoSlug, host, ref, err) - } - } + owner, repo, err := splitRepositoryPackageSlug(repoSpec.RepoSlug) + if err != nil { + return nil, err } + ref := resolveRepositoryPackageRef(ctx, repoSpec, host) packagePath := strings.Trim(repoSpec.PackagePath, "/") - manifestPath, manifestContent, err := loadRepositoryPackageManifestFile(ctx, owner, repo, packagePath, ref, host) if err != nil { return nil, err } - manifest, warnings, err := parseRepositoryPackageManifest(manifestPath, manifestContent) if err != nil { return nil, err } + return buildResolvedRepositoryPackage(ctx, owner, repo, packagePath, ref, host, repoSpec, manifest, manifestPath, warnings) +} - includeInstallablePaths, includeSkillDirs, includeAgentFiles := splitManifestIncludePaths(manifest.Includes) - includeInstallablePaths = append(includeInstallablePaths, manifest.Files...) +func splitRepositoryPackageSlug(repoSlug string) (string, string, error) { + parts := strings.SplitN(repoSlug, "/", 2) + if len(parts) != 2 { + return "", "", fmt.Errorf("invalid repository slug: %s", repoSlug) + } + return parts[0], parts[1], nil +} - installationSources := normalizePackageInstallablePaths(includeInstallablePaths, packagePath) - if len(installationSources) == 0 { - installationSources, err = scanRepositoryPackageInstallablePaths(ctx, owner, repo, packagePath, ref, host) - if err != nil { - return nil, err +func resolveRepositoryPackageRef(ctx context.Context, repoSpec *RepoSpec, host string) string { + ref := repositoryPackageEffectiveRef(repoSpec, nil) + if ref != "" { + return ref + } + if isGhAwRepository(repoSpec.RepoSlug) { + if latestRelease, err := getRepositoryPackageLatestRelease(ctx, repoSpec.RepoSlug, host); err == nil { + return latestRelease } + addPackageManifestLog.Printf("failed to resolve latest release for %s (host=%q)", repoSpec.RepoSlug, host) } - if err := validateUniqueManifestWorkflowFilenames(installationSources, manifestPath); err != nil { - return nil, err + ref = "main" + if defaultBranch, err := getRepositoryPackageDefaultBranch(ctx, repoSpec.RepoSlug, host); err == nil { + return defaultBranch } + addPackageManifestLog.Printf("failed to resolve default branch for %s (host=%q), falling back to %q", repoSpec.RepoSlug, host, ref) + return ref +} +func buildResolvedRepositoryPackage(ctx context.Context, owner, repo, packagePath, ref, host string, repoSpec *RepoSpec, manifest *repositoryPackageManifest, manifestPath string, warnings []string) (*resolvedRepositoryPackage, error) { + installationSources, skillDirs, agentFiles, err := resolveRepositoryPackageContentPaths(ctx, owner, repo, packagePath, ref, host, manifest, manifestPath) + if err != nil { + return nil, err + } docsPath, err := resolveRepositoryPackageDocsPath(ctx, owner, repo, packagePath, ref, host) if err != nil { return nil, err } - - // Resolve skill files: explicit from manifest or auto-scanned. - explicitSkillDirs := append([]string{}, manifest.Skills...) - explicitSkillDirs = append(explicitSkillDirs, includeSkillDirs...) - skillFiles, skillWarnings, err := resolvePackageSkillFiles(ctx, owner, repo, packagePath, ref, host, explicitSkillDirs) + skillFiles, skillWarnings, err := resolvePackageSkillFiles(ctx, owner, repo, packagePath, ref, host, skillDirs) if err != nil { return nil, err } - warnings = append(warnings, skillWarnings...) - - // Resolve agent files: explicit from manifest or auto-scanned. - explicitAgentFiles := append([]string{}, manifest.Agents...) - explicitAgentFiles = append(explicitAgentFiles, includeAgentFiles...) - agentFiles, agentWarnings, err := resolvePackageAgentFiles(ctx, owner, repo, packagePath, ref, host, explicitAgentFiles) + agentFiles, agentWarnings, err := resolvePackageAgentFiles(ctx, owner, repo, packagePath, ref, host, agentFiles) if err != nil { return nil, err } + warnings = append(warnings, skillWarnings...) warnings = append(warnings, agentWarnings...) - if len(installationSources) == 0 && len(skillFiles) == 0 && len(agentFiles) == 0 { return nil, fmt.Errorf("repository %q does not contain any installable workflows, skills, or agents (either explicitly declared or auto-discovered)", repositoryPackageIdentifier(repoSpec.RepoSlug, packagePath)) } + return &resolvedRepositoryPackage{ManifestPath: manifestPath, ResolvedRef: ref, Name: manifest.Name, Emoji: manifest.Emoji, Description: manifest.Description, License: manifest.License, DocsPath: docsPath, InstallationSource: installationSources, Bootstrap: manifest.Bootstrap, SkillFiles: skillFiles, AgentFiles: agentFiles, Warnings: warnings}, nil +} - return &resolvedRepositoryPackage{ - ManifestPath: manifestPath, - ResolvedRef: ref, - Name: manifest.Name, - Emoji: manifest.Emoji, - Description: manifest.Description, - License: manifest.License, - DocsPath: docsPath, - InstallationSource: installationSources, - Bootstrap: manifest.Bootstrap, - SkillFiles: skillFiles, - AgentFiles: agentFiles, - Warnings: warnings, - }, nil +func resolveRepositoryPackageContentPaths(ctx context.Context, owner, repo, packagePath, ref, host string, manifest *repositoryPackageManifest, manifestPath string) ([]string, []string, []string, error) { + includeInstallablePaths, includeSkillDirs, includeAgentFiles := splitManifestIncludePaths(manifest.Includes) + includeInstallablePaths = append(includeInstallablePaths, manifest.Files...) + installationSources := normalizePackageInstallablePaths(includeInstallablePaths, packagePath) + if len(installationSources) == 0 { + var err error + installationSources, err = scanRepositoryPackageInstallablePaths(ctx, owner, repo, packagePath, ref, host) + if err != nil { + return nil, nil, nil, err + } + } + if err := validateUniqueManifestWorkflowFilenames(installationSources, manifestPath); err != nil { + return nil, nil, nil, err + } + return installationSources, append([]string{}, append(manifest.Skills, includeSkillDirs...)...), append([]string{}, append(manifest.Agents, includeAgentFiles...)...), nil } func loadRepositoryPackageManifestFile(ctx context.Context, owner, repo, packagePath, ref, host string) (string, []byte, error) { @@ -205,106 +196,129 @@ type repositoryPackageManifest struct { } func parseRepositoryPackageManifest(manifestPath string, content []byte) (*repositoryPackageManifest, []string, error) { + root, name, err := parseRepositoryPackageManifestRoot(manifestPath, content) + if err != nil { + return nil, nil, err + } + manifest := &repositoryPackageManifest{Name: strings.TrimSpace(name)} + warnings := make([]string, 0) + applyRepositoryManifestVersionFields(manifest, root) + if warnings, err = applyRepositoryManifestMinVersion(manifest, root, manifestPath, warnings); err != nil { + return nil, nil, err + } + warnings = applyRepositoryManifestMetadata(manifest, root, manifestPath, warnings) + warnings, err = applyRepositoryManifestCollections(manifest, root, manifestPath, warnings) + if err != nil { + return nil, nil, err + } + return manifest, warnings, nil +} + +func parseRepositoryPackageManifestRoot(manifestPath string, content []byte) (map[string]any, string, error) { var raw any if err := yaml.Unmarshal(content, &raw); err != nil { - return nil, nil, fmt.Errorf("invalid Agentic Workflow manifest %q: %s", manifestPath, parser.FormatYAMLError(err, 1, string(content))) + return nil, "", fmt.Errorf("invalid Agentic Workflow manifest %q: %s", manifestPath, parser.FormatYAMLError(err, 1, string(content))) } - root, ok := raw.(map[string]any) if !ok { - return nil, nil, fmt.Errorf("invalid Agentic Workflow manifest %q: top-level document must be a mapping", manifestPath) + return nil, "", fmt.Errorf("invalid Agentic Workflow manifest %q: top-level document must be a mapping", manifestPath) } - - // Validate name before schema validation to provide a clear error message for - // the most common manifest authoring error (missing or empty name). name, ok := stringValue(root["name"]) if !ok || strings.TrimSpace(name) == "" { - return nil, nil, fmt.Errorf("invalid Agentic Workflow manifest %q: name must be a non-empty string", manifestPath) + return nil, "", fmt.Errorf("invalid Agentic Workflow manifest %q: name must be a non-empty string", manifestPath) } - if err := parser.ValidateRepositoryPackageManifestWithSchemaAndLocation(root, manifestPath); err != nil { - return nil, nil, fmt.Errorf("invalid Agentic Workflow manifest %q: %w", manifestPath, err) - } - - manifest := &repositoryPackageManifest{ - Name: strings.TrimSpace(name), + return nil, "", fmt.Errorf("invalid Agentic Workflow manifest %q: %w", manifestPath, err) } - var warnings []string + return root, name, nil +} +func applyRepositoryManifestVersionFields(manifest *repositoryPackageManifest, root map[string]any) { if manifestVersion, ok := stringValue(root["manifest-version"]); ok { manifest.ManifestVersion = strings.TrimSpace(manifestVersion) } else { manifest.ManifestVersion = repositoryPackageManifestVersion } +} - if minVersion, ok := stringValue(root["min-version"]); ok { - manifest.MinVersion = strings.TrimSpace(minVersion) - if !isSupportedManifestMinVersion(manifest.MinVersion) { - return nil, nil, fmt.Errorf("invalid Agentic Workflow manifest %q: min-version must use vMAJOR.minor.patch, got %q", manifestPath, minVersion) - } - currentVersion := GetVersion() - if !semverutil.IsValid(currentVersion) { - return nil, nil, fmt.Errorf("invalid Agentic Workflow manifest %q: min-version validation requires a semantic-versioned compiler, but the current compiler version %q is not a valid semantic version (this indicates a build issue)", manifestPath, currentVersion) - } - currentVersion = semverutil.NormalizeGitDescribeSemver(currentVersion) - if semverutil.Compare(currentVersion, manifest.MinVersion) < 0 { - return nil, nil, fmt.Errorf("invalid Agentic Workflow manifest %q: min-version %q requires gh-aw %s or newer (current: %s)", manifestPath, manifest.MinVersion, manifest.MinVersion, currentVersion) - } +func applyRepositoryManifestMinVersion(manifest *repositoryPackageManifest, root map[string]any, manifestPath string, warnings []string) ([]string, error) { + minVersion, ok := stringValue(root["min-version"]) + if !ok { + return warnings, nil + } + manifest.MinVersion = strings.TrimSpace(minVersion) + if !isSupportedManifestMinVersion(manifest.MinVersion) { + return nil, fmt.Errorf("invalid Agentic Workflow manifest %q: min-version must use vMAJOR.minor.patch, got %q", manifestPath, minVersion) + } + currentVersion := GetVersion() + if !semverutil.IsValid(currentVersion) { + return nil, fmt.Errorf("invalid Agentic Workflow manifest %q: min-version validation requires a semantic-versioned compiler, but the current compiler version %q is not a valid semantic version (this indicates a build issue)", manifestPath, currentVersion) } + currentVersion = semverutil.NormalizeGitDescribeSemver(currentVersion) + if semverutil.Compare(currentVersion, manifest.MinVersion) < 0 { + return nil, fmt.Errorf("invalid Agentic Workflow manifest %q: min-version %q requires gh-aw %s or newer (current: %s)", manifestPath, manifest.MinVersion, manifest.MinVersion, currentVersion) + } + return warnings, nil +} +func applyRepositoryManifestMetadata(manifest *repositoryPackageManifest, root map[string]any, manifestPath string, warnings []string) []string { if description, ok := stringValue(root["description"]); ok { manifest.Description = description if len(description) > 255 { warnings = append(warnings, fmt.Sprintf("Manifest %s description exceeds the 255-character marketplace display limit", manifestPath)) } } - if emoji, ok := stringValue(root["emoji"]); ok { manifest.Emoji = emoji } - if license, ok := stringValue(root["license"]); ok { manifest.License = license } + return warnings +} +func applyRepositoryManifestCollections(manifest *repositoryPackageManifest, root map[string]any, manifestPath string, warnings []string) ([]string, error) { if includesValue, ok := root["includes"]; ok { includes, includeWarnings := extractManifestIncludes(includesValue, manifestPath) manifest.Includes = includes warnings = append(warnings, includeWarnings...) } - if filesValue, ok := root["files"]; ok { files, fileWarnings := extractManifestFiles(filesValue, manifestPath) manifest.Files = files warnings = append(warnings, fileWarnings...) if len(files) > 0 { - warnings = append(warnings, fmt.Sprintf("Field 'files' in %s is deprecated; use 'includes' instead.", manifestPath)) - warnings = append(warnings, "Codemod suggestion:\n"+formatIncludesCodemodSuggestion(codemodManifestFilesToIncludes(files))) + warnings = append( + warnings, + fmt.Sprintf("Field 'files' in %s is deprecated; use 'includes' instead.", manifestPath), + "Codemod suggestion:\n"+formatIncludesCodemodSuggestion(codemodManifestFilesToIncludes(files)), + ) + } + } + warnings = applyRepositoryManifestSkillAgentCollections(manifest, root, manifestPath, warnings) + if configValue, ok := root["config"]; ok { + warnings = append(warnings, "Using experimental feature: config") + bootstrap, err := extractManifestConfig(configValue, manifestPath) + if err != nil { + return nil, err } + manifest.Bootstrap = bootstrap } + return warnings, nil +} +func applyRepositoryManifestSkillAgentCollections(manifest *repositoryPackageManifest, root map[string]any, manifestPath string, warnings []string) []string { if skillsValue, ok := root["skills"]; ok { skills, skillWarnings := extractManifestSkillDirs(skillsValue, manifestPath) manifest.Skills = skills warnings = append(warnings, skillWarnings...) } - if agentsValue, ok := root["agents"]; ok { agents, agentWarnings := extractManifestAgentFiles(agentsValue, manifestPath) manifest.Agents = agents warnings = append(warnings, agentWarnings...) } - - if configValue, ok := root["config"]; ok { - warnings = append(warnings, "Using experimental feature: config") - bootstrap, err := extractManifestConfig(configValue, manifestPath) - if err != nil { - return nil, nil, err - } - manifest.Bootstrap = bootstrap - } - - return manifest, warnings, nil + return warnings } func extractManifestIncludes(value any, manifestPath string) ([]string, []string) { @@ -542,85 +556,111 @@ func agentDirectoryRoot(cleaned string) string { // contain a SKILL.md file but are not already covered by the manifest. Each skill folder // is traversed recursively so that all nested files are included. func resolvePackageSkillFiles(ctx context.Context, owner, repo, packagePath, ref, host string, explicitSkillDirs []string) ([]resolvedPackageSkillFile, []string, error) { - // seenSkillDirs tracks full skill directories already added so that auto-scanned - // duplicates of manifest-specified skills are not added a second time. - seenSkillDirs := make(map[string]struct{}) - var warnings []string + manifestSkillDirs := packageManifestSkillDirs(packagePath, explicitSkillDirs) + autoScanned, warnings, err := autoScanPackageSkillDirs(ctx, owner, repo, packagePath, ref, host, len(manifestSkillDirs) > 0) + if err != nil { + return nil, nil, err + } + skillDirs := uniquePackageSkillDirs(manifestSkillDirs, autoScanned) + manifestSkillDirSet := packageSkillDirSet(manifestSkillDirs) + skillFiles, fileWarnings, err := collectRemotePackageSkillFiles(ctx, owner, repo, ref, host, skillDirs, manifestSkillDirSet) + if err != nil { + return nil, nil, err + } + return skillFiles, append(warnings, fileWarnings...), nil +} - // Step 1: resolve manifest skills first (explicit dirs). - var manifestSkillDirs []string +func packageManifestSkillDirs(packagePath string, explicitSkillDirs []string) []string { + manifestSkillDirs := make([]string, 0, len(explicitSkillDirs)) for _, dir := range explicitSkillDirs { manifestSkillDirs = append(manifestSkillDirs, joinRepositoryPackagePath(packagePath, dir)) } + return manifestSkillDirs +} - // Step 2: always auto-scan and append any skills not already in the manifest. +func autoScanPackageSkillDirs(ctx context.Context, owner, repo, packagePath, ref, host string, hasManifestSkills bool) ([]string, []string, error) { autoScanned, err := scanPackageSkillDirs(ctx, owner, repo, packagePath, ref, host) - if err != nil { - // Auto-scan is supplementary for manifest-declared skills; preserve manifest - // resolution even when scan fails transiently. - if len(manifestSkillDirs) > 0 { - warnings = append(warnings, fmt.Sprintf("failed to auto-scan skills directory, proceeding with manifest skills only: %v", err)) - } else { - return nil, nil, err - } + if err == nil { + return autoScanned, nil, nil + } + if hasManifestSkills { + return nil, []string{fmt.Sprintf("failed to auto-scan skills directory, proceeding with manifest skills only: %v", err)}, nil } + return nil, nil, err +} - // Build the final ordered list: manifest skills first, then auto-scanned extras. - var skillDirs []string - appendIfNew := func(dir string) { - if _, exists := seenSkillDirs[dir]; !exists { - seenSkillDirs[dir] = struct{}{} - skillDirs = append(skillDirs, dir) +func uniquePackageSkillDirs(manifestSkillDirs, autoScanned []string) []string { + seenSkillDirs := make(map[string]struct{}) + skillDirs := make([]string, 0, len(manifestSkillDirs)+len(autoScanned)) + for _, dir := range append(append([]string{}, manifestSkillDirs...), autoScanned...) { + if _, exists := seenSkillDirs[dir]; exists { + continue } + seenSkillDirs[dir] = struct{}{} + skillDirs = append(skillDirs, dir) } - for _, dir := range manifestSkillDirs { - appendIfNew(dir) - } - for _, dir := range autoScanned { - appendIfNew(dir) - } + return skillDirs +} - // manifestSkillDirSet is used to know which dirs require a SKILL.md marker check. +func packageSkillDirSet(manifestSkillDirs []string) map[string]struct{} { manifestSkillDirSet := make(map[string]struct{}, len(manifestSkillDirs)) - for _, d := range manifestSkillDirs { - manifestSkillDirSet[d] = struct{}{} + for _, dir := range manifestSkillDirs { + manifestSkillDirSet[dir] = struct{}{} } + return manifestSkillDirSet +} +func collectRemotePackageSkillFiles(ctx context.Context, owner, repo, ref, host string, skillDirs []string, manifestSkillDirSet map[string]struct{}) ([]resolvedPackageSkillFile, []string, error) { + var warnings []string var skillFiles []resolvedPackageSkillFile for _, skillDir := range skillDirs { - // For skills that came from the manifest, validate that the SKILL.md marker - // exists so that typos in the manifest surface as clear warnings. - if _, fromManifest := manifestSkillDirSet[skillDir]; fromManifest { - markerPath := joinRepositoryPackagePath(skillDir, packageSkillMarkerFile) - if _, err := downloadPackageFileFromGitHubForHost(ctx, owner, repo, markerPath, ref, host); err != nil { - if isRepositoryFileNotFound(err) { - warnings = append(warnings, fmt.Sprintf("Skill directory %q is missing required %s marker file", skillDir, packageSkillMarkerFile)) - continue - } - return nil, nil, fmt.Errorf("failed to validate skill marker %q: %w", markerPath, err) - } - } - skillName := filepath.Base(skillDir) - // Use recursive listing so that the entire skill folder (including any - // subdirectories) is copied, not just the top-level files. - files, err := listPackageDirFilesRecursivelyForHost(ctx, owner, repo, ref, skillDir, host) + files, fileWarnings, err := collectRemotePackageSkillDirFiles(ctx, owner, repo, ref, host, skillDir, manifestSkillDirSet) if err != nil { - if isRepositoryFileNotFound(err) { - warnings = append(warnings, fmt.Sprintf("Skill directory %q not found in package, skipping", skillDir)) - continue - } - return nil, nil, fmt.Errorf("failed to list files in skill directory %q: %w", skillDir, err) - } - for _, file := range files { - skillFiles = append(skillFiles, resolvedPackageSkillFile{ - SourcePath: file, - SkillName: skillName, - }) + return nil, nil, err } + warnings = append(warnings, fileWarnings...) + skillFiles = append(skillFiles, files...) } return skillFiles, warnings, nil } +func collectRemotePackageSkillDirFiles(ctx context.Context, owner, repo, ref, host, skillDir string, manifestSkillDirSet map[string]struct{}) ([]resolvedPackageSkillFile, []string, error) { + if err := validateRemotePackageSkillMarker(ctx, owner, repo, ref, host, skillDir, manifestSkillDirSet); err != nil { + if errors.Is(err, errRepositoryPackageFileNotFound) { + return nil, []string{fmt.Sprintf("Skill directory %q is missing required %s marker file", skillDir, packageSkillMarkerFile)}, nil + } + return nil, nil, err + } + files, err := listPackageDirFilesRecursivelyForHost(ctx, owner, repo, ref, skillDir, host) + if err != nil { + if isRepositoryFileNotFound(err) { + return nil, []string{fmt.Sprintf("Skill directory %q not found in package, skipping", skillDir)}, nil + } + return nil, nil, fmt.Errorf("failed to list files in skill directory %q: %w", skillDir, err) + } + skillName := filepath.Base(skillDir) + skillFiles := make([]resolvedPackageSkillFile, 0, len(files)) + for _, file := range files { + skillFiles = append(skillFiles, resolvedPackageSkillFile{SourcePath: file, SkillName: skillName}) + } + return skillFiles, nil, nil +} + +func validateRemotePackageSkillMarker(ctx context.Context, owner, repo, ref, host, skillDir string, manifestSkillDirSet map[string]struct{}) error { + if _, fromManifest := manifestSkillDirSet[skillDir]; !fromManifest { + return nil + } + markerPath := joinRepositoryPackagePath(skillDir, packageSkillMarkerFile) + if _, err := downloadPackageFileFromGitHubForHost(ctx, owner, repo, markerPath, ref, host); err != nil { + if isRepositoryFileNotFound(err) { + return errRepositoryPackageFileNotFound + } + return fmt.Errorf("failed to validate skill marker %q: %w", markerPath, err) + } + return nil +} + +// resolvePackageAgentFiles returns the list of agent file source paths for a package. // resolvePackageAgentFiles returns the list of agent file source paths for a package. // If explicitAgentFiles is non-empty it is used; otherwise the agents/ directory is // auto-scanned for .md files. diff --git a/pkg/cli/add_skill_rewrite.go b/pkg/cli/add_skill_rewrite.go index 6b5cafd164c..f192ae65248 100644 --- a/pkg/cli/add_skill_rewrite.go +++ b/pkg/cli/add_skill_rewrite.go @@ -167,95 +167,82 @@ func rewriteLocalSkillRefsInContent(content, repoSlug, headSHA string) (string, // - object block keys: " skill: .github/skills/my-skill" (rare form) func rewriteSkillsInFrontmatterLines(lines []string, repoSlug, headSHA string) []string { newLines := make([]string, 0, len(lines)) - inSkills := false - skillsBaseIndent := -1 - + state := skillRewriteState{skillsBaseIndent: -1} for _, line := range lines { - if line == "" { - newLines = append(newLines, line) - continue - } - - trimmed := strings.TrimSpace(line) - indent := countLeadingSpacesSkill(line) + newLines = append(newLines, rewriteSkillsFrontmatterLine(line, repoSlug, headSHA, &state)) + } + return newLines +} - if !inSkills { - if isSkillsKeyLine(trimmed) { - inSkills = true - skillsBaseIndent = indent - } else if strings.HasPrefix(trimmed, "skills:") { - // Flow-sequence form: skills: [item1, item2, ...] - line = rewriteFlowSkillsLine(line, trimmed, indent, repoSlug, headSHA) - } - newLines = append(newLines, line) - continue - } +type skillRewriteState struct { + inSkills bool + skillsBaseIndent int +} - // A non-list, non-empty line at or below the "skills:" indent level - // means we have left the skills block. - if indent <= skillsBaseIndent && !strings.HasPrefix(trimmed, "-") { - inSkills = false - newLines = append(newLines, line) - continue - } +func rewriteSkillsFrontmatterLine(line, repoSlug, headSHA string, state *skillRewriteState) string { + if line == "" { + return line + } + trimmed := strings.TrimSpace(line) + indent := countLeadingSpacesSkill(line) + if !state.inSkills { + return rewriteSkillsEntryLine(line, trimmed, indent, repoSlug, headSHA, state) + } + if indent <= state.skillsBaseIndent && !strings.HasPrefix(trimmed, "-") { + state.inSkills = false + return line + } + return rewriteSkillsListLine(line, trimmed, indent, repoSlug, headSHA, state.skillsBaseIndent) +} - // Handle list items: " - " or " - skill: " - if strings.HasPrefix(trimmed, "- ") { - itemContent := trimmed[2:] // content after "- " - leadingSpace := strings.Repeat(" ", indent) +func rewriteSkillsEntryLine(line, trimmed string, indent int, repoSlug, headSHA string, state *skillRewriteState) string { + if isSkillsKeyLine(trimmed) { + state.inSkills = true + state.skillsBaseIndent = indent + return line + } + if strings.HasPrefix(trimmed, "skills:") { + return rewriteFlowSkillsLine(line, trimmed, indent, repoSlug, headSHA) + } + return line +} - if rest, ok := strings.CutPrefix(itemContent, "skill:"); ok { - // Object form: "- skill: " - rawVal := strings.TrimSpace(rest) - valPart, comment := splitYAMLValueAndComment(rawVal) - unquoted := trimYAMLQuotesSkill(valPart) - if isLocalSkillRef(unquoted) { - qualified := buildQualifiedSkillRef(unquoted, repoSlug, headSHA) - suffix := "" - if comment != "" { - suffix = " " + comment - } - line = leadingSpace + "- skill: " + qualified + suffix - skillRewriteLog.Printf("Rewrote local skill ref (object form): %q -> %q", unquoted, qualified) - } - } else { - // String form: "- " - valPart, comment := splitYAMLValueAndComment(itemContent) - unquoted := trimYAMLQuotesSkill(valPart) - if isLocalSkillRef(unquoted) { - qualified := buildQualifiedSkillRef(unquoted, repoSlug, headSHA) - suffix := "" - if comment != "" { - suffix = " " + comment - } - line = leadingSpace + "- " + qualified + suffix - skillRewriteLog.Printf("Rewrote local skill ref (string form): %q -> %q", unquoted, qualified) - } - } - } else if indent > skillsBaseIndent && strings.HasPrefix(trimmed, "skill:") { - // Object block key on its own line (rare YAML form where the list - // item marker was on the previous line and skill: is indented): - // - - // skill: .github/skills/my-skill - rawVal := strings.TrimSpace(strings.TrimPrefix(trimmed, "skill:")) - valPart, comment := splitYAMLValueAndComment(rawVal) - unquoted := trimYAMLQuotesSkill(valPart) - if isLocalSkillRef(unquoted) { - qualified := buildQualifiedSkillRef(unquoted, repoSlug, headSHA) - leadingSpace := strings.Repeat(" ", indent) - suffix := "" - if comment != "" { - suffix = " " + comment - } - line = leadingSpace + "skill: " + qualified + suffix - skillRewriteLog.Printf("Rewrote local skill ref (block object form): %q -> %q", unquoted, qualified) - } - } +func rewriteSkillsListLine(line, trimmed string, indent int, repoSlug, headSHA string, skillsBaseIndent int) string { + switch { + case strings.HasPrefix(trimmed, "- "): + return rewriteInlineSkillListLine(trimmed[2:], indent, repoSlug, headSHA) + case indent > skillsBaseIndent && strings.HasPrefix(trimmed, "skill:"): + return rewriteBlockSkillLine(strings.TrimSpace(strings.TrimPrefix(trimmed, "skill:")), indent, repoSlug, headSHA) + default: + return line + } +} - newLines = append(newLines, line) +func rewriteInlineSkillListLine(itemContent string, indent int, repoSlug, headSHA string) string { + leadingSpace := strings.Repeat(" ", indent) + if rest, ok := strings.CutPrefix(itemContent, "skill:"); ok { + return rewriteQualifiedSkillLine(strings.TrimSpace(rest), leadingSpace+"- skill: ", "object form", repoSlug, headSHA) } + return rewriteQualifiedSkillLine(itemContent, leadingSpace+"- ", "string form", repoSlug, headSHA) +} - return newLines +func rewriteBlockSkillLine(rawVal string, indent int, repoSlug, headSHA string) string { + return rewriteQualifiedSkillLine(rawVal, strings.Repeat(" ", indent)+"skill: ", "block object form", repoSlug, headSHA) +} + +func rewriteQualifiedSkillLine(rawVal, prefix, form, repoSlug, headSHA string) string { + valPart, comment := splitYAMLValueAndComment(rawVal) + unquoted := trimYAMLQuotesSkill(valPart) + if !isLocalSkillRef(unquoted) { + return prefix + rawVal + } + qualified := buildQualifiedSkillRef(unquoted, repoSlug, headSHA) + suffix := "" + if comment != "" { + suffix = " " + comment + } + skillRewriteLog.Printf("Rewrote local skill ref (%s): %q -> %q", form, unquoted, qualified) + return prefix + qualified + suffix } // trimYAMLQuotesSkill strips a single layer of matching single or double diff --git a/pkg/cli/add_wizard_command.go b/pkg/cli/add_wizard_command.go index b0613619592..d1989cca63b 100644 --- a/pkg/cli/add_wizard_command.go +++ b/pkg/cli/add_wizard_command.go @@ -13,7 +13,15 @@ var addWizardLog = logger.New("cli:add_wizard_command") // NewAddWizardCommand creates the add-wizard command, which is always interactive. func NewAddWizardCommand(validateEngine func(string) error) *cobra.Command { - cmd := &cobra.Command{ + cmd := newAddWizardCommand(validateEngine) + configureAddWizardFlags(cmd) + RegisterEngineFlagCompletion(cmd) + RegisterDirFlagCompletion(cmd, "dir") + return cmd +} + +func newAddWizardCommand(validateEngine func(string) error) *cobra.Command { + return &cobra.Command{ Use: "add-wizard ...", Short: "Interactively add one or more agentic workflows with guided setup", Long: `Interactively add one or more agentic workflows with guided setup. @@ -55,85 +63,79 @@ Note: To create a new workflow from scratch, use the 'new' command instead.`, ` + string(constants.CLIExtensionPrefix) + ` add-wizard githubnext/agentics/ci-doctor --append "custom footer" # Append custom content ` + string(constants.CLIExtensionPrefix) + ` add-wizard githubnext/agentics/ci-doctor --no-security-scanner # Skip security scan `, - Args: func(cmd *cobra.Command, args []string) error { - if len(args) < 1 { - return errors.New("missing workflow specification\n\nRun 'gh aw add-wizard --help' for usage information") - } - return nil - }, - RunE: func(cmd *cobra.Command, args []string) error { - workflows := args - engineOverride, _ := cmd.Flags().GetString("engine") - verbose, _ := cmd.Flags().GetBool("verbose") - noGitattributes, _ := cmd.Flags().GetBool("no-gitattributes") - workflowDir, _ := cmd.Flags().GetString("dir") - noStopAfter, _ := cmd.Flags().GetBool("no-stop-after") - stopAfter, _ := cmd.Flags().GetString("stop-after") - noSecret, _ := cmd.Flags().GetBool("no-secret") - skipSecretLegacy, _ := cmd.Flags().GetBool("skip-secret") - skipSecret := noSecret || skipSecretLegacy - appendText, _ := cmd.Flags().GetString("append") - disableSecurityScanner := resolveDeprecatedBoolFlag(cmd, "no-security-scanner", "disable-security-scanner") + Args: validateAddWizardArgs, + RunE: runAddWizardCommand(validateEngine), + } +} - addWizardLog.Printf("Starting add-wizard: workflows=%v, engine=%s, verbose=%v", workflows, engineOverride, verbose) +func validateAddWizardArgs(cmd *cobra.Command, args []string) error { + if len(args) < 1 { + return errors.New("missing workflow specification\n\nRun 'gh aw add-wizard --help' for usage information") + } + return nil +} - if err := validateEngine(engineOverride); err != nil { - return err - } +func runAddWizardCommand(validateEngine func(string) error) func(*cobra.Command, []string) error { + return func(cmd *cobra.Command, args []string) error { + config := addWizardConfigFromFlags(cmd, args) + addWizardLog.Printf("Starting add-wizard: workflows=%v, engine=%s, verbose=%v", config.WorkflowSpecs, config.EngineOverride, config.Verbose) + if err := validateEngine(config.EngineOverride); err != nil { + return err + } + if err := validateAddWizardInteractiveTerminal(); err != nil { + return err + } + return RunAddInteractive(cmd.Context(), config) + } +} + +func addWizardConfigFromFlags(cmd *cobra.Command, workflows []string) *AddInteractiveConfig { + engineOverride, _ := cmd.Flags().GetString("engine") + verbose, _ := cmd.Flags().GetBool("verbose") + noGitattributes, _ := cmd.Flags().GetBool("no-gitattributes") + workflowDir, _ := cmd.Flags().GetString("dir") + noStopAfter, _ := cmd.Flags().GetBool("no-stop-after") + stopAfter, _ := cmd.Flags().GetString("stop-after") + appendText, _ := cmd.Flags().GetString("append") + return &AddInteractiveConfig{ + WorkflowSpecs: workflows, + Verbose: verbose, + EngineOverride: engineOverride, + NoGitattributes: noGitattributes, + WorkflowDir: workflowDir, + NoStopAfter: noStopAfter, + StopAfter: stopAfter, + SkipSecret: resolveAddWizardSkipSecret(cmd), + AppendText: appendText, + DisableSecurityScanner: resolveDeprecatedBoolFlag(cmd, "no-security-scanner", "disable-security-scanner"), + } +} - // add-wizard requires an interactive terminal - isTerminal := tty.IsStdoutTerminal() - isCIEnv := IsRunningInCI() - addWizardLog.Printf("Terminal check: is_terminal=%v, is_ci=%v", isTerminal, isCIEnv) - if !isTerminal || isCIEnv { - return errors.New("add-wizard requires an interactive terminal; use 'add' for non-interactive environments") - } +func resolveAddWizardSkipSecret(cmd *cobra.Command) bool { + noSecret, _ := cmd.Flags().GetBool("no-secret") + skipSecretLegacy, _ := cmd.Flags().GetBool("skip-secret") + return noSecret || skipSecretLegacy +} - return RunAddInteractive(cmd.Context(), &AddInteractiveConfig{ - WorkflowSpecs: workflows, - Verbose: verbose, - EngineOverride: engineOverride, - NoGitattributes: noGitattributes, - WorkflowDir: workflowDir, - NoStopAfter: noStopAfter, - StopAfter: stopAfter, - SkipSecret: skipSecret, - AppendText: appendText, - DisableSecurityScanner: disableSecurityScanner, - }) - }, +func validateAddWizardInteractiveTerminal() error { + isTerminal := tty.IsStdoutTerminal() + isCIEnv := IsRunningInCI() + addWizardLog.Printf("Terminal check: is_terminal=%v, is_ci=%v", isTerminal, isCIEnv) + if !isTerminal || isCIEnv { + return errors.New("add-wizard requires an interactive terminal; use 'add' for non-interactive environments") } + return nil +} - // Add AI engine flag +func configureAddWizardFlags(cmd *cobra.Command) { addEngineFlag(cmd) - - // Add no-gitattributes flag cmd.Flags().Bool("no-gitattributes", false, "Skip updating .gitattributes file") - - // Add workflow directory flag cmd.Flags().StringP("dir", "d", "", "Workflow directory (default: $GH_AW_WORKFLOWS_DIR or .github/workflows)") - - // Add no-stop-after flag cmd.Flags().Bool("no-stop-after", false, "Remove any stop-after field from the workflow") - - // Add stop-after flag cmd.Flags().String("stop-after", "", "Override stop-after value in the workflow (e.g., '+48h', '2025-12-31 23:59:59')") - - // Add no-secret flag (--skip-secret is kept as an undocumented alias) cmd.Flags().Bool("no-secret", false, "Skip the API secret prompt (use when the secret is already set at the org or repo level)") cmd.Flags().Bool("skip-secret", false, "Skip the API secret prompt (use when the secret is already set at the org or repo level)") _ = cmd.Flags().MarkHidden("skip-secret") - - // Add append flag (matches --append in add command) cmd.Flags().String("append", "", "Append extra content to the end of the agentic workflow on installation") - - // Add no-security-scanner flag (--disable-security-scanner is kept as a deprecated alias - // for consistency with add and other install entry points) addSecurityScannerFlag(cmd) - - // Register completions - RegisterEngineFlagCompletion(cmd) - RegisterDirFlagCompletion(cmd, "dir") - - return cmd } diff --git a/pkg/cli/add_workflow_pr.go b/pkg/cli/add_workflow_pr.go index 1bd62c28be0..17930344ec2 100644 --- a/pkg/cli/add_workflow_pr.go +++ b/pkg/cli/add_workflow_pr.go @@ -52,103 +52,88 @@ func sanitizeBranchName(name string) string { // addWorkflowsWithPR handles workflow addition with PR creation using pre-resolved workflows. func addWorkflowsWithPR(ctx context.Context, workflows []*ResolvedWorkflow, opts AddOptions) (int, string, error) { addWorkflowPRLog.Printf("Adding %d workflow(s) with PR creation (resolved)", len(workflows)) + currentBranch, branchName, tracker, err := prepareWorkflowPRBranch(workflows, opts) + if err != nil { + return 0, "", err + } + defer restoreWorkflowPRBranch(currentBranch, opts.Verbose) + if err := stageWorkflowsForPR(ctx, workflows, tracker, opts); err != nil { + return 0, "", err + } + commitMessage, prTitle, prBody := workflowPRMessages(workflows) + if err := commitWorkflowPRChanges(commitMessage, branchName, prTitle, opts.Verbose); err != nil { + return 0, "", err + } + if err := pushWorkflowPRBranch(branchName, prTitle, opts.Verbose); err != nil { + return 0, "", err + } + return createWorkflowPullRequest(ctx, tracker, currentBranch, branchName, prTitle, prBody, opts.Verbose) +} - // Get current branch for restoration later +func prepareWorkflowPRBranch(workflows []*ResolvedWorkflow, opts AddOptions) (string, string, *FileTracker, error) { currentBranch, err := getCurrentBranch() if err != nil { addWorkflowPRLog.Printf("Failed to get current branch: %v", err) - return 0, "", fmt.Errorf("failed to get current branch: %w", err) + return "", "", nil, fmt.Errorf("failed to get current branch: %w", err) } - - addWorkflowPRLog.Printf("Current branch: %s", currentBranch) - - // Create temporary branch with random 4-digit number - // Use sanitized workflow name to avoid invalid git ref characters - randomNum := rand.Intn(9000) + 1000 // Generate number between 1000-9999 - sanitizedName := sanitizeBranchName(workflows[0].Spec.WorkflowPath) - branchName := fmt.Sprintf("add-workflow-%s-%04d", sanitizedName, randomNum) - + branchName := fmt.Sprintf("add-workflow-%s-%04d", sanitizeBranchName(workflows[0].Spec.WorkflowPath), rand.Intn(9000)+1000) addWorkflowPRLog.Printf("Creating temporary branch: %s", branchName) - if err := createAndSwitchBranch(branchName, opts.Verbose); err != nil { - return 0, "", fmt.Errorf("failed to create branch %s: %w", branchName, err) + return "", "", nil, fmt.Errorf("failed to create branch %s: %w", branchName, err) } + return currentBranch, branchName, NewFileTracker(), nil +} - // Create file tracker for rollback capability - tracker := NewFileTracker() - - // Ensure we switch back to original branch on exit - defer func() { - if switchErr := switchBranch(currentBranch, opts.Verbose); switchErr != nil && opts.Verbose { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Failed to switch back to branch %s: %v", currentBranch, switchErr))) - } - }() +func restoreWorkflowPRBranch(currentBranch string, verbose bool) { + if switchErr := switchBranch(currentBranch, verbose); switchErr != nil && verbose { + fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Failed to switch back to branch %s: %v", currentBranch, switchErr))) + } +} - // Add workflows using the resolved workflow path +func stageWorkflowsForPR(ctx context.Context, workflows []*ResolvedWorkflow, tracker *FileTracker, opts AddOptions) error { addWorkflowPRLog.Print("Adding workflows to repository") - prOpts := opts - if err := addWorkflowsWithTracking(ctx, workflows, tracker, prOpts); err != nil { + if err := addWorkflowsWithTracking(ctx, workflows, tracker, opts); err != nil { addWorkflowPRLog.Printf("Failed to add workflows: %v", err) - return 0, "", fmt.Errorf("failed to add workflows: %w", err) + return fmt.Errorf("failed to add workflows: %w", err) } - - // Stage all files before creating PR addWorkflowPRLog.Print("Staging workflow files") if err := tracker.StageAllFiles(opts.Verbose); err != nil { - if rollbackErr := tracker.RollbackAllFiles(opts.Verbose); rollbackErr != nil && opts.Verbose { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Failed to rollback files: %v", rollbackErr))) - } - return 0, "", fmt.Errorf("failed to stage workflow files: %w", err) + rollbackWorkflowPRFiles(tracker, opts.Verbose) + return fmt.Errorf("failed to stage workflow files: %w", err) } - - // Update .gitattributes and stage it if changed if err := stageGitAttributesIfChanged(); err != nil && opts.Verbose { fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Failed to stage .gitattributes: %v", err))) } + return nil +} + +func rollbackWorkflowPRFiles(tracker *FileTracker, verbose bool) { + if rollbackErr := tracker.RollbackAllFiles(verbose); rollbackErr != nil && verbose { + fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Failed to rollback files: %v", rollbackErr))) + } +} - // Commit changes - var commitMessage, prTitle, prBody, joinedNames string +func workflowPRMessages(workflows []*ResolvedWorkflow) (string, string, string) { if len(workflows) == 1 { - joinedNames = workflows[0].Spec.WorkflowName - commitMessage = "Add agentic workflow " + joinedNames - prTitle = "Add agentic workflow " + joinedNames - prBody = "Add agentic workflow " + joinedNames - } else { - workflowNames := sliceutil.Map(workflows, func(wf *ResolvedWorkflow) string { - return wf.Spec.WorkflowName - }) - joinedNames = strings.Join(workflowNames, ", ") - commitMessage = "Add agentic workflows: " + joinedNames - prTitle = "Add agentic workflows: " + joinedNames - prBody = "Add agentic workflows: " + joinedNames - } - - if err := commitChanges(commitMessage, opts.Verbose); err != nil { - // Don't rollback - leave the workflow files on disk for manual recovery. - // Return a richly formatted error with clear instructions so the user can - // commit and push manually. The top-level error handler will print this. - return 0, "", fmt.Errorf( - "failed to commit workflow files: %w\n\n"+ - "The workflow files have been written to disk and staged in git.\n"+ - "Please commit the files manually, then either push them to the\n"+ - "repository or create a pull request:\n\n"+ - " git commit -m %q\n"+ - " git push\n\n"+ - "Or to create a pull request:\n\n"+ - " git checkout -b %s\n"+ - " git commit -m %q\n"+ - " git push -u origin %s\n"+ - " gh pr create --title %q", - err, commitMessage, branchName, commitMessage, branchName, prTitle, - ) - } - - // Push branch + message := "Add agentic workflow " + workflows[0].Spec.WorkflowName + return message, message, message + } + workflowNames := sliceutil.Map(workflows, func(wf *ResolvedWorkflow) string { return wf.Spec.WorkflowName }) + message := "Add agentic workflows: " + strings.Join(workflowNames, ", ") + return message, message, message +} + +func commitWorkflowPRChanges(commitMessage, branchName, prTitle string, verbose bool) error { + if err := commitChanges(commitMessage, verbose); err != nil { + return fmt.Errorf("failed to commit workflow files: %w\n\nThe workflow files have been written to disk and staged in git.\nPlease commit the files manually, then either push them to the\nrepository or create a pull request:\n\n git commit -m %q\n git push\n\nOr to create a pull request:\n\n git checkout -b %s\n git commit -m %q\n git push -u origin %s\n gh pr create --title %q", err, commitMessage, branchName, commitMessage, branchName, prTitle) + } + return nil +} + +func pushWorkflowPRBranch(branchName, prTitle string, verbose bool) error { addWorkflowPRLog.Printf("Pushing branch %s to remote", branchName) - if err := pushBranch(branchName, opts.Verbose); err != nil { + if err := pushBranch(branchName, verbose); err != nil { addWorkflowPRLog.Printf("Failed to push branch: %v", err) - // Treat push failure as a warning: keep the files and commit intact so the - // user can push manually. Do NOT rollback. fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Failed to push branch %s: %v", branchName, err))) fmt.Fprintln(os.Stderr, console.FormatInfoMessage( "The workflow files have been committed to local branch "+branchName+".\n"+ @@ -156,27 +141,23 @@ func addWorkflowsWithPR(ctx context.Context, workflows []*ResolvedWorkflow, opts " git push -u origin "+branchName+"\n"+ " gh pr create --title "+fmt.Sprintf("%q", prTitle), )) - return 0, "", fmt.Errorf("failed to push branch %s: %w", branchName, err) + return fmt.Errorf("failed to push branch %s: %w", branchName, err) } + return nil +} - // Create PR +func createWorkflowPullRequest(ctx context.Context, tracker *FileTracker, currentBranch, branchName, prTitle, prBody string, verbose bool) (int, string, error) { addWorkflowPRLog.Printf("Creating pull request: %s", prTitle) - prNumber, prURL, err := createPR(ctx, branchName, prTitle, prBody, opts.Verbose) + prNumber, prURL, err := createPR(ctx, branchName, prTitle, prBody, verbose) if err != nil { addWorkflowPRLog.Printf("Failed to create PR: %v", err) - if rollbackErr := tracker.RollbackAllFiles(opts.Verbose); rollbackErr != nil && opts.Verbose { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Failed to rollback files: %v", rollbackErr))) - } + rollbackWorkflowPRFiles(tracker, verbose) return 0, "", fmt.Errorf("failed to create PR: %w", err) } - addWorkflowPRLog.Printf("Successfully created PR #%d: %s", prNumber, prURL) - - // Switch back to original branch - if err := switchBranch(currentBranch, opts.Verbose); err != nil { + if err := switchBranch(currentBranch, verbose); err != nil { return prNumber, prURL, fmt.Errorf("failed to switch back to branch %s: %w", currentBranch, err) } - fmt.Fprintln(os.Stderr, console.FormatSuccessMessage("Created pull request "+prURL)) return prNumber, prURL, nil } diff --git a/pkg/cli/add_workflow_resolution.go b/pkg/cli/add_workflow_resolution.go index 900d357f3ba..0922ce60de1 100644 --- a/pkg/cli/add_workflow_resolution.go +++ b/pkg/cli/add_workflow_resolution.go @@ -437,23 +437,17 @@ func resolveLocalRepositoryPackage(source string) (*resolvedRepositoryPackage, e if !isLocalWorkflowPath(source) { return nil, nil } - manifestPath, packageDir, err := localRepositoryPackageManifest(source) + if shouldIgnoreMissingLocalPackageManifest(manifestPath, err) { + return nil, nil + } if err != nil { - if errors.Is(err, os.ErrNotExist) { - return nil, nil - } return nil, err } - if manifestPath == "" { - return nil, nil - } - content, err := os.ReadFile(manifestPath) if err != nil { return nil, fmt.Errorf("failed to read Agentic Workflow manifest %q: %w", manifestPath, err) } - manifest, warnings, err := parseRepositoryPackageManifest(manifestPath, content) if err != nil { return nil, err @@ -461,49 +455,49 @@ func resolveLocalRepositoryPackage(source string) (*resolvedRepositoryPackage, e if err := validateLocalRepositoryPackageContents(manifestPath); err != nil { return nil, err } + return buildResolvedLocalRepositoryPackage(packageDir, manifestPath, manifest, warnings) +} - includeInstallablePaths, includeSkillDirs, includeAgentFiles := splitManifestIncludePaths(manifest.Includes) - includeInstallablePaths = append(includeInstallablePaths, manifest.Files...) - installationSources := normalizeLocalPackageInstallablePaths(includeInstallablePaths, packageDir) - if len(installationSources) == 0 { - installationSources, err = scanLocalRepositoryPackageInstallablePaths(packageDir) - if err != nil { - return nil, err - } - } - if err := validateUniqueManifestWorkflowFilenames(installationSources, manifestPath); err != nil { +func shouldIgnoreMissingLocalPackageManifest(manifestPath string, err error) bool { + return (err != nil && errors.Is(err, os.ErrNotExist)) || manifestPath == "" +} + +func buildResolvedLocalRepositoryPackage(packageDir, manifestPath string, manifest *repositoryPackageManifest, warnings []string) (*resolvedRepositoryPackage, error) { + installationSources, skillDirs, agentFiles, err := resolveLocalRepositoryPackageContent(packageDir, manifest, manifestPath) + if err != nil { return nil, err } - - skillFiles, skillWarnings, err := resolveLocalPackageSkillFiles(packageDir, append(append([]string{}, manifest.Skills...), includeSkillDirs...)) + skillFiles, skillWarnings, err := resolveLocalPackageSkillFiles(packageDir, skillDirs) if err != nil { return nil, err } - warnings = append(warnings, skillWarnings...) - - agentFiles, agentWarnings, err := resolveLocalPackageAgentFiles(packageDir, append(append([]string{}, manifest.Agents...), includeAgentFiles...)) + agentFiles, agentWarnings, err := resolveLocalPackageAgentFiles(packageDir, agentFiles) if err != nil { return nil, err } + warnings = append(warnings, skillWarnings...) warnings = append(warnings, agentWarnings...) - if len(installationSources) == 0 && len(skillFiles) == 0 && len(agentFiles) == 0 { return nil, fmt.Errorf("repository package at %q does not contain any installable workflows, skills, or agents (either explicitly declared or auto-discovered)", packageDir) } + return &resolvedRepositoryPackage{ManifestPath: manifestPath, Name: manifest.Name, Emoji: manifest.Emoji, Description: manifest.Description, License: manifest.License, DocsPath: filepath.Join(packageDir, "README.md"), InstallationSource: installationSources, Bootstrap: manifest.Bootstrap, SkillFiles: skillFiles, AgentFiles: agentFiles, Warnings: warnings}, nil +} - return &resolvedRepositoryPackage{ - ManifestPath: manifestPath, - Name: manifest.Name, - Emoji: manifest.Emoji, - Description: manifest.Description, - License: manifest.License, - DocsPath: filepath.Join(packageDir, "README.md"), - InstallationSource: installationSources, - Bootstrap: manifest.Bootstrap, - SkillFiles: skillFiles, - AgentFiles: agentFiles, - Warnings: warnings, - }, nil +func resolveLocalRepositoryPackageContent(packageDir string, manifest *repositoryPackageManifest, manifestPath string) ([]string, []string, []string, error) { + includeInstallablePaths, includeSkillDirs, includeAgentFiles := splitManifestIncludePaths(manifest.Includes) + includeInstallablePaths = append(includeInstallablePaths, manifest.Files...) + installationSources := normalizeLocalPackageInstallablePaths(includeInstallablePaths, packageDir) + if len(installationSources) == 0 { + var err error + installationSources, err = scanLocalRepositoryPackageInstallablePaths(packageDir) + if err != nil { + return nil, nil, nil, err + } + } + if err := validateUniqueManifestWorkflowFilenames(installationSources, manifestPath); err != nil { + return nil, nil, nil, err + } + return installationSources, append([]string{}, append(manifest.Skills, includeSkillDirs...)...), append([]string{}, append(manifest.Agents, includeAgentFiles...)...), nil } func localRepositoryPackageManifest(source string) (string, string, error) { @@ -586,70 +580,112 @@ func appendLocalRepositoryPackageWorkflowSpecs(parsedSpecs []*WorkflowSpec, pkg } func resolveLocalPackageSkillFiles(packageDir string, explicitSkillDirs []string) ([]resolvedPackageSkillFile, []string, error) { - seenSkillDirs := make(map[string]struct{}) - var warnings []string + skillDirs := uniqueLocalPackageSkillDirs(packageDir, explicitSkillDirs) + autoScanned, warnings, err := autoScanLocalPackageSkillDirs(packageDir, len(skillDirs) > 0) + if err != nil { + return nil, nil, err + } + skillDirs = uniqueLocalSkillDirList(skillDirs, autoScanned) + manifestSkillDirSet := localManifestSkillDirSet(packageDir, explicitSkillDirs) + skillFiles, fileWarnings, err := collectLocalPackageSkillFiles(skillDirs, manifestSkillDirSet) + if err != nil { + return nil, nil, err + } + return skillFiles, append(warnings, fileWarnings...), nil +} - var skillDirs []string - appendIfNew := func(dir string) { +func uniqueLocalPackageSkillDirs(packageDir string, explicitSkillDirs []string) []string { + skillDirs := make([]string, 0, len(explicitSkillDirs)) + for _, dir := range explicitSkillDirs { + skillDirs = append(skillDirs, filepath.Clean(filepath.Join(packageDir, filepath.FromSlash(dir)))) + } + return skillDirs +} + +func autoScanLocalPackageSkillDirs(packageDir string, hasManifestSkills bool) ([]string, []string, error) { + autoScanned, err := scanLocalPackageSkillDirs(packageDir) + if err == nil { + return autoScanned, nil, nil + } + if hasManifestSkills { + return nil, []string{fmt.Sprintf("failed to auto-scan skills directory, proceeding with manifest skills only: %v", err)}, nil + } + return nil, nil, err +} + +func uniqueLocalSkillDirList(existing, autoScanned []string) []string { + seenSkillDirs := make(map[string]struct{}) + skillDirs := make([]string, 0, len(existing)+len(autoScanned)) + for _, dir := range append(append([]string{}, existing...), autoScanned...) { cleaned := filepath.Clean(dir) if _, exists := seenSkillDirs[cleaned]; exists { - return + continue } seenSkillDirs[cleaned] = struct{}{} skillDirs = append(skillDirs, cleaned) } + return skillDirs +} +func localManifestSkillDirSet(packageDir string, explicitSkillDirs []string) map[string]struct{} { + manifestSkillDirSet := make(map[string]struct{}, len(explicitSkillDirs)) for _, dir := range explicitSkillDirs { - appendIfNew(filepath.Join(packageDir, filepath.FromSlash(dir))) + manifestSkillDirSet[filepath.Clean(filepath.Join(packageDir, filepath.FromSlash(dir)))] = struct{}{} } - autoScanned, err := scanLocalPackageSkillDirs(packageDir) - if err != nil { - if len(skillDirs) == 0 { + return manifestSkillDirSet +} + +func collectLocalPackageSkillFiles(skillDirs []string, manifestSkillDirSet map[string]struct{}) ([]resolvedPackageSkillFile, []string, error) { + var warnings []string + var skillFiles []resolvedPackageSkillFile + for _, skillDir := range skillDirs { + files, fileWarnings, err := collectLocalSkillDirFiles(skillDir, manifestSkillDirSet) + if err != nil { return nil, nil, err } - warnings = append(warnings, fmt.Sprintf("failed to auto-scan skills directory, proceeding with manifest skills only: %v", err)) - } - for _, dir := range autoScanned { - appendIfNew(dir) + warnings = append(warnings, fileWarnings...) + skillFiles = append(skillFiles, files...) } + return skillFiles, warnings, nil +} - manifestSkillDirSet := make(map[string]struct{}, len(explicitSkillDirs)) - for _, dir := range explicitSkillDirs { - manifestSkillDirSet[filepath.Clean(filepath.Join(packageDir, filepath.FromSlash(dir)))] = struct{}{} +func collectLocalSkillDirFiles(skillDir string, manifestSkillDirSet map[string]struct{}) ([]resolvedPackageSkillFile, []string, error) { + if err := validateLocalSkillMarker(skillDir, manifestSkillDirSet); err != nil { + if errors.Is(err, os.ErrNotExist) { + return nil, []string{fmt.Sprintf("Skill directory %q is missing required %s marker file", skillDir, packageSkillMarkerFile)}, nil + } + return nil, nil, err } - + skillName := filepath.Base(skillDir) var skillFiles []resolvedPackageSkillFile - for _, skillDir := range skillDirs { - if _, fromManifest := manifestSkillDirSet[skillDir]; fromManifest { - markerPath := filepath.Join(skillDir, packageSkillMarkerFile) - if _, err := os.Stat(markerPath); err != nil { - if errors.Is(err, os.ErrNotExist) { - warnings = append(warnings, fmt.Sprintf("Skill directory %q is missing required %s marker file", skillDir, packageSkillMarkerFile)) - continue - } - return nil, nil, fmt.Errorf("failed to validate skill marker %q: %w", markerPath, err) - } + err := filepath.WalkDir(skillDir, func(currentPath string, d os.DirEntry, walkErr error) error { + if walkErr != nil { + return walkErr } - skillName := filepath.Base(skillDir) - err := filepath.WalkDir(skillDir, func(currentPath string, d os.DirEntry, walkErr error) error { - if walkErr != nil { - return walkErr - } - if d.IsDir() { - return nil - } - skillFiles = append(skillFiles, resolvedPackageSkillFile{ - SourcePath: currentPath, - SkillName: skillName, - }) + if d.IsDir() { return nil - }) - if err != nil { - return nil, nil, fmt.Errorf("failed to list files in skill directory %q: %w", skillDir, err) } + skillFiles = append(skillFiles, resolvedPackageSkillFile{SourcePath: currentPath, SkillName: skillName}) + return nil + }) + if err != nil { + return nil, nil, fmt.Errorf("failed to list files in skill directory %q: %w", skillDir, err) } + return skillFiles, nil, nil +} - return skillFiles, warnings, nil +func validateLocalSkillMarker(skillDir string, manifestSkillDirSet map[string]struct{}) error { + if _, fromManifest := manifestSkillDirSet[skillDir]; !fromManifest { + return nil + } + markerPath := filepath.Join(skillDir, packageSkillMarkerFile) + if _, err := os.Stat(markerPath); err != nil { + if errors.Is(err, os.ErrNotExist) { + return err + } + return fmt.Errorf("failed to validate skill marker %q: %w", markerPath, err) + } + return nil } func resolveLocalPackageAgentFiles(packageDir string, explicitAgentFiles []string) ([]string, []string, error) { @@ -711,135 +747,112 @@ func appendRepositoryPackageWorkflowSpecs(parsedSpecs []*WorkflowSpec, repoSpec } host := explicitHostForRepo(repoSpec.RepoSlug) effectiveVersion := repositoryPackageEffectiveRef(repoSpec, pkg) + parsedSpecs = appendRepositoryPackageInstallableSpecs(parsedSpecs, repoSpec, pkg, host, effectiveVersion) + parsedSpecs = appendRepositoryPackageSkillSpecs(parsedSpecs, repoSpec, pkg, host, effectiveVersion) + return appendRepositoryPackageAgentSpecs(parsedSpecs, repoSpec, pkg, host, effectiveVersion) +} + +func appendRepositoryPackageInstallableSpecs(parsedSpecs []*WorkflowSpec, repoSpec *RepoSpec, pkg *resolvedRepositoryPackage, host, effectiveVersion string) []*WorkflowSpec { for _, installationSource := range pkg.InstallationSource { - // installationSource is guaranteed by isSupportedPackageInstallablePath to be - // either a .md agentic workflow or a .yml action workflow file; no other - // extensions can reach this point. base := filepath.Base(installationSource) - // Use filepath.Ext for case-insensitive extension removal (e.g. ".YML" or ".MD"). workflowName := strings.TrimSuffix(base, filepath.Ext(base)) - parsedSpecs = append(parsedSpecs, &WorkflowSpec{ - RepoSpec: RepoSpec{ - RepoSlug: repoSpec.RepoSlug, - Version: effectiveVersion, - PackagePath: repoSpec.PackagePath, - }, - WorkflowPath: installationSource, - WorkflowName: workflowName, - Host: host, - FromRepositoryManifest: true, - }) + parsedSpecs = append(parsedSpecs, repositoryPackageWorkflowSpec(repoSpec, host, effectiveVersion, installationSource, workflowName)) } + return parsedSpecs +} - // Append skill file specs. Each spec carries IsPackageSkillFile=true and the SkillName - // so that the installation step can route the file to the correct skill directory. +func appendRepositoryPackageSkillSpecs(parsedSpecs []*WorkflowSpec, repoSpec *RepoSpec, pkg *resolvedRepositoryPackage, host, effectiveVersion string) []*WorkflowSpec { for _, skillFile := range pkg.SkillFiles { base := filepath.Base(skillFile.SourcePath) - // WorkflowName is unused for skill files but set to a stable value for logging. workflowName := skillFile.SkillName + "/" + strings.TrimSuffix(base, filepath.Ext(base)) - parsedSpecs = append(parsedSpecs, &WorkflowSpec{ - RepoSpec: RepoSpec{ - RepoSlug: repoSpec.RepoSlug, - Version: effectiveVersion, - PackagePath: repoSpec.PackagePath, - }, - WorkflowPath: skillFile.SourcePath, - WorkflowName: workflowName, - Host: host, - IsPackageSkillFile: true, - SkillName: skillFile.SkillName, - }) + spec := repositoryPackageWorkflowSpec(repoSpec, host, effectiveVersion, skillFile.SourcePath, workflowName) + spec.IsPackageSkillFile = true + spec.SkillName = skillFile.SkillName + parsedSpecs = append(parsedSpecs, spec) } + return parsedSpecs +} - // Append agent file specs. Each spec carries IsPackageAgentFile=true so the installation - // step routes the file to the correct agents directory. +func appendRepositoryPackageAgentSpecs(parsedSpecs []*WorkflowSpec, repoSpec *RepoSpec, pkg *resolvedRepositoryPackage, host, effectiveVersion string) []*WorkflowSpec { for _, agentFile := range pkg.AgentFiles { base := filepath.Base(agentFile) workflowName := strings.TrimSuffix(base, filepath.Ext(base)) - parsedSpecs = append(parsedSpecs, &WorkflowSpec{ - RepoSpec: RepoSpec{ - RepoSlug: repoSpec.RepoSlug, - Version: effectiveVersion, - PackagePath: repoSpec.PackagePath, - }, - WorkflowPath: agentFile, - WorkflowName: workflowName, - Host: host, - IsPackageAgentFile: true, - }) + spec := repositoryPackageWorkflowSpec(repoSpec, host, effectiveVersion, agentFile, workflowName) + spec.IsPackageAgentFile = true + parsedSpecs = append(parsedSpecs, spec) } - return parsedSpecs } +func repositoryPackageWorkflowSpec(repoSpec *RepoSpec, host, effectiveVersion, workflowPath, workflowName string) *WorkflowSpec { + return &WorkflowSpec{RepoSpec: RepoSpec{RepoSlug: repoSpec.RepoSlug, Version: effectiveVersion, PackagePath: repoSpec.PackagePath}, WorkflowPath: workflowPath, WorkflowName: workflowName, Host: host, FromRepositoryManifest: true} +} + func resolveAddWorkflowSpecAndContent(ctx context.Context, initialSpec *WorkflowSpec, verbose bool) (*WorkflowSpec, *FetchedWorkflow, error) { currentSpec := *initialSpec visited := make(map[string]struct{}) followedRedirect := false - for range maxRedirectDepth { - // Fetch workflow content - handles both local and remote. fetched, err := fetchWorkflowFromSourceWithContextFn(ctx, ¤tSpec, verbose) if err != nil { return nil, nil, err } - - // Redirects only apply to remote workflows. if fetched.IsLocal { return ¤tSpec, fetched, nil } - - currentRef := currentSpec.Version - if currentRef == "" { - currentRef = "main" - } - locationKey := fmt.Sprintf("%s/%s@%s", currentSpec.RepoSlug, currentSpec.WorkflowPath, currentRef) - if _, exists := visited[locationKey]; exists { - return nil, nil, fmt.Errorf("redirect loop detected at %s", locationKey) + locationKey := workflowRedirectLocationKey(¤tSpec) + if err := trackWorkflowRedirectVisit(visited, locationKey); err != nil { + return nil, nil, err } - visited[locationKey] = struct{}{} - redirect, err := extractRedirectFromContent(string(fetched.Content)) if err != nil { return nil, nil, err } if redirect == "" { - // Preserve the original WorkflowName from the user's request only when - // one or more redirects were followed, so the final local file keeps - // the requested name. - // Without redirects, keep any name derived during fetch, such as JSON - // imports where conversion picks a better filename from `name`. if followedRedirect { currentSpec.WorkflowName = initialSpec.WorkflowName } return ¤tSpec, fetched, nil } - - redirectedSource, err := normalizeRedirectToSourceSpec(redirect) + currentSpec, err = followAddWorkflowRedirect(redirect, locationKey, currentSpec.Host, verbose) if err != nil { - return nil, nil, fmt.Errorf("invalid redirect %q in %s: %w", redirect, locationKey, err) - } - - nextSpec := &WorkflowSpec{ - RepoSpec: RepoSpec{ - RepoSlug: redirectedSource.Repo, - Version: redirectedSource.Ref, - }, - WorkflowPath: redirectedSource.Path, - WorkflowName: normalizeWorkflowID(redirectedSource.Path), - Host: currentSpec.Host, - } - resolutionLog.Printf("Following redirect for add: from=%s to=%s", locationKey, nextSpec.String()) - if verbose { - fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Workflow redirect: %s -> %s", locationKey, nextSpec.String()))) + return nil, nil, err } followedRedirect = true - currentSpec = *nextSpec } - return nil, nil, fmt.Errorf("redirect chain exceeded maximum depth (%d) for workflow '%s'", maxRedirectDepth, initialSpec.String()) } +func workflowRedirectLocationKey(spec *WorkflowSpec) string { + currentRef := spec.Version + if currentRef == "" { + currentRef = "main" + } + return fmt.Sprintf("%s/%s@%s", spec.RepoSlug, spec.WorkflowPath, currentRef) +} + +func trackWorkflowRedirectVisit(visited map[string]struct{}, locationKey string) error { + if _, exists := visited[locationKey]; exists { + return fmt.Errorf("redirect loop detected at %s", locationKey) + } + visited[locationKey] = struct{}{} + return nil +} + +func followAddWorkflowRedirect(redirect, locationKey, host string, verbose bool) (WorkflowSpec, error) { + redirectedSource, err := normalizeRedirectToSourceSpec(redirect) + if err != nil { + return WorkflowSpec{}, fmt.Errorf("invalid redirect %q in %s: %w", redirect, locationKey, err) + } + nextSpec := WorkflowSpec{RepoSpec: RepoSpec{RepoSlug: redirectedSource.Repo, Version: redirectedSource.Ref}, WorkflowPath: redirectedSource.Path, WorkflowName: normalizeWorkflowID(redirectedSource.Path), Host: host} + resolutionLog.Printf("Following redirect for add: from=%s to=%s", locationKey, nextSpec.String()) + if verbose { + fmt.Fprintln(os.Stderr, console.FormatWarningMessage(fmt.Sprintf("Workflow redirect: %s -> %s", locationKey, nextSpec.String()))) + } + return nextSpec, nil +} + +// expandLocalWildcardWorkflows expands wildcard workflow specifications for local workflows only. // expandLocalWildcardWorkflows expands wildcard workflow specifications for local workflows only. func expandLocalWildcardWorkflows(specs []*WorkflowSpec, verbose bool) ([]*WorkflowSpec, error) { expandedWorkflows := []*WorkflowSpec{} diff --git a/pkg/cli/audit_agentic_analysis.go b/pkg/cli/audit_agentic_analysis.go index ad7afe3c636..59ee36dcef3 100644 --- a/pkg/cli/audit_agentic_analysis.go +++ b/pkg/cli/audit_agentic_analysis.go @@ -98,82 +98,81 @@ func mergeMCPToolUsageInfo(toolUsage []ToolUsageInfo, mcpToolUsage *MCPToolUsage if mcpToolUsage == nil { return toolUsage } + toolStats := cloneToolUsageInfoMap(toolUsage) + mergeMCPToolUsageSummaries(toolStats, mcpToolUsage) + return sortedToolUsageInfos(toolStats) +} +func cloneToolUsageInfoMap(toolUsage []ToolUsageInfo) map[string]*ToolUsageInfo { toolStats := make(map[string]*ToolUsageInfo) for _, info := range toolUsage { cloned := info toolStats[info.Name] = &cloned } + return toolStats +} - addOrUpdateToolUsage := func(name string, callCount, maxInputSize, maxOutputSize int, maxDuration string) { - normalizedName := strings.TrimSpace(name) - if normalizedName == "" { - return - } - displayKey := workflow.PrettifyToolName(normalizedName) - if existing, exists := toolStats[displayKey]; exists { - existing.CallCount += callCount - if maxInputSize > existing.MaxInputSize { - existing.MaxInputSize = maxInputSize - } - if maxOutputSize > existing.MaxOutputSize { - existing.MaxOutputSize = maxOutputSize - } - if maxDuration != "" { - maxDurationValue := parseDurationString(maxDuration) - if existing.MaxDuration == "" { - existing.MaxDuration = maxDuration - } else { - existingMaxDurationValue := parseDurationString(existing.MaxDuration) - if maxDurationValue > existingMaxDurationValue { - existing.MaxDuration = maxDuration - } - } - } - return +func mergeMCPToolUsageSummaries(toolStats map[string]*ToolUsageInfo, mcpToolUsage *MCPToolUsageData) { + if len(mcpToolUsage.Summary) > 0 { + for _, summary := range mcpToolUsage.Summary { + mergeToolUsageEntry(toolStats, buildMCPToolUsageName(summary.ServerName, summary.ToolName), summary.CallCount, summary.MaxInputSize, summary.MaxOutputSize, summary.MaxDuration) } + return + } + for _, call := range mcpToolUsage.ToolCalls { + mergeToolUsageEntry(toolStats, buildMCPToolUsageName(call.ServerName, call.ToolName), 1, call.InputSize, call.OutputSize, call.Duration) + } +} - toolStats[displayKey] = &ToolUsageInfo{ - Name: displayKey, - CallCount: callCount, - MaxInputSize: maxInputSize, - MaxOutputSize: maxOutputSize, - MaxDuration: maxDuration, - } +func buildMCPToolUsageName(serverName, toolName string) string { + switch { + case serverName != "" && toolName != "": + return serverName + "." + toolName + default: + return toolName } +} - if len(mcpToolUsage.Summary) > 0 { - for _, summary := range mcpToolUsage.Summary { - switch { - case summary.ServerName != "" && summary.ToolName != "": - addOrUpdateToolUsage(summary.ServerName+"."+summary.ToolName, summary.CallCount, summary.MaxInputSize, summary.MaxOutputSize, summary.MaxDuration) - case summary.ToolName != "": - addOrUpdateToolUsage(summary.ToolName, summary.CallCount, summary.MaxInputSize, summary.MaxOutputSize, summary.MaxDuration) - } - } - } else { - for _, call := range mcpToolUsage.ToolCalls { - switch { - case call.ServerName != "" && call.ToolName != "": - addOrUpdateToolUsage(call.ServerName+"."+call.ToolName, 1, call.InputSize, call.OutputSize, call.Duration) - case call.ToolName != "": - addOrUpdateToolUsage(call.ToolName, 1, call.InputSize, call.OutputSize, call.Duration) - } - } +func mergeToolUsageEntry(toolStats map[string]*ToolUsageInfo, name string, callCount, maxInputSize, maxOutputSize int, maxDuration string) { + normalizedName := strings.TrimSpace(name) + if normalizedName == "" { + return + } + displayKey := workflow.PrettifyToolName(normalizedName) + if existing, exists := toolStats[displayKey]; exists { + mergeToolUsageIntoExisting(existing, callCount, maxInputSize, maxOutputSize, maxDuration) + return } + toolStats[displayKey] = &ToolUsageInfo{Name: displayKey, CallCount: callCount, MaxInputSize: maxInputSize, MaxOutputSize: maxOutputSize, MaxDuration: maxDuration} +} + +func mergeToolUsageIntoExisting(existing *ToolUsageInfo, callCount, maxInputSize, maxOutputSize int, maxDuration string) { + existing.CallCount += callCount + if maxInputSize > existing.MaxInputSize { + existing.MaxInputSize = maxInputSize + } + if maxOutputSize > existing.MaxOutputSize { + existing.MaxOutputSize = maxOutputSize + } + if maxDuration == "" { + return + } + if existing.MaxDuration == "" || parseDurationString(maxDuration) > parseDurationString(existing.MaxDuration) { + existing.MaxDuration = maxDuration + } +} +func sortedToolUsageInfos(toolStats map[string]*ToolUsageInfo) []ToolUsageInfo { mergedToolUsage := make([]ToolUsageInfo, 0, len(toolStats)) for _, info := range toolStats { mergedToolUsage = append(mergedToolUsage, *info) } - slices.SortFunc(mergedToolUsage, func(a, b ToolUsageInfo) int { if a.CallCount != b.CallCount { return b.CallCount - a.CallCount } return strings.Compare(a.Name, b.Name) }) - return mergedToolUsage } @@ -310,96 +309,97 @@ func buildAgenticAssessments(processedRun ProcessedRun, metrics MetricsData, too return nil } auditAgenticLog.Printf("Building agentic assessments: run_id=%d domain=%s resource=%s execution=%s", processedRun.Run.DatabaseID, domain.Name, fingerprint.ResourceProfile, fingerprint.ExecutionStyle) - + inputs := agenticAssessmentInputs{processedRun: processedRun, metrics: metrics, toolTypes: len(toolUsage), frictionEvents: len(processedRun.MissingTools) + len(processedRun.MCPFailures) + len(processedRun.MissingData), writeCount: len(createdItems) + processedRun.Run.SafeItemsCount, domain: domain, fingerprint: fingerprint, awContext: awContext} assessments := make([]AgenticAssessment, 0, 4) - toolTypes := len(toolUsage) - frictionEvents := len(processedRun.MissingTools) + len(processedRun.MCPFailures) + len(processedRun.MissingData) - writeCount := len(createdItems) + processedRun.Run.SafeItemsCount + assessments = appendHeavyResourceAssessment(assessments, inputs) + assessments = appendOverkillAssessment(assessments, inputs) + assessments = appendPoorControlAssessment(assessments, inputs) + assessments = appendPartiallyReducibleAssessment(assessments, inputs) + assessments = appendModelDowngradeAssessment(assessments, inputs) + assessments = appendDelegatedContextAssessment(assessments, inputs) + auditAgenticLog.Printf("Built %d agentic assessments", len(assessments)) + return assessments +} - if fingerprint.ResourceProfile == "heavy" { - severity := "medium" - if metrics.Turns >= 14 || toolTypes >= 7 || processedRun.Run.Duration >= 20*time.Minute { - severity = "high" - } - assessments = append(assessments, AgenticAssessment{ - Kind: "resource_heavy_for_domain", - Severity: severity, - Summary: fmt.Sprintf("This %s run consumed a heavy execution profile for its task shape.", domain.Label), - Evidence: fmt.Sprintf("turns=%d tool_types=%d duration=%s write_actions=%d", metrics.Turns, toolTypes, formatAssessmentDuration(processedRun.Run.Duration), writeCount), - Recommendation: "Compare this run to similar successful runs and trim unnecessary turns, tools, or write actions.", - }) +type agenticAssessmentInputs struct { + processedRun ProcessedRun + metrics MetricsData + toolTypes int + frictionEvents int + writeCount int + domain *TaskDomainInfo + fingerprint *BehaviorFingerprint + awContext *AwContext +} + +func appendHeavyResourceAssessment(assessments []AgenticAssessment, inputs agenticAssessmentInputs) []AgenticAssessment { + if inputs.fingerprint.ResourceProfile != "heavy" { + return assessments } + severity := "medium" + if inputs.metrics.Turns >= 14 || inputs.toolTypes >= 7 || inputs.processedRun.Run.Duration >= 20*time.Minute { + severity = "high" + } + return append(assessments, AgenticAssessment{Kind: "resource_heavy_for_domain", Severity: severity, Summary: fmt.Sprintf("This %s run consumed a heavy execution profile for its task shape.", inputs.domain.Label), Evidence: fmt.Sprintf("turns=%d tool_types=%d duration=%s write_actions=%d", inputs.metrics.Turns, inputs.toolTypes, formatAssessmentDuration(inputs.processedRun.Run.Duration), inputs.writeCount), Recommendation: "Compare this run to similar successful runs and trim unnecessary turns, tools, or write actions."}) +} - if (domain.Name == "triage" || domain.Name == "repo_maintenance" || domain.Name == "issue_response") && fingerprint.ResourceProfile == "lean" && fingerprint.ExecutionStyle == "directed" && fingerprint.ToolBreadth == "narrow" { - assessments = append(assessments, AgenticAssessment{ - Kind: "overkill_for_agentic", - Severity: "low", - Summary: fmt.Sprintf("This %s run looks stable enough that deterministic automation may be a simpler fit.", domain.Label), - Evidence: fmt.Sprintf("turns=%d tool_types=%d actuation=%s", metrics.Turns, toolTypes, fingerprint.ActuationStyle), - Recommendation: "Consider whether a scripted rule or deterministic workflow step could replace this agentic path.", - }) +func appendOverkillAssessment(assessments []AgenticAssessment, inputs agenticAssessmentInputs) []AgenticAssessment { + if !isSimpleAgenticDomain(inputs.domain.Name) || inputs.fingerprint.ResourceProfile != "lean" || inputs.fingerprint.ExecutionStyle != "directed" || inputs.fingerprint.ToolBreadth != "narrow" { + return assessments } + return append(assessments, AgenticAssessment{Kind: "overkill_for_agentic", Severity: "low", Summary: fmt.Sprintf("This %s run looks stable enough that deterministic automation may be a simpler fit.", inputs.domain.Label), Evidence: fmt.Sprintf("turns=%d tool_types=%d actuation=%s", inputs.metrics.Turns, inputs.toolTypes, inputs.fingerprint.ActuationStyle), Recommendation: "Consider whether a scripted rule or deterministic workflow step could replace this agentic path."}) +} - if frictionEvents >= 3 || (frictionEvents > 0 && writeCount >= 3) || ((domain.Name == "triage" || domain.Name == "repo_maintenance" || domain.Name == "issue_response") && fingerprint.ExecutionStyle == "exploratory") { - severity := "medium" - if frictionEvents >= 4 || (frictionEvents > 0 && fingerprint.ActuationStyle == "write_heavy") { - severity = "high" - } - assessments = append(assessments, AgenticAssessment{ - Kind: "poor_agentic_control", - Severity: severity, - Summary: "The run showed signs of broad or weakly controlled agentic behavior.", - Evidence: fmt.Sprintf("friction=%d execution=%s actuation=%s", frictionEvents, fingerprint.ExecutionStyle, fingerprint.ActuationStyle), - Recommendation: "Tighten instructions, reduce unnecessary tools, or delay write actions until the workflow has stronger evidence.", - }) +func appendPoorControlAssessment(assessments []AgenticAssessment, inputs agenticAssessmentInputs) []AgenticAssessment { + if !hasPoorAgenticControl(inputs) { + return assessments } + severity := "medium" + if inputs.frictionEvents >= 4 || (inputs.frictionEvents > 0 && inputs.fingerprint.ActuationStyle == "write_heavy") { + severity = "high" + } + return append(assessments, AgenticAssessment{Kind: "poor_agentic_control", Severity: severity, Summary: "The run showed signs of broad or weakly controlled agentic behavior.", Evidence: fmt.Sprintf("friction=%d execution=%s actuation=%s", inputs.frictionEvents, inputs.fingerprint.ExecutionStyle, inputs.fingerprint.ActuationStyle), Recommendation: "Tighten instructions, reduce unnecessary tools, or delay write actions until the workflow has stronger evidence."}) +} - // Partially reducible: the workflow has a low agentic fraction, meaning - // many turns are data-gathering that could be moved to deterministic steps: - // or post-steps: in the frontmatter. Only flag when there's substantive work - // (not lean/directed runs which overkill_for_agentic already covers). - if fingerprint.AgenticFraction > 0 && fingerprint.AgenticFraction < 0.6 && - fingerprint.ResourceProfile != "lean" { - severity := "low" - if fingerprint.AgenticFraction < 0.4 { - severity = "medium" - } - deterministicPct := int((1.0 - fingerprint.AgenticFraction) * 100) - assessments = append(assessments, AgenticAssessment{ - Kind: "partially_reducible", - Severity: severity, - Summary: fmt.Sprintf("About %d%% of this run's turns appear to be data-gathering that could move to deterministic steps.", deterministicPct), - Evidence: fmt.Sprintf("agentic_fraction=%.2f turns=%d", fingerprint.AgenticFraction, metrics.Turns), - Recommendation: "Move data-fetching work to frontmatter steps: (pre-agent) writing to /tmp/gh-aw/agent/ or post-steps: (post-agent) to reduce inference cost. See the DeterministicOps guide.", - }) +func hasPoorAgenticControl(inputs agenticAssessmentInputs) bool { + return inputs.frictionEvents >= 3 || + (inputs.frictionEvents > 0 && inputs.writeCount >= 3) || + (isSimpleAgenticDomain(inputs.domain.Name) && inputs.fingerprint.ExecutionStyle == "exploratory") +} + +func appendPartiallyReducibleAssessment(assessments []AgenticAssessment, inputs agenticAssessmentInputs) []AgenticAssessment { + if inputs.fingerprint.AgenticFraction <= 0 || inputs.fingerprint.AgenticFraction >= 0.6 || inputs.fingerprint.ResourceProfile == "lean" { + return assessments + } + severity := "low" + if inputs.fingerprint.AgenticFraction < 0.4 { + severity = "medium" } + deterministicPct := int((1.0 - inputs.fingerprint.AgenticFraction) * 100) + return append(assessments, AgenticAssessment{Kind: "partially_reducible", Severity: severity, Summary: fmt.Sprintf("About %d%% of this run's turns appear to be data-gathering that could move to deterministic steps.", deterministicPct), Evidence: fmt.Sprintf("agentic_fraction=%.2f turns=%d", inputs.fingerprint.AgenticFraction, inputs.metrics.Turns), Recommendation: "Move data-fetching work to frontmatter steps: (pre-agent) writing to /tmp/gh-aw/agent/ or post-steps: (post-agent) to reduce inference cost. See the DeterministicOps guide."}) +} - // Model downgrade suggestion: the run uses a heavy resource profile but - // the task domain is simple enough that a smaller model would likely suffice. - if fingerprint.ResourceProfile != "lean" && - (domain.Name == "triage" || domain.Name == "repo_maintenance" || domain.Name == "issue_response") && - fingerprint.ActuationStyle != "write_heavy" { - assessments = append(assessments, AgenticAssessment{ - Kind: "model_downgrade_available", - Severity: "low", - Summary: fmt.Sprintf("This %s run may not need a frontier model. A smaller model (e.g. gpt-4.1-mini, claude-haiku-4-5) could handle the task at lower cost.", domain.Label), - Evidence: fmt.Sprintf("domain=%s resource_profile=%s actuation=%s", domain.Name, fingerprint.ResourceProfile, fingerprint.ActuationStyle), - Recommendation: "Try engine.model: gpt-4.1-mini or claude-haiku-4-5 in the workflow frontmatter.", - }) +func appendModelDowngradeAssessment(assessments []AgenticAssessment, inputs agenticAssessmentInputs) []AgenticAssessment { + if inputs.fingerprint.ResourceProfile == "lean" || !isSimpleAgenticDomain(inputs.domain.Name) || inputs.fingerprint.ActuationStyle == "write_heavy" { + return assessments } + return append(assessments, AgenticAssessment{Kind: "model_downgrade_available", Severity: "low", Summary: fmt.Sprintf("This %s run may not need a frontier model. A smaller model (e.g. gpt-4.1-mini, claude-haiku-4-5) could handle the task at lower cost.", inputs.domain.Label), Evidence: fmt.Sprintf("domain=%s resource_profile=%s actuation=%s", inputs.domain.Name, inputs.fingerprint.ResourceProfile, inputs.fingerprint.ActuationStyle), Recommendation: "Try engine.model: gpt-4.1-mini or claude-haiku-4-5 in the workflow frontmatter."}) +} - if awContext != nil { - assessments = append(assessments, AgenticAssessment{ - Kind: "delegated_context_present", - Severity: "info", - Summary: "The run preserved upstream dispatch context, which helps trace multi-workflow episodes.", - Evidence: fmt.Sprintf("workflow_call_id=%s event_type=%s", awContext.WorkflowCallID, awContext.EventType), - Recommendation: "Use this context when comparing downstream runs so follow-up workflows are evaluated as part of one task chain.", - }) +func appendDelegatedContextAssessment(assessments []AgenticAssessment, inputs agenticAssessmentInputs) []AgenticAssessment { + if inputs.awContext == nil { + return assessments } + return append(assessments, AgenticAssessment{Kind: "delegated_context_present", Severity: "info", Summary: "The run preserved upstream dispatch context, which helps trace multi-workflow episodes.", Evidence: fmt.Sprintf("workflow_call_id=%s event_type=%s", inputs.awContext.WorkflowCallID, inputs.awContext.EventType), Recommendation: "Use this context when comparing downstream runs so follow-up workflows are evaluated as part of one task chain."}) +} - auditAgenticLog.Printf("Built %d agentic assessments", len(assessments)) - return assessments +func isSimpleAgenticDomain(domain string) bool { + switch domain { + case "triage", "repo_maintenance", "issue_response": + return true + default: + return false + } } func generateAgenticAssessmentFindings(assessments []AgenticAssessment) []Finding { diff --git a/pkg/cli/audit_comparison.go b/pkg/cli/audit_comparison.go index 0ee05986d40..52b8528bcaa 100644 --- a/pkg/cli/audit_comparison.go +++ b/pkg/cli/audit_comparison.go @@ -340,100 +340,91 @@ func buildAuditComparison(currentConclusion string, current auditComparisonSnaps if baselineRun == nil || baseline == nil { return &AuditComparisonData{BaselineFound: false} } - - reasonCodes := make([]string, 0, 4) currentConclusion = strings.TrimSpace(strings.ToLower(currentConclusion)) - currentRunUnsuccessful := currentConclusion != "" && currentConclusion != "success" - delta := &AuditComparisonDelta{ - Turns: AuditComparisonIntDelta{ - Before: baseline.Turns, - After: current.Turns, - Changed: baseline.Turns != current.Turns, - }, - Posture: AuditComparisonStringDelta{ - Before: baseline.Posture, - After: current.Posture, - Changed: baseline.Posture != current.Posture, - }, - BlockedRequests: AuditComparisonIntDelta{ - Before: baseline.BlockedRequests, - After: current.BlockedRequests, - Changed: baseline.BlockedRequests != current.BlockedRequests, - }, + delta, newMCPFailure, mcpFailuresResolved := buildAuditComparisonDelta(current, *baseline) + reasonCodes := collectAuditComparisonReasonCodes(currentConclusion, current, *baseline, newMCPFailure, mcpFailuresResolved) + label := classifyAuditComparisonLabel(currentConclusion, delta, baseline.BlockedRequests, current.BlockedRequests, newMCPFailure, mcpFailuresResolved, reasonCodes) + return &AuditComparisonData{ + BaselineFound: true, + Baseline: buildAuditComparisonBaseline(baselineRun), + Delta: delta, + Classification: &AuditComparisonClassification{Label: label, ReasonCodes: reasonCodes}, + Recommendation: &AuditComparisonRecommendation{Action: recommendAuditComparisonAction(label, currentConclusion, delta)}, } +} - if current.Turns > baseline.Turns { - reasonCodes = append(reasonCodes, "turns_increase") - } else if current.Turns < baseline.Turns { - reasonCodes = append(reasonCodes, "turns_decrease") +func buildAuditComparisonDelta(current, baseline auditComparisonSnapshot) (*AuditComparisonDelta, bool, bool) { + newMCPFailure := len(baseline.MCPFailures) == 0 && len(current.MCPFailures) > 0 + mcpFailuresResolved := len(baseline.MCPFailures) > 0 && len(current.MCPFailures) == 0 + delta := &AuditComparisonDelta{ + Turns: AuditComparisonIntDelta{Before: baseline.Turns, After: current.Turns, Changed: baseline.Turns != current.Turns}, + Posture: AuditComparisonStringDelta{Before: baseline.Posture, After: current.Posture, Changed: baseline.Posture != current.Posture}, + BlockedRequests: AuditComparisonIntDelta{Before: baseline.BlockedRequests, After: current.BlockedRequests, Changed: baseline.BlockedRequests != current.BlockedRequests}, + } + if newMCPFailure || len(baseline.MCPFailures) > 0 || len(current.MCPFailures) > 0 { + delta.MCPFailure = &AuditComparisonMCPFailureDelta{Before: baseline.MCPFailures, After: current.MCPFailures, NewlyPresent: newMCPFailure} } + return delta, newMCPFailure, mcpFailuresResolved +} + +func collectAuditComparisonReasonCodes(currentConclusion string, current, baseline auditComparisonSnapshot, newMCPFailure, mcpFailuresResolved bool) []string { + reasonCodes := make([]string, 0, 4) + reasonCodes = appendAuditComparisonTurnReason(reasonCodes, baseline.Turns, current.Turns) if baseline.Posture != current.Posture { reasonCodes = append(reasonCodes, "posture_changed") } - if current.BlockedRequests > baseline.BlockedRequests { - reasonCodes = append(reasonCodes, "blocked_requests_increase") - } else if current.BlockedRequests < baseline.BlockedRequests { - reasonCodes = append(reasonCodes, "blocked_requests_decrease") - } - if currentRunUnsuccessful { + reasonCodes = appendAuditComparisonBlockedReason(reasonCodes, baseline.BlockedRequests, current.BlockedRequests) + if currentConclusion != "" && currentConclusion != "success" { reasonCodes = append(reasonCodes, "run_unsuccessful") } - - newMCPFailure := len(baseline.MCPFailures) == 0 && len(current.MCPFailures) > 0 - mcpFailuresResolved := len(baseline.MCPFailures) > 0 && len(current.MCPFailures) == 0 - if newMCPFailure || len(baseline.MCPFailures) > 0 || len(current.MCPFailures) > 0 { - delta.MCPFailure = &AuditComparisonMCPFailureDelta{ - Before: baseline.MCPFailures, - After: current.MCPFailures, - NewlyPresent: newMCPFailure, - } - } if newMCPFailure { reasonCodes = append(reasonCodes, "new_mcp_failure") } else if mcpFailuresResolved { reasonCodes = append(reasonCodes, "mcp_failures_resolved") } + return reasonCodes +} - label := "stable" +func appendAuditComparisonTurnReason(reasonCodes []string, before, after int) []string { switch { - case currentRunUnsuccessful: - label = "risky" - case delta.Posture.Before == "read_only" && delta.Posture.After == "write_capable": - label = "risky" - case newMCPFailure: - label = "risky" - case current.BlockedRequests > baseline.BlockedRequests: - label = "risky" - case delta.Posture.Before != "" && delta.Posture.After != "" && delta.Posture.Before != delta.Posture.After: - label = "changed" - case mcpFailuresResolved: - label = "changed" - case current.BlockedRequests < baseline.BlockedRequests: - label = "changed" - case len(reasonCodes) > 0: - label = "changed" + case after > before: + return append(reasonCodes, "turns_increase") + case after < before: + return append(reasonCodes, "turns_decrease") + default: + return reasonCodes } +} - return &AuditComparisonData{ - BaselineFound: true, - Baseline: &AuditComparisonBaseline{ - RunID: baselineRun.DatabaseID, - WorkflowName: baselineRun.WorkflowName, - Conclusion: baselineRun.Conclusion, - CreatedAt: baselineRun.CreatedAt.Format("2006-01-02T15:04:05Z07:00"), - Selection: "latest_success", - }, - Delta: delta, - Classification: &AuditComparisonClassification{ - Label: label, - ReasonCodes: reasonCodes, - }, - Recommendation: &AuditComparisonRecommendation{ - Action: recommendAuditComparisonAction(label, currentConclusion, delta), - }, +func appendAuditComparisonBlockedReason(reasonCodes []string, before, after int) []string { + switch { + case after > before: + return append(reasonCodes, "blocked_requests_increase") + case after < before: + return append(reasonCodes, "blocked_requests_decrease") + default: + return reasonCodes + } +} + +func classifyAuditComparisonLabel(currentConclusion string, delta *AuditComparisonDelta, baselineBlocked, currentBlocked int, newMCPFailure, mcpFailuresResolved bool, reasonCodes []string) string { + currentRunUnsuccessful := currentConclusion != "" && currentConclusion != "success" + switch { + case currentRunUnsuccessful, delta.Posture.Before == "read_only" && delta.Posture.After == "write_capable", newMCPFailure, currentBlocked > baselineBlocked: + return "risky" + case delta.Posture.Before != "" && delta.Posture.After != "" && delta.Posture.Before != delta.Posture.After: + return "changed" + case mcpFailuresResolved, currentBlocked < baselineBlocked, len(reasonCodes) > 0: + return "changed" + default: + return "stable" } } +func buildAuditComparisonBaseline(baselineRun *WorkflowRun) *AuditComparisonBaseline { + return &AuditComparisonBaseline{RunID: baselineRun.DatabaseID, WorkflowName: baselineRun.WorkflowName, Conclusion: baselineRun.Conclusion, CreatedAt: baselineRun.CreatedAt.Format("2006-01-02T15:04:05Z07:00"), Selection: "latest_success"} +} + func recommendAuditComparisonAction(label, currentConclusion string, delta *AuditComparisonDelta) string { if currentConclusion != "" && currentConclusion != "success" { if currentConclusion == "failure" { diff --git a/pkg/cli/audit_expanded.go b/pkg/cli/audit_expanded.go index 42d0fac19a6..c85d1e5c0f2 100644 --- a/pkg/cli/audit_expanded.go +++ b/pkg/cli/audit_expanded.go @@ -293,84 +293,84 @@ func extractPromptAnalysis(logsPath string) *PromptAnalysis { // buildSessionAnalysis creates session performance metrics from available data func buildSessionAnalysis(processedRun ProcessedRun, metrics LogMetrics) *SessionAnalysis { run := processedRun.Run + session := &SessionAnalysis{TurnCount: metrics.Turns, NoopCount: run.NoopCount} + applySessionWallTime(session, run.Duration) + applySessionTurnTimings(session, metrics, run.Duration) + applySessionTokensPerMinute(session, metrics.TokenUsage, run.Duration) + session.TimeoutDetected = sessionTimeoutDetected(processedRun) + auditExpandedLog.Printf("Built session analysis: turns=%d, wall_time=%s, avg_tbt=%s, max_tbt=%s, timeout=%v", + session.TurnCount, session.WallTime, session.AvgTimeBetweenTurns, session.MaxTimeBetweenTurns, session.TimeoutDetected) + return session +} - session := &SessionAnalysis{ - TurnCount: metrics.Turns, - NoopCount: run.NoopCount, +func applySessionWallTime(session *SessionAnalysis, duration time.Duration) { + if duration > 0 { + session.WallTime = timeutil.FormatDuration(duration) } +} - // Wall time from run duration - if run.Duration > 0 { - session.WallTime = timeutil.FormatDuration(run.Duration) +func applySessionTurnTimings(session *SessionAnalysis, metrics LogMetrics, duration time.Duration) { + if metrics.Turns > 0 && duration > 0 { + session.AvgTurnDuration = timeutil.FormatDuration(duration / time.Duration(metrics.Turns)) + } + if metrics.AvgTimeBetweenTurns > 0 { + applyObservedSessionTBT(session, metrics) + return } + applyEstimatedSessionTBT(session, metrics.Turns, duration) +} - // Average turn duration - if metrics.Turns > 0 && run.Duration > 0 { - avgTurnDuration := run.Duration / (time.Duration(metrics.Turns)) - session.AvgTurnDuration = timeutil.FormatDuration(avgTurnDuration) +func applyObservedSessionTBT(session *SessionAnalysis, metrics LogMetrics) { + session.AvgTimeBetweenTurns = timeutil.FormatDuration(metrics.AvgTimeBetweenTurns) + if metrics.MaxTimeBetweenTurns > 0 { + session.MaxTimeBetweenTurns = timeutil.FormatDuration(metrics.MaxTimeBetweenTurns) } + session.CacheWarning = sessionCacheWarning(metrics.AvgTimeBetweenTurns, metrics.MaxTimeBetweenTurns) +} - // Time Between Turns (TBT): prefer precise per-turn timestamps from log metrics; - // fall back to wall-time / turns when timestamps are unavailable. - // TBT measures the gap between consecutive LLM API calls (tool execution overhead). - // Anthropic's prompt cache TTL is 5 minutes — if TBT exceeds this, cache entries - // expire and every turn incurs full prompt re-processing costs. - const anthropicCacheTTL = 5 * time.Minute - if metrics.AvgTimeBetweenTurns > 0 { - session.AvgTimeBetweenTurns = timeutil.FormatDuration(metrics.AvgTimeBetweenTurns) - if metrics.MaxTimeBetweenTurns > 0 { - session.MaxTimeBetweenTurns = timeutil.FormatDuration(metrics.MaxTimeBetweenTurns) - } - // Warn when the maximum observed TBT exceeds the Anthropic cache TTL. - if metrics.MaxTimeBetweenTurns > anthropicCacheTTL { - session.CacheWarning = fmt.Sprintf( - "Max TBT (%s) exceeds Anthropic 5-min cache TTL — prompt cache will expire between turns, increasing cost", - timeutil.FormatDuration(metrics.MaxTimeBetweenTurns), - ) - } else if metrics.AvgTimeBetweenTurns > anthropicCacheTTL { - session.CacheWarning = fmt.Sprintf( - "Avg TBT (%s) exceeds Anthropic 5-min cache TTL — prompt cache likely expiring between turns", - timeutil.FormatDuration(metrics.AvgTimeBetweenTurns), - ) - } - } else if metrics.Turns > 1 && run.Duration > 0 { - // Fallback: estimate TBT from wall time over turns-1 intervals. - avgTBT := run.Duration / time.Duration(metrics.Turns-1) - session.AvgTimeBetweenTurns = timeutil.FormatDuration(avgTBT) + " (estimated)" - if avgTBT > anthropicCacheTTL { - session.CacheWarning = fmt.Sprintf( - "Estimated avg TBT (%s) exceeds Anthropic 5-min cache TTL — prompt cache likely expiring between turns", - timeutil.FormatDuration(avgTBT), - ) - } +func applyEstimatedSessionTBT(session *SessionAnalysis, turns int, duration time.Duration) { + if turns <= 1 || duration <= 0 { + return } + avgTBT := duration / time.Duration(turns-1) + session.AvgTimeBetweenTurns = timeutil.FormatDuration(avgTBT) + " (estimated)" + if avgTBT > 5*time.Minute { + session.CacheWarning = fmt.Sprintf("Estimated avg TBT (%s) exceeds Anthropic 5-min cache TTL — prompt cache likely expiring between turns", timeutil.FormatDuration(avgTBT)) + } +} - // Tokens per minute - if metrics.TokenUsage > 0 && run.Duration > 0 { - minutes := run.Duration.Minutes() - if minutes > 0 { - session.TokensPerMinute = float64(metrics.TokenUsage) / minutes - } +func sessionCacheWarning(avgTBT, maxTBT time.Duration) string { + const anthropicCacheTTL = 5 * time.Minute + switch { + case maxTBT > anthropicCacheTTL: + return fmt.Sprintf("Max TBT (%s) exceeds Anthropic 5-min cache TTL — prompt cache will expire between turns, increasing cost", timeutil.FormatDuration(maxTBT)) + case avgTBT > anthropicCacheTTL: + return fmt.Sprintf("Avg TBT (%s) exceeds Anthropic 5-min cache TTL — prompt cache likely expiring between turns", timeutil.FormatDuration(avgTBT)) + default: + return "" } +} - // Timeout detection: check if the run was cancelled (typically indicates timeout) - if run.Conclusion == "cancelled" || run.Conclusion == "timed_out" { - session.TimeoutDetected = true +func applySessionTokensPerMinute(session *SessionAnalysis, tokenUsage int, duration time.Duration) { + if tokenUsage <= 0 || duration <= 0 || duration.Minutes() <= 0 { + return } + session.TokensPerMinute = float64(tokenUsage) / duration.Minutes() +} - // Check for timeout patterns in job conclusions +func sessionTimeoutDetected(processedRun ProcessedRun) bool { + if processedRun.Run.Conclusion == "cancelled" || processedRun.Run.Conclusion == "timed_out" { + return true + } for _, job := range processedRun.JobDetails { if job.Conclusion == "cancelled" || job.Conclusion == "timed_out" { - session.TimeoutDetected = true - break + return true } } - - auditExpandedLog.Printf("Built session analysis: turns=%d, wall_time=%s, avg_tbt=%s, max_tbt=%s, timeout=%v", - session.TurnCount, session.WallTime, session.AvgTimeBetweenTurns, session.MaxTimeBetweenTurns, session.TimeoutDetected) - return session + return false } +// buildSafeOutputSummary creates a summary of safe output items by type // buildSafeOutputSummary creates a summary of safe output items by type func buildSafeOutputSummary(items []CreatedItemReport, chainMetrics SafeOutputChainMetrics) *SafeOutputSummary { if len(items) == 0 && chainMetrics.TemporaryIDMapStatus == "" { @@ -470,82 +470,74 @@ func buildMCPServerHealth(mcpToolUsage *MCPToolUsageData, mcpFailures []MCPFailu if mcpToolUsage == nil && len(mcpFailures) == 0 { return nil } - health := &MCPServerHealth{} + failedServers := collectFailedMCPServers(mcpFailures) + health.FailedSvrs = len(failedServers) + populateMCPServerHealth(health, mcpToolUsage, failedServers) + appendMissingFailedMCPServers(health, failedServers) + finalizeMCPServerHealth(health) + auditExpandedLog.Printf("Built MCP server health: %s, total_requests=%d, error_rate=%.1f%%", health.Summary, health.TotalRequests, health.ErrorRate) + return health +} - // Track failed servers from MCPFailures - failedServers := make(map[string]struct { - }) +func collectFailedMCPServers(mcpFailures []MCPFailureReport) map[string]struct{} { + failedServers := make(map[string]struct{}) for _, failure := range mcpFailures { - failedServers[failure.ServerName] = struct { - }{} + failedServers[failure.ServerName] = struct{}{} } - health.FailedSvrs = len(failedServers) - - // Process server statistics from mcpToolUsage - if mcpToolUsage != nil { - for _, server := range mcpToolUsage.Servers { - health.TotalRequests += server.RequestCount - health.TotalErrors += server.ErrorCount - - errorRate := safePercent(server.ErrorCount, server.RequestCount) - - status := "✅ healthy" - if _, isFailed := failedServers[server.ServerName]; isFailed { - status = "❌ failed" - } else if errorRate > 10 { - status = "⚠️ degraded" - } + return failedServers +} - health.Servers = append(health.Servers, MCPServerHealthDetail{ - ServerName: server.ServerName, - RequestCount: server.RequestCount, - ToolCalls: server.ToolCallCount, - ErrorCount: server.ErrorCount, - ErrorRate: errorRate, - ErrorRateStr: fmt.Sprintf("%.1f%%", errorRate), - AvgLatency: server.AvgDuration, - Status: status, - }) - } +func populateMCPServerHealth(health *MCPServerHealth, mcpToolUsage *MCPToolUsageData, failedServers map[string]struct{}) { + if mcpToolUsage == nil { + return + } + for _, server := range mcpToolUsage.Servers { + health.TotalRequests += server.RequestCount + health.TotalErrors += server.ErrorCount + health.Servers = append(health.Servers, buildMCPServerHealthDetail(server, failedServers)) + } + health.SlowestCalls = buildSlowestToolCalls(mcpToolUsage.ToolCalls, 5) +} - // Build slowest tool calls from individual call records (top 5) - health.SlowestCalls = buildSlowestToolCalls(mcpToolUsage.ToolCalls, 5) +func buildMCPServerHealthDetail(server MCPServerStats, failedServers map[string]struct{}) MCPServerHealthDetail { + errorRate := safePercent(server.ErrorCount, server.RequestCount) + status := "✅ healthy" + if _, isFailed := failedServers[server.ServerName]; isFailed { + status = "❌ failed" + } else if errorRate > 10 { + status = "⚠️ degraded" } + return MCPServerHealthDetail{ServerName: server.ServerName, RequestCount: server.RequestCount, ToolCalls: server.ToolCallCount, ErrorCount: server.ErrorCount, ErrorRate: errorRate, ErrorRateStr: fmt.Sprintf("%.1f%%", errorRate), AvgLatency: server.AvgDuration, Status: status} +} - // Add failed servers that don't appear in stats +func appendMissingFailedMCPServers(health *MCPServerHealth, failedServers map[string]struct{}) { for serverName := range failedServers { - found := false - for _, s := range health.Servers { - if s.ServerName == serverName { - found = true - break - } + if hasMCPServerHealthDetail(health.Servers, serverName) { + continue } - if !found { - health.Servers = append(health.Servers, MCPServerHealthDetail{ - ServerName: serverName, - Status: "❌ failed", - }) + health.Servers = append(health.Servers, MCPServerHealthDetail{ServerName: serverName, Status: "❌ failed"}) + } +} + +func hasMCPServerHealthDetail(details []MCPServerHealthDetail, serverName string) bool { + for _, detail := range details { + if detail.ServerName == serverName { + return true } } + return false +} +func finalizeMCPServerHealth(health *MCPServerHealth) { health.TotalServers = len(health.Servers) - - // Count servers by status for accurate summary - degradedCount := 0 - for _, s := range health.Servers { - if strings.Contains(s.Status, "degraded") { - degradedCount++ + for _, server := range health.Servers { + if strings.Contains(server.Status, "degraded") { + health.DegradedSvrs++ } } - health.DegradedSvrs = degradedCount health.HealthySvrs = health.TotalServers - health.FailedSvrs - health.DegradedSvrs - - // Calculate overall error rate health.ErrorRate = safePercent(health.TotalErrors, health.TotalRequests) - - // Sort servers by request count (highest first) slices.SortFunc(health.Servers, func(a, b MCPServerHealthDetail) int { if a.RequestCount > b.RequestCount { return -1 @@ -555,16 +547,10 @@ func buildMCPServerHealth(mcpToolUsage *MCPToolUsageData, mcpFailures []MCPFailu } return 0 }) - - // Build summary string - health.Summary = fmt.Sprintf("%d server(s), %d healthy, %d degraded, %d failed", - health.TotalServers, health.HealthySvrs, health.DegradedSvrs, health.FailedSvrs) - - auditExpandedLog.Printf("Built MCP server health: %s, total_requests=%d, error_rate=%.1f%%", - health.Summary, health.TotalRequests, health.ErrorRate) - return health + health.Summary = fmt.Sprintf("%d server(s), %d healthy, %d degraded, %d failed", health.TotalServers, health.HealthySvrs, health.DegradedSvrs, health.FailedSvrs) } +// buildSlowestToolCalls extracts the N slowest tool calls from the call records // buildSlowestToolCalls extracts the N slowest tool calls from the call records func buildSlowestToolCalls(calls []MCPToolCall, topN int) []MCPSlowestToolCall { if len(calls) == 0 { diff --git a/pkg/cli/audit_report_experiments.go b/pkg/cli/audit_report_experiments.go index 7c2cc479795..532f100f0da 100644 --- a/pkg/cli/audit_report_experiments.go +++ b/pkg/cli/audit_report_experiments.go @@ -64,64 +64,70 @@ func extractExperimentData(logsPath string) *ExperimentData { if logsPath == "" { return nil } - experimentDataLog.Printf("Extracting experiment data from: %s", logsPath) + if data := extractExperimentDataFromState(logsPath); data != nil { + return data + } + if data := extractExperimentDataFromUsageSummary(logsPath); data != nil { + return data + } + experimentDataLog.Print("No experiment data found") + return nil +} +func extractExperimentDataFromState(logsPath string) *ExperimentData { statePath := findExperimentStatePath(logsPath) - if statePath != "" { - experimentDataLog.Printf("Reading experiment state from: %s", statePath) - raw, err := os.ReadFile(statePath) - if err == nil { - state := parseExperimentState(raw) - if len(state.Counts) > 0 { - experimentDataLog.Printf("Found %d experiment(s) in state file", len(state.Counts)) - - // When per-run records are available, use the most recent run's assignments directly - // instead of inferring them from cumulative counts. - if len(state.Runs) > 0 { - lastRun := state.Runs[len(state.Runs)-1] - if len(lastRun.Assignments) > 0 { - experimentDataLog.Printf("Using run record from run_id=%s (timestamp=%s)", lastRun.RunID, lastRun.Timestamp) - return &ExperimentData{ - Assignments: lastRun.Assignments, - CumulativeCounts: state.Counts, - } - } - } - - // Derive this-run assignments: the variant selected on the most-recent run is - // the one with the maximum count (ties resolved by sorted order). - assignments := make(map[string]string, len(state.Counts)) - names := sliceutil.SortedKeys(state.Counts) - for _, name := range names { - variantCounts := state.Counts[name] - selected := deriveLastSelectedVariant(variantCounts) - assignments[name] = selected - experimentDataLog.Printf("Experiment %q: selected variant=%q", name, selected) - } - return &ExperimentData{ - Assignments: assignments, - CumulativeCounts: state.Counts, - } - } - } + if statePath == "" { + return nil } + experimentDataLog.Printf("Reading experiment state from: %s", statePath) + raw, err := os.ReadFile(statePath) + if err != nil { + return nil + } + state := parseExperimentState(raw) + if len(state.Counts) == 0 { + return nil + } + experimentDataLog.Printf("Found %d experiment(s) in state file", len(state.Counts)) + if data := extractExperimentDataFromRuns(state); data != nil { + return data + } + return deriveExperimentDataFromCounts(state.Counts) +} - // Fall back to the usage activity summary (written by the conclusion job). - // This is available when the experiment artifact was not downloaded separately, - // and the conclusion job was run with pick_experiment.cjs v2+ (JSONL ledger). - usageSummary, err := loadUsageActivitySummary(logsPath) - if err == nil && usageSummary != nil && usageSummary.Experiments != nil { - if len(usageSummary.Experiments.Assignments) > 0 { - experimentDataLog.Printf("Loaded experiment assignments from usage activity summary (%d experiment(s))", len(usageSummary.Experiments.Assignments)) - return &ExperimentData{ - Assignments: usageSummary.Experiments.Assignments, - } - } +func extractExperimentDataFromRuns(state *ExperimentState) *ExperimentData { + if state == nil || len(state.Runs) == 0 { + return nil + } + lastRun := state.Runs[len(state.Runs)-1] + if len(lastRun.Assignments) == 0 { + return nil } + experimentDataLog.Printf("Using run record from run_id=%s (timestamp=%s)", lastRun.RunID, lastRun.Timestamp) + return &ExperimentData{Assignments: lastRun.Assignments, CumulativeCounts: state.Counts} +} - experimentDataLog.Print("No experiment data found") - return nil +func deriveExperimentDataFromCounts(counts map[string]map[string]int) *ExperimentData { + assignments := make(map[string]string, len(counts)) + for _, name := range sliceutil.SortedKeys(counts) { + selected := deriveLastSelectedVariant(counts[name]) + assignments[name] = selected + experimentDataLog.Printf("Experiment %q: selected variant=%q", name, selected) + } + return &ExperimentData{Assignments: assignments, CumulativeCounts: counts} +} + +func extractExperimentDataFromUsageSummary(logsPath string) *ExperimentData { + usageSummary, err := loadUsageActivitySummary(logsPath) + if err != nil || usageSummary == nil || usageSummary.Experiments == nil { + return nil + } + if len(usageSummary.Experiments.Assignments) == 0 { + return nil + } + experimentDataLog.Printf("Loaded experiment assignments from usage activity summary (%d experiment(s))", len(usageSummary.Experiments.Assignments)) + return &ExperimentData{Assignments: usageSummary.Experiments.Assignments} } // formatExperimentLabel returns a compact, human-readable label summarising the diff --git a/pkg/workflow/safe_outputs_actions.go b/pkg/workflow/safe_outputs_actions.go index 9b62b1e892c..ae7ea0abf4f 100644 --- a/pkg/workflow/safe_outputs_actions.go +++ b/pkg/workflow/safe_outputs_actions.go @@ -166,6 +166,10 @@ func parseActionUsesField(uses string) (*actionRef, error) { // When available, the action reference is pinned to a commit SHA for security; // if no pin is available, later step generation falls back to the original config.Uses. func (c *Compiler) fetchAndParseActionYAML(actionName string, config *SafeOutputActionConfig, markdownPath string, data *WorkflowData) { + c.fetchAndParseActionYAMLBody(actionName, config, markdownPath, data) +} + +func (c *Compiler) fetchAndParseActionYAMLBody(actionName string, config *SafeOutputActionConfig, markdownPath string, data *WorkflowData) { if config.Uses == "" { return } @@ -380,6 +384,10 @@ func isGitHubExpressionDefault(input *ActionYAMLInput) bool { // generateActionToolDefinition creates an MCP tool definition for a custom safe output action. // The tool name is the normalized action name. Inputs are derived from the action.yml. func generateActionToolDefinition(actionName string, config *SafeOutputActionConfig) map[string]any { + return generateActionToolDefinitionBody(actionName, config) +} + +func generateActionToolDefinitionBody(actionName string, config *SafeOutputActionConfig) map[string]any { normalizedName := stringutil.NormalizeSafeOutputIdentifier(actionName) description := config.Description @@ -500,6 +508,10 @@ func actionOutputKey(normalizedName string) string { // - Uses the resolved action reference // - Has a `with:` block populated from parsed payload output via fromJSON func (c *Compiler) buildActionSteps(data *WorkflowData) []string { + return c.buildActionStepsBody(data) +} + +func (c *Compiler) buildActionStepsBody(data *WorkflowData) []string { if data.SafeOutputs == nil || len(data.SafeOutputs.Actions) == 0 { return nil } diff --git a/pkg/workflow/safe_outputs_app_config.go b/pkg/workflow/safe_outputs_app_config.go index 9c84b9bd30a..038b1e80249 100644 --- a/pkg/workflow/safe_outputs_app_config.go +++ b/pkg/workflow/safe_outputs_app_config.go @@ -33,6 +33,10 @@ type GitHubAppConfig struct { // parseAppConfig parses the app configuration from a map func parseAppConfig(appMap map[string]any) *GitHubAppConfig { + return parseAppConfigBody(appMap) +} + +func parseAppConfigBody(appMap map[string]any) *GitHubAppConfig { safeOutputsAppLog.Print("Parsing GitHub App configuration") appConfig := &GitHubAppConfig{} @@ -372,6 +376,10 @@ func (c *Compiler) buildGitHubAppTokenMintStepForRepository(app *GitHubAppConfig } func (c *Compiler) buildGitHubAppTokenMintStepWithMeta(app *GitHubAppConfig, permissions *Permissions, fallbackRepoExpr string, ownerSourceRepository string, stepName string, stepID string) []string { + return c.buildGitHubAppTokenMintStepWithMetaBody(app, permissions, fallbackRepoExpr, ownerSourceRepository, stepName, stepID) +} + +func (c *Compiler) buildGitHubAppTokenMintStepWithMetaBody(app *GitHubAppConfig, permissions *Permissions, fallbackRepoExpr string, ownerSourceRepository string, stepName string, stepID string) []string { safeOutputsAppLog.Printf("Building GitHub App token mint step: owner=%s, repos=%d", app.Owner, len(app.Repositories)) var steps []string @@ -475,6 +483,10 @@ func (c *Compiler) buildGitHubAppTokenMintStepWithMeta(app *GitHubAppConfig, per // GetExplicit() so that only scopes the user actually declared are forwarded — a "read-all" // shorthand must never accidentally grant broad GitHub App-only permissions. func convertPermissionsToAppTokenFields(permissions *Permissions) map[string]string { + return convertPermissionsToAppTokenFieldsBody(permissions) +} + +func convertPermissionsToAppTokenFieldsBody(permissions *Permissions) map[string]string { fields := make(map[string]string) // Map GitHub Actions permissions to GitHub App permissions diff --git a/pkg/workflow/safe_outputs_config_base.go b/pkg/workflow/safe_outputs_config_base.go index 8ce74319b5b..9731dfe75ab 100644 --- a/pkg/workflow/safe_outputs_config_base.go +++ b/pkg/workflow/safe_outputs_config_base.go @@ -9,6 +9,10 @@ import ( // before parsing the max field from configMap. Supports both integer values and GitHub // Actions expression strings (e.g. "${{ inputs.max }}"). func (c *Compiler) parseBaseSafeOutputConfig(configMap map[string]any, config *BaseSafeOutputConfig, defaultMax int) { + c.parseBaseSafeOutputConfigBody(configMap, config, defaultMax) +} + +func (c *Compiler) parseBaseSafeOutputConfigBody(configMap map[string]any, config *BaseSafeOutputConfig, defaultMax int) { // Set default max if provided if defaultMax > 0 { safeOutputsConfigLog.Printf("Setting default max: %d", defaultMax) diff --git a/pkg/workflow/safe_outputs_config_extraction.go b/pkg/workflow/safe_outputs_config_extraction.go index 3137cac317f..4d7425044f0 100644 --- a/pkg/workflow/safe_outputs_config_extraction.go +++ b/pkg/workflow/safe_outputs_config_extraction.go @@ -43,6 +43,10 @@ package workflow // extractSafeOutputsConfig extracts output configuration from frontmatter func (c *Compiler) extractSafeOutputsConfig(frontmatter map[string]any) *SafeOutputsConfig { + return c.extractSafeOutputsConfigBody(frontmatter) +} + +func (c *Compiler) extractSafeOutputsConfigBody(frontmatter map[string]any) *SafeOutputsConfig { safeOutputsConfigLog.Print("Extracting safe-outputs configuration from frontmatter") var config *SafeOutputsConfig diff --git a/pkg/workflow/safe_outputs_config_generation.go b/pkg/workflow/safe_outputs_config_generation.go index 181b65f17c2..41439a21fff 100644 --- a/pkg/workflow/safe_outputs_config_generation.go +++ b/pkg/workflow/safe_outputs_config_generation.go @@ -28,6 +28,10 @@ import ( // MCP server. Standard handler configs are sourced from handlerRegistry to ensure // they stay in sync with GH_AW_SAFE_OUTPUTS_HANDLER_CONFIG. func generateSafeOutputsConfig(data *WorkflowData) (string, error) { + return generateSafeOutputsConfigBody(data) +} + +func generateSafeOutputsConfigBody(data *WorkflowData) (string, error) { if data.SafeOutputs == nil { safeOutputsConfigLog.Print("No safe outputs configuration found, returning empty config") return "", nil @@ -242,6 +246,10 @@ func getEngineAgentFileInfoFromWorkflowData(data *WorkflowData) (manifestFiles [ // generateCustomJobToolDefinition creates an MCP tool definition for a custom safe-output job. // Returns a map representing the tool definition in MCP format with name, description, and inputSchema. func generateCustomJobToolDefinition(jobName string, jobConfig *SafeJobConfig) map[string]any { + return generateCustomJobToolDefinitionBody(jobName, jobConfig) +} + +func generateCustomJobToolDefinitionBody(jobName string, jobConfig *SafeJobConfig) map[string]any { safeOutputsConfigLog.Printf("Generating tool definition for custom job: %s", jobName) description := jobConfig.Description diff --git a/pkg/workflow/safe_outputs_config_global.go b/pkg/workflow/safe_outputs_config_global.go index 74e3f2bf902..b912aa7e0cd 100644 --- a/pkg/workflow/safe_outputs_config_global.go +++ b/pkg/workflow/safe_outputs_config_global.go @@ -10,6 +10,10 @@ import ( // extractGlobalConfigFields parses safe-outputs fields that apply across handlers, // keeping extractSafeOutputsConfig focused on routing handler-specific configuration. func (c *Compiler) extractGlobalConfigFields(outputMap map[string]any, config *SafeOutputsConfig) { + c.extractGlobalConfigFieldsBody(outputMap, config) +} + +func (c *Compiler) extractGlobalConfigFieldsBody(outputMap map[string]any, config *SafeOutputsConfig) { // Parse allowed-domains configuration (additional domains, unioned with network.allowed; supports ecosystem identifiers) if allowedDomains, exists := outputMap["allowed-domains"]; exists { if domainsArray, ok := allowedDomains.([]any); ok { diff --git a/pkg/workflow/safe_outputs_config_runtime.go b/pkg/workflow/safe_outputs_config_runtime.go index 406ef2f71e4..c80ada3317b 100644 --- a/pkg/workflow/safe_outputs_config_runtime.go +++ b/pkg/workflow/safe_outputs_config_runtime.go @@ -26,6 +26,10 @@ type SafeOutputStepConfig struct { } func (c *Compiler) addHandlerManagerConfigEnvVar(steps *[]string, data *WorkflowData) { + c.addHandlerManagerConfigEnvVarBody(steps, data) +} + +func (c *Compiler) addHandlerManagerConfigEnvVarBody(steps *[]string, data *WorkflowData) { if data.SafeOutputs == nil { safeOutputsConfigLog.Print("No safe-outputs configuration, skipping handler manager config") return diff --git a/pkg/workflow/safe_outputs_data_schema.go b/pkg/workflow/safe_outputs_data_schema.go index 61e58a4d2a5..5164b202b2a 100644 --- a/pkg/workflow/safe_outputs_data_schema.go +++ b/pkg/workflow/safe_outputs_data_schema.go @@ -110,6 +110,10 @@ func resolveSafeOutputsDataSchema(config *SafeOutputsConfig) (bool, map[string]a } func simplifyDataSchemaNode(raw any, path string, allowShorthand bool) (map[string]any, error) { + return simplifyDataSchemaNodeBody(raw, path, allowShorthand) +} + +func simplifyDataSchemaNodeBody(raw any, path string, allowShorthand bool) (map[string]any, error) { if typeName, ok := raw.(string); ok { if !allowShorthand { return nil, fmt.Errorf("%s: string shorthand is not allowed here", path) diff --git a/pkg/workflow/safe_outputs_handler_registry.go b/pkg/workflow/safe_outputs_handler_registry.go index cd0c2aa36f9..f95e170b668 100644 --- a/pkg/workflow/safe_outputs_handler_registry.go +++ b/pkg/workflow/safe_outputs_handler_registry.go @@ -504,80 +504,7 @@ var handlerRegistry = map[string]handlerBuilder{ AddTemplatableBool("staged", templatableBoolPtrToStringPtr(c.Staged)). Build() }, - "create_pull_request": func(cfg *SafeOutputsConfig) map[string]any { - if cfg.CreatePullRequests == nil { - return nil - } - c := cfg.CreatePullRequests - protectedFilesPolicy := "request_review" - if c.ManifestFilesPolicy != nil { - protectedFilesPolicy = *c.ManifestFilesPolicy - } - maxPatchSize := 4096 // default 4096 KB - if cfg.MaximumPatchSize > 0 { - maxPatchSize = cfg.MaximumPatchSize - } - if c.MaxPatchSize > 0 { - maxPatchSize = c.MaxPatchSize - } - maxPatchFiles := 100 // default 100 unique files - if cfg.MaximumPatchFiles > 0 { - maxPatchFiles = cfg.MaximumPatchFiles - } - if c.MaxPatchFiles > 0 { - maxPatchFiles = c.MaxPatchFiles - } - builder := newHandlerConfigBuilder(). - AddTemplatableInt("max", c.Max). - AddIfTrue("require_temporary_id", c.RequireTemporaryID). - AddIfNotEmpty("branch_prefix", c.BranchPrefix). - AddIfNotEmpty("title_prefix", c.TitlePrefix). - AddTemplatableStringSlice("labels", c.Labels). - AddStringSlice("fallback_labels", c.FallbackLabels). - AddTemplatableStringSlice("reviewers", c.Reviewers). - AddTemplatableStringSlice("team_reviewers", c.TeamReviewers). - AddTemplatableStringSlice("assignees", c.Assignees). - AddTemplatableBool("draft", c.Draft). - AddIfNotEmpty("if_no_changes", c.IfNoChanges). - AddTemplatableBool("allow_empty", c.AllowEmpty). - AddTemplatableBool("auto_merge", c.AutoMerge). - AddIfPositive("expires", c.Expires). - AddIfNotEmpty("target-repo", c.TargetRepoSlug). - AddIfNotEmpty("head-repo", c.HeadRepoSlug). - AddTemplatableStringSlice("allowed_repos", c.AllowedRepos). - AddTemplatableStringSlice("allowed_base_branches", c.AllowedBaseBranches). - AddTemplatableStringSlice("allowed_branches", c.AllowedBranches). - AddDefault("max_patch_size", maxPatchSize). - AddDefault("max_patch_files", maxPatchFiles). - AddIfNotEmpty("github-token", resolveHandlerGitHubToken(c.GitHubApp, "create-pull-request", c.GitHubToken)). - AddTemplatableBool("footer", getEffectiveFooterForTemplatable(c.Footer, cfg.Footer)). - AddBoolPtr("normalize_closing_keywords", c.NormalizeClosingKeywords). - AddBoolPtr("fallback_as_issue", c.FallbackAsIssue). - AddTemplatableBool("auto_close_issue", c.AutoCloseIssue). - AddIfNotEmpty("base_branch", c.BaseBranch). - AddDefault("protected_files_policy", protectedFilesPolicy). - AddStringSlice("protected_files", getAllManifestFiles()). - AddStringSlice("protected_path_prefixes", getProtectedPathPrefixes()). - AddDefault("protect_top_level_dot_folders", true). - AddStringSlice("_protected_files_exclude", c.ProtectedFilesExclude). - AddStringSlice("allowed_files", c.AllowedFiles). - AddStringSlice("excluded_files", c.ExcludedFiles). - AddIfTrue("preserve_branch_name", c.PreserveBranchName). - AddIfTrue("recreate_ref", c.RecreateRef). - AddIfNotEmpty("patch_format", c.PatchFormat). - AddBoolPtr("signed_commits", c.SignedCommits). - AddTemplatableBool("close_older_pull_requests", c.CloseOlderPullRequests). - AddIfNotEmpty("close_older_key", c.CloseOlderKey). - AddTemplatableBool("staged", templatableBoolPtrToStringPtr(c.Staged)) - // Use app-minted token if head-github-app is configured; fall back to head-github-token. - if c.HeadGitHubApp != nil { - //nolint:gosec // G101: False positive - this is a GitHub Actions expression template, not a hardcoded credential - builder.AddIfNotEmpty("head-github-token", "${{ steps.safe-outputs-head-app-token.outputs.token }}") - } else { - builder.AddIfNotEmpty("head-github-token", c.HeadGitHubToken) - } - return builder.Build() - }, + "create_pull_request": createPullRequestHandlerConfig, "push_to_pull_request_branch": func(cfg *SafeOutputsConfig) map[string]any { if cfg.PushToPullRequestBranch == nil { return nil @@ -1039,3 +966,97 @@ var handlerRegistry = map[string]handlerBuilder{ return config }, } + +func createPullRequestHandlerConfig(cfg *SafeOutputsConfig) map[string]any { + if cfg.CreatePullRequests == nil { + return nil + } + c := cfg.CreatePullRequests + builder := newHandlerConfigBuilder(). + AddTemplatableInt("max", c.Max). + AddIfTrue("require_temporary_id", c.RequireTemporaryID). + AddIfNotEmpty("branch_prefix", c.BranchPrefix). + AddIfNotEmpty("title_prefix", c.TitlePrefix). + AddTemplatableStringSlice("labels", c.Labels). + AddStringSlice("fallback_labels", c.FallbackLabels). + AddTemplatableStringSlice("reviewers", c.Reviewers). + AddTemplatableStringSlice("team_reviewers", c.TeamReviewers). + AddTemplatableStringSlice("assignees", c.Assignees). + AddTemplatableBool("draft", c.Draft). + AddIfNotEmpty("if_no_changes", c.IfNoChanges). + AddTemplatableBool("allow_empty", c.AllowEmpty). + AddTemplatableBool("auto_merge", c.AutoMerge). + AddIfPositive("expires", c.Expires). + AddIfNotEmpty("target-repo", c.TargetRepoSlug). + AddIfNotEmpty("head-repo", c.HeadRepoSlug). + AddTemplatableStringSlice("allowed_repos", c.AllowedRepos). + AddTemplatableStringSlice("allowed_base_branches", c.AllowedBaseBranches). + AddTemplatableStringSlice("allowed_branches", c.AllowedBranches) + addCreatePullRequestHandlerDefaults(builder, cfg, c) + addCreatePullRequestHeadToken(builder, c) + return builder.Build() +} + +func addCreatePullRequestHandlerDefaults(builder *handlerConfigBuilder, cfg *SafeOutputsConfig, c *CreatePullRequestsConfig) { + builder. + AddDefault("max_patch_size", resolveCreatePullRequestMaxPatchSize(cfg, c)). + AddDefault("max_patch_files", resolveCreatePullRequestMaxPatchFiles(cfg, c)). + AddIfNotEmpty("github-token", resolveHandlerGitHubToken(c.GitHubApp, "create-pull-request", c.GitHubToken)). + AddTemplatableBool("footer", getEffectiveFooterForTemplatable(c.Footer, cfg.Footer)). + AddBoolPtr("normalize_closing_keywords", c.NormalizeClosingKeywords). + AddBoolPtr("fallback_as_issue", c.FallbackAsIssue). + AddTemplatableBool("auto_close_issue", c.AutoCloseIssue). + AddIfNotEmpty("base_branch", c.BaseBranch). + AddDefault("protected_files_policy", resolveCreatePullRequestProtectedFilesPolicy(c)). + AddStringSlice("protected_files", getAllManifestFiles()). + AddStringSlice("protected_path_prefixes", getProtectedPathPrefixes()). + AddDefault("protect_top_level_dot_folders", true). + AddStringSlice("_protected_files_exclude", c.ProtectedFilesExclude). + AddStringSlice("allowed_files", c.AllowedFiles). + AddStringSlice("excluded_files", c.ExcludedFiles). + AddIfTrue("preserve_branch_name", c.PreserveBranchName). + AddIfTrue("recreate_ref", c.RecreateRef). + AddIfNotEmpty("patch_format", c.PatchFormat). + AddBoolPtr("signed_commits", c.SignedCommits). + AddTemplatableBool("close_older_pull_requests", c.CloseOlderPullRequests). + AddIfNotEmpty("close_older_key", c.CloseOlderKey). + AddTemplatableBool("staged", templatableBoolPtrToStringPtr(c.Staged)) +} + +func addCreatePullRequestHeadToken(builder *handlerConfigBuilder, c *CreatePullRequestsConfig) { + if c.HeadGitHubApp != nil { + //nolint:gosec // G101: False positive - this is a GitHub Actions expression template, not a hardcoded credential + builder.AddIfNotEmpty("head-github-token", "${{ steps.safe-outputs-head-app-token.outputs.token }}") + return + } + builder.AddIfNotEmpty("head-github-token", c.HeadGitHubToken) +} + +func resolveCreatePullRequestProtectedFilesPolicy(c *CreatePullRequestsConfig) string { + if c.ManifestFilesPolicy != nil { + return *c.ManifestFilesPolicy + } + return "request_review" +} + +func resolveCreatePullRequestMaxPatchSize(cfg *SafeOutputsConfig, c *CreatePullRequestsConfig) int { + maxPatchSize := 4096 + if cfg.MaximumPatchSize > 0 { + maxPatchSize = cfg.MaximumPatchSize + } + if c.MaxPatchSize > 0 { + maxPatchSize = c.MaxPatchSize + } + return maxPatchSize +} + +func resolveCreatePullRequestMaxPatchFiles(cfg *SafeOutputsConfig, c *CreatePullRequestsConfig) int { + maxPatchFiles := 100 + if cfg.MaximumPatchFiles > 0 { + maxPatchFiles = cfg.MaximumPatchFiles + } + if c.MaxPatchFiles > 0 { + maxPatchFiles = c.MaxPatchFiles + } + return maxPatchFiles +} diff --git a/pkg/workflow/safe_outputs_jobs.go b/pkg/workflow/safe_outputs_jobs.go index 7762cd7c123..d425c095d41 100644 --- a/pkg/workflow/safe_outputs_jobs.go +++ b/pkg/workflow/safe_outputs_jobs.go @@ -50,6 +50,10 @@ type SafeOutputJobConfig struct { // 3. Invoke buildGitHubScriptStep // 4. Create Job with standard metadata func (c *Compiler) buildSafeOutputJob(data *WorkflowData, config SafeOutputJobConfig) (*Job, error) { + return c.buildSafeOutputJobBody(data, config) +} + +func (c *Compiler) buildSafeOutputJobBody(data *WorkflowData, config SafeOutputJobConfig) (*Job, error) { safeOutputsJobsLog.Printf("Building safe output job: %s (actionMode=%s)", config.JobName, c.actionMode) var steps []string diff --git a/pkg/workflow/safe_outputs_max_validation.go b/pkg/workflow/safe_outputs_max_validation.go index 1b2301ebc21..87ebf5246ca 100644 --- a/pkg/workflow/safe_outputs_max_validation.go +++ b/pkg/workflow/safe_outputs_max_validation.go @@ -56,6 +56,10 @@ func checkMaxField(toolName string, maxPtr *string) error { // it is on the hot path and called on every compilation. The field ordering matches // the sorted safeOutputFieldMapping keys for deterministic error reporting. func validateSafeOutputsMax(config *SafeOutputsConfig) error { + return validateSafeOutputsMaxBody(config) +} + +func validateSafeOutputsMaxBody(config *SafeOutputsConfig) error { if config == nil { return nil } diff --git a/pkg/workflow/safe_outputs_messages_config.go b/pkg/workflow/safe_outputs_messages_config.go index d449ecb063b..642626e9dd3 100644 --- a/pkg/workflow/safe_outputs_messages_config.go +++ b/pkg/workflow/safe_outputs_messages_config.go @@ -65,6 +65,10 @@ func parseMessagesConfig(messagesMap map[string]any) *SafeOutputMessagesConfig { // - true: always allows mentions (error in strict mode) // - object: detailed configuration with allowed-collaborators, allow-context, allowed, max func parseMentionsConfig(mentions any) *MentionsConfig { + return parseMentionsConfigBody(mentions) +} + +func parseMentionsConfigBody(mentions any) *MentionsConfig { safeOutputMessagesLog.Printf("Parsing mentions configuration: type=%T", mentions) config := &MentionsConfig{} diff --git a/pkg/workflow/safe_outputs_permissions.go b/pkg/workflow/safe_outputs_permissions.go index 4c1f19e102f..777cba32972 100644 --- a/pkg/workflow/safe_outputs_permissions.go +++ b/pkg/workflow/safe_outputs_permissions.go @@ -110,6 +110,10 @@ func ComputePermissionsForSafeOutputs(safeOutputs *SafeOutputsConfig) *Permissio } func computePermissionsForSafeOutputs(safeOutputs *SafeOutputsConfig, excludePerHandlerApps bool) *Permissions { + return computePermissionsForSafeOutputsBody(safeOutputs, excludePerHandlerApps) +} + +func computePermissionsForSafeOutputsBody(safeOutputs *SafeOutputsConfig, excludePerHandlerApps bool) *Permissions { if safeOutputs == nil { safeOutputsPermissionsLog.Print("No safe outputs configured, returning empty permissions") return NewPermissions() diff --git a/pkg/workflow/safe_outputs_steps_shell_expansion_validation.go b/pkg/workflow/safe_outputs_steps_shell_expansion_validation.go index 8d55385b41c..3b522ce6edc 100644 --- a/pkg/workflow/safe_outputs_steps_shell_expansion_validation.go +++ b/pkg/workflow/safe_outputs_steps_shell_expansion_validation.go @@ -117,6 +117,10 @@ func validateSafeOutputsStepsShellExpansion(config *SafeOutputsConfig) error { // validateRunScriptForShellExpansion checks a single run: script for dangerous // bash expansion patterns. stepIndex is 0-based and is included in error messages. func validateRunScriptForShellExpansion(stepIndex int, script string) error { + return validateRunScriptForShellExpansionBody(stepIndex, script) +} + +func validateRunScriptForShellExpansionBody(stepIndex int, script string) error { // Fast path: no '$' or backtick character means no expansion pattern can be present. if !strings.ContainsAny(script, "$`") { return nil diff --git a/pkg/workflow/safe_outputs_tools_computation.go b/pkg/workflow/safe_outputs_tools_computation.go index 420c62fac25..4b7b599a1d9 100644 --- a/pkg/workflow/safe_outputs_tools_computation.go +++ b/pkg/workflow/safe_outputs_tools_computation.go @@ -8,6 +8,11 @@ var safeOutputsToolsComputationLog = logger.New("workflow:safe_outputs_tools_com // by the workflow's SafeOutputsConfig. Dynamic tools (dispatch-workflow, custom jobs, // call-workflow) are excluded because they are generated separately. func computeEnabledToolNames(data *WorkflowData) map[string]struct { +} { + return computeEnabledToolNamesBody(data) +} + +func computeEnabledToolNamesBody(data *WorkflowData) map[string]struct { } { enabledTools := make(map[string]struct { }) diff --git a/pkg/workflow/safe_outputs_tools_generation.go b/pkg/workflow/safe_outputs_tools_generation.go index 83e22df5dfa..3a2c34e234c 100644 --- a/pkg/workflow/safe_outputs_tools_generation.go +++ b/pkg/workflow/safe_outputs_tools_generation.go @@ -30,6 +30,10 @@ import ( // These tools are not in safe_outputs_tools.json and must be generated from // the workflow configuration at compile time. func generateDynamicTools(data *WorkflowData, markdownPath string) ([]map[string]any, error) { + return generateDynamicToolsBody(data, markdownPath) +} + +func generateDynamicToolsBody(data *WorkflowData, markdownPath string) ([]map[string]any, error) { var dynamicTools []map[string]any // Add custom job tools from SafeOutputs.Jobs @@ -366,6 +370,10 @@ func computePropertyInjections(safeOutputs *SafeOutputsConfig) map[string]map[st // the actions folder, applies the meta overrides from tools_meta.json, and writes // the final ${RUNNER_TEMP}/gh-aw/safeoutputs/tools.json. func generateToolsMetaJSON(data *WorkflowData, markdownPath string) (string, error) { + return generateToolsMetaJSONBody(data, markdownPath) +} + +func generateToolsMetaJSONBody(data *WorkflowData, markdownPath string) (string, error) { if data.SafeOutputs == nil { empty := ToolsMeta{ DescriptionSuffixes: map[string]string{}, diff --git a/pkg/workflow/safe_outputs_tools_repo_params.go b/pkg/workflow/safe_outputs_tools_repo_params.go index 57883598ff2..1677f175e6a 100644 --- a/pkg/workflow/safe_outputs_tools_repo_params.go +++ b/pkg/workflow/safe_outputs_tools_repo_params.go @@ -5,6 +5,10 @@ import "fmt" // addRepoParameterIfNeeded adds a "repo" parameter to the tool's inputSchema // if the safe output configuration has allowed-repos entries or a wildcard "*" target-repo func addRepoParameterIfNeeded(tool map[string]any, toolName string, safeOutputs *SafeOutputsConfig) { + addRepoParameterIfNeededBody(tool, toolName, safeOutputs) +} + +func addRepoParameterIfNeededBody(tool map[string]any, toolName string, safeOutputs *SafeOutputsConfig) { safeOutputsConfigLog.Printf("Checking if repo parameter needed for tool: %s", toolName) if safeOutputs == nil { return diff --git a/pkg/workflow/safe_outputs_validation.go b/pkg/workflow/safe_outputs_validation.go index 44cad1349e6..69aee123542 100644 --- a/pkg/workflow/safe_outputs_validation.go +++ b/pkg/workflow/safe_outputs_validation.go @@ -79,6 +79,10 @@ var safeOutputsTargetValidationLog = logger.New("workflow:safe_outputs_target_va // - A positive integer as a string (e.g., "123") // - A GitHub Actions expression (e.g., "${{ github.event.issue.number }}") func validateSafeOutputsTarget(config *SafeOutputsConfig) error { + return validateSafeOutputsTargetBody(config) +} + +func validateSafeOutputsTargetBody(config *SafeOutputsConfig) error { if config == nil { return nil } diff --git a/pkg/workflow/safe_outputs_validation_config.go b/pkg/workflow/safe_outputs_validation_config.go index 7537adf311b..a519eeba8b4 100644 --- a/pkg/workflow/safe_outputs_validation_config.go +++ b/pkg/workflow/safe_outputs_validation_config.go @@ -493,6 +493,10 @@ var validationConfigJSONCache sync.Map // key: string → value: string // GetValidationConfigJSONWithDataSchema behaves like GetValidationConfigJSONWithDataSchema and additionally // injects a normalized data schema into body-bearing safe-output types. func GetValidationConfigJSONWithDataSchema(enabledTypes []string, mentions map[string]any, dataEnabled bool, dataSchema map[string]any) (string, error) { + return GetValidationConfigJSONWithDataSchemaBody(enabledTypes, mentions, dataEnabled, dataSchema) +} + +func GetValidationConfigJSONWithDataSchemaBody(enabledTypes []string, mentions map[string]any, dataEnabled bool, dataSchema map[string]any) (string, error) { safeOutputValidationLog.Printf("Getting validation config JSON for %d types (mentions=%t)", len(enabledTypes), len(mentions) > 0) // Cache only the schema-only path; mentions are workflow-specific and cheap to remarshal. diff --git a/pkg/workflow/tools_parser.go b/pkg/workflow/tools_parser.go index 96650d7e0a9..30a80cfb71a 100644 --- a/pkg/workflow/tools_parser.go +++ b/pkg/workflow/tools_parser.go @@ -107,6 +107,10 @@ var knownTools = map[string]struct{}{ } func NewTools(toolsMap map[string]any) *Tools { + return NewToolsBody(toolsMap) +} + +func NewToolsBody(toolsMap map[string]any) *Tools { toolsParserLog.Printf("Creating tools configuration from map with %d entries", len(toolsMap)) if toolsMap == nil { return &Tools{ @@ -188,6 +192,10 @@ func NewTools(toolsMap map[string]any) *Tools { // parseGitHubTool converts raw github tool configuration to GitHubToolConfig func parseGitHubTool(val any) *GitHubToolConfig { + return parseGitHubToolBody(val) +} + +func parseGitHubToolBody(val any) *GitHubToolConfig { if val == nil { toolsParserLog.Print("GitHub tool enabled with default configuration") return &GitHubToolConfig{ @@ -687,6 +695,10 @@ func parseStartupTimeoutTool(val any) *TemplatableInt32 { // parseMCPServerConfig converts raw MCP server configuration to MCPServerConfig func parseMCPServerConfig(val any) MCPServerConfig { + return parseMCPServerConfigBody(val) +} + +func parseMCPServerConfigBody(val any) MCPServerConfig { config := MCPServerConfig{ CustomFields: make(map[string]any), } From 284b28058a664dadeef703b13c41dcaf226b7d78 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Sun, 2 Aug 2026 18:56:31 +0000 Subject: [PATCH 4/4] Refine audit diff metrics check Co-authored-by: pelikhan <4175913+pelikhan@users.noreply.github.com> --- pkg/cli/audit_diff.go | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/pkg/cli/audit_diff.go b/pkg/cli/audit_diff.go index 972c246431b..212204585a7 100644 --- a/pkg/cli/audit_diff.go +++ b/pkg/cli/audit_diff.go @@ -594,11 +594,11 @@ func runSummaryTurnCount(summary *RunSummary) int { } func hasRunMetricsData(inputs runMetricsInputs) bool { - return !(inputs.run1Tokens == 0 && inputs.run2Tokens == 0 && - inputs.run1Duration == 0 && inputs.run2Duration == 0 && - inputs.run1Turns == 0 && inputs.run2Turns == 0 && - inputs.tokenUsage1 == nil && inputs.tokenUsage2 == nil && - inputs.rateLimit1 == nil && inputs.rateLimit2 == nil) + return inputs.run1Tokens != 0 || inputs.run2Tokens != 0 || + inputs.run1Duration != 0 || inputs.run2Duration != 0 || + inputs.run1Turns != 0 || inputs.run2Turns != 0 || + inputs.tokenUsage1 != nil || inputs.tokenUsage2 != nil || + inputs.rateLimit1 != nil || inputs.rateLimit2 != nil } func newRunMetricsDiff(inputs runMetricsInputs) *RunMetricsDiff {