Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 commits
Commits
Show all changes
18 commits
Select commit Hold shift + click to select a range
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions pkg/workflow/safe_outputs_tools_generation_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -234,6 +234,23 @@ func TestAddRepoParameterIfNeededClosePullRequestWithAllowedRepos(t *testing.T)
assert.Contains(t, repoProp["description"].(string), "org/default-repo", "description should include default repo")
}

func TestRepoTargetAccessorsCoverRepoTargetTools(t *testing.T) {
expectedTools := []string{
"create_issue", "create_discussion", "add_comment", "create_pull_request",
"create_pull_request_review_comment", "reply_to_pull_request_review_comment",
"dismiss_pull_request_review", "create_agent_session", "close_issue", "update_issue",
"close_discussion", "update_discussion", "close_pull_request", "update_pull_request",
"merge_pull_request", "add_labels", "remove_labels", "replace_label", "hide_comment",
"link_sub_issue", "mark_pull_request_as_ready_for_review", "add_reviewer",
"assign_milestone", "assign_to_agent", "assign_to_user", "unassign_from_user",
"set_issue_type", "set_issue_field",
}
assert.Len(t, repoTargetAccessors, len(expectedTools))
for _, toolName := range expectedTools {
assert.Contains(t, repoTargetAccessors, toolName)
}
}

func TestParseUpdateIssuesConfigWithWildcardTargetRepo(t *testing.T) {
compiler := &Compiler{}
outputMap := map[string]any{
Comment thread
github-actions[bot] marked this conversation as resolved.
Expand Down
338 changes: 189 additions & 149 deletions pkg/workflow/safe_outputs_tools_repo_params.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,184 @@ package workflow

import "fmt"

type repoTargetConfig struct {
allowedRepos []string
targetRepoSlug string
}

type repoTargetAccessor func(*SafeOutputsConfig) *repoTargetConfig

var repoTargetAccessors = map[string]repoTargetAccessor{
Comment thread
github-actions[bot] marked this conversation as resolved.
Outdated
"create_issue": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.CreateIssues; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"create_discussion": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.CreateDiscussions; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"add_comment": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.AddComments; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"create_pull_request": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.CreatePullRequests; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"create_pull_request_review_comment": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.CreatePullRequestReviewComments; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"reply_to_pull_request_review_comment": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.ReplyToPullRequestReviewComment; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"dismiss_pull_request_review": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.DismissPullRequestReview; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"create_agent_session": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.CreateAgentSessions; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"close_issue": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.CloseIssues; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"update_issue": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.UpdateIssues; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"close_discussion": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.CloseDiscussions; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"update_discussion": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.UpdateDiscussions; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"close_pull_request": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.ClosePullRequests; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"update_pull_request": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.UpdatePullRequests; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"merge_pull_request": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.MergePullRequest; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"add_labels": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.AddLabels; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"remove_labels": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.RemoveLabels; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"replace_label": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.ReplaceLabel; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"hide_comment": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.HideComment; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"link_sub_issue": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.LinkSubIssue; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"mark_pull_request_as_ready_for_review": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.MarkPullRequestAsReadyForReview; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"add_reviewer": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.AddReviewer; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"assign_milestone": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.AssignMilestone; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"assign_to_agent": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.AssignToAgent; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"assign_to_user": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.AssignToUser; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"unassign_from_user": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.UnassignFromUser; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"set_issue_type": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.SetIssueType; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
"set_issue_field": func(config *SafeOutputsConfig) *repoTargetConfig {
if output := config.SetIssueField; output != nil {
return &repoTargetConfig{output.AllowedRepos, output.TargetRepoSlug}
}
return nil
},
}

// 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) {
Expand All @@ -10,155 +188,17 @@ func addRepoParameterIfNeeded(tool map[string]any, toolName string, safeOutputs
return
}

// Determine if this tool should have a repo parameter based on allowed-repos and target-repo configuration (including wildcard "*")
var hasAllowedRepos bool
var targetRepoSlug string

switch toolName {
case "create_issue":
if config := safeOutputs.CreateIssues; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "create_discussion":
if config := safeOutputs.CreateDiscussions; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "add_comment":
if config := safeOutputs.AddComments; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "create_pull_request":
if config := safeOutputs.CreatePullRequests; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "create_pull_request_review_comment":
if config := safeOutputs.CreatePullRequestReviewComments; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "reply_to_pull_request_review_comment":
if config := safeOutputs.ReplyToPullRequestReviewComment; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "dismiss_pull_request_review":
if config := safeOutputs.DismissPullRequestReview; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "create_agent_session":
if config := safeOutputs.CreateAgentSessions; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "close_issue", "update_issue":
if config := safeOutputs.CloseIssues; config != nil && toolName == "close_issue" {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
} else if config := safeOutputs.UpdateIssues; config != nil && toolName == "update_issue" {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "close_discussion", "update_discussion":
if config := safeOutputs.CloseDiscussions; config != nil && toolName == "close_discussion" {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
} else if config := safeOutputs.UpdateDiscussions; config != nil && toolName == "update_discussion" {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "close_pull_request", "update_pull_request":
if config := safeOutputs.ClosePullRequests; config != nil && toolName == "close_pull_request" {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
} else if config := safeOutputs.UpdatePullRequests; config != nil && toolName == "update_pull_request" {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "merge_pull_request":
if config := safeOutputs.MergePullRequest; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "add_labels", "remove_labels", "replace_label", "hide_comment", "link_sub_issue", "mark_pull_request_as_ready_for_review",
"add_reviewer", "assign_milestone", "assign_to_agent", "assign_to_user", "unassign_from_user",
"set_issue_type", "set_issue_field":
// These use SafeOutputTargetConfig - check the appropriate config
switch toolName {
case "add_labels":
if config := safeOutputs.AddLabels; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "remove_labels":
if config := safeOutputs.RemoveLabels; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "replace_label":
if config := safeOutputs.ReplaceLabel; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "hide_comment":
if config := safeOutputs.HideComment; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "link_sub_issue":
if config := safeOutputs.LinkSubIssue; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "mark_pull_request_as_ready_for_review":
if config := safeOutputs.MarkPullRequestAsReadyForReview; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "add_reviewer":
if config := safeOutputs.AddReviewer; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "assign_milestone":
if config := safeOutputs.AssignMilestone; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "assign_to_agent":
if config := safeOutputs.AssignToAgent; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "assign_to_user":
if config := safeOutputs.AssignToUser; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "unassign_from_user":
if config := safeOutputs.UnassignFromUser; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "set_issue_type":
if config := safeOutputs.SetIssueType; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
case "set_issue_field":
if config := safeOutputs.SetIssueField; config != nil {
hasAllowedRepos = len(config.AllowedRepos) > 0
targetRepoSlug = config.TargetRepoSlug
}
}
accessor, ok := repoTargetAccessors[toolName]
if !ok {
return
}
targetConfig := accessor(safeOutputs)
if targetConfig == nil {
return
}

// Only add repo parameter if allowed-repos has entries or target-repo is wildcard ("*")
if !hasAllowedRepos && targetRepoSlug != "*" {
if len(targetConfig.allowedRepos) == 0 && targetConfig.targetRepoSlug != "*" {
safeOutputsConfigLog.Printf("Skipping repo parameter for tool %s: no allowed-repos and target-repo is not wildcard", toolName)
return
}
Expand All @@ -176,10 +216,10 @@ func addRepoParameterIfNeeded(tool map[string]any, toolName string, safeOutputs

// Build repo parameter description
var repoDescription string
if targetRepoSlug == "*" {
if targetConfig.targetRepoSlug == "*" {
repoDescription = "Target repository for this operation in 'owner/repo' format. Any repository can be targeted."
} else if targetRepoSlug != "" {
repoDescription = fmt.Sprintf("Target repository for this operation in 'owner/repo' format. Default is %q. Must be the target-repo or in the allowed-repos list.", targetRepoSlug)
} else if targetConfig.targetRepoSlug != "" {
repoDescription = fmt.Sprintf("Target repository for this operation in 'owner/repo' format. Default is %q. Must be the target-repo or in the allowed-repos list.", targetConfig.targetRepoSlug)
} else {
repoDescription = "Target repository for this operation in 'owner/repo' format. Must be the target-repo or in the allowed-repos list."
}
Expand Down
Loading