Files
1Panel/agent/utils/terminal/ai/command_generator.go
glmgbj233andssongliu 8ed33fd06a feat: validate generated terminal commands (#12826)
* Fix taint path 1 detected by Codex

Repository: 1Panel-dev_1Panel
Result: /home/hejunjie/llm_web_serve/find_github_project/codex_find_taint/batch_results_go_llm_full_mini/1Panel-dev_1Panel.json
Sink kind: command_exec

* Fix taint path 0 detected by Codex

Repository: 1Panel-dev_1Panel
Result: /home/hejunjie/llm_web_serve/find_github_project/codex_find_taint/batch_results_go_llm_full_mini/1Panel-dev_1Panel.json
Sink kind: command_exec

* Validate generated terminal commands

Keep generated output behind the intended trust boundary while preserving the normal safe workflow.

Add regression coverage for the unsafe flow and the expected safe behavior.

* Narrow terminal AI input validation

* refine terminal ai command validation

---------

Co-authored-by: ssongliu <sloooop1x@gmail.com>
2026-05-27 11:08:52 +08:00

200 lines
5.1 KiB
Go

package ai
import (
"context"
"fmt"
"strings"
"unicode"
)
type CommandGenerator struct {
client Client
}
type CommandGenerateRequest struct {
Input string
Shell string
WorkingDir string
OS string
RecentCommands []string
DirectoryHints []string
}
type CommandGenerateResponse struct {
Command string
Model string
Provider string
RawText string
Usage ResponseUsage
}
func NewCommandGeneratorFromConfig(cfg GeneratorConfig) (*CommandGenerator, error) {
client, err := NewClient(ClientConfig{
Provider: cfg.Provider,
BaseURL: cfg.BaseURL,
APIKey: cfg.APIKey,
Model: cfg.Model,
APIType: cfg.APIType,
MaxTokens: cfg.MaxTokens,
})
if err != nil {
return nil, err
}
return NewCommandGenerator(client)
}
func NewCommandGenerator(client Client) (*CommandGenerator, error) {
if client == nil {
return nil, fmt.Errorf("client is required")
}
return &CommandGenerator{client: client}, nil
}
func (g *CommandGenerator) Generate(ctx context.Context, req CommandGenerateRequest) (*CommandGenerateResponse, error) {
if strings.TrimSpace(req.Input) == "" {
return nil, fmt.Errorf("input is required")
}
resp, err := g.client.ChatCompletion(ctx, ChatCompletionRequest{
Messages: []ChatMessage{
{Role: "system", Content: buildCommandSystemPrompt()},
{Role: "user", Content: buildCommandUserPrompt(req)},
},
})
if err != nil {
return nil, err
}
command := sanitizeCommand(resp.Content)
if command == "" {
return nil, fmt.Errorf("model returned empty command")
}
if err := validateGeneratedCommand(command); err != nil {
return nil, err
}
return &CommandGenerateResponse{
Command: command,
Model: resp.Model,
Provider: providerNameFromModel(resp.Model),
RawText: resp.RawText,
Usage: resp.Usage,
}, nil
}
func buildCommandSystemPrompt() string {
return strings.Join([]string{
"You are a shell command generator.",
"Return exactly one command suitable for direct execution in the user's shell.",
"Do not include markdown, code fences, explanations, numbering, comments, or backticks.",
"If multiple commands are required, join them with shell operators in a single line.",
"Prefer safe, non-destructive commands unless the user explicitly asks for destructive behavior.",
"Preserve the user's language when filenames or arguments are ambiguous, but output only the command.",
}, "\n")
}
func buildCommandUserPrompt(req CommandGenerateRequest) string {
var sections []string
sections = append(sections, "Task:\n"+strings.TrimSpace(req.Input))
var env []string
if shell := strings.TrimSpace(req.Shell); shell != "" {
env = append(env, "Shell: "+shell)
}
if wd := strings.TrimSpace(req.WorkingDir); wd != "" {
env = append(env, "Working directory: "+wd)
}
if osName := strings.TrimSpace(req.OS); osName != "" {
env = append(env, "Operating system: "+osName)
}
if len(env) > 0 {
sections = append(sections, "Environment:\n"+strings.Join(env, "\n"))
}
if block := formatBulletBlock(req.DirectoryHints); block != "" {
sections = append(sections, "Directory hints:\n"+block)
}
if block := formatBulletBlock(req.RecentCommands); block != "" {
sections = append(sections, "Recent commands:\n"+block)
}
sections = append(sections, "Output requirement:\nReturn one shell command only.")
return strings.Join(sections, "\n\n")
}
func formatBulletBlock(values []string) string {
var lines []string
for _, value := range values {
value = strings.TrimSpace(value)
if value == "" {
continue
}
lines = append(lines, "- "+value)
}
return strings.Join(lines, "\n")
}
func sanitizeCommand(raw string) string {
command := strings.TrimSpace(raw)
if command == "" {
return ""
}
command = strings.TrimPrefix(command, "```sh")
command = strings.TrimPrefix(command, "```bash")
command = strings.TrimPrefix(command, "```zsh")
command = strings.TrimPrefix(command, "```shell")
command = strings.TrimPrefix(command, "```")
command = strings.TrimSuffix(command, "```")
command = strings.TrimSpace(command)
lines := strings.Split(command, "\n")
for _, line := range lines {
line = strings.TrimSpace(line)
if line == "" {
continue
}
if strings.HasPrefix(line, "#") {
continue
}
if strings.HasPrefix(strings.ToLower(line), "command:") {
line = strings.TrimSpace(line[len("command:"):])
}
return strings.Trim(line, "` ")
}
return ""
}
func validateGeneratedCommand(command string) error {
command = strings.TrimSpace(command)
if command == "" {
return fmt.Errorf("model returned empty command")
}
if !isSingleLinePrintableCommand(command) {
return fmt.Errorf("model returned unsafe command")
}
return nil
}
func isSingleLinePrintableCommand(command string) bool {
if strings.ContainsAny(command, "\x00\r\n\x1b") {
return false
}
for _, r := range command {
if unicode.IsControl(r) || unicode.In(r, unicode.Cf) {
return false
}
}
return true
}
func providerNameFromModel(model string) string {
model = strings.TrimSpace(model)
if model == "" {
return ""
}
if parts := strings.SplitN(model, "/", 2); len(parts) == 2 {
return parts[0]
}
return ""
}