mirror of
https://github.com/Ed1s0nZ/CyberStrikeAI.git
synced 2026-08-01 00:27:35 +02:00
Add files via upload
This commit is contained in:
@@ -172,6 +172,8 @@ func permissionForRequest(method, fullPath string) string {
|
||||
return "workflow:read"
|
||||
case strings.HasPrefix(path, "/workflow-package-inspections"), strings.HasPrefix(path, "/workflow-package-imports"):
|
||||
return "workflow:write"
|
||||
case path == "/workflows/generate-draft":
|
||||
return "workflow:write"
|
||||
case strings.HasPrefix(path, "/workflows"):
|
||||
if path == "/workflows/validate" || path == "/workflows/dry-run" || strings.HasSuffix(path, "/resume") {
|
||||
return "workflow:execute"
|
||||
@@ -265,7 +267,7 @@ func isProcessGlobalMutationPath(path string) bool {
|
||||
}
|
||||
if strings.HasPrefix(path, "/workflows") {
|
||||
// Workflow runs inherit conversation access; definitions are global.
|
||||
return !strings.HasPrefix(path, "/workflows/runs/") && path != "/workflows/validate" && path != "/workflows/dry-run"
|
||||
return !strings.HasPrefix(path, "/workflows/runs/") && path != "/workflows/validate" && path != "/workflows/dry-run" && path != "/workflows/generate-draft"
|
||||
}
|
||||
if strings.HasPrefix(path, "/workflow-package-inspections") || strings.HasPrefix(path, "/workflow-package-imports") {
|
||||
return true
|
||||
|
||||
@@ -157,9 +157,15 @@ func TestWorkflowRunPermissionIsSeparateFromDefinitionManagement(t *testing.T) {
|
||||
if got := permissionForRequest(http.MethodPost, "/api/workflows/runs/run-1/resume"); got != "workflow:execute" {
|
||||
t.Fatalf("resume permission = %q, want workflow:execute", got)
|
||||
}
|
||||
if got := permissionForRequest(http.MethodPost, "/api/workflows/generate-draft"); got != "workflow:write" {
|
||||
t.Fatalf("generate draft permission = %q, want workflow:write", got)
|
||||
}
|
||||
if got := permissionForRequest(http.MethodPut, "/api/workflows/workflow-1"); got != "workflow:write" {
|
||||
t.Fatalf("definition permission = %q, want workflow:write", got)
|
||||
}
|
||||
if isProcessGlobalMutationPath("/workflows/generate-draft") {
|
||||
t.Fatalf("generate draft should not be treated as a process-global mutation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRBACDenyHookReceivesDeniedDecision(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,782 @@
|
||||
package workflow
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"hash/fnv"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"cyberstrike-ai/internal/config"
|
||||
"cyberstrike-ai/internal/openai"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
type DraftTool struct {
|
||||
Key string `json:"key"`
|
||||
Name string `json:"name,omitempty"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
type DraftOptions struct {
|
||||
IncludeObjective bool `json:"include_objective"`
|
||||
AllowSchedule bool `json:"allow_schedule"`
|
||||
AllowHighRisk bool `json:"allow_high_risk"`
|
||||
}
|
||||
|
||||
type DraftRequest struct {
|
||||
Prompt string `json:"prompt"`
|
||||
Options DraftOptions `json:"options"`
|
||||
AvailableTools []DraftTool `json:"available_tools,omitempty"`
|
||||
}
|
||||
|
||||
type DraftMeta struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
type DraftCapability struct {
|
||||
Label string `json:"label"`
|
||||
ToolName string `json:"tool_name,omitempty"`
|
||||
ToolCandidates []string `json:"tool_candidates,omitempty"`
|
||||
}
|
||||
|
||||
type DraftAudit struct {
|
||||
Savable bool `json:"savable"`
|
||||
Validation []string `json:"validation,omitempty"`
|
||||
MissingFields []string `json:"missing_fields,omitempty"`
|
||||
RiskWarnings []string `json:"risk_warnings,omitempty"`
|
||||
Assumptions []string `json:"assumptions,omitempty"`
|
||||
HighRisk bool `json:"high_risk"`
|
||||
NeedsHITL bool `json:"needs_hitl"`
|
||||
}
|
||||
|
||||
type DraftResult struct {
|
||||
Graph *graphDef `json:"graph"`
|
||||
Meta DraftMeta `json:"meta"`
|
||||
Generator string `json:"generator"`
|
||||
Audit DraftAudit `json:"audit"`
|
||||
Capabilities []DraftCapability `json:"capabilities,omitempty"`
|
||||
Stats map[string]int `json:"stats"`
|
||||
}
|
||||
|
||||
type llmDraftEnvelope struct {
|
||||
Graph graphDef `json:"graph"`
|
||||
Meta DraftMeta `json:"meta"`
|
||||
Capabilities []DraftCapability `json:"capabilities,omitempty"`
|
||||
Audit DraftAudit `json:"audit,omitempty"`
|
||||
}
|
||||
|
||||
type draftToolHint struct {
|
||||
Label string
|
||||
Keywords []string
|
||||
Tools []string
|
||||
}
|
||||
|
||||
var draftToolHints = []draftToolHint{
|
||||
{Label: "子域名发现", Keywords: []string{"子域名", "subdomain", "subfinder", "amass"}, Tools: []string{"subfinder", "amass"}},
|
||||
{Label: "端口扫描", Keywords: []string{"端口", "port", "nmap", "rustscan", "masscan"}, Tools: []string{"nmap", "rustscan", "masscan"}},
|
||||
{Label: "漏洞扫描", Keywords: []string{"漏洞", "vuln", "漏洞扫描", "nuclei", "nikto", "zap"}, Tools: []string{"nuclei", "nikto", "zap"}},
|
||||
{Label: "暴露面探测", Keywords: []string{"目录", "路径", "暴露页面", "dir", "ffuf", "gobuster", "feroxbuster"}, Tools: []string{"ffuf", "gobuster", "feroxbuster", "dirsearch"}},
|
||||
{Label: "证书与域名线索收集", Keywords: []string{"证书", "certificate", "crt"}, Tools: []string{"subfinder"}},
|
||||
{Label: "云配置审计", Keywords: []string{"云", "cloud", "配置审计", "prowler", "scout"}, Tools: []string{"prowler", "scout-suite"}},
|
||||
{Label: "容器安全检查", Keywords: []string{"容器", "镜像", "k8s", "kubernetes", "trivy", "kube"}, Tools: []string{"trivy", "kube-bench", "kube-hunter"}},
|
||||
{Label: "威胁情报收集", Keywords: []string{"情报", "威胁情报", "threat", "ioc", "virustotal", "shodan", "fofa"}, Tools: []string{"virustotal_search", "shodan_search", "fofa_search"}},
|
||||
}
|
||||
|
||||
var highRiskDraftRE = regexp.MustCompile(`(?i)(隔离|封禁|加固|修复|执行|命令|脚本|删除|清理|阻断|封锁|攻击|利用|getshell|shell|payload|exploit|isolate|block|execute|script|delete|exploit|payload)`)
|
||||
|
||||
func GenerateDraftFromNaturalLanguage(ctx context.Context, req DraftRequest) (*DraftResult, error) {
|
||||
prompt := strings.TrimSpace(req.Prompt)
|
||||
if prompt == "" {
|
||||
return nil, fmt.Errorf("工作流需求不能为空")
|
||||
}
|
||||
capabilities := detectDraftCapabilities(prompt, req.AvailableTools)
|
||||
wantsApproval := containsAnyFold(prompt, "审批", "确认", "审核", "负责人", "人工", "review", "approve", "approval", "human")
|
||||
wantsReport := containsAnyFold(prompt, "报告", "汇总", "输出", "通知", "任务", "工单", "report", "summary", "notify", "ticket")
|
||||
wantsCondition := containsAnyFold(prompt, "如果", "发现", "存在", "高危", "新增", "失败", "通过", "否则", "if", "when", "high", "critical", "new", "fail")
|
||||
highRisk := highRiskDraftRE.MatchString(prompt)
|
||||
|
||||
builder := &draftGraphBuilder{x: 120, y: 150}
|
||||
assumptions := make([]string, 0)
|
||||
riskWarnings := make([]string, 0)
|
||||
missingFields := make([]string, 0)
|
||||
|
||||
start := builder.add("start", "开始", map[string]any{"input_keys": "message, conversationId, projectId, target"}, 0)
|
||||
previous := start
|
||||
for _, capability := range capabilities {
|
||||
hasTool := strings.TrimSpace(capability.ToolName) != ""
|
||||
var id string
|
||||
if hasTool {
|
||||
id = builder.add("tool", capability.Label, map[string]any{
|
||||
"tool_name": capability.ToolName,
|
||||
"arguments": `{"target":"{{inputs.target}}","message":"{{inputs.message}}"}`,
|
||||
"timeout_seconds": "120",
|
||||
"join_strategy": "all_merge",
|
||||
}, 0)
|
||||
} else {
|
||||
id = builder.add("agent", capability.Label, map[string]any{
|
||||
"agent_mode": "eino_single",
|
||||
"input_binding": map[string]any{"from": "previous", "field": "output"},
|
||||
"instruction": capability.Label + "。根据用户需求执行安全流程步骤,并输出结构化结果:" + prompt,
|
||||
"output_key": "agent_result",
|
||||
"join_strategy": "all_merge",
|
||||
"missing_tool_candidates": strings.Join(capability.ToolCandidates, ", "),
|
||||
}, 0)
|
||||
if len(capability.ToolCandidates) > 0 {
|
||||
assumptions = append(assumptions, capability.Label+" 未匹配到已启用工具,已生成 Agent 草稿节点。")
|
||||
missingFields = append(missingFields, capability.Label+": 选择或启用对应 MCP 工具")
|
||||
}
|
||||
}
|
||||
builder.connect(previous, id, "", nil)
|
||||
previous = id
|
||||
}
|
||||
|
||||
openConditionID := ""
|
||||
if wantsCondition {
|
||||
expr := `{{previous.output}} != ""`
|
||||
label := "是否满足触发条件"
|
||||
if highRisk {
|
||||
expr = `{{previous.output}} contains "高危"`
|
||||
label = "是否需要高风险处置"
|
||||
}
|
||||
condition := builder.add("condition", label, map[string]any{"expression": expr, "join_strategy": "all_merge"}, 0)
|
||||
builder.connect(previous, condition, "", nil)
|
||||
openConditionID = condition
|
||||
report := builder.add("output", draftOutputLabel(wantsReport), map[string]any{
|
||||
"output_key": "result",
|
||||
"source_binding": map[string]any{"from": "previous", "field": "output"},
|
||||
"static_value": "",
|
||||
"join_strategy": "all_merge",
|
||||
}, 130)
|
||||
builder.connect(condition, report, "否", map[string]any{"condition": `{{previous.matched}} == "false"`, "branch": "false"})
|
||||
previous = condition
|
||||
}
|
||||
|
||||
insertedHITL := false
|
||||
if highRisk {
|
||||
if !req.Options.AllowHighRisk || wantsApproval {
|
||||
approval := builder.add("hitl", "人工审批", map[string]any{
|
||||
"prompt": "请确认是否允许继续执行高风险处置:" + prompt,
|
||||
"prompt_binding": map[string]any{"from": "previous", "field": "output"},
|
||||
"reviewer": "human",
|
||||
"join_strategy": "all_merge",
|
||||
"risk_level": "high",
|
||||
}, 0)
|
||||
builder.connect(previous, approval, branchLabel(previous, openConditionID), branchConfig(previous, openConditionID, true))
|
||||
if previous == openConditionID {
|
||||
openConditionID = ""
|
||||
}
|
||||
previous = approval
|
||||
insertedHITL = true
|
||||
}
|
||||
action := builder.add("agent", "执行受控处置", map[string]any{
|
||||
"agent_mode": "eino_single",
|
||||
"input_binding": map[string]any{"from": "previous", "field": "output"},
|
||||
"instruction": "仅在授权范围内生成处置步骤草稿;实际执行前必须由人工确认。用户需求:" + prompt,
|
||||
"output_key": "remediation_plan",
|
||||
"join_strategy": "all_merge",
|
||||
"risk_level": "high",
|
||||
"requires_human_confirmation": "true",
|
||||
}, 0)
|
||||
builder.connect(previous, action, branchLabel(previous, openConditionID), branchConfig(previous, openConditionID, true))
|
||||
if previous == openConditionID {
|
||||
openConditionID = ""
|
||||
}
|
||||
previous = action
|
||||
if insertedHITL {
|
||||
riskWarnings = append(riskWarnings, "检测到高风险动作,已加入人工审批与 requires_human_confirmation 标记。")
|
||||
} else {
|
||||
riskWarnings = append(riskWarnings, "检测到高风险动作,已保留为草稿并添加 requires_human_confirmation 标记。")
|
||||
}
|
||||
} else if wantsApproval {
|
||||
approval := builder.add("hitl", "人工审批", map[string]any{
|
||||
"prompt": "请审核工作流阶段结果:" + prompt,
|
||||
"prompt_binding": map[string]any{"from": "previous", "field": "output"},
|
||||
"reviewer": "human",
|
||||
"join_strategy": "all_merge",
|
||||
}, 0)
|
||||
builder.connect(previous, approval, "", nil)
|
||||
previous = approval
|
||||
insertedHITL = true
|
||||
}
|
||||
|
||||
output := builder.add("output", draftOutputLabel(wantsReport), map[string]any{
|
||||
"output_key": "result",
|
||||
"source_binding": map[string]any{"from": "previous", "field": "output"},
|
||||
"static_value": "",
|
||||
"join_strategy": "all_merge",
|
||||
}, 0)
|
||||
builder.connect(previous, output, branchLabel(previous, openConditionID), branchConfig(previous, openConditionID, true))
|
||||
|
||||
graph := &graphDef{Nodes: builder.nodes, Edges: builder.edges, Config: map[string]any{
|
||||
"schema_version": 1,
|
||||
"generated_by": "natural_language",
|
||||
"source_prompt": prompt,
|
||||
}}
|
||||
if req.Options.IncludeObjective {
|
||||
graph.Config["objective"] = prompt
|
||||
}
|
||||
if req.Options.AllowSchedule && containsAnyFold(prompt, "每天", "每周", "定时", "周期", "持续", "daily", "weekly", "schedule", "monitor") {
|
||||
if containsAnyFold(prompt, "每天", "daily") {
|
||||
graph.Config["trigger_suggestion"] = "daily"
|
||||
} else {
|
||||
graph.Config["trigger_suggestion"] = "scheduled"
|
||||
}
|
||||
assumptions = append(assumptions, "已记录定时触发建议;保存后仍需在触发器或角色绑定处配置。")
|
||||
}
|
||||
|
||||
raw, _ := json.Marshal(graph)
|
||||
validation := make([]string, 0)
|
||||
if err := ValidateGraphJSON(ctx, string(raw)); err != nil {
|
||||
validation = append(validation, err.Error())
|
||||
}
|
||||
return &DraftResult{
|
||||
Graph: graph,
|
||||
Meta: DraftMeta{ID: draftSlug(prompt), Name: draftName(prompt), Description: prompt, Enabled: true},
|
||||
Generator: "deterministic",
|
||||
Audit: DraftAudit{
|
||||
Savable: len(validation) == 0,
|
||||
Validation: validation,
|
||||
MissingFields: missingFields,
|
||||
RiskWarnings: riskWarnings,
|
||||
Assumptions: assumptions,
|
||||
HighRisk: highRisk,
|
||||
NeedsHITL: insertedHITL,
|
||||
},
|
||||
Capabilities: capabilities,
|
||||
Stats: map[string]int{"nodes": len(graph.Nodes), "edges": len(graph.Edges)},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func GenerateDraftFromLLM(ctx context.Context, req DraftRequest, oa config.OpenAIConfig, logger *zap.Logger) (*DraftResult, error) {
|
||||
prompt := strings.TrimSpace(req.Prompt)
|
||||
if prompt == "" {
|
||||
return nil, fmt.Errorf("工作流需求不能为空")
|
||||
}
|
||||
if strings.TrimSpace(oa.APIKey) == "" || strings.TrimSpace(oa.Model) == "" {
|
||||
return nil, fmt.Errorf("AI 通道未配置 api_key 或 model")
|
||||
}
|
||||
if logger == nil {
|
||||
logger = zap.NewNop()
|
||||
}
|
||||
callCtx, cancel := context.WithTimeout(ctx, 90*time.Second)
|
||||
defer cancel()
|
||||
toolJSON, _ := json.Marshal(req.AvailableTools)
|
||||
systemPrompt := `你是 CyberStrikeAI 的工作流编排助手。你必须把用户的一句话需求转换为可保存的工作流草稿 JSON。
|
||||
只返回 JSON 对象,不要 Markdown,不要解释。JSON 必须符合:
|
||||
{
|
||||
"meta": {"id":"kebab-case-id","name":"短名称","description":"用户需求","enabled":true},
|
||||
"graph": {
|
||||
"nodes": [{"id":"start-1","type":"start","label":"显示名","position":{"x":120,"y":150},"config":{}}],
|
||||
"edges": [{"id":"edge-1","source":"start-1","target":"node-2","label":"","config":{}}],
|
||||
"config": {"schema_version":1,"generated_by":"llm","source_prompt":"用户原文"}
|
||||
},
|
||||
"capabilities": [{"label":"能力名","tool_name":"已匹配工具名","tool_candidates":["候选工具"]}],
|
||||
"audit": {"assumptions":[],"missing_fields":[],"risk_warnings":[]}
|
||||
}
|
||||
硬性规则:
|
||||
- 只能输出一个合法 JSON object;不要输出 JSON Schema、注释、解释文字、Markdown 代码块或多余前后缀。
|
||||
- 不要在 JSON 字符串值中使用竖线枚举写法;type 字段一次只能填写一个节点类型字符串。
|
||||
- 至少 1 个 start 和 1 个 output;output/end 不能有出边。
|
||||
- 节点 type 只能从这些字符串中选择:start、tool、agent、condition、hitl、output、end。
|
||||
- 每个 agent、tool、output 节点都必须配置唯一的 output_key;output 节点默认使用 result。
|
||||
- agent 节点必须配置 instruction 或 input_binding;默认 input_binding 为 {"from":"previous","field":"output"}。
|
||||
- output 节点必须配置 source_binding 或 static_value;默认 source_binding 为 {"from":"previous","field":"output"}。
|
||||
- tool 节点必须配置 tool_name、arguments、timeout_seconds;arguments 必须是合法 JSON 字符串。
|
||||
- 所有非 start 且可能有多个上游的节点必须配置 join_strategy:"all_merge"。
|
||||
- condition 最多 2 条出边,必须用 branch true/false,并用 label 是/否。
|
||||
- tool 节点只有在 available_tools 中存在启用工具时才使用,否则用 agent 节点并在 audit.missing_fields 写明缺失工具。
|
||||
- 高风险动作(执行脚本、隔离、封禁、删除、利用、payload、命令执行等)必须加入 hitl 审批,或在高风险节点 config 中标记 requires_human_confirmation:"true"、risk_level:"high"。
|
||||
- 不要生成会真实执行攻击的参数;工具参数使用 {{inputs.target}}、{{inputs.message}} 占位。
|
||||
- 所有节点 config 加 generated_by:"llm" 和 needs_review:"true"。`
|
||||
userPrompt := fmt.Sprintf("用户需求:%s\n\n选项:%+v\n\n可用工具 JSON:%s", prompt, req.Options, string(toolJSON))
|
||||
requestBody := map[string]interface{}{
|
||||
"model": strings.TrimSpace(oa.Model),
|
||||
"messages": []map[string]interface{}{
|
||||
{"role": "system", "content": systemPrompt},
|
||||
{"role": "user", "content": userPrompt},
|
||||
},
|
||||
"temperature": 0,
|
||||
"max_completion_tokens": 4096,
|
||||
"response_format": map[string]interface{}{"type": "json_object"},
|
||||
"thinking": map[string]interface{}{"type": "disabled"},
|
||||
}
|
||||
var apiResponse struct {
|
||||
Choices []struct {
|
||||
Message struct {
|
||||
Content string `json:"content"`
|
||||
ReasoningContent string `json:"reasoning_content"`
|
||||
} `json:"message"`
|
||||
} `json:"choices"`
|
||||
}
|
||||
client := openai.NewClient(&oa, nil, logger)
|
||||
if err := client.ChatCompletion(callCtx, requestBody, &apiResponse); err != nil {
|
||||
return nil, fmt.Errorf("调用大模型失败: %w", err)
|
||||
}
|
||||
if len(apiResponse.Choices) == 0 {
|
||||
return nil, fmt.Errorf("大模型未返回候选结果")
|
||||
}
|
||||
raw := strings.TrimSpace(apiResponse.Choices[0].Message.Content)
|
||||
if raw == "" {
|
||||
raw = strings.TrimSpace(apiResponse.Choices[0].Message.ReasoningContent)
|
||||
}
|
||||
env, err := parseLLMDraftEnvelope(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := normalizeLLMDraft(prompt, req, env)
|
||||
graphRaw, _ := json.Marshal(result.Graph)
|
||||
validation := make([]string, 0)
|
||||
if err := ValidateGraphJSON(ctx, string(graphRaw)); err != nil {
|
||||
validation = append(validation, err.Error())
|
||||
}
|
||||
result.Audit.Validation = validation
|
||||
result.Audit.Savable = len(validation) == 0
|
||||
if !result.Audit.Savable {
|
||||
return nil, fmt.Errorf("大模型生成的工作流未通过校验: %s", strings.Join(validation, ";"))
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
type draftGraphBuilder struct {
|
||||
nodes []graphNode
|
||||
edges []graphEdge
|
||||
x float64
|
||||
y float64
|
||||
nodeSeq int
|
||||
edgeSeq int
|
||||
}
|
||||
|
||||
func (b *draftGraphBuilder) add(nodeType, label string, config map[string]any, yOffset float64) string {
|
||||
b.nodeSeq++
|
||||
id := fmt.Sprintf("%s-%d", nodeType, b.nodeSeq)
|
||||
if config == nil {
|
||||
config = make(map[string]any)
|
||||
}
|
||||
config["generated_by"] = "natural_language"
|
||||
config["needs_review"] = "true"
|
||||
b.nodes = append(b.nodes, graphNode{
|
||||
ID: id,
|
||||
Type: nodeType,
|
||||
Label: label,
|
||||
Position: graphPosition{X: b.x, Y: b.y + yOffset},
|
||||
Config: config,
|
||||
})
|
||||
b.x += 210
|
||||
return id
|
||||
}
|
||||
|
||||
func (b *draftGraphBuilder) connect(source, target, label string, config map[string]any) {
|
||||
b.edgeSeq++
|
||||
if config == nil {
|
||||
config = make(map[string]any)
|
||||
}
|
||||
b.edges = append(b.edges, graphEdge{ID: fmt.Sprintf("edge-ai-%d", b.edgeSeq), Source: source, Target: target, Label: label, Config: config})
|
||||
}
|
||||
|
||||
func parseLLMDraftEnvelope(raw string) (llmDraftEnvelope, error) {
|
||||
var lastErr error
|
||||
for _, candidate := range jsonObjectCandidates(raw) {
|
||||
var env llmDraftEnvelope
|
||||
if err := json.Unmarshal([]byte(candidate), &env); err == nil {
|
||||
if len(env.Graph.Nodes) == 0 {
|
||||
lastErr = fmt.Errorf("大模型 JSON 缺少 graph.nodes")
|
||||
continue
|
||||
}
|
||||
return env, nil
|
||||
} else {
|
||||
lastErr = err
|
||||
}
|
||||
}
|
||||
if lastErr == nil {
|
||||
lastErr = fmt.Errorf("大模型响应为空")
|
||||
}
|
||||
return llmDraftEnvelope{}, fmt.Errorf("解析大模型工作流 JSON 失败: %w", lastErr)
|
||||
}
|
||||
|
||||
func jsonObjectCandidates(raw string) []string {
|
||||
s := strings.TrimSpace(raw)
|
||||
s = strings.TrimPrefix(s, "```json")
|
||||
s = strings.TrimPrefix(s, "```")
|
||||
s = strings.TrimSuffix(s, "```")
|
||||
s = strings.TrimSpace(s)
|
||||
candidates := []string{s}
|
||||
if start := strings.Index(s, "{"); start >= 0 {
|
||||
if end := strings.LastIndex(s, "}"); end > start {
|
||||
candidates = append(candidates, s[start:end+1])
|
||||
}
|
||||
}
|
||||
return candidates
|
||||
}
|
||||
|
||||
func normalizeLLMDraft(prompt string, req DraftRequest, env llmDraftEnvelope) *DraftResult {
|
||||
g := env.Graph
|
||||
if g.Config == nil {
|
||||
g.Config = make(map[string]any)
|
||||
}
|
||||
g.Config["schema_version"] = 1
|
||||
g.Config["generated_by"] = "llm"
|
||||
g.Config["source_prompt"] = prompt
|
||||
if req.Options.IncludeObjective {
|
||||
g.Config["objective"] = prompt
|
||||
}
|
||||
enabledTools := enabledDraftToolNames(req.AvailableTools)
|
||||
usedOutputKeys := make(map[string]bool)
|
||||
nodeTypes := make(map[string]string, len(g.Nodes))
|
||||
for i := range g.Nodes {
|
||||
if strings.TrimSpace(g.Nodes[i].ID) == "" {
|
||||
g.Nodes[i].ID = fmt.Sprintf("%s-%d", firstNonEmpty(g.Nodes[i].Type, "node"), i+1)
|
||||
}
|
||||
if strings.TrimSpace(g.Nodes[i].Type) == "" {
|
||||
g.Nodes[i].Type = "agent"
|
||||
}
|
||||
if strings.TrimSpace(g.Nodes[i].Label) == "" {
|
||||
g.Nodes[i].Label = displayNodeType(g.Nodes[i].Type)
|
||||
}
|
||||
if g.Nodes[i].Position.X == 0 && g.Nodes[i].Position.Y == 0 {
|
||||
g.Nodes[i].Position = graphPosition{X: 120 + float64(i)*210, Y: 150}
|
||||
}
|
||||
if g.Nodes[i].Config == nil {
|
||||
g.Nodes[i].Config = make(map[string]any)
|
||||
}
|
||||
g.Nodes[i].Config["generated_by"] = "llm"
|
||||
g.Nodes[i].Config["needs_review"] = "true"
|
||||
normalizeLLMNodeConfig(prompt, &g.Nodes[i], enabledTools, usedOutputKeys)
|
||||
nodeTypes[g.Nodes[i].ID] = strings.ToLower(strings.TrimSpace(g.Nodes[i].Type))
|
||||
}
|
||||
conditionBranchCounts := make(map[string]int)
|
||||
for i := range g.Edges {
|
||||
if strings.TrimSpace(g.Edges[i].ID) == "" {
|
||||
g.Edges[i].ID = fmt.Sprintf("edge-llm-%d", i+1)
|
||||
}
|
||||
if g.Edges[i].Config == nil {
|
||||
g.Edges[i].Config = make(map[string]any)
|
||||
}
|
||||
normalizeLLMEdgeConfig(&g.Edges[i], nodeTypes, conditionBranchCounts)
|
||||
}
|
||||
audit := env.Audit
|
||||
highRisk := highRiskDraftRE.MatchString(prompt) || graphHasHighRisk(g)
|
||||
audit.HighRisk = highRisk
|
||||
audit.NeedsHITL = graphHasNodeType(g, "hitl")
|
||||
if highRisk && !audit.NeedsHITL && !graphHasConfirmation(g) {
|
||||
audit.RiskWarnings = append(audit.RiskWarnings, "大模型生成包含高风险语义,请补充人工审批或确认标记后再运行。")
|
||||
}
|
||||
if len(audit.RiskWarnings) == 0 && highRisk {
|
||||
audit.RiskWarnings = append(audit.RiskWarnings, "检测到高风险动作,已标记为需要重点审计。")
|
||||
}
|
||||
meta := env.Meta
|
||||
if strings.TrimSpace(meta.Description) == "" {
|
||||
meta.Description = prompt
|
||||
}
|
||||
if strings.TrimSpace(meta.Name) == "" {
|
||||
meta.Name = draftName(prompt)
|
||||
}
|
||||
if strings.TrimSpace(meta.ID) == "" {
|
||||
meta.ID = draftSlug(prompt)
|
||||
}
|
||||
meta.Enabled = true
|
||||
return &DraftResult{
|
||||
Graph: &g,
|
||||
Meta: meta,
|
||||
Generator: "llm",
|
||||
Audit: audit,
|
||||
Capabilities: env.Capabilities,
|
||||
Stats: map[string]int{"nodes": len(g.Nodes), "edges": len(g.Edges)},
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeLLMEdgeConfig(edge *graphEdge, nodeTypes map[string]string, conditionBranchCounts map[string]int) {
|
||||
if nodeTypes[strings.TrimSpace(edge.Source)] != "condition" {
|
||||
return
|
||||
}
|
||||
if conditionBranchHint(*edge) != "" {
|
||||
return
|
||||
}
|
||||
conditionBranchCounts[edge.Source]++
|
||||
branch := "true"
|
||||
label := "是"
|
||||
if conditionBranchCounts[edge.Source] > 1 {
|
||||
branch = "false"
|
||||
label = "否"
|
||||
}
|
||||
edge.Label = label
|
||||
edge.Config["branch"] = branch
|
||||
}
|
||||
|
||||
func normalizeLLMNodeConfig(prompt string, node *graphNode, enabledTools map[string]bool, usedOutputKeys map[string]bool) {
|
||||
nodeType := strings.ToLower(strings.TrimSpace(node.Type))
|
||||
switch nodeType {
|
||||
case "start":
|
||||
if cfgString(node.Config, "input_keys") == "" {
|
||||
node.Config["input_keys"] = "message, conversationId, projectId, target"
|
||||
}
|
||||
case "tool":
|
||||
toolName := cfgString(node.Config, "tool_name")
|
||||
if toolName == "" || !enabledTools[strings.ToLower(toolName)] {
|
||||
node.Type = "agent"
|
||||
node.Config["missing_tool_name"] = toolName
|
||||
normalizeAgentDraftConfig(prompt, node, usedOutputKeys)
|
||||
return
|
||||
}
|
||||
if cfgString(node.Config, "arguments") == "" {
|
||||
node.Config["arguments"] = `{"target":"{{inputs.target}}","message":"{{inputs.message}}"}`
|
||||
}
|
||||
if cfgString(node.Config, "timeout_seconds") == "" {
|
||||
node.Config["timeout_seconds"] = "120"
|
||||
}
|
||||
ensureNodeOutputKey(node, usedOutputKeys, draftOutputKeyBase(node, "tool_result"))
|
||||
ensureJoinStrategy(node)
|
||||
case "agent":
|
||||
normalizeAgentDraftConfig(prompt, node, usedOutputKeys)
|
||||
case "condition":
|
||||
if cfgString(node.Config, "expression") == "" {
|
||||
node.Config["expression"] = `{{previous.output}} != ""`
|
||||
}
|
||||
ensureJoinStrategy(node)
|
||||
case "hitl":
|
||||
if cfgString(node.Config, "prompt") == "" {
|
||||
node.Config["prompt"] = "请审核工作流阶段结果:" + prompt
|
||||
}
|
||||
if cfgString(node.Config, "reviewer") == "" {
|
||||
node.Config["reviewer"] = "human"
|
||||
}
|
||||
ensureJoinStrategy(node)
|
||||
case "output":
|
||||
ensureNodeOutputKey(node, usedOutputKeys, "result")
|
||||
if cfgString(node.Config, "static_value") == "" {
|
||||
if _, ok := parseFieldBinding(node.Config, "source_binding"); !ok {
|
||||
node.Config["source_binding"] = map[string]any{"from": "previous", "field": "output"}
|
||||
}
|
||||
}
|
||||
ensureJoinStrategy(node)
|
||||
case "end":
|
||||
ensureJoinStrategy(node)
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeAgentDraftConfig(prompt string, node *graphNode, usedOutputKeys map[string]bool) {
|
||||
if cfgString(node.Config, "agent_mode") == "" {
|
||||
node.Config["agent_mode"] = "eino_single"
|
||||
}
|
||||
if cfgString(node.Config, "instruction") == "" {
|
||||
node.Config["instruction"] = node.Label + "。根据用户需求执行安全流程步骤,并输出结构化结果:" + prompt
|
||||
}
|
||||
if _, ok := parseFieldBinding(node.Config, "input_binding"); !ok {
|
||||
node.Config["input_binding"] = map[string]any{"from": "previous", "field": "output"}
|
||||
}
|
||||
ensureNodeOutputKey(node, usedOutputKeys, draftOutputKeyBase(node, "agent_result"))
|
||||
ensureJoinStrategy(node)
|
||||
}
|
||||
|
||||
func ensureJoinStrategy(node *graphNode) {
|
||||
if cfgString(node.Config, "join_strategy") == "" {
|
||||
node.Config["join_strategy"] = "all_merge"
|
||||
}
|
||||
}
|
||||
|
||||
func ensureNodeOutputKey(node *graphNode, used map[string]bool, fallback string) {
|
||||
current := sanitizeOutputKey(cfgString(node.Config, "output_key"))
|
||||
if current == "" {
|
||||
current = sanitizeOutputKey(fallback)
|
||||
}
|
||||
if current == "" {
|
||||
current = "result"
|
||||
}
|
||||
base := current
|
||||
for i := 2; used[current]; i++ {
|
||||
current = fmt.Sprintf("%s_%d", base, i)
|
||||
}
|
||||
node.Config["output_key"] = current
|
||||
used[current] = true
|
||||
}
|
||||
|
||||
func draftOutputKeyBase(node *graphNode, fallback string) string {
|
||||
if name := cfgString(node.Config, "tool_name"); name != "" {
|
||||
return name + "_result"
|
||||
}
|
||||
if node.ID != "" {
|
||||
return node.ID + "_result"
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func sanitizeOutputKey(value string) string {
|
||||
value = strings.ToLower(strings.TrimSpace(value))
|
||||
var b strings.Builder
|
||||
lastUnderscore := false
|
||||
for _, r := range value {
|
||||
if (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') {
|
||||
b.WriteRune(r)
|
||||
lastUnderscore = false
|
||||
continue
|
||||
}
|
||||
if b.Len() > 0 && !lastUnderscore {
|
||||
b.WriteByte('_')
|
||||
lastUnderscore = true
|
||||
}
|
||||
}
|
||||
return strings.Trim(b.String(), "_")
|
||||
}
|
||||
|
||||
func enabledDraftToolNames(tools []DraftTool) map[string]bool {
|
||||
names := make(map[string]bool, len(tools)*2)
|
||||
for _, tool := range tools {
|
||||
if !tool.Enabled {
|
||||
continue
|
||||
}
|
||||
if key := strings.ToLower(strings.TrimSpace(tool.Key)); key != "" {
|
||||
names[key] = true
|
||||
}
|
||||
if name := strings.ToLower(strings.TrimSpace(tool.Name)); name != "" {
|
||||
names[name] = true
|
||||
}
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
func graphHasNodeType(g graphDef, nodeType string) bool {
|
||||
for _, node := range g.Nodes {
|
||||
if strings.EqualFold(node.Type, nodeType) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func graphHasConfirmation(g graphDef) bool {
|
||||
for _, node := range g.Nodes {
|
||||
if cfgString(node.Config, "requires_human_confirmation") == "true" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func graphHasHighRisk(g graphDef) bool {
|
||||
for _, node := range g.Nodes {
|
||||
if cfgString(node.Config, "risk_level") == "high" || cfgString(node.Config, "requires_human_confirmation") == "true" {
|
||||
return true
|
||||
}
|
||||
if highRiskDraftRE.MatchString(node.Label) || highRiskDraftRE.MatchString(cfgString(node.Config, "instruction")) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func detectDraftCapabilities(prompt string, tools []DraftTool) []DraftCapability {
|
||||
capabilities := make([]DraftCapability, 0)
|
||||
for _, hint := range draftToolHints {
|
||||
if containsAnyFold(prompt, hint.Keywords...) {
|
||||
capabilities = append(capabilities, DraftCapability{
|
||||
Label: hint.Label,
|
||||
ToolName: matchDraftTool(hint.Tools, tools),
|
||||
ToolCandidates: append([]string(nil), hint.Tools...),
|
||||
})
|
||||
}
|
||||
}
|
||||
if len(capabilities) == 0 {
|
||||
capabilities = append(capabilities, DraftCapability{Label: "节点能力", ToolCandidates: nil})
|
||||
}
|
||||
return capabilities
|
||||
}
|
||||
|
||||
func matchDraftTool(candidates []string, tools []DraftTool) string {
|
||||
if len(candidates) == 0 || len(tools) == 0 {
|
||||
return ""
|
||||
}
|
||||
for _, enabledOnly := range []bool{true, false} {
|
||||
for _, candidate := range candidates {
|
||||
candidate = strings.ToLower(strings.TrimSpace(candidate))
|
||||
for _, tool := range tools {
|
||||
if enabledOnly && !tool.Enabled {
|
||||
continue
|
||||
}
|
||||
key := strings.ToLower(strings.TrimSpace(firstNonEmpty(tool.Key, tool.Name)))
|
||||
if key != "" && strings.Contains(key, candidate) {
|
||||
return firstNonEmpty(tool.Key, tool.Name)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func containsAnyFold(text string, needles ...string) bool {
|
||||
lower := strings.ToLower(text)
|
||||
for _, needle := range needles {
|
||||
if strings.Contains(lower, strings.ToLower(needle)) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func draftOutputLabel(wantsReport bool) string {
|
||||
if wantsReport {
|
||||
return "输出报告"
|
||||
}
|
||||
return "输出"
|
||||
}
|
||||
|
||||
func branchLabel(source, conditionID string) string {
|
||||
if source == conditionID && conditionID != "" {
|
||||
return "是"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func branchConfig(source, conditionID string, yes bool) map[string]any {
|
||||
if source != conditionID || conditionID == "" {
|
||||
return nil
|
||||
}
|
||||
if yes {
|
||||
return map[string]any{"condition": `{{previous.matched}} == "true"`, "branch": "true"}
|
||||
}
|
||||
return map[string]any{"condition": `{{previous.matched}} == "false"`, "branch": "false"}
|
||||
}
|
||||
|
||||
func draftName(prompt string) string {
|
||||
runes := []rune(strings.TrimSpace(prompt))
|
||||
if len(runes) > 22 {
|
||||
return string(runes[:22]) + "..."
|
||||
}
|
||||
return string(runes)
|
||||
}
|
||||
|
||||
func draftSlug(prompt string) string {
|
||||
lower := strings.ToLower(strings.TrimSpace(prompt))
|
||||
var b strings.Builder
|
||||
lastDash := false
|
||||
for _, r := range lower {
|
||||
if r >= 'a' && r <= 'z' || r >= '0' && r <= '9' {
|
||||
b.WriteRune(r)
|
||||
lastDash = false
|
||||
continue
|
||||
}
|
||||
if !lastDash && b.Len() > 0 {
|
||||
b.WriteByte('-')
|
||||
lastDash = true
|
||||
}
|
||||
}
|
||||
slug := strings.Trim(b.String(), "-")
|
||||
if slug != "" {
|
||||
if len(slug) > 48 {
|
||||
return strings.Trim(slug[:48], "-")
|
||||
}
|
||||
return slug
|
||||
}
|
||||
h := fnv.New32a()
|
||||
_, _ = h.Write([]byte(lower))
|
||||
if !utf8.ValidString(lower) || lower == "" {
|
||||
lower = "workflow"
|
||||
}
|
||||
return fmt.Sprintf("ai-workflow-%x", h.Sum32())
|
||||
}
|
||||
@@ -0,0 +1,213 @@
|
||||
package workflow
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"cyberstrike-ai/internal/config"
|
||||
)
|
||||
|
||||
func TestGenerateDraftFromNaturalLanguageHighRiskAddsHITLAndValidGraph(t *testing.T) {
|
||||
result, err := GenerateDraftFromNaturalLanguage(context.Background(), DraftRequest{
|
||||
Prompt: "对目标资产做端口扫描,如果发现高危端口就执行加固脚本,最后输出报告",
|
||||
Options: DraftOptions{
|
||||
IncludeObjective: true,
|
||||
AllowSchedule: false,
|
||||
AllowHighRisk: false,
|
||||
},
|
||||
AvailableTools: []DraftTool{{Key: "nmap", Name: "nmap", Enabled: true}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateDraftFromNaturalLanguage: %v", err)
|
||||
}
|
||||
raw, _ := json.Marshal(result.Graph)
|
||||
if err := ValidateGraphJSON(context.Background(), string(raw)); err != nil {
|
||||
t.Fatalf("generated graph should validate: %v\n%s", err, raw)
|
||||
}
|
||||
if !result.Audit.HighRisk || !result.Audit.NeedsHITL || len(result.Audit.RiskWarnings) == 0 {
|
||||
t.Fatalf("audit did not flag high-risk HITL path: %#v", result.Audit)
|
||||
}
|
||||
var hasTool, hasHITL, hasConfirmation bool
|
||||
for _, node := range result.Graph.Nodes {
|
||||
if node.Type == "tool" && cfgString(node.Config, "tool_name") == "nmap" {
|
||||
hasTool = true
|
||||
}
|
||||
if node.Type == "hitl" {
|
||||
hasHITL = true
|
||||
}
|
||||
if cfgString(node.Config, "requires_human_confirmation") == "true" {
|
||||
hasConfirmation = true
|
||||
}
|
||||
}
|
||||
if !hasTool || !hasHITL || !hasConfirmation {
|
||||
t.Fatalf("expected nmap tool, HITL, and confirmation marker; tool=%v hitl=%v confirmation=%v", hasTool, hasHITL, hasConfirmation)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateDraftAllowHighRiskStillLabelsConditionBranch(t *testing.T) {
|
||||
result, err := GenerateDraftFromNaturalLanguage(context.Background(), DraftRequest{
|
||||
Prompt: "如果漏洞扫描发现高危漏洞,允许生成执行修复脚本的草稿并输出报告",
|
||||
Options: DraftOptions{
|
||||
AllowHighRisk: true,
|
||||
},
|
||||
AvailableTools: []DraftTool{{Key: "nuclei", Name: "nuclei", Enabled: true}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateDraftFromNaturalLanguage: %v", err)
|
||||
}
|
||||
raw, _ := json.Marshal(result.Graph)
|
||||
if err := ValidateGraphJSON(context.Background(), string(raw)); err != nil {
|
||||
t.Fatalf("generated graph should validate: %v\n%s", err, raw)
|
||||
}
|
||||
branches := map[string]bool{}
|
||||
for _, edge := range result.Graph.Edges {
|
||||
if branch := cfgString(edge.Config, "branch"); branch != "" {
|
||||
branches[branch] = true
|
||||
}
|
||||
}
|
||||
if !branches["true"] || !branches["false"] {
|
||||
t.Fatalf("condition branches = %#v, want true and false", branches)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateDraftFromLLMUsesOpenAICompatibleEndpoint(t *testing.T) {
|
||||
called := false
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
called = true
|
||||
if r.URL.Path != "/chat/completions" {
|
||||
t.Fatalf("path = %s, want /chat/completions", r.URL.Path)
|
||||
}
|
||||
if got := r.Header.Get("Authorization"); got != "Bearer test-key" {
|
||||
t.Fatalf("authorization = %q", got)
|
||||
}
|
||||
var payload struct {
|
||||
Temperature float64 `json:"temperature"`
|
||||
ResponseFormat struct {
|
||||
Type string `json:"type"`
|
||||
} `json:"response_format"`
|
||||
Messages []struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
} `json:"messages"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&payload); err != nil {
|
||||
t.Fatalf("decode request: %v", err)
|
||||
}
|
||||
if payload.Temperature != 0 || payload.ResponseFormat.Type != "json_object" {
|
||||
t.Fatalf("unexpected structured output controls: temperature=%v response_format=%#v", payload.Temperature, payload.ResponseFormat)
|
||||
}
|
||||
if len(payload.Messages) == 0 || strings.Contains(payload.Messages[0].Content, "start|tool|agent") {
|
||||
t.Fatalf("system prompt still contains pipe enum: %q", payload.Messages[0].Content)
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{\"meta\":{\"id\":\"llm-port-scan\",\"name\":\"端口扫描\",\"description\":\"端口扫描\",\"enabled\":true},\"graph\":{\"nodes\":[{\"id\":\"start-1\",\"type\":\"start\",\"label\":\"开始\",\"position\":{\"x\":120,\"y\":150},\"config\":{\"input_keys\":\"message, target\"}},{\"id\":\"tool-2\",\"type\":\"tool\",\"label\":\"端口扫描\",\"position\":{\"x\":330,\"y\":150},\"config\":{\"tool_name\":\"nmap\",\"arguments\":\"{\\\"target\\\":\\\"{{inputs.target}}\\\"}\",\"timeout_seconds\":\"120\",\"join_strategy\":\"all_merge\"}},{\"id\":\"output-3\",\"type\":\"output\",\"label\":\"输出报告\",\"position\":{\"x\":540,\"y\":150},\"config\":{\"source_binding\":{\"from\":\"previous\",\"field\":\"output\"},\"join_strategy\":\"all_merge\"}}],\"edges\":[{\"id\":\"edge-1\",\"source\":\"start-1\",\"target\":\"tool-2\"},{\"id\":\"edge-2\",\"source\":\"tool-2\",\"target\":\"output-3\"}],\"config\":{\"schema_version\":1}},\"capabilities\":[{\"label\":\"端口扫描\",\"tool_name\":\"nmap\",\"tool_candidates\":[\"nmap\"]}],\"audit\":{\"assumptions\":[]}}"}}]}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
result, err := GenerateDraftFromLLM(context.Background(), DraftRequest{
|
||||
Prompt: "对目标做端口扫描并输出报告",
|
||||
AvailableTools: []DraftTool{{Key: "nmap", Name: "nmap", Enabled: true}},
|
||||
}, config.OpenAIConfig{APIKey: "test-key", BaseURL: srv.URL, Model: "test-model"}, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateDraftFromLLM: %v", err)
|
||||
}
|
||||
if !called {
|
||||
t.Fatal("expected LLM endpoint to be called")
|
||||
}
|
||||
if result.Generator != "llm" || !result.Audit.Savable || result.Meta.ID != "llm-port-scan" {
|
||||
t.Fatalf("unexpected result: %#v", result)
|
||||
}
|
||||
for _, node := range result.Graph.Nodes {
|
||||
if node.Type == "output" && cfgString(node.Config, "output_key") != "result" {
|
||||
t.Fatalf("output_key = %q, want result", cfgString(node.Config, "output_key"))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateDraftFromLLMReturnsErrorOnMalformedJSON(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{\"meta\":{\"id\":\"bad\"},|\"graph\":{\"nodes\":[]}}"}}]}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
_, err := GenerateDraftFromLLM(context.Background(), DraftRequest{
|
||||
Prompt: "随便生成一个工作流,要求所有节点都用到输出变量",
|
||||
Options: DraftOptions{
|
||||
IncludeObjective: true,
|
||||
},
|
||||
}, config.OpenAIConfig{APIKey: "test-key", BaseURL: srv.URL, Model: "test-model"}, nil)
|
||||
if err == nil {
|
||||
t.Fatal("expected malformed JSON error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "解析大模型工作流 JSON 失败") {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeLLMDraftRepairsMissingRequiredConfig(t *testing.T) {
|
||||
result := normalizeLLMDraft("随便生成一个工作流", DraftRequest{}, llmDraftEnvelope{
|
||||
Graph: graphDef{
|
||||
Nodes: []graphNode{
|
||||
{ID: "start-1", Type: "start", Label: "开始", Config: map[string]any{}},
|
||||
{ID: "agent-1", Type: "agent", Label: "分析", Config: map[string]any{}},
|
||||
{ID: "out-1", Type: "output", Label: "输出结果", Config: map[string]any{}},
|
||||
},
|
||||
Edges: []graphEdge{
|
||||
{ID: "e1", Source: "start-1", Target: "agent-1"},
|
||||
{ID: "e2", Source: "agent-1", Target: "out-1"},
|
||||
},
|
||||
},
|
||||
})
|
||||
raw, _ := json.Marshal(result.Graph)
|
||||
if err := ValidateGraphJSON(context.Background(), string(raw)); err != nil {
|
||||
t.Fatalf("normalized graph should validate: %v\n%s", err, raw)
|
||||
}
|
||||
var agentKey, outputKey string
|
||||
for _, node := range result.Graph.Nodes {
|
||||
switch node.Type {
|
||||
case "agent":
|
||||
agentKey = cfgString(node.Config, "output_key")
|
||||
case "output":
|
||||
outputKey = cfgString(node.Config, "output_key")
|
||||
}
|
||||
}
|
||||
if agentKey == "" || outputKey != "result" {
|
||||
t.Fatalf("agentKey=%q outputKey=%q", agentKey, outputKey)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeLLMDraftRepairsConditionBranches(t *testing.T) {
|
||||
result := normalizeLLMDraft("如果发现异常则输出详情,否则输出正常", DraftRequest{}, llmDraftEnvelope{
|
||||
Graph: graphDef{
|
||||
Nodes: []graphNode{
|
||||
{ID: "start-1", Type: "start", Label: "开始", Config: map[string]any{}},
|
||||
{ID: "cond-1", Type: "condition", Label: "判断", Config: map[string]any{"expression": `{{inputs.message}} != ""`}},
|
||||
{ID: "out-yes", Type: "output", Label: "异常", Config: map[string]any{}},
|
||||
{ID: "out-no", Type: "output", Label: "正常", Config: map[string]any{}},
|
||||
},
|
||||
Edges: []graphEdge{
|
||||
{ID: "e1", Source: "start-1", Target: "cond-1"},
|
||||
{ID: "e2", Source: "cond-1", Target: "out-yes"},
|
||||
{ID: "e3", Source: "cond-1", Target: "out-no"},
|
||||
},
|
||||
},
|
||||
})
|
||||
raw, _ := json.Marshal(result.Graph)
|
||||
if err := ValidateGraphJSON(context.Background(), string(raw)); err != nil {
|
||||
t.Fatalf("normalized graph should validate: %v\n%s", err, raw)
|
||||
}
|
||||
branches := map[string]bool{}
|
||||
for _, edge := range result.Graph.Edges {
|
||||
if edge.Source == "cond-1" {
|
||||
branches[cfgString(edge.Config, "branch")] = true
|
||||
}
|
||||
}
|
||||
if !branches["true"] || !branches["false"] {
|
||||
t.Fatalf("branches = %#v, want true and false", branches)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user