diff --git a/internal/multiagent/eino_middleware.go b/internal/multiagent/eino_middleware.go new file mode 100644 index 00000000..1f2c057a --- /dev/null +++ b/internal/multiagent/eino_middleware.go @@ -0,0 +1,435 @@ +package multiagent + +import ( + "context" + "fmt" + "os" + "path/filepath" + "strings" + + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/mcp/builtin" + + localbk "github.com/cloudwego/eino-ext/adk/backend/local" + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/adk/middlewares/dynamictool/toolsearch" + "github.com/cloudwego/eino/adk/middlewares/patchtoolcalls" + "github.com/cloudwego/eino/adk/middlewares/plantask" + "github.com/cloudwego/eino/adk/middlewares/reduction" + "github.com/cloudwego/eino/components/tool" + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" +) + +// einoMWPlacement controls which optional middleware runs on orchestrator vs sub-agents. +type einoMWPlacement int + +const ( + einoMWMain einoMWPlacement = iota // Deep / Supervisor main chat agent + einoMWSub // Specialist ChatModelAgent +) + +func sanitizeEinoPathSegment(s string) string { + s = strings.TrimSpace(s) + if s == "" { + return "default" + } + s = strings.ReplaceAll(s, string(filepath.Separator), "-") + s = strings.ReplaceAll(s, "/", "-") + s = strings.ReplaceAll(s, "\\", "-") + s = strings.ReplaceAll(s, "..", "__") + if len(s) > 180 { + s = s[:180] + } + return s +} + +func splitToolsForToolSearch(all []tool.BaseTool, alwaysVisible int) (static []tool.BaseTool, dynamic []tool.BaseTool, ok bool) { + if alwaysVisible <= 0 || len(all) <= alwaysVisible+1 { + return all, nil, false + } + return append([]tool.BaseTool(nil), all[:alwaysVisible]...), append([]tool.BaseTool(nil), all[alwaysVisible:]...), true +} + +func splitToolsForToolSearchByNames(all []tool.BaseTool, names []string, fallbackAlwaysVisible int) (static []tool.BaseTool, dynamic []tool.BaseTool, ok bool) { + nameSet := expandAlwaysVisibleNameSet(names) + if len(nameSet) == 0 { + return splitToolsForToolSearch(all, fallbackAlwaysVisible) + } + static = make([]tool.BaseTool, 0, len(all)) + dynamic = make([]tool.BaseTool, 0, len(all)) + for _, t := range all { + if t == nil { + continue + } + info, err := t.Info(context.Background()) + name := "" + if err == nil && info != nil { + name = info.Name + } + if toolMatchesAlwaysVisible(name, nameSet) { + static = append(static, t) + continue + } + dynamic = append(dynamic, t) + } + if len(static) == 0 || len(dynamic) == 0 { + // fallback: preserve previous behavior when whitelist misses all or includes all. + return splitToolsForToolSearch(all, fallbackAlwaysVisible) + } + return static, dynamic, true +} + +func mergeAlwaysVisibleToolNames(configured []string) []string { + merged := make([]string, 0, len(configured)+32) + seen := make(map[string]struct{}, len(configured)+32) + add := func(name string) { + n := strings.TrimSpace(strings.ToLower(name)) + if n == "" { + return + } + if _, ok := seen[n]; ok { + return + } + seen[n] = struct{}{} + merged = append(merged, n) + } + for _, n := range configured { + add(n) + } + // Always include hardcoded backend builtin MCP tools from constants. + for _, n := range builtin.GetAllBuiltinTools() { + add(n) + } + return merged +} + +func reductionCacheRootDir(configuredBase, projectID, conversationID string) string { + base := strings.TrimSpace(configuredBase) + if base == "" { + base = filepath.Join("tmp", "reduction") + } + if pid := strings.TrimSpace(projectID); pid != "" { + return filepath.Join(base, "projects", sanitizeEinoPathSegment(pid)) + } + conv := strings.TrimSpace(conversationID) + if conv == "" { + conv = "default" + } + return filepath.Join(base, "conversations", sanitizeEinoPathSegment(conv)) +} + +func buildReductionMiddleware(ctx context.Context, mw config.MultiAgentEinoMiddlewareConfig, projectID, convID string, loc *localbk.Local, logger *zap.Logger) (adk.ChatModelAgentMiddleware, error) { + if loc == nil { + return nil, fmt.Errorf("reduction: local backend nil") + } + root := reductionCacheRootDir(mw.ReductionRootDir, projectID, convID) + if err := os.MkdirAll(root, 0o755); err != nil { + return nil, fmt.Errorf("reduction root: %w", err) + } + excl := append([]string(nil), mw.ReductionClearExclude...) + defaultExcl := []string{ + "task", "transfer_to_agent", "exit", "write_todos", "skill", "tool_search", + "TaskCreate", "TaskGet", "TaskUpdate", "TaskList", + } + excl = append(excl, defaultExcl...) + redMW, err := reduction.New(ctx, &reduction.Config{ + Backend: loc, + RootDir: root, + ReadFileToolName: "read_file", + ClearExcludeTools: excl, + MaxLengthForTrunc: mw.ReductionMaxLengthForTruncEffective(), + MaxTokensForClear: int64(mw.ReductionMaxTokensForClearEffective()), + }) + if err != nil { + return nil, err + } + if logger != nil { + logger.Info("eino middleware: reduction enabled", zap.String("root", root)) + } + return redMW, nil +} + +func buildAgenticReductionMiddleware( + ctx context.Context, + mw config.MultiAgentEinoMiddlewareConfig, + projectID, convID string, + loc *localbk.Local, + logger *zap.Logger, +) (adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], error) { + if loc == nil { + return nil, fmt.Errorf("agentic reduction: local backend nil") + } + root := reductionCacheRootDir(mw.ReductionRootDir, projectID, convID) + if err := os.MkdirAll(root, 0o755); err != nil { + return nil, fmt.Errorf("agentic reduction root: %w", err) + } + excl := append([]string(nil), mw.ReductionClearExclude...) + defaultExcl := []string{ + "task", "transfer_to_agent", "exit", "write_todos", "skill", "tool_search", + "TaskCreate", "TaskGet", "TaskUpdate", "TaskList", + } + excl = append(excl, defaultExcl...) + redMW, err := reduction.NewTyped[*schema.AgenticMessage](ctx, &reduction.TypedConfig[*schema.AgenticMessage]{ + Backend: loc, + RootDir: root, + ReadFileToolName: "read_file", + ClearExcludeTools: excl, + MaxLengthForTrunc: mw.ReductionMaxLengthForTruncEffective(), + MaxTokensForClear: int64(mw.ReductionMaxTokensForClearEffective()), + }) + if err != nil { + return nil, err + } + if logger != nil { + logger.Info("eino middleware: agentic reduction enabled", zap.String("root", root)) + } + return redMW, nil +} + +// prependEinoMiddlewares returns handlers to prepend (outermost first) and optionally replaces tools when tool_search is used. +// toolSearchActive is true when the toolsearch middleware was mounted (dynamic tools split off); callers should pass this to +// injectToolNamesOnlyInstruction — tool_search is not part of the pre-middleware tools list, so name-scanning alone cannot detect it. +func prependEinoMiddlewares( + ctx context.Context, + mw *config.MultiAgentEinoMiddlewareConfig, + place einoMWPlacement, + tools []tool.BaseTool, + einoLoc *localbk.Local, + skillsRoot string, + conversationID string, + projectID string, + logger *zap.Logger, +) (outTools []tool.BaseTool, extraHandlers []adk.ChatModelAgentMiddleware, toolSearchActive bool, err error) { + if mw == nil { + return tools, nil, false, nil + } + outTools = tools + + if mw.PatchToolCallsEffective() { + patchMW, perr := patchtoolcalls.New(ctx, &patchtoolcalls.Config{}) + if perr != nil { + return nil, nil, false, fmt.Errorf("patchtoolcalls: %w", perr) + } + extraHandlers = append(extraHandlers, patchMW) + } + + if mw.ReductionEnable && einoLoc != nil { + if place == einoMWSub && !mw.ReductionSubAgents { + // skip + } else { + redMW, rerr := buildReductionMiddleware(ctx, *mw, projectID, conversationID, einoLoc, logger) + if rerr != nil { + return nil, nil, false, rerr + } + extraHandlers = append(extraHandlers, redMW) + } + } + + minTools := mw.ToolSearchMinTools + if minTools <= 0 { + minTools = 20 + } + alwaysVis := mw.ToolSearchAlwaysVisible + if alwaysVis <= 0 { + alwaysVis = 12 + } + if mw.ToolSearchEnable && len(tools) >= minTools { + static, dynamic, split := splitToolsForToolSearchByNames(tools, mergeAlwaysVisibleToolNames(mw.ToolSearchAlwaysVisibleTools), alwaysVis) + if split && len(dynamic) > 0 { + ts, terr := toolsearch.New(ctx, &toolsearch.Config{DynamicTools: dynamic}) + if terr != nil { + return nil, nil, false, fmt.Errorf("toolsearch: %w", terr) + } + extraHandlers = append(extraHandlers, ts) + outTools = static + toolSearchActive = true + if logger != nil { + logger.Info("eino middleware: tool_search enabled", + zap.Int("static_tools", len(static)), + zap.Int("dynamic_tools", len(dynamic))) + } + } + } + + if place == einoMWMain && mw.PlantaskEnable { + if einoLoc == nil || strings.TrimSpace(skillsRoot) == "" { + if logger != nil { + logger.Warn("eino middleware: plantask_enable ignored (need eino_skills + skills_dir)") + } + } else { + rel := strings.TrimSpace(mw.PlantaskRelDir) + if rel == "" { + rel = ".eino/plantask" + } + baseDir := filepath.Join(skillsRoot, rel, sanitizeEinoPathSegment(conversationID)) + if mk := os.MkdirAll(baseDir, 0o755); mk != nil { + return nil, nil, toolSearchActive, fmt.Errorf("plantask mkdir: %w", mk) + } + ptBE := newLocalPlantaskBackend(einoLoc) + pt, perr := plantask.New(ctx, &plantask.Config{Backend: ptBE, BaseDir: baseDir}) + if perr != nil { + return nil, nil, toolSearchActive, fmt.Errorf("plantask: %w", perr) + } + extraHandlers = append(extraHandlers, pt) + if logger != nil { + logger.Info("eino middleware: plantask enabled", zap.String("baseDir", baseDir)) + } + } + } + + return outTools, extraHandlers, toolSearchActive, nil +} + +func prependEinoAgenticMiddlewares( + ctx context.Context, + mw *config.MultiAgentEinoMiddlewareConfig, + place einoMWPlacement, + tools []tool.BaseTool, + einoLoc *localbk.Local, + skillsRoot string, + conversationID string, + projectID string, + logger *zap.Logger, +) (outTools []tool.BaseTool, extraHandlers []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], toolSearchActive bool, err error) { + if mw == nil { + return tools, nil, false, nil + } + outTools = tools + + if mw.PatchToolCallsEffective() { + patchMW, perr := patchtoolcalls.NewTyped[*schema.AgenticMessage](ctx, &patchtoolcalls.Config{}) + if perr != nil { + return nil, nil, false, fmt.Errorf("agentic patchtoolcalls: %w", perr) + } + extraHandlers = append(extraHandlers, patchMW) + } + + if mw.ReductionEnable && einoLoc != nil { + if place == einoMWSub && !mw.ReductionSubAgents { + // skip + } else { + redMW, rerr := buildAgenticReductionMiddleware(ctx, *mw, projectID, conversationID, einoLoc, logger) + if rerr != nil { + return nil, nil, false, rerr + } + extraHandlers = append(extraHandlers, redMW) + } + } + + minTools := mw.ToolSearchMinTools + if minTools <= 0 { + minTools = 20 + } + alwaysVis := mw.ToolSearchAlwaysVisible + if alwaysVis <= 0 { + alwaysVis = 12 + } + if mw.ToolSearchEnable && len(tools) >= minTools { + static, dynamic, split := splitToolsForToolSearchByNames(tools, mergeAlwaysVisibleToolNames(mw.ToolSearchAlwaysVisibleTools), alwaysVis) + if split && len(dynamic) > 0 { + ts, terr := toolsearch.NewTyped[*schema.AgenticMessage](ctx, &toolsearch.Config{DynamicTools: dynamic}) + if terr != nil { + return nil, nil, false, fmt.Errorf("agentic toolsearch: %w", terr) + } + extraHandlers = append(extraHandlers, ts) + outTools = static + toolSearchActive = true + if logger != nil { + logger.Info("eino middleware: agentic tool_search enabled", + zap.Int("static_tools", len(static)), + zap.Int("dynamic_tools", len(dynamic))) + } + } + } + + if place == einoMWMain && mw.PlantaskEnable { + if einoLoc == nil || strings.TrimSpace(skillsRoot) == "" { + if logger != nil { + logger.Warn("eino middleware: agentic plantask_enable ignored (need eino_skills + skills_dir)") + } + } else { + rel := strings.TrimSpace(mw.PlantaskRelDir) + if rel == "" { + rel = ".eino/plantask" + } + baseDir := filepath.Join(skillsRoot, rel, sanitizeEinoPathSegment(conversationID)) + if mk := os.MkdirAll(baseDir, 0o755); mk != nil { + return nil, nil, toolSearchActive, fmt.Errorf("agentic plantask mkdir: %w", mk) + } + ptBE := newLocalPlantaskBackend(einoLoc) + pt, perr := plantask.NewTyped[*schema.AgenticMessage](ctx, &plantask.Config{Backend: ptBE, BaseDir: baseDir}) + if perr != nil { + return nil, nil, toolSearchActive, fmt.Errorf("agentic plantask: %w", perr) + } + extraHandlers = append(extraHandlers, pt) + if logger != nil { + logger.Info("eino middleware: agentic plantask enabled", zap.String("baseDir", baseDir)) + } + } + } + + return outTools, extraHandlers, toolSearchActive, nil +} + +func deepExtrasFromConfig(ma *config.MultiAgentConfig) (outputKey string, taskDesc func(context.Context, []adk.Agent) (string, error)) { + if ma == nil { + return "", nil + } + mw := ma.EinoMiddleware + if k := strings.TrimSpace(mw.DeepOutputKey); k != "" { + outputKey = k + } + prefix := strings.TrimSpace(mw.TaskToolDescriptionPrefix) + if prefix != "" { + taskDesc = func(ctx context.Context, agents []adk.Agent) (string, error) { + _ = ctx + var names []string + for _, a := range agents { + if a == nil { + continue + } + n := strings.TrimSpace(a.Name(ctx)) + if n != "" { + names = append(names, n) + } + } + if len(names) == 0 { + return prefix, nil + } + return prefix + "\n可用子代理(按名称 transfer / task 调用):" + strings.Join(names, "、"), nil + } + } + return outputKey, taskDesc +} + +func deepAgenticExtrasFromConfig(ma *config.MultiAgentConfig) (outputKey string, taskDesc func(context.Context, []adk.TypedAgent[*schema.AgenticMessage]) (string, error)) { + if ma == nil { + return "", nil + } + mw := ma.EinoMiddleware + if k := strings.TrimSpace(mw.DeepOutputKey); k != "" { + outputKey = k + } + prefix := strings.TrimSpace(mw.TaskToolDescriptionPrefix) + if prefix != "" { + taskDesc = func(ctx context.Context, agents []adk.TypedAgent[*schema.AgenticMessage]) (string, error) { + _ = ctx + var names []string + for _, a := range agents { + if a == nil { + continue + } + n := strings.TrimSpace(a.Name(ctx)) + if n != "" { + names = append(names, n) + } + } + if len(names) == 0 { + return prefix, nil + } + return prefix + "\n可用子代理(按名称 transfer / task 调用):" + strings.Join(names, "、"), nil + } + } + return outputKey, taskDesc +} diff --git a/internal/multiagent/eino_middleware_test.go b/internal/multiagent/eino_middleware_test.go new file mode 100644 index 00000000..45842d86 --- /dev/null +++ b/internal/multiagent/eino_middleware_test.go @@ -0,0 +1,220 @@ +package multiagent + +import ( + "context" + "fmt" + "path/filepath" + "strings" + "testing" + + "cyberstrike-ai/internal/config" + + localbk "github.com/cloudwego/eino-ext/adk/backend/local" + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/components/tool" + "github.com/cloudwego/eino/schema" +) + +func TestReductionCacheRootDir(t *testing.T) { + got := reductionCacheRootDir("", "proj-1", "conv-1") + want := filepath.Join("tmp", "reduction", "projects", "proj-1") + if got != want { + t.Fatalf("project scope: got %q want %q", got, want) + } + got = reductionCacheRootDir("", "", "conv-abc") + want = filepath.Join("tmp", "reduction", "conversations", "conv-abc") + if got != want { + t.Fatalf("conversation scope: got %q want %q", got, want) + } + custom := reductionCacheRootDir("/data/cache", "p1", "c1") + if !strings.HasSuffix(custom, filepath.Join("projects", "p1")) { + t.Fatalf("custom base should still scope by project, got %q", custom) + } +} + +func TestBuildAgenticReductionMiddlewareClearsOldAgenticToolResult(t *testing.T) { + ctx := context.Background() + loc, err := localbk.NewBackend(ctx, &localbk.Config{}) + if err != nil { + t.Fatalf("NewBackend: %v", err) + } + root := t.TempDir() + mw, err := buildAgenticReductionMiddleware(ctx, config.MultiAgentEinoMiddlewareConfig{ + ReductionRootDir: root, + ReductionMaxTokensForClear: 1, + }, "", "conv-1", loc, nil) + if err != nil { + t.Fatalf("buildAgenticReductionMiddleware: %v", err) + } + oldText := strings.Repeat("old-tool-output-", 20) + newText := strings.Repeat("new-tool-output-", 20) + state := &adk.TypedChatModelAgentState[*schema.AgenticMessage]{ + Messages: []*schema.AgenticMessage{ + agenticAssistantToolCall("old-call", "execute", `{"command":"old"}`), + agenticToolResult("old-call", "execute", oldText), + agenticAssistantToolCall("new-call", "execute", `{"command":"new"}`), + agenticToolResult("new-call", "execute", newText), + }, + } + _, out, err := mw.BeforeModelRewriteState(ctx, state, nil) + if err != nil { + t.Fatalf("BeforeModelRewriteState: %v", err) + } + oldGot := out.Messages[1].ContentBlocks[0].FunctionToolResult.Content[0].Text.Text + newGot := out.Messages[3].ContentBlocks[0].FunctionToolResult.Content[0].Text.Text + if oldGot == oldText { + t.Fatal("agentic reduction did not clear old oversized tool result") + } + if !strings.Contains(oldGot, "read_file") { + t.Fatalf("cleared content should mention read_file, got %q", oldGot) + } + if newGot != newText { + t.Fatalf("latest tool result should be retained, got %q", newGot) + } +} + +func agenticAssistantToolCall(callID, name, arguments string) *schema.AgenticMessage { + return &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.FunctionToolCall{ + CallID: callID, + Name: name, + Arguments: arguments, + })}, + } +} + +func agenticToolResult(callID, name, text string) *schema.AgenticMessage { + return &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeUser, + ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.FunctionToolResult{ + CallID: callID, + Name: name, + Content: []*schema.FunctionToolResultContentBlock{{ + Type: schema.FunctionToolResultContentBlockTypeText, + Text: &schema.UserInputText{Text: text}, + }}, + })}, + } +} + +func TestBuildAgenticReductionMiddlewareHandlesSingleAgenticToolResult(t *testing.T) { + ctx := context.Background() + loc, err := localbk.NewBackend(ctx, &localbk.Config{}) + if err != nil { + t.Fatalf("NewBackend: %v", err) + } + mw, err := buildAgenticReductionMiddleware(ctx, config.MultiAgentEinoMiddlewareConfig{ + ReductionRootDir: t.TempDir(), + ReductionMaxTokensForClear: 1, + }, "", "conv-1", loc, nil) + if err != nil { + t.Fatalf("buildAgenticReductionMiddleware: %v", err) + } + state := &adk.TypedChatModelAgentState[*schema.AgenticMessage]{ + Messages: []*schema.AgenticMessage{ + { + Role: schema.AgenticRoleTypeUser, + ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.FunctionToolResult{ + CallID: "call-1", + Name: "execute", + Content: []*schema.FunctionToolResultContentBlock{{ + Type: schema.FunctionToolResultContentBlockTypeText, + Text: &schema.UserInputText{Text: strings.Repeat("tool-output-", 20)}, + }}, + })}, + }, + }, + } + _, out, err := mw.BeforeModelRewriteState(ctx, state, nil) + if err != nil { + t.Fatalf("BeforeModelRewriteState: %v", err) + } + got := out.Messages[0].ContentBlocks[0].FunctionToolResult.Content[0].Text.Text + if got != strings.Repeat("tool-output-", 20) { + t.Fatalf("single retained tool result should not be cleared, got %q", got) + } +} + +func TestPrependEinoAgenticMiddlewaresRespectsReductionPlacement(t *testing.T) { + ctx := context.Background() + loc, err := localbk.NewBackend(ctx, &localbk.Config{}) + if err != nil { + t.Fatalf("NewBackend: %v", err) + } + patchToolCalls := false + mw := &config.MultiAgentEinoMiddlewareConfig{ + ReductionEnable: true, + ReductionRootDir: t.TempDir(), + ReductionMaxTokensForClear: 100, + PatchToolCalls: &patchToolCalls, + } + _, mainHandlers, _, err := prependEinoAgenticMiddlewares(ctx, mw, einoMWMain, nil, loc, "", "conv-1", "", nil) + if err != nil { + t.Fatalf("prepend main: %v", err) + } + if len(mainHandlers) != 1 { + t.Fatalf("main handlers = %d, want reduction", len(mainHandlers)) + } + _, subHandlers, _, err := prependEinoAgenticMiddlewares(ctx, mw, einoMWSub, nil, loc, "", "conv-1", "", nil) + if err != nil { + t.Fatalf("prepend sub: %v", err) + } + if len(subHandlers) != 0 { + t.Fatalf("sub handlers = %d, want skipped when reduction_sub_agents=false", len(subHandlers)) + } + mw.ReductionSubAgents = true + _, subHandlers, _, err = prependEinoAgenticMiddlewares(ctx, mw, einoMWSub, nil, loc, "", "conv-1", "", nil) + if err != nil { + t.Fatalf("prepend sub enabled: %v", err) + } + if len(subHandlers) != 1 { + t.Fatalf("sub handlers = %d, want reduction when reduction_sub_agents=true", len(subHandlers)) + } +} + +func TestPrependEinoAgenticMiddlewaresMountsToolSearchAndPatchToolCalls(t *testing.T) { + ctx := context.Background() + mw := &config.MultiAgentEinoMiddlewareConfig{ + ToolSearchEnable: true, + ToolSearchMinTools: 20, + ToolSearchAlwaysVisible: 5, + } + outTools, handlers, toolSearchActive, err := prependEinoAgenticMiddlewares(ctx, mw, einoMWMain, stubTools(25), nil, "", "conv-test", "", nil) + if err != nil { + t.Fatalf("prependEinoAgenticMiddlewares: %v", err) + } + if !toolSearchActive { + t.Fatal("agentic tool_search should be active") + } + if len(outTools) != 5 { + t.Fatalf("mounted tools = %d, want static visible tools only", len(outTools)) + } + if len(handlers) != 2 { + t.Fatalf("handlers = %d, want patchtoolcalls + toolsearch", len(handlers)) + } +} + +type stubTool struct{ name string } + +func (s stubTool) Info(_ context.Context) (*schema.ToolInfo, error) { + return &schema.ToolInfo{Name: s.name}, nil +} + +func TestSplitToolsForToolSearch(t *testing.T) { + mk := func(n int) []tool.BaseTool { + out := make([]tool.BaseTool, n) + for i := 0; i < n; i++ { + out[i] = stubTool{name: fmt.Sprintf("t%d", i)} + } + return out + } + static, dynamic, ok := splitToolsForToolSearch(mk(4), 3) + if ok || len(static) != 4 || dynamic != nil { + t.Fatalf("expected no split when len<=alwaysVisible+1, got ok=%v static=%d dynamic=%v", ok, len(static), dynamic) + } + static, dynamic, ok = splitToolsForToolSearch(mk(20), 5) + if !ok || len(static) != 5 || len(dynamic) != 15 { + t.Fatalf("expected split 5+15, got ok=%v static=%d dynamic=%d", ok, len(static), len(dynamic)) + } +} diff --git a/internal/multiagent/eino_model_facing_persistence_test.go b/internal/multiagent/eino_model_facing_persistence_test.go new file mode 100644 index 00000000..eddf1bed --- /dev/null +++ b/internal/multiagent/eino_model_facing_persistence_test.go @@ -0,0 +1,34 @@ +package multiagent + +import ( + "strings" + "testing" + + "cyberstrike-ai/internal/agent" + + "github.com/cloudwego/eino/schema" +) + +func TestBuildEinoRunResultNeverPersistsRawAccumulationWithoutModelFacingTrace(t *testing.T) { + raw := []schema.Message{*schema.ToolMessage(strings.Repeat("raw-tool-output", 1000), "call-1")} + rawMsgs := make([]*schema.Message, len(raw)) + for i := range raw { + rawMsgs[i] = &raw[i] + } + result := buildEinoRunResultFromAccumulated("deep", rawMsgs, nil, "", "", "empty", nil, true) + if result.LastAgentTraceInput != "" { + t.Fatalf("pre-model raw accumulation must not be persisted: %d bytes", len(result.LastAgentTraceInput)) + } + + modelFacing := []*schema.Message{schema.UserMessage("bounded-model-view")} + result = buildEinoRunResultFromAccumulated("deep", rawMsgs, modelFacing, "ok", "", "empty", nil, false) + if !strings.Contains(result.LastAgentTraceInput, "bounded-model-view") { + t.Fatalf("model-facing trace missing: %s", result.LastAgentTraceInput) + } + if strings.Contains(result.LastAgentTraceInput, "raw-tool-output") { + t.Fatal("raw accumulation leaked into persisted model-facing trace") + } + if !agent.IsModelFacingTraceJSON(result.LastAgentTraceInput) { + t.Fatal("persisted model-facing trace is missing its version marker") + } +} diff --git a/internal/multiagent/eino_model_facing_trace.go b/internal/multiagent/eino_model_facing_trace.go new file mode 100644 index 00000000..33d8d011 --- /dev/null +++ b/internal/multiagent/eino_model_facing_trace.go @@ -0,0 +1,124 @@ +package multiagent + +import ( + "context" + "encoding/json" + "sync" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +// modelFacingTraceHolder 保存「即将送入 ChatModel」的消息快照(已走 summarization / reduction / orphan 修剪等), +// 用于 last_react_input 落库,使续跑与「上下文压缩后」的模型视角一致,而非仅依赖事件流 append 的 runAccumulatedMsgs。 +type modelFacingTraceHolder struct { + mu sync.Mutex + // msgs 为深拷贝后的切片,避免框架后续原地修改污染快照 + msgs []adk.Message +} + +func newModelFacingTraceHolder() *modelFacingTraceHolder { + return &modelFacingTraceHolder{} +} + +// Snapshot 返回当前快照的再一次深拷贝(供序列化落库,避免与 holder 互斥长期持锁)。 +func (h *modelFacingTraceHolder) Snapshot() []adk.Message { + if h == nil { + return nil + } + h.mu.Lock() + defer h.mu.Unlock() + return cloneADKMessagesForTrace(h.msgs) +} + +func (h *modelFacingTraceHolder) storeFromState(state *adk.ChatModelAgentState) { + if h == nil || state == nil || len(state.Messages) == 0 { + return + } + cloned := cloneADKMessagesForTrace(state.Messages) + if len(cloned) == 0 { + return + } + h.mu.Lock() + h.msgs = cloned + h.mu.Unlock() +} + +func (h *modelFacingTraceHolder) storeFromAgenticState(state *adk.TypedChatModelAgentState[*schema.AgenticMessage]) { + if h == nil || state == nil || len(state.Messages) == 0 { + return + } + cloned := cloneADKMessagesForTrace(AgenticMessagesToEino(state.Messages)) + if len(cloned) == 0 { + return + } + h.mu.Lock() + h.msgs = cloned + h.mu.Unlock() +} + +func cloneADKMessagesForTrace(msgs []adk.Message) []adk.Message { + if len(msgs) == 0 { + return nil + } + b, err := json.Marshal(msgs) + if err != nil { + return nil + } + var out []adk.Message + if err := json.Unmarshal(b, &out); err != nil { + return nil + } + return out +} + +// modelFacingTraceMiddleware 必须在 Handlers 链中处于 **BeforeModel 最后**(telemetry 之后), +// 此时 state.Messages 即为本次 LLM 调用的最终入参。 +type modelFacingTraceMiddleware struct { + adk.BaseChatModelAgentMiddleware + holder *modelFacingTraceHolder +} + +func newModelFacingTraceMiddleware(holder *modelFacingTraceHolder) adk.ChatModelAgentMiddleware { + if holder == nil { + return nil + } + return &modelFacingTraceMiddleware{holder: holder} +} + +func (m *modelFacingTraceMiddleware) BeforeModelRewriteState( + ctx context.Context, + state *adk.ChatModelAgentState, + mc *adk.ModelContext, +) (context.Context, *adk.ChatModelAgentState, error) { + if m.holder != nil && state != nil { + m.holder.storeFromState(state) + } + return ctx, state, nil +} + +type agenticModelFacingTraceMiddleware struct { + *adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage] + holder *modelFacingTraceHolder +} + +func newAgenticModelFacingTraceMiddleware(holder *modelFacingTraceHolder) adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] { + if holder == nil { + return nil + } + return &agenticModelFacingTraceMiddleware{ + TypedBaseChatModelAgentMiddleware: &adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage]{}, + holder: holder, + } +} + +func (m *agenticModelFacingTraceMiddleware) BeforeModelRewriteState( + ctx context.Context, + state *adk.TypedChatModelAgentState[*schema.AgenticMessage], + mc *adk.TypedModelContext[*schema.AgenticMessage], +) (context.Context, *adk.TypedChatModelAgentState[*schema.AgenticMessage], error) { + if m.holder != nil && state != nil { + m.holder.storeFromAgenticState(state) + } + return ctx, state, nil +} diff --git a/internal/multiagent/eino_model_resilience.go b/internal/multiagent/eino_model_resilience.go new file mode 100644 index 00000000..bf85bbc2 --- /dev/null +++ b/internal/multiagent/eino_model_resilience.go @@ -0,0 +1,500 @@ +package multiagent + +import ( + "context" + "errors" + "fmt" + "net" + "net/http" + "strings" + "sync" + "time" + + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/openai" + "cyberstrike-ai/internal/reasoning" + + agenticopenai "github.com/cloudwego/eino-ext/components/model/agenticopenai" + einoopenai "github.com/cloudwego/eino-ext/components/model/openai" + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" +) + +type einoModelMode string + +const ( + einoModelModeNormal einoModelMode = "normal" + einoModelModePlanner einoModelMode = "planner" +) + +type einoModelFactory func(ctx context.Context, oa config.OpenAIConfig, mode einoModelMode) (model.ToolCallingChatModel, error) +type einoAgenticModelConfigFactory func(ctx context.Context, oa config.OpenAIConfig, mode einoModelMode) (model.AgenticModel, error) + +func newEinoBaseHTTPClient() *http.Client { + return &http.Client{ + Timeout: 30 * time.Minute, + Transport: &http.Transport{ + DialContext: (&net.Dialer{ + Timeout: 300 * time.Second, + KeepAlive: 300 * time.Second, + }).DialContext, + MaxIdleConns: 100, + MaxIdleConnsPerHost: 10, + IdleConnTimeout: 90 * time.Second, + TLSHandshakeTimeout: 30 * time.Second, + ResponseHeaderTimeout: 60 * time.Minute, + }, + } +} + +func newEinoOpenAIChatModelFactory( + baseHTTPClient *http.Client, + reasoningClient *reasoning.ClientIntent, + logger *zap.Logger, +) einoModelFactory { + if baseHTTPClient == nil { + baseHTTPClient = newEinoBaseHTTPClient() + } + return func(ctx context.Context, oa config.OpenAIConfig, mode einoModelMode) (model.ToolCallingChatModel, error) { + httpClient := openai.NewEinoHTTPClient(&oa, baseHTTPClient) + openai.AttachSummarizationDiagTransport(httpClient, logger) + maxCompletionTokens := oa.MaxCompletionTokensEffective() + modelCfg := &einoopenai.ChatModelConfig{ + APIKey: oa.APIKey, + BaseURL: strings.TrimSuffix(oa.BaseURL, "/"), + Model: oa.Model, + HTTPClient: httpClient, + MaxCompletionTokens: &maxCompletionTokens, + } + if mode == einoModelModePlanner { + reasoning.ApplyPlanExecutePlannerModelConfig(modelCfg, &oa) + } else { + reasoning.ApplyToEinoChatModelConfig(modelCfg, &oa, reasoningClient) + } + baseModel, err := einoopenai.NewChatModel(ctx, modelCfg) + if err != nil { + return nil, err + } + return newStreamToolCallIndexRepairModel(baseModel), nil + } +} + +func newEinoOpenAIAgenticChatModelFactory( + baseHTTPClient *http.Client, + reasoningClient *reasoning.ClientIntent, + logger *zap.Logger, +) einoAgenticModelConfigFactory { + if baseHTTPClient == nil { + baseHTTPClient = newEinoBaseHTTPClient() + } + return func(ctx context.Context, oa config.OpenAIConfig, mode einoModelMode) (model.AgenticModel, error) { + if !supportsEinoAgenticOpenAIBackend(oa) { + return nil, fmt.Errorf("eino agentic model: provider %q is not enabled for agenticopenai backend", strings.TrimSpace(oa.Provider)) + } + httpClient := openai.NewEinoHTTPClient(&oa, baseHTTPClient) + openai.AttachSummarizationDiagTransport(httpClient, logger) + maxCompletionTokens := oa.MaxCompletionTokensEffective() + modelCfg := &agenticopenai.ChatConfig{ + APIKey: oa.APIKey, + BaseURL: strings.TrimSuffix(oa.BaseURL, "/"), + Model: oa.Model, + HTTPClient: httpClient, + MaxCompletionTokens: &maxCompletionTokens, + ExtraFields: reasoning.AgenticOpenAIExtraFields(&oa, reasoningClient), + } + if mode == einoModelModePlanner { + modelCfg.ExtraFields = reasoning.AgenticOpenAIPlannerExtraFields(&oa) + } + return agenticopenai.NewChatModel(ctx, modelCfg) + } +} + +func supportsEinoAgenticOpenAIBackend(oa config.OpenAIConfig) bool { + provider := strings.ToLower(strings.TrimSpace(oa.Provider)) + return provider == "" || provider == "openai" || provider == "openai_compatible" +} + +func agenticModelGateFactory(factory einoAgenticModelConfigFactory, oa config.OpenAIConfig, mode einoModelMode) einoAgenticModelFactory { + if factory == nil { + return nil + } + return func(ctx context.Context) (model.AgenticModel, error) { + return factory(ctx, oa, mode) + } +} + +func newEinoModelRetryConfig( + mw *config.MultiAgentEinoMiddlewareConfig, + logger *zap.Logger, + scope string, +) *adk.ModelRetryConfig { + maxRetries := RunRetryMaxAttemptsFromConfig(mw) + maxBackoff := einoRunRetryMaxBackoffFromConfig(mw) + return &adk.ModelRetryConfig{ + MaxRetries: maxRetries, + BackoffFunc: func(_ context.Context, attempt int) time.Duration { + return einoTransientRetryBackoff(attempt-1, maxBackoff) + }, + ShouldRetry: func(ctx context.Context, retryCtx *adk.RetryContext) *adk.RetryDecision { + if retryCtx == nil || ctx.Err() != nil { + return &adk.RetryDecision{} + } + if retryCtx.Err != nil { + if !isEinoTransientRunError(retryCtx.Err) { + return &adk.RetryDecision{} + } + if logger != nil { + kind, summary := einoTransientRunErrorUserDetail(retryCtx.Err) + logger.Warn("eino native model retry", + zap.String("scope", scope), + zap.Int("attempt", retryCtx.RetryAttempt), + zap.Int("maxRetries", maxRetries), + zap.String("errorKind", kind), + zap.String("errorSummary", summary), + ) + } + return &adk.RetryDecision{Retry: true, RejectReason: "transient_model_error"} + } + if isRetryableEmptyModelOutput(retryCtx.OutputMessage) { + if logger != nil { + logger.Warn("eino native model retry: empty model output", + zap.String("scope", scope), + zap.Int("attempt", retryCtx.RetryAttempt), + zap.Int("maxRetries", maxRetries), + ) + } + return &adk.RetryDecision{Retry: true, RejectReason: "empty_model_output"} + } + return &adk.RetryDecision{} + }, + } +} + +func newEinoAgenticModelRetryConfig( + mw *config.MultiAgentEinoMiddlewareConfig, + logger *zap.Logger, + scope string, +) *adk.TypedModelRetryConfig[*schema.AgenticMessage] { + maxRetries := RunRetryMaxAttemptsFromConfig(mw) + maxBackoff := einoRunRetryMaxBackoffFromConfig(mw) + return &adk.TypedModelRetryConfig[*schema.AgenticMessage]{ + MaxRetries: maxRetries, + BackoffFunc: func(_ context.Context, attempt int) time.Duration { + return einoTransientRetryBackoff(attempt-1, maxBackoff) + }, + ShouldRetry: func(ctx context.Context, retryCtx *adk.TypedRetryContext[*schema.AgenticMessage]) *adk.TypedRetryDecision[*schema.AgenticMessage] { + if retryCtx == nil || ctx.Err() != nil { + return &adk.TypedRetryDecision[*schema.AgenticMessage]{} + } + if retryCtx.Err != nil { + if !isEinoTransientRunError(retryCtx.Err) { + return &adk.TypedRetryDecision[*schema.AgenticMessage]{} + } + if logger != nil { + kind, summary := einoTransientRunErrorUserDetail(retryCtx.Err) + logger.Warn("eino native agentic model retry", + zap.String("scope", scope), + zap.Int("attempt", retryCtx.RetryAttempt), + zap.Int("maxRetries", maxRetries), + zap.String("errorKind", kind), + zap.String("errorSummary", summary), + ) + } + return &adk.TypedRetryDecision[*schema.AgenticMessage]{Retry: true, RejectReason: "transient_model_error"} + } + if isRetryableEmptyAgenticModelOutput(retryCtx.OutputMessage) { + if logger != nil { + logger.Warn("eino native agentic model retry: empty model output", + zap.String("scope", scope), + zap.Int("attempt", retryCtx.RetryAttempt), + zap.Int("maxRetries", maxRetries), + ) + } + return &adk.TypedRetryDecision[*schema.AgenticMessage]{Retry: true, RejectReason: "empty_model_output"} + } + return &adk.TypedRetryDecision[*schema.AgenticMessage]{} + }, + } +} + +func newEinoModelFailoverConfig( + ctx context.Context, + appCfg *config.Config, + mw *config.MultiAgentEinoMiddlewareConfig, + mode einoModelMode, + factory einoModelFactory, + logger *zap.Logger, + scope string, + progress func(eventType, message string, data interface{}), + orchestration string, + conversationID string, +) (*adk.ModelFailoverConfig[*schema.Message], error) { + channels := resolveEinoFailoverChannels(appCfg, mw) + if len(channels) == 0 { + return nil, nil + } + if factory == nil { + return nil, fmt.Errorf("eino model failover: 模型工厂为空") + } + + maxRetries := len(channels) + if mw != nil && mw.ModelFailoverMaxRetries > 0 && mw.ModelFailoverMaxRetries < maxRetries { + maxRetries = mw.ModelFailoverMaxRetries + } + channels = channels[:maxRetries] + + cache := make(map[string]model.BaseModel[*schema.Message], len(channels)) + var mu sync.Mutex + return &adk.ModelFailoverConfig[*schema.Message]{ + MaxRetries: uint(maxRetries), + ShouldFailover: func(ctx context.Context, _ *schema.Message, err error) bool { + if ctx.Err() != nil || err == nil { + return false + } + err = unwrapEinoRetryExhausted(err) + return isEinoTransientRunError(err) + }, + GetFailoverModel: func(ctx context.Context, failoverCtx *adk.FailoverContext[*schema.Message]) (model.BaseModel[*schema.Message], []*schema.Message, error) { + if failoverCtx == nil || failoverCtx.FailoverAttempt == 0 { + return nil, nil, fmt.Errorf("eino model failover: invalid failover attempt") + } + idx := int(failoverCtx.FailoverAttempt) - 1 + if idx < 0 || idx >= len(channels) { + return nil, nil, fmt.Errorf("eino model failover: no channel for attempt %d", failoverCtx.FailoverAttempt) + } + ch := channels[idx] + mu.Lock() + cached := cache[ch.id] + mu.Unlock() + if cached != nil { + emitEinoModelFailoverEvent(progress, conversationID, orchestration, scope, ch.id, ch.cfg.Model, failoverCtx.FailoverAttempt) + if logger != nil { + logger.Warn("eino native model failover", + zap.String("scope", scope), + zap.String("channel", ch.id), + zap.String("model", ch.cfg.Model), + zap.Uint("attempt", failoverCtx.FailoverAttempt), + ) + } + return cached, nil, nil + } + m, err := factory(ctx, ch.cfg, mode) + if err != nil { + return nil, nil, fmt.Errorf("eino model failover channel %q: %w", ch.id, err) + } + mu.Lock() + cache[ch.id] = m + mu.Unlock() + emitEinoModelFailoverEvent(progress, conversationID, orchestration, scope, ch.id, ch.cfg.Model, failoverCtx.FailoverAttempt) + if logger != nil { + logger.Warn("eino native model failover", + zap.String("scope", scope), + zap.String("channel", ch.id), + zap.String("model", ch.cfg.Model), + zap.Uint("attempt", failoverCtx.FailoverAttempt), + ) + } + return m, nil, nil + }, + }, nil +} + +func newEinoAgenticModelFailoverConfig( + ctx context.Context, + appCfg *config.Config, + mw *config.MultiAgentEinoMiddlewareConfig, + mode einoModelMode, + factory einoAgenticModelConfigFactory, + logger *zap.Logger, + scope string, + progress func(eventType, message string, data interface{}), + orchestration string, + conversationID string, +) (*adk.ModelFailoverConfig[*schema.AgenticMessage], error) { + channels := resolveEinoFailoverChannels(appCfg, mw) + if len(channels) == 0 { + return nil, nil + } + if factory == nil { + return nil, fmt.Errorf("eino agentic model failover: 模型工厂为空") + } + + maxRetries := len(channels) + if mw != nil && mw.ModelFailoverMaxRetries > 0 && mw.ModelFailoverMaxRetries < maxRetries { + maxRetries = mw.ModelFailoverMaxRetries + } + channels = channels[:maxRetries] + + cache := make(map[string]model.BaseModel[*schema.AgenticMessage], len(channels)) + var mu sync.Mutex + return &adk.ModelFailoverConfig[*schema.AgenticMessage]{ + MaxRetries: uint(maxRetries), + ShouldFailover: func(ctx context.Context, _ *schema.AgenticMessage, err error) bool { + if ctx.Err() != nil || err == nil { + return false + } + err = unwrapEinoRetryExhausted(err) + return isEinoTransientRunError(err) + }, + GetFailoverModel: func(ctx context.Context, failoverCtx *adk.FailoverContext[*schema.AgenticMessage]) (model.BaseModel[*schema.AgenticMessage], []*schema.AgenticMessage, error) { + if failoverCtx == nil || failoverCtx.FailoverAttempt == 0 { + return nil, nil, fmt.Errorf("eino agentic model failover: invalid failover attempt") + } + idx := int(failoverCtx.FailoverAttempt) - 1 + if idx < 0 || idx >= len(channels) { + return nil, nil, fmt.Errorf("eino agentic model failover: no channel for attempt %d", failoverCtx.FailoverAttempt) + } + ch := channels[idx] + mu.Lock() + cached := cache[ch.id] + mu.Unlock() + if cached != nil { + emitEinoModelFailoverEvent(progress, conversationID, orchestration, scope, ch.id, ch.cfg.Model, failoverCtx.FailoverAttempt) + if logger != nil { + logger.Warn("eino native agentic model failover", + zap.String("scope", scope), + zap.String("channel", ch.id), + zap.String("model", ch.cfg.Model), + zap.Uint("attempt", failoverCtx.FailoverAttempt), + ) + } + return cached, nil, nil + } + m, err := factory(ctx, ch.cfg, mode) + if err != nil { + return nil, nil, fmt.Errorf("eino agentic model failover channel %q: %w", ch.id, err) + } + mu.Lock() + cache[ch.id] = m + mu.Unlock() + emitEinoModelFailoverEvent(progress, conversationID, orchestration, scope, ch.id, ch.cfg.Model, failoverCtx.FailoverAttempt) + if logger != nil { + logger.Warn("eino native agentic model failover", + zap.String("scope", scope), + zap.String("channel", ch.id), + zap.String("model", ch.cfg.Model), + zap.Uint("attempt", failoverCtx.FailoverAttempt), + ) + } + return m, nil, nil + }, + }, nil +} + +type resolvedEinoFailoverChannel struct { + id string + cfg config.OpenAIConfig +} + +func resolveEinoFailoverChannels(appCfg *config.Config, mw *config.MultiAgentEinoMiddlewareConfig) []resolvedEinoFailoverChannel { + if appCfg == nil || mw == nil || len(mw.ModelFailoverChannels) == 0 { + return nil + } + primary := appCfg.OpenAI + seen := map[string]struct{}{} + out := make([]resolvedEinoFailoverChannel, 0, len(mw.ModelFailoverChannels)) + for _, raw := range mw.ModelFailoverChannels { + id := config.NormalizeAIChannelID(raw) + if id == "" { + continue + } + if _, ok := seen[id]; ok { + continue + } + oa, resolvedID, ok := appCfg.AI.ResolveChannel(id) + if !ok { + continue + } + if sameOpenAIModelEndpoint(primary, oa) { + continue + } + seen[resolvedID] = struct{}{} + out = append(out, resolvedEinoFailoverChannel{id: resolvedID, cfg: oa}) + } + return out +} + +func sameOpenAIModelEndpoint(a, b config.OpenAIConfig) bool { + return strings.EqualFold(strings.TrimSpace(a.Provider), strings.TrimSpace(b.Provider)) && + strings.TrimRight(strings.TrimSpace(a.BaseURL), "/") == strings.TrimRight(strings.TrimSpace(b.BaseURL), "/") && + strings.TrimSpace(a.APIKey) == strings.TrimSpace(b.APIKey) && + strings.TrimSpace(a.Model) == strings.TrimSpace(b.Model) +} + +func isRetryableEmptyModelOutput(msg *schema.Message) bool { + if msg == nil { + return true + } + return strings.TrimSpace(msg.Content) == "" && + strings.TrimSpace(msg.ReasoningContent) == "" && + len(msg.ToolCalls) == 0 && + len(msg.MultiContent) == 0 && + len(msg.UserInputMultiContent) == 0 && + len(msg.AssistantGenMultiContent) == 0 +} + +func isRetryableEmptyAgenticModelOutput(msg *schema.AgenticMessage) bool { + if msg == nil { + return true + } + for _, block := range msg.ContentBlocks { + if block == nil { + continue + } + switch { + case block.Reasoning != nil: + if strings.TrimSpace(block.Reasoning.Text) != "" { + return false + } + case block.UserInputText != nil: + if strings.TrimSpace(block.UserInputText.Text) != "" { + return false + } + case block.AssistantGenText != nil: + if strings.TrimSpace(block.AssistantGenText.Text) != "" { + return false + } + default: + return false + } + } + return true +} + +func unwrapEinoRetryExhausted(err error) error { + var retryErr *adk.RetryExhaustedError + if errors.As(err, &retryErr) && retryErr.LastErr != nil { + return retryErr.LastErr + } + return err +} + +func isEinoNativeWillRetry(err error) (*adk.WillRetryError, bool) { + var willRetry *adk.WillRetryError + if errors.As(err, &willRetry) { + return willRetry, true + } + return nil, false +} + +func emitEinoModelFailoverEvent( + progress func(eventType, message string, data interface{}), + conversationID, orchestration, scope, channelID, modelName string, + attempt uint, +) { + if progress == nil { + return + } + msg := fmt.Sprintf("主模型重试耗尽,正在切换备用模型 %s。", modelName) + progress("eino_model_failover", msg, map[string]interface{}{ + "conversationId": conversationID, + "source": "eino", + "orchestration": orchestration, + "scope": scope, + "channel": channelID, + "model": modelName, + "attempt": attempt, + }) +} diff --git a/internal/multiagent/eino_model_resilience_test.go b/internal/multiagent/eino_model_resilience_test.go new file mode 100644 index 00000000..e0651c40 --- /dev/null +++ b/internal/multiagent/eino_model_resilience_test.go @@ -0,0 +1,376 @@ +package multiagent + +import ( + "context" + "errors" + "testing" + "time" + + "cyberstrike-ai/internal/config" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/schema" +) + +func TestNewEinoModelRetryConfigUsesNativeFieldsFirst(t *testing.T) { + t.Parallel() + mw := &config.MultiAgentEinoMiddlewareConfig{ + ModelRetryMaxRetries: 2, + ModelRetryMaxBackoffSec: 7, + RunRetryMaxAttempts: 9, + RunRetryMaxBackoffSec: 11, + } + cfg := newEinoModelRetryConfig(mw, nil, "test") + if cfg.MaxRetries != 2 { + t.Fatalf("MaxRetries = %d, want 2", cfg.MaxRetries) + } + backoff := cfg.BackoffFunc(context.Background(), 1) + if backoff < 500*time.Millisecond || backoff > 2*time.Second { + t.Fatalf("attempt 1 backoff = %v, want first equal-jitter window", backoff) + } + if got := einoRunRetryMaxBackoffFromConfig(mw); got != 7*time.Second { + t.Fatalf("backoff from config = %v, want 7s", got) + } +} + +func TestEinoModelRetryPolicyRetriesTransientAndEmptyOutput(t *testing.T) { + t.Parallel() + cfg := newEinoModelRetryConfig(&config.MultiAgentEinoMiddlewareConfig{ModelRetryMaxRetries: 1}, nil, "test") + if got := cfg.ShouldRetry(context.Background(), &adk.RetryContext{Err: errors.New("HTTP 429 Too Many Requests")}); got == nil || !got.Retry { + t.Fatal("transient model error should retry") + } + if got := cfg.ShouldRetry(context.Background(), &adk.RetryContext{OutputMessage: schema.AssistantMessage("", nil)}); got == nil || !got.Retry { + t.Fatal("empty assistant output should retry") + } + if got := cfg.ShouldRetry(context.Background(), &adk.RetryContext{OutputMessage: schema.AssistantMessage("", []schema.ToolCall{{ID: "call_1"}})}); got == nil || got.Retry { + t.Fatal("assistant tool call output should not be treated as empty") + } + if got := cfg.ShouldRetry(context.Background(), &adk.RetryContext{Err: errors.New("invalid api key")}); got == nil || got.Retry { + t.Fatal("permanent auth error should not retry") + } +} + +func TestEinoAgenticModelRetryPolicyRetriesTransientAndEmptyOutput(t *testing.T) { + t.Parallel() + cfg := newEinoAgenticModelRetryConfig(&config.MultiAgentEinoMiddlewareConfig{ModelRetryMaxRetries: 1}, nil, "agentic") + if got := cfg.ShouldRetry(context.Background(), &adk.TypedRetryContext[*schema.AgenticMessage]{Err: errors.New("HTTP 429 Too Many Requests")}); got == nil || !got.Retry { + t.Fatal("transient agentic model error should retry") + } + if got := cfg.ShouldRetry(context.Background(), &adk.TypedRetryContext[*schema.AgenticMessage]{ + OutputMessage: &schema.AgenticMessage{Role: schema.AgenticRoleTypeAssistant}, + }); got == nil || !got.Retry { + t.Fatal("empty agentic assistant output should retry") + } + if got := cfg.ShouldRetry(context.Background(), &adk.TypedRetryContext[*schema.AgenticMessage]{ + OutputMessage: &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.AssistantGenText{Text: "ok"})}, + }, + }); got == nil || got.Retry { + t.Fatal("agentic assistant text should not be treated as empty") + } + if got := cfg.ShouldRetry(context.Background(), &adk.TypedRetryContext[*schema.AgenticMessage]{ + OutputMessage: &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.FunctionToolCall{ + CallID: "call_1", Name: "search", Arguments: `{"q":"x"}`, + })}, + }, + }); got == nil || got.Retry { + t.Fatal("agentic tool call output should not be treated as empty") + } + if got := cfg.ShouldRetry(context.Background(), &adk.TypedRetryContext[*schema.AgenticMessage]{Err: errors.New("invalid api key")}); got == nil || got.Retry { + t.Fatal("permanent auth error should not retry") + } +} + +func TestResolveEinoFailoverChannelsSkipsPrimaryDuplicateAndUnknown(t *testing.T) { + t.Parallel() + appCfg := &config.Config{ + OpenAI: config.OpenAIConfig{Provider: "openai", APIKey: "k1", BaseURL: "https://api.example/v1", Model: "primary"}, + AI: config.AIConfig{Channels: map[string]config.AIChannelConfig{ + "same": {Provider: "openai", APIKey: "k1", BaseURL: "https://api.example/v1", Model: "primary"}, + "fb1": {Provider: "openai", APIKey: "k2", BaseURL: "https://api.example/v1", Model: "fallback-1"}, + "fb2": {Provider: "claude", APIKey: "k3", BaseURL: "https://api.anthropic.com/v1", Model: "claude-sonnet"}, + }}, + } + got := resolveEinoFailoverChannels(appCfg, &config.MultiAgentEinoMiddlewareConfig{ + ModelFailoverChannels: []string{"same", "missing", "fb1", "fb1", "fb2"}, + ModelFailoverMaxRetries: 1, + }) + if len(got) != 2 { + t.Fatalf("resolved channels len = %d, want 2 before max cap is applied by config builder", len(got)) + } + if got[0].id != "fb1" || got[1].id != "fb2" { + t.Fatalf("resolved channel order = %#v", got) + } +} + +func TestNewEinoModelFailoverConfigBuildsDistinctFallbackModel(t *testing.T) { + t.Parallel() + appCfg := &config.Config{ + OpenAI: config.OpenAIConfig{APIKey: "k1", BaseURL: "https://api.example/v1", Model: "primary"}, + AI: config.AIConfig{Channels: map[string]config.AIChannelConfig{ + "fb1": {APIKey: "k2", BaseURL: "https://api.example/v1", Model: "fallback-1"}, + "fb2": {APIKey: "k3", BaseURL: "https://api.example/v1", Model: "fallback-2"}, + }}, + } + var built []string + cfg, err := newEinoModelFailoverConfig( + context.Background(), + appCfg, + &config.MultiAgentEinoMiddlewareConfig{ + ModelFailoverChannels: []string{"fb1", "fb2"}, + ModelFailoverMaxRetries: 1, + }, + einoModelModeNormal, + func(_ context.Context, oa config.OpenAIConfig, _ einoModelMode) (model.ToolCallingChatModel, error) { + built = append(built, oa.Model) + return &streamToolCallIndexFakeModel{}, nil + }, + nil, + "test", + nil, + "deep", + "conv-1", + ) + if err != nil { + t.Fatalf("newEinoModelFailoverConfig: %v", err) + } + if cfg == nil || cfg.MaxRetries != 1 { + t.Fatalf("failover cfg = %#v, want max retries 1", cfg) + } + m, msgs, err := cfg.GetFailoverModel(context.Background(), &adk.FailoverContext[*schema.Message]{FailoverAttempt: 1}) + if err != nil || m == nil || msgs != nil { + t.Fatalf("GetFailoverModel = (%v, %v, %v)", m, msgs, err) + } + if len(built) != 1 || built[0] != "fallback-1" { + t.Fatalf("built models = %v, want [fallback-1]", built) + } + if !cfg.ShouldFailover(context.Background(), nil, &adk.RetryExhaustedError{LastErr: errors.New("upstream returned 503"), TotalRetries: 4}) { + t.Fatal("retry-exhausted transient error should fail over") + } + if cfg.ShouldFailover(context.Background(), nil, &adk.RetryExhaustedError{LastErr: errors.New("invalid api key"), TotalRetries: 4}) { + t.Fatal("retry-exhausted permanent error should not fail over") + } +} + +func TestNewEinoModelFailoverConfigEmitsProgressEvent(t *testing.T) { + t.Parallel() + appCfg := &config.Config{ + OpenAI: config.OpenAIConfig{APIKey: "k1", BaseURL: "https://api.example/v1", Model: "primary"}, + AI: config.AIConfig{Channels: map[string]config.AIChannelConfig{ + "fb1": {APIKey: "k2", BaseURL: "https://api.example/v1", Model: "fallback-1"}, + }}, + } + var events []struct { + eventType string + message string + data interface{} + } + cfg, err := newEinoModelFailoverConfig( + context.Background(), + appCfg, + &config.MultiAgentEinoMiddlewareConfig{ModelFailoverChannels: []string{"fb1"}}, + einoModelModeNormal, + func(_ context.Context, _ config.OpenAIConfig, _ einoModelMode) (model.ToolCallingChatModel, error) { + return &streamToolCallIndexFakeModel{}, nil + }, + nil, + "test", + func(eventType, message string, data interface{}) { + events = append(events, struct { + eventType string + message string + data interface{} + }{eventType: eventType, message: message, data: data}) + }, + "deep", + "conv-1", + ) + if err != nil { + t.Fatalf("newEinoModelFailoverConfig: %v", err) + } + if _, _, err := cfg.GetFailoverModel(context.Background(), &adk.FailoverContext[*schema.Message]{FailoverAttempt: 1}); err != nil { + t.Fatalf("GetFailoverModel: %v", err) + } + if len(events) != 1 || events[0].eventType != "eino_model_failover" { + t.Fatalf("events = %#v, want one eino_model_failover", events) + } + payload, ok := events[0].data.(map[string]interface{}) + if !ok { + t.Fatalf("event payload type = %T", events[0].data) + } + if payload["conversationId"] != "conv-1" || payload["orchestration"] != "deep" || payload["channel"] != "fb1" || payload["model"] != "fallback-1" { + t.Fatalf("payload = %#v", payload) + } +} + +func TestNewEinoAgenticModelFailoverConfigBuildsDistinctFallbackModel(t *testing.T) { + t.Parallel() + appCfg := &config.Config{ + OpenAI: config.OpenAIConfig{Provider: "openai", APIKey: "k1", BaseURL: "https://api.example/v1", Model: "primary"}, + AI: config.AIConfig{Channels: map[string]config.AIChannelConfig{ + "fb1": {Provider: "openai", APIKey: "k2", BaseURL: "https://api.example/v1", Model: "fallback-1"}, + "fb2": {Provider: "openai", APIKey: "k3", BaseURL: "https://api.example/v1", Model: "fallback-2"}, + }}, + } + var built []string + cfg, err := newEinoAgenticModelFailoverConfig( + context.Background(), + appCfg, + &config.MultiAgentEinoMiddlewareConfig{ + ModelFailoverChannels: []string{"fb1", "fb2"}, + ModelFailoverMaxRetries: 1, + }, + einoModelModeNormal, + func(_ context.Context, oa config.OpenAIConfig, _ einoModelMode) (model.AgenticModel, error) { + built = append(built, oa.Model) + return &fakeAgenticGateModel{}, nil + }, + nil, + "agentic", + nil, + "eino_single_agentic", + "conv-1", + ) + if err != nil { + t.Fatalf("newEinoAgenticModelFailoverConfig: %v", err) + } + if cfg == nil || cfg.MaxRetries != 1 { + t.Fatalf("agentic failover cfg = %#v, want max retries 1", cfg) + } + m, msgs, err := cfg.GetFailoverModel(context.Background(), &adk.FailoverContext[*schema.AgenticMessage]{FailoverAttempt: 1}) + if err != nil || m == nil || msgs != nil { + t.Fatalf("GetFailoverModel = (%v, %v, %v)", m, msgs, err) + } + if len(built) != 1 || built[0] != "fallback-1" { + t.Fatalf("built models = %v, want [fallback-1]", built) + } + if !cfg.ShouldFailover(context.Background(), nil, &adk.RetryExhaustedError{LastErr: errors.New("upstream returned 503"), TotalRetries: 4}) { + t.Fatal("retry-exhausted transient agentic error should fail over") + } + if cfg.ShouldFailover(context.Background(), nil, &adk.RetryExhaustedError{LastErr: errors.New("invalid api key"), TotalRetries: 4}) { + t.Fatal("retry-exhausted permanent agentic error should not fail over") + } +} + +func TestNewEinoAgenticModelFailoverConfigEmitsProgressEvent(t *testing.T) { + t.Parallel() + appCfg := &config.Config{ + OpenAI: config.OpenAIConfig{Provider: "openai", APIKey: "k1", BaseURL: "https://api.example/v1", Model: "primary"}, + AI: config.AIConfig{Channels: map[string]config.AIChannelConfig{ + "fb1": {Provider: "openai", APIKey: "k2", BaseURL: "https://api.example/v1", Model: "fallback-1"}, + }}, + } + var events []struct { + eventType string + message string + data interface{} + } + cfg, err := newEinoAgenticModelFailoverConfig( + context.Background(), + appCfg, + &config.MultiAgentEinoMiddlewareConfig{ModelFailoverChannels: []string{"fb1"}}, + einoModelModeNormal, + func(_ context.Context, _ config.OpenAIConfig, _ einoModelMode) (model.AgenticModel, error) { + return &fakeAgenticGateModel{}, nil + }, + nil, + "agentic", + func(eventType, message string, data interface{}) { + events = append(events, struct { + eventType string + message string + data interface{} + }{eventType: eventType, message: message, data: data}) + }, + "eino_single_agentic", + "conv-1", + ) + if err != nil { + t.Fatalf("newEinoAgenticModelFailoverConfig: %v", err) + } + if _, _, err := cfg.GetFailoverModel(context.Background(), &adk.FailoverContext[*schema.AgenticMessage]{FailoverAttempt: 1}); err != nil { + t.Fatalf("GetFailoverModel: %v", err) + } + if len(events) != 1 || events[0].eventType != "eino_model_failover" { + t.Fatalf("events = %#v, want one eino_model_failover", events) + } + payload, ok := events[0].data.(map[string]interface{}) + if !ok { + t.Fatalf("event payload type = %T", events[0].data) + } + if payload["conversationId"] != "conv-1" || payload["orchestration"] != "eino_single_agentic" || payload["channel"] != "fb1" || payload["model"] != "fallback-1" { + t.Fatalf("payload = %#v", payload) + } +} + +func TestNewEinoOpenAIAgenticChatModelFactoryBuildsBackend(t *testing.T) { + t.Parallel() + factory := newEinoOpenAIAgenticChatModelFactory(newEinoBaseHTTPClient(), nil, nil) + m, err := factory(context.Background(), config.OpenAIConfig{ + Provider: "openai", + APIKey: "test-key", + BaseURL: "https://api.example/v1", + Model: "gpt-4o-mini", + Reasoning: config.OpenAIReasoningConfig{ + Profile: "openai_compat", + Mode: "on", + Effort: "high", + }, + }, einoModelModeNormal) + if err != nil { + t.Fatalf("agentic factory: %v", err) + } + if m == nil { + t.Fatal("agentic factory returned nil model") + } + gate := evaluateEinoAgenticModelGate(agenticModelGateFactory(factory, config.OpenAIConfig{ + Provider: "openai", + APIKey: "test-key", + BaseURL: "https://api.example/v1", + Model: "gpt-4o-mini", + }, einoModelModeNormal), einoAgenticRuntimeSupportV0914()) + if !gate.Ready { + t.Fatalf("gate = %#v, want ready with buildable agentic backend", gate) + } +} + +func TestNewEinoOpenAIAgenticChatModelFactoryRejectsUnsupportedProvider(t *testing.T) { + t.Parallel() + factory := newEinoOpenAIAgenticChatModelFactory(newEinoBaseHTTPClient(), nil, nil) + if _, err := factory(context.Background(), config.OpenAIConfig{ + Provider: "claude", + APIKey: "test-key", + BaseURL: "https://api.anthropic.com/v1", + Model: "claude-sonnet-4", + }, einoModelModeNormal); err == nil { + t.Fatal("expected unsupported provider error") + } + gate := evaluateEinoAgenticModelGate(agenticModelGateFactory(factory, config.OpenAIConfig{ + Provider: "claude", + APIKey: "test-key", + BaseURL: "https://api.anthropic.com/v1", + Model: "claude-sonnet-4", + }, einoModelModeNormal), einoAgenticRuntimeSupportV0914()) + if gate.Ready || !containsString(gate.Missing, "model.AgenticModel backend") { + t.Fatalf("gate = %#v, want backend missing for unsupported provider", gate) + } +} + +func TestEinoNativeRetryErrorsDoNotTriggerRunLevelTransientRetry(t *testing.T) { + t.Parallel() + err := &adk.WillRetryError{ErrStr: "HTTP 429 Too Many Requests", RetryAttempt: 1} + if isEinoTransientRunError(err) { + t.Fatal("WillRetryError should be observed, not treated as a run-level transient failure") + } + exhausted := &adk.RetryExhaustedError{LastErr: errors.New("HTTP 429 Too Many Requests"), TotalRetries: 4} + if isEinoTransientRunError(exhausted) { + t.Fatal("RetryExhaustedError should not trigger a second run-level retry layer") + } + if got := unwrapEinoRetryExhausted(exhausted); got == exhausted { + t.Fatal("unwrapEinoRetryExhausted should return the underlying model error") + } +} diff --git a/internal/multiagent/eino_model_rewrite_pipeline.go b/internal/multiagent/eino_model_rewrite_pipeline.go new file mode 100644 index 00000000..aabd3c1d --- /dev/null +++ b/internal/multiagent/eino_model_rewrite_pipeline.go @@ -0,0 +1,38 @@ +package multiagent + +import ( + "context" + "fmt" + + "github.com/cloudwego/eino/adk" +) + +func applyBeforeModelRewriteHandlers( + ctx context.Context, + msgs []adk.Message, + handlers []adk.ChatModelAgentMiddleware, +) ([]adk.Message, error) { + if len(msgs) == 0 || len(handlers) == 0 { + return msgs, nil + } + state := &adk.ChatModelAgentState{Messages: msgs} + modelCtx := &adk.ModelContext{} + curCtx := ctx + for _, h := range handlers { + if h == nil { + continue + } + nextCtx, nextState, err := h.BeforeModelRewriteState(curCtx, state, modelCtx) + if err != nil { + return nil, fmt.Errorf("before model rewrite: %w", err) + } + if nextCtx != nil { + curCtx = nextCtx + } + if nextState != nil { + state = nextState + } + } + return state.Messages, nil +} + diff --git a/internal/multiagent/eino_native_cancel.go b/internal/multiagent/eino_native_cancel.go new file mode 100644 index 00000000..605390b6 --- /dev/null +++ b/internal/multiagent/eino_native_cancel.go @@ -0,0 +1,102 @@ +package multiagent + +import ( + "context" + "errors" + "time" + + "github.com/cloudwego/eino/adk" +) + +const ( + einoNativeCancelImmediateWait = 1200 * time.Millisecond + einoNativeCancelSafePointWait = 3500 * time.Millisecond + einoNativeCancelSafePointTTL = 3 * time.Second +) + +type agentRuntimeCancelRegistrarKey struct{} +type agentTurnLoopInterruptRegistrarKey struct{} + +// AgentRuntimeCancelRegistrar binds the currently active Eino ADK cancel hook +// into the host task manager. The hook returns true when Eino accepted and +// handled the cancel request, so the host can avoid canceling the parent context. +type AgentRuntimeCancelRegistrar func(cancel func(error) bool) (unregister func()) + +// WithAgentRuntimeCancelRegistrar lets the HTTP/task layer trigger Eino's native +// Agent Cancel before falling back to the existing context cancellation path. +func WithAgentRuntimeCancelRegistrar(ctx context.Context, registrar AgentRuntimeCancelRegistrar) context.Context { + if ctx == nil || registrar == nil { + return ctx + } + return context.WithValue(ctx, agentRuntimeCancelRegistrarKey{}, registrar) +} + +func agentRuntimeCancelRegistrarFromContext(ctx context.Context) AgentRuntimeCancelRegistrar { + if ctx == nil { + return nil + } + if v, ok := ctx.Value(agentRuntimeCancelRegistrarKey{}).(AgentRuntimeCancelRegistrar); ok { + return v + } + return nil +} + +// AgentTurnLoopInterruptRegistrar binds a conversation-level TurnLoop interrupt +// pusher into the host task manager. The pusher receives the user supplied note +// and returns true when the note was accepted by the loop. +type AgentTurnLoopInterruptRegistrar func(push func(note string) bool) (unregister func()) + +// WithAgentTurnLoopInterruptRegistrar lets the HTTP/task layer enqueue a user +// supplement into an active Eino TurnLoop before falling back to cancellation. +func WithAgentTurnLoopInterruptRegistrar(ctx context.Context, registrar AgentTurnLoopInterruptRegistrar) context.Context { + if ctx == nil || registrar == nil { + return ctx + } + return context.WithValue(ctx, agentTurnLoopInterruptRegistrarKey{}, registrar) +} + +func agentTurnLoopInterruptRegistrarFromContext(ctx context.Context) AgentTurnLoopInterruptRegistrar { + if ctx == nil { + return nil + } + if v, ok := ctx.Value(agentTurnLoopInterruptRegistrarKey{}).(AgentTurnLoopInterruptRegistrar); ok { + return v + } + return nil +} + +func requestEinoNativeAgentCancel(cancelFn adk.AgentCancelFunc, cause error) (waitErr error, submitted bool, handled bool) { + if cancelFn == nil { + return nil, false, false + } + opts, waitFor := einoNativeCancelOptions(cause) + handle, submitted := cancelFn(opts...) + if !submitted || handle == nil { + return nil, submitted, false + } + waitCh := make(chan error, 1) + go func() { + waitCh <- handle.Wait() + }() + select { + case err := <-waitCh: + handled := err == nil || errors.Is(err, adk.ErrCancelTimeout) || errors.Is(err, adk.ErrExecutionEnded) + return err, submitted, handled + case <-time.After(waitFor): + return context.DeadlineExceeded, submitted, false + } +} + +func einoNativeCancelOptions(cause error) ([]adk.AgentCancelOption, time.Duration) { + if errors.Is(cause, ErrInterruptContinue) { + return []adk.AgentCancelOption{ + adk.WithAgentCancelMode(adk.CancelAfterChatModel | adk.CancelAfterToolCalls), + adk.WithAgentCancelTimeout(einoNativeCancelSafePointTTL), + adk.WithRecursive(), + }, einoNativeCancelSafePointWait + } + return []adk.AgentCancelOption{ + adk.WithAgentCancelMode(adk.CancelImmediate), + adk.WithRecursive(), + }, einoNativeCancelImmediateWait +} diff --git a/internal/multiagent/eino_native_cancel_test.go b/internal/multiagent/eino_native_cancel_test.go new file mode 100644 index 00000000..e3b4c5a7 --- /dev/null +++ b/internal/multiagent/eino_native_cancel_test.go @@ -0,0 +1,27 @@ +package multiagent + +import ( + "context" + "testing" +) + +func TestEinoNativeCancelOptionsByCause(t *testing.T) { + fullStopOpts, fullStopWait := einoNativeCancelOptions(context.Canceled) + if len(fullStopOpts) != 2 { + t.Fatalf("full stop options: got %d want 2", len(fullStopOpts)) + } + if fullStopWait != einoNativeCancelImmediateWait { + t.Fatalf("full stop wait: got %v want %v", fullStopWait, einoNativeCancelImmediateWait) + } + + interruptOpts, interruptWait := einoNativeCancelOptions(ErrInterruptContinue) + if len(interruptOpts) != 3 { + t.Fatalf("interrupt options: got %d want 3", len(interruptOpts)) + } + if interruptWait != einoNativeCancelSafePointWait { + t.Fatalf("interrupt wait: got %v want %v", interruptWait, einoNativeCancelSafePointWait) + } + if interruptWait <= einoNativeCancelSafePointTTL { + t.Fatalf("interrupt wait must allow the Eino safe-point timeout to elapse: wait=%v ttl=%v", interruptWait, einoNativeCancelSafePointTTL) + } +} diff --git a/internal/multiagent/eino_native_model_retry_progress.go b/internal/multiagent/eino_native_model_retry_progress.go new file mode 100644 index 00000000..c41890e9 --- /dev/null +++ b/internal/multiagent/eino_native_model_retry_progress.go @@ -0,0 +1,41 @@ +package multiagent + +import ( + "fmt" + + "github.com/cloudwego/eino/adk" + "go.uber.org/zap" +) + +func emitEinoNativeModelRetryProgress( + conversationID, orchMode string, + willRetry *adk.WillRetryError, + progress func(eventType, message string, data interface{}), + logger *zap.Logger, + runErr error, +) bool { + if willRetry == nil { + return false + } + if progress != nil { + reason := "" + if willRetry.RejectReason() != nil { + reason = fmt.Sprint(willRetry.RejectReason()) + } + progress("eino_model_retry", "模型调用遇到临时问题,Eino 正在原生重试…", map[string]interface{}{ + "conversationId": conversationID, + "source": "eino", + "orchestration": orchMode, + "attempt": willRetry.RetryAttempt, + "reason": reason, + "error": willRetry.Error(), + }) + } + if logger != nil { + logger.Warn("eino native model retry event", + zap.String("orchestration", orchMode), + zap.Int("attempt", willRetry.RetryAttempt), + zap.Error(runErr)) + } + return true +} diff --git a/internal/multiagent/eino_native_model_retry_progress_test.go b/internal/multiagent/eino_native_model_retry_progress_test.go new file mode 100644 index 00000000..c3030b2c --- /dev/null +++ b/internal/multiagent/eino_native_model_retry_progress_test.go @@ -0,0 +1,85 @@ +package multiagent + +import ( + "testing" + + "github.com/cloudwego/eino/adk" + "go.uber.org/zap" + "go.uber.org/zap/zaptest/observer" +) + +func TestEmitEinoNativeModelRetryProgress(t *testing.T) { + willRetry := &adk.WillRetryError{ + ErrStr: "HTTP 429 Too Many Requests", + RetryAttempt: 2, + } + var gotType, gotMessage string + var gotData map[string]interface{} + called := emitEinoNativeModelRetryProgress("conv-1", "deep_agent", willRetry, func(eventType, message string, data interface{}) { + gotType = eventType + gotMessage = message + var ok bool + gotData, ok = data.(map[string]interface{}) + if !ok { + t.Fatalf("progress data type = %T, want map[string]interface{}", data) + } + }, nil, willRetry) + if !called { + t.Fatal("called = false, want true") + } + if gotType != "eino_model_retry" { + t.Fatalf("event type = %q, want eino_model_retry", gotType) + } + if gotMessage != "模型调用遇到临时问题,Eino 正在原生重试…" { + t.Fatalf("message = %q", gotMessage) + } + assertNativeRetryMapValue(t, gotData, "conversationId", "conv-1") + assertNativeRetryMapValue(t, gotData, "source", "eino") + assertNativeRetryMapValue(t, gotData, "orchestration", "deep_agent") + assertNativeRetryMapValue(t, gotData, "attempt", 2) + assertNativeRetryMapValue(t, gotData, "reason", "") + assertNativeRetryMapValue(t, gotData, "error", "HTTP 429 Too Many Requests") +} + +func TestEmitEinoNativeModelRetryProgressNilSafe(t *testing.T) { + calledProgress := false + called := emitEinoNativeModelRetryProgress("conv-1", "deep_agent", nil, func(string, string, interface{}) { + calledProgress = true + }, nil, nil) + if called { + t.Fatal("called = true, want false") + } + if calledProgress { + t.Fatal("progress called for nil willRetry") + } +} + +func TestEmitEinoNativeModelRetryProgressLogsEvent(t *testing.T) { + core, logs := observer.New(zap.WarnLevel) + logger := zap.New(core) + willRetry := &adk.WillRetryError{ + ErrStr: "HTTP 500", + RetryAttempt: 3, + } + + emitEinoNativeModelRetryProgress("conv-1", "single_agent", willRetry, nil, logger, willRetry) + + entry := logs.FilterMessage("eino native model retry event").TakeAll() + if len(entry) != 1 { + t.Fatalf("log count = %d, want 1", len(entry)) + } + fields := entry[0].ContextMap() + if fields["orchestration"] != "single_agent" { + t.Fatalf("orchestration field = %v", fields["orchestration"]) + } + if fields["attempt"] != int64(3) { + t.Fatalf("attempt field = %v", fields["attempt"]) + } +} + +func assertNativeRetryMapValue(t *testing.T, data map[string]interface{}, key string, want interface{}) { + t.Helper() + if got := data[key]; got != want { + t.Fatalf("%s = %v, want %v", key, got, want) + } +} diff --git a/internal/multiagent/eino_orchestration.go b/internal/multiagent/eino_orchestration.go new file mode 100644 index 00000000..051984e2 --- /dev/null +++ b/internal/multiagent/eino_orchestration.go @@ -0,0 +1,408 @@ +package multiagent + +import ( + "context" + "fmt" + "strings" + + "cyberstrike-ai/internal/agent" + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/database" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/adk/prebuilt/planexecute" + "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" +) + +// PlanExecuteRootArgs 构建 Eino adk/prebuilt/planexecute 根 Agent 所需参数。 +type PlanExecuteRootArgs struct { + MainToolCallingModel model.ToolCallingChatModel + AgenticExecModel model.AgenticModel + OrchInstruction string + ToolsCfg adk.ToolsConfig + ExecMaxIter int + LoopMaxIter int + // AppCfg / Logger 非空时为 Executor 挂载与 Deep/Supervisor 一致的 Eino summarization 中间件。 + AppCfg *config.Config + MwCfg *config.MultiAgentEinoMiddlewareConfig + // ConversationID is used for transcript/isolation paths in middleware. + ConversationID string + DB *database.DB + ProjectID string + Logger *zap.Logger + // ModelName is used for model input token estimation logs. + ModelName string + // AgenticExecPreMiddlewares 是由 prependEinoAgenticMiddlewares 构建的前置中间件。 + AgenticExecPreMiddlewares []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] + AgenticSkillMiddleware adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] + AgenticFilesystemMiddleware adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] + // PlannerReplannerRewriteHandlers applies BeforeModelRewriteState pipeline for planner/replanner input. + PlannerReplannerRewriteHandlers []adk.ChatModelAgentMiddleware + // ModelFacingTrace 可选:由 Executor Handlers 链末尾写入,供 last_react 与 summarization 后上下文对齐。 + ModelFacingTrace *modelFacingTraceHolder + AgenticModelRetryConfig *adk.TypedModelRetryConfig[*schema.AgenticMessage] + AgenticModelFailoverConfig *adk.ModelFailoverConfig[*schema.AgenticMessage] +} + +// NewPlanExecuteRoot 返回 plan → execute → replan 预置编排根节点(与 Deep / Supervisor 并列)。 +func NewPlanExecuteRoot(ctx context.Context, a *PlanExecuteRootArgs) (adk.ResumableAgent, error) { + if a == nil { + return nil, fmt.Errorf("plan_execute: args 为空") + } + if a.MainToolCallingModel == nil || a.AgenticExecModel == nil { + return nil, fmt.Errorf("plan_execute: 模型为空") + } + tcm, ok := interface{}(a.MainToolCallingModel).(model.ToolCallingChatModel) + if !ok { + return nil, fmt.Errorf("plan_execute: 主模型需实现 ToolCallingChatModel") + } + plannerCfg := &planexecute.PlannerConfig{ + ToolCallingChatModel: tcm, + NewPlan: newLenientPlan, + } + if fn := planExecutePlannerGenInput(a.OrchInstruction, a.AppCfg, a.MwCfg, a.Logger, a.ModelName, a.ConversationID, a.PlannerReplannerRewriteHandlers); fn != nil { + plannerCfg.GenInputFn = fn + } + planner, err := planexecute.NewPlanner(ctx, plannerCfg) + if err != nil { + return nil, fmt.Errorf("plan_execute planner: %w", err) + } + replanner, err := planexecute.NewReplanner(ctx, &planexecute.ReplannerConfig{ + ChatModel: tcm, + GenInputFn: planExecuteReplannerGenInput(a.OrchInstruction, a.AppCfg, a.MwCfg, a.Logger, a.ModelName, a.ConversationID, a.PlannerReplannerRewriteHandlers), + NewPlan: newLenientPlan, + }) + if err != nil { + return nil, fmt.Errorf("plan_execute replanner: %w", err) + } + + var executor adk.Agent + agenticExecHandlers, herr := buildPlanExecuteAgenticExecutorHandlers(ctx, a) + if herr != nil { + return nil, herr + } + executor, err = newPlanExecuteAgenticExecutor(ctx, &planexecute.ExecutorConfig{ + ToolsConfig: a.ToolsCfg, + MaxIterations: a.ExecMaxIter, + GenInputFn: planExecuteExecutorGenInput(a.OrchInstruction, a.AppCfg, a.MwCfg, a.Logger, a.ModelName, a.ConversationID), + }, a.AgenticExecModel, agenticExecHandlers, a.AgenticModelRetryConfig, a.AgenticModelFailoverConfig) + if err != nil { + return nil, fmt.Errorf("plan_execute executor: %w", err) + } + loopMax := a.LoopMaxIter + if loopMax <= 0 { + loopMax = 10 + } + return planexecute.New(ctx, &planexecute.Config{ + Planner: planner, + Executor: executor, + Replanner: replanner, + MaxIterations: loopMax, + }) +} + +func buildPlanExecuteAgenticExecutorHandlers(ctx context.Context, a *PlanExecuteRootArgs) ([]adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], error) { + if a == nil { + return nil, fmt.Errorf("plan_execute: args 为空") + } + var execHandlers []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] + if len(a.AgenticExecPreMiddlewares) > 0 { + execHandlers = append(execHandlers, a.AgenticExecPreMiddlewares...) + } + if a.AgenticFilesystemMiddleware != nil { + execHandlers = append(execHandlers, a.AgenticFilesystemMiddleware) + } + if a.AgenticSkillMiddleware != nil { + execHandlers = append(execHandlers, a.AgenticSkillMiddleware) + } + if a.AppCfg != nil { + sumMw, sumErr := newEinoAgenticSummarizationMiddleware(ctx, a.AgenticExecModel, a.AppCfg, a.MwCfg, a.ConversationID, a.DB, a.ProjectID, a.Logger) + if sumErr != nil { + return nil, fmt.Errorf("plan_execute agentic executor summarization: %w", sumErr) + } + execHandlers = appendEinoAgenticChatModelTailMiddlewares(execHandlers, einoChatModelTailConfig{ + logger: a.Logger, + phase: "plan_execute_executor", + agenticSummarization: sumMw, + modelName: a.ModelName, + maxTotalTokens: a.AppCfg.OpenAI.MaxTotalTokens, + toolMaxBytes: toolMaxBytesFromMW(a.MwCfg), + conversationID: a.ConversationID, + trace: a.ModelFacingTrace, + middlewareConfig: a.MwCfg, + }) + } + return execHandlers, nil +} + +// planExecutePlannerGenInput 将 orchestrator instruction 作为 SystemMessage 注入 planner 输入。 +// 返回 nil 时 Eino 使用内置默认 planner prompt。 +func planExecutePlannerGenInput( + orchInstruction string, + appCfg *config.Config, + mwCfg *config.MultiAgentEinoMiddlewareConfig, + logger *zap.Logger, + modelName string, + conversationID string, + rewriteHandlers []adk.ChatModelAgentMiddleware, +) planexecute.GenPlannerModelInputFn { + oi := strings.TrimSpace(orchInstruction) + if oi == "" && appCfg == nil { + return nil + } + return func(ctx context.Context, userInput []adk.Message) ([]adk.Message, error) { + userInput = capPlanExecuteUserInputMessages(userInput, appCfg, mwCfg) + msgs := make([]adk.Message, 0, len(userInput)) + msgs = append(msgs, userInput...) + if rewritten, rerr := applyBeforeModelRewriteHandlers(ctx, msgs, rewriteHandlers); rerr == nil && len(rewritten) > 0 { + msgs = rewritten + } + msgs = normalizeSingleLeadingSystemMessage(msgs, oi) + logPlanExecuteModelInputEstimate(logger, modelName, conversationID, "plan_execute_planner", msgs) + return msgs, nil + } +} + +func planExecuteExecutorGenInput( + orchInstruction string, + appCfg *config.Config, + mwCfg *config.MultiAgentEinoMiddlewareConfig, + logger *zap.Logger, + modelName string, + conversationID string, +) planexecute.GenModelInputFn { + oi := strings.TrimSpace(orchInstruction) + return func(ctx context.Context, in *planexecute.ExecutionContext) ([]adk.Message, error) { + planContent, err := in.Plan.MarshalJSON() + if err != nil { + return nil, err + } + userMsgs, err := planexecute.ExecutorPrompt.Format(ctx, map[string]any{ + "input": planExecuteFormatInput(capPlanExecuteUserInputMessages(in.UserInput, appCfg, mwCfg)), + "plan": string(planContent), + "executed_steps": planExecuteFormatExecutedSteps(in.ExecutedSteps, appCfg, mwCfg), + "step": in.Plan.FirstStep(), + }) + if err != nil { + return nil, err + } + userMsgs = normalizeSingleLeadingSystemMessage(userMsgs, oi) + logPlanExecuteModelInputEstimate(logger, modelName, conversationID, "plan_execute_executor_gen_input", userMsgs) + return userMsgs, nil + } +} + +func planExecuteFormatInput(input []adk.Message) string { + var sb strings.Builder + for _, msg := range input { + sb.WriteString(msg.Content) + sb.WriteString("\n") + } + return sb.String() +} + +func planExecuteFormatExecutedSteps(results []planexecute.ExecutedStep, appCfg *config.Config, mwCfg *config.MultiAgentEinoMiddlewareConfig) string { + capped := capPlanExecuteExecutedStepsWithConfig(results, mwCfg) + return renderPlanExecuteStepsByBudget(capped, appCfg, mwCfg) +} + +// planExecuteReplannerGenInput 与 Eino 默认 Replanner 输入一致,但 executed_steps 经 cap 后再写入 prompt, +// 且在 orchInstruction 非空时 prepend SystemMessage 使 replanner 也能接收全局指令。 +func planExecuteReplannerGenInput( + orchInstruction string, + appCfg *config.Config, + mwCfg *config.MultiAgentEinoMiddlewareConfig, + logger *zap.Logger, + modelName string, + conversationID string, + rewriteHandlers []adk.ChatModelAgentMiddleware, +) planexecute.GenModelInputFn { + oi := strings.TrimSpace(orchInstruction) + return func(ctx context.Context, in *planexecute.ExecutionContext) ([]adk.Message, error) { + planContent, err := in.Plan.MarshalJSON() + if err != nil { + return nil, err + } + msgs, err := planexecute.ReplannerPrompt.Format(ctx, map[string]any{ + "plan": string(planContent), + "input": planExecuteFormatInput(capPlanExecuteUserInputMessages(in.UserInput, appCfg, mwCfg)), + "executed_steps": planExecuteFormatExecutedSteps(in.ExecutedSteps, appCfg, mwCfg), + "plan_tool": planexecute.PlanToolInfo.Name, + "respond_tool": planexecute.RespondToolInfo.Name, + }) + if err != nil { + return nil, err + } + if rewritten, rerr := applyBeforeModelRewriteHandlers(ctx, msgs, rewriteHandlers); rerr == nil && len(rewritten) > 0 { + msgs = rewritten + } + msgs = normalizeSingleLeadingSystemMessage(msgs, oi) + logPlanExecuteModelInputEstimate(logger, modelName, conversationID, "plan_execute_replanner", msgs) + return msgs, nil + } +} + +// normalizeSingleLeadingSystemMessage enforces a provider-friendly message shape: +// exactly one system message at index 0 (when any system context exists). +// For strict OpenAI-compatible backends (e.g. qwen/vllm templates), this avoids +// "System message must be at the beginning" caused by multiple/disordered system messages. +func normalizeSingleLeadingSystemMessage(msgs []adk.Message, extraSystem string) []adk.Message { + extraSystem = strings.TrimSpace(extraSystem) + if len(msgs) == 0 { + if extraSystem == "" { + return msgs + } + return []adk.Message{schema.SystemMessage(extraSystem)} + } + + systemParts := make([]string, 0, 2) + if extraSystem != "" { + systemParts = append(systemParts, extraSystem) + } + nonSystem := make([]adk.Message, 0, len(msgs)) + for _, msg := range msgs { + if msg == nil { + continue + } + if msg.Role == schema.System { + if s := strings.TrimSpace(msg.Content); s != "" { + systemParts = append(systemParts, s) + } + continue + } + nonSystem = append(nonSystem, msg) + } + if len(systemParts) == 0 { + return nonSystem + } + out := make([]adk.Message, 0, len(nonSystem)+1) + out = append(out, schema.SystemMessage(strings.Join(systemParts, "\n\n"))) + out = append(out, nonSystem...) + return out +} + +func capPlanExecuteUserInputMessages(input []adk.Message, appCfg *config.Config, mwCfg *config.MultiAgentEinoMiddlewareConfig) []adk.Message { + if len(input) == 0 { + return input + } + maxTotal := 120000 + modelName := "gpt-4o" + if appCfg != nil { + if appCfg.OpenAI.MaxTotalTokens > 0 { + maxTotal = appCfg.OpenAI.MaxTotalTokens + } + if m := strings.TrimSpace(appCfg.OpenAI.Model); m != "" { + modelName = m + } + } + // Reserve most tokens for planner/replanner prompt and tool schema. + ratio := 0.35 + if mwCfg != nil { + ratio = mwCfg.PlanExecuteUserInputBudgetRatioEffective() + } + budget := int(float64(maxTotal) * ratio) + if budget < 4096 { + budget = 4096 + } + tc := agent.NewTikTokenCounter() + out := make([]adk.Message, 0, len(input)) + used := 0 + for i := len(input) - 1; i >= 0; i-- { + msg := input[i] + if msg == nil { + continue + } + n, err := tc.Count(modelName, string(msg.Role)+"\n"+msg.Content) + if err != nil { + n = (len(msg.Content) + 3) / 4 + } + if n <= 0 { + n = 1 + } + if used+n > budget { + break + } + used += n + out = append(out, msg) + } + for i, j := 0, len(out)-1; i < j; i, j = i+1, j-1 { + out[i], out[j] = out[j], out[i] + } + if len(out) == 0 { + // Keep the latest user message at least. + return []adk.Message{input[len(input)-1]} + } + return out +} + +func renderPlanExecuteStepsByBudget(steps []planexecute.ExecutedStep, appCfg *config.Config, mwCfg *config.MultiAgentEinoMiddlewareConfig) string { + if len(steps) == 0 { + return "" + } + maxTotal := 120000 + modelName := "gpt-4o" + if appCfg != nil { + if appCfg.OpenAI.MaxTotalTokens > 0 { + maxTotal = appCfg.OpenAI.MaxTotalTokens + } + if m := strings.TrimSpace(appCfg.OpenAI.Model); m != "" { + modelName = m + } + } + ratio := 0.2 + if mwCfg != nil { + ratio = mwCfg.PlanExecuteExecutedStepsBudgetRatioEffective() + } + budget := int(float64(maxTotal) * ratio) + if budget < 3072 { + budget = 3072 + } + tc := agent.NewTikTokenCounter() + var kept []string + used := 0 + skipped := 0 + for i := len(steps) - 1; i >= 0; i-- { + block := fmt.Sprintf("Step: %s\nResult: %s\n\n", steps[i].Step, steps[i].Result) + n, err := tc.Count(modelName, block) + if err != nil { + n = (len(block) + 3) / 4 + } + if n <= 0 { + n = 1 + } + if used+n > budget { + skipped = i + 1 + break + } + used += n + kept = append(kept, block) + } + var sb strings.Builder + if skipped > 0 { + sb.WriteString(fmt.Sprintf("Earlier executed steps omitted due to context budget: %d steps.\n\n", skipped)) + } + for i := len(kept) - 1; i >= 0; i-- { + sb.WriteString(kept[i]) + } + return sb.String() +} + +// planExecuteStreamsMainAssistant 将规划/执行/重规划各阶段助手流式输出映射到主对话区。 +func planExecuteStreamsMainAssistant(agent string) bool { + if agent == "" { + return true + } + switch agent { + case "planner", "executor", "replanner", "execute_replan", "plan_execute_replan": + return true + default: + return false + } +} + +func planExecuteEinoRoleTag(agent string) string { + _ = agent + return "orchestrator" +} diff --git a/internal/multiagent/eino_orchestration_system_message_test.go b/internal/multiagent/eino_orchestration_system_message_test.go new file mode 100644 index 00000000..2cb32cfc --- /dev/null +++ b/internal/multiagent/eino_orchestration_system_message_test.go @@ -0,0 +1,45 @@ +package multiagent + +import ( + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestNormalizeSingleLeadingSystemMessage_MergesMultipleSystems(t *testing.T) { + in := []adk.Message{ + schema.SystemMessage("sys-1"), + schema.UserMessage("u1"), + schema.SystemMessage("sys-2"), + schema.AssistantMessage("a1", nil), + } + out := normalizeSingleLeadingSystemMessage(in, "orch") + if len(out) != 3 { + t.Fatalf("unexpected output length: got %d want 3", len(out)) + } + if out[0].Role != schema.System { + t.Fatalf("first message role must be system, got %s", out[0].Role) + } + if got := out[0].Content; got != "orch\n\nsys-1\n\nsys-2" { + t.Fatalf("unexpected merged system content: %q", got) + } + if out[1].Role != schema.User || out[2].Role != schema.Assistant { + t.Fatalf("non-system message order changed unexpectedly") + } +} + +func TestNormalizeSingleLeadingSystemMessage_NoSystemKeepsFlow(t *testing.T) { + in := []adk.Message{ + schema.UserMessage("u1"), + schema.AssistantMessage("a1", nil), + } + out := normalizeSingleLeadingSystemMessage(in, "") + if len(out) != 2 { + t.Fatalf("unexpected output length: got %d want 2", len(out)) + } + if out[0].Role != schema.User || out[1].Role != schema.Assistant { + t.Fatalf("message order changed unexpectedly") + } +} + diff --git a/internal/multiagent/eino_pending_tool_calls.go b/internal/multiagent/eino_pending_tool_calls.go new file mode 100644 index 00000000..2693ab93 --- /dev/null +++ b/internal/multiagent/eino_pending_tool_calls.go @@ -0,0 +1,125 @@ +package multiagent + +import ( + "fmt" + "strings" + "sync" +) + +type einoPendingToolCalls struct { + conversationID string + progress func(eventType, message string, data interface{}) + + mu sync.Mutex + byID map[string]toolCallPendingInfo + queueByAgent map[string][]string +} + +func newEinoPendingToolCalls(conversationID string, progress func(eventType, message string, data interface{})) *einoPendingToolCalls { + return &einoPendingToolCalls{ + conversationID: conversationID, + progress: progress, + byID: make(map[string]toolCallPendingInfo), + queueByAgent: make(map[string][]string), + } +} + +func (p *einoPendingToolCalls) Mark(tc toolCallPendingInfo) { + if p == nil || strings.TrimSpace(tc.ToolCallID) == "" { + return + } + p.mu.Lock() + defer p.mu.Unlock() + p.byID[tc.ToolCallID] = tc + p.queueByAgent[tc.EinoAgent] = append(p.queueByAgent[tc.EinoAgent], tc.ToolCallID) +} + +func (p *einoPendingToolCalls) PopNextForAgent(agentName string) (toolCallPendingInfo, bool) { + if p == nil { + return toolCallPendingInfo{}, false + } + p.mu.Lock() + defer p.mu.Unlock() + q := p.queueByAgent[agentName] + for len(q) > 0 { + id := q[0] + q = q[1:] + p.queueByAgent[agentName] = q + if tc, ok := p.byID[id]; ok { + delete(p.byID, id) + return tc, true + } + } + return toolCallPendingInfo{}, false +} + +func (p *einoPendingToolCalls) RemoveByID(toolCallID string) { + if p == nil || strings.TrimSpace(toolCallID) == "" { + return + } + p.mu.Lock() + defer p.mu.Unlock() + delete(p.byID, toolCallID) +} + +func (p *einoPendingToolCalls) PopAny() (toolCallPendingInfo, bool) { + if p == nil { + return toolCallPendingInfo{}, false + } + p.mu.Lock() + defer p.mu.Unlock() + for id, tc := range p.byID { + delete(p.byID, id) + return tc, true + } + return toolCallPendingInfo{}, false +} + +func (p *einoPendingToolCalls) Count() int { + if p == nil { + return 0 + } + p.mu.Lock() + defer p.mu.Unlock() + return len(p.byID) +} + +func (p *einoPendingToolCalls) FlushAsFailed(err error) { + if p == nil { + return + } + p.mu.Lock() + pendingSnapshot := make([]toolCallPendingInfo, 0, len(p.byID)) + for _, tc := range p.byID { + pendingSnapshot = append(pendingSnapshot, tc) + } + p.byID = make(map[string]toolCallPendingInfo) + p.queueByAgent = make(map[string][]string) + p.mu.Unlock() + + if p.progress == nil { + return + } + msg := "" + if err != nil { + msg = err.Error() + } + for _, tc := range pendingSnapshot { + toolName := tc.ToolName + if strings.TrimSpace(toolName) == "" { + toolName = "unknown" + } + p.progress("tool_result", fmt.Sprintf("工具结果 (%s)", toolName), map[string]interface{}{ + "toolName": toolName, + "success": false, + "isError": true, + "result": msg, + "resultPreview": msg, + "toolCallId": tc.ToolCallID, + "conversationId": p.conversationID, + "einoAgent": tc.EinoAgent, + "einoRole": tc.EinoRole, + "source": "eino", + }) + } +} diff --git a/internal/multiagent/eino_pending_tool_calls_test.go b/internal/multiagent/eino_pending_tool_calls_test.go new file mode 100644 index 00000000..c241cf11 --- /dev/null +++ b/internal/multiagent/eino_pending_tool_calls_test.go @@ -0,0 +1,81 @@ +package multiagent + +import ( + "errors" + "testing" +) + +func TestEinoPendingToolCallsPopNextForAgentSkipsRemovedIDs(t *testing.T) { + p := newEinoPendingToolCalls("conv", nil) + p.Mark(toolCallPendingInfo{ToolCallID: "call-1", ToolName: "first", EinoAgent: "agent"}) + p.Mark(toolCallPendingInfo{ToolCallID: "call-2", ToolName: "second", EinoAgent: "agent"}) + p.RemoveByID("call-1") + + got, ok := p.PopNextForAgent("agent") + if !ok { + t.Fatal("expected pending tool call") + } + if got.ToolCallID != "call-2" { + t.Fatalf("toolCallID = %q, want call-2", got.ToolCallID) + } + if p.Count() != 0 { + t.Fatalf("pending count = %d, want 0", p.Count()) + } +} + +func TestEinoPendingToolCallsPopAny(t *testing.T) { + p := newEinoPendingToolCalls("conv", nil) + p.Mark(toolCallPendingInfo{ToolCallID: "call-1", ToolName: "tool"}) + + got, ok := p.PopAny() + if !ok || got.ToolCallID != "call-1" { + t.Fatalf("PopAny = %#v ok=%v", got, ok) + } + if _, ok := p.PopAny(); ok { + t.Fatal("PopAny should be empty after first pop") + } +} + +func TestEinoPendingToolCallsFlushAsFailedEmitsAndClears(t *testing.T) { + var events []struct { + eventType string + message string + data map[string]interface{} + } + p := newEinoPendingToolCalls("conv-1", func(eventType, message string, data interface{}) { + m, _ := data.(map[string]interface{}) + events = append(events, struct { + eventType string + message string + data map[string]interface{} + }{eventType: eventType, message: message, data: m}) + }) + p.Mark(toolCallPendingInfo{ + ToolCallID: "call-err", + ToolName: "", + EinoAgent: "agent", + EinoRole: "sub", + }) + + p.FlushAsFailed(errors.New("boom")) + + if p.Count() != 0 { + t.Fatalf("pending count = %d, want 0", p.Count()) + } + if len(events) != 1 { + t.Fatalf("events = %#v, want one", events) + } + ev := events[0] + if ev.eventType != "tool_result" || ev.message != "工具结果 (unknown)" { + t.Fatalf("event = %#v", ev) + } + if ev.data["toolCallId"] != "call-err" || + ev.data["conversationId"] != "conv-1" || + ev.data["einoAgent"] != "agent" || + ev.data["einoRole"] != "sub" || + ev.data["isError"] != true || + ev.data["success"] != false || + ev.data["result"] != "boom" { + t.Fatalf("payload = %#v", ev.data) + } +} diff --git a/internal/multiagent/eino_reasoning_stream_emitter.go b/internal/multiagent/eino_reasoning_stream_emitter.go new file mode 100644 index 00000000..e77afa6c --- /dev/null +++ b/internal/multiagent/eino_reasoning_stream_emitter.go @@ -0,0 +1,111 @@ +package multiagent + +import ( + "strings" + + "cyberstrike-ai/internal/openai" +) + +type einoReasoningStreamEmitter struct { + progress func(eventType, message string, data interface{}) + conversation string + orchMode string + agentName string + einoRole string + nextStreamID func() string + + streamID string + rawBuf string + displayPrev string +} + +func newEinoReasoningStreamEmitter( + conversationID, orchMode, agentName, einoRole string, + progress func(eventType, message string, data interface{}), + nextStreamID func() string, +) *einoReasoningStreamEmitter { + return &einoReasoningStreamEmitter{ + progress: progress, + conversation: conversationID, + orchMode: orchMode, + agentName: agentName, + einoRole: einoRole, + nextStreamID: nextStreamID, + } +} + +func (e *einoReasoningStreamEmitter) EmitDelta(reasoningContent string) bool { + if e == nil || strings.TrimSpace(reasoningContent) == "" { + return false + } + var rawDelta string + e.rawBuf, rawDelta = normalizeStreamingDelta(e.rawBuf, reasoningContent) + if rawDelta == "" || e.progress == nil { + return false + } + fullDisplay := openai.DisplayReasoningContent(e.rawBuf) + displayDelta := fullDisplay + if strings.HasPrefix(fullDisplay, e.displayPrev) { + displayDelta = fullDisplay[len(e.displayPrev):] + } + e.displayPrev = fullDisplay + if displayDelta == "" { + return false + } + if e.streamID == "" { + if e.nextStreamID != nil { + e.streamID = e.nextStreamID() + } + if e.streamID == "" { + e.streamID = "eino-reasoning" + } + e.progress("reasoning_chain_stream_start", " ", map[string]interface{}{ + "streamId": e.streamID, + "source": "eino", + "einoAgent": e.agentName, + "einoRole": e.einoRole, + "orchestration": e.orchMode, + }) + } + e.progress("reasoning_chain_stream_delta", displayDelta, openai.WithSSEAccumulated(map[string]interface{}{ + "streamId": e.streamID, + }, fullDisplay)) + return true +} + +func (e *einoReasoningStreamEmitter) Finish() string { + if e == nil { + return "" + } + display := openai.DisplayReasoningContent(strings.TrimSpace(e.rawBuf)) + if display == "" || e.streamID == "" || e.progress == nil { + return display + } + e.progress("reasoning_chain_stream_end", display, map[string]interface{}{ + "streamId": e.streamID, + "conversationId": e.conversation, + "source": "eino", + "einoAgent": e.agentName, + "einoRole": e.einoRole, + "orchestration": e.orchMode, + }) + return display +} + +func (e *einoReasoningStreamEmitter) EmitComplete(reasoningContent string) bool { + if e == nil || e.progress == nil { + return false + } + display := openai.DisplayReasoningContent(strings.TrimSpace(reasoningContent)) + if display == "" { + return false + } + e.progress("reasoning_chain", display, map[string]interface{}{ + "conversationId": e.conversation, + "source": "eino", + "einoAgent": e.agentName, + "einoRole": e.einoRole, + "orchestration": e.orchMode, + }) + return true +} diff --git a/internal/multiagent/eino_run_cancellation_handler.go b/internal/multiagent/eino_run_cancellation_handler.go new file mode 100644 index 00000000..d4e6cec6 --- /dev/null +++ b/internal/multiagent/eino_run_cancellation_handler.go @@ -0,0 +1,56 @@ +package multiagent + +import "context" + +type einoRunCancellationHandler struct { + ctx context.Context + conversationID string + progress func(eventType, message string, data interface{}) + pending *einoPendingToolCalls + takePartial einoPartialResultFunc +} + +type einoRunCancellationHandlerConfig struct { + Context context.Context + ConversationID string + Progress func(eventType, message string, data interface{}) + Pending *einoPendingToolCalls + TakePartial einoPartialResultFunc +} + +func newEinoRunCancellationHandler(cfg einoRunCancellationHandlerConfig) *einoRunCancellationHandler { + return &einoRunCancellationHandler{ + ctx: cfg.Context, + conversationID: cfg.ConversationID, + progress: cfg.Progress, + pending: cfg.Pending, + takePartial: cfg.TakePartial, + } +} + +func (h *einoRunCancellationHandler) Handle(runErr error) (*RunResult, error) { + if h == nil { + return nil, runErr + } + if h.pending != nil { + h.pending.FlushAsFailed(runErr) + } + if h.progress != nil { + if isInterruptContinue(h.ctx) { + h.progress("progress", "已暂停当前输出,正在合并用户补充并继续…", map[string]interface{}{ + "conversationId": h.conversationID, + "source": "eino", + "kind": "interrupt_continue", + }) + } else if runErr != nil { + h.progress("error", runErr.Error(), map[string]interface{}{ + "conversationId": h.conversationID, + "source": "eino", + }) + } + } + if h.takePartial == nil { + return nil, runErr + } + return h.takePartial(runErr) +} diff --git a/internal/multiagent/eino_run_cancellation_handler_test.go b/internal/multiagent/eino_run_cancellation_handler_test.go new file mode 100644 index 00000000..406aaef0 --- /dev/null +++ b/internal/multiagent/eino_run_cancellation_handler_test.go @@ -0,0 +1,97 @@ +package multiagent + +import ( + "context" + "errors" + "testing" +) + +func TestEinoRunCancellationHandlerFlushesPendingAndEmitsError(t *testing.T) { + runErr := errors.New("context canceled") + var events []struct { + eventType string + data map[string]interface{} + } + progress := func(eventType, _ string, data interface{}) { + m, _ := data.(map[string]interface{}) + events = append(events, struct { + eventType string + data map[string]interface{} + }{eventType: eventType, data: m}) + } + pending := newEinoPendingToolCalls("conv-1", progress) + pending.Mark(toolCallPendingInfo{ToolCallID: "call-1", ToolName: "execute", EinoAgent: "lead", EinoRole: "orchestrator"}) + want := &RunResult{Response: "partial"} + + result, err := newEinoRunCancellationHandler(einoRunCancellationHandlerConfig{ + Context: context.Background(), + ConversationID: "conv-1", + Progress: progress, + Pending: pending, + TakePartial: func(got error) (*RunResult, error) { + if !errors.Is(got, runErr) { + t.Fatalf("partial err = %v", got) + } + return want, got + }, + }).Handle(runErr) + + if result != want || !errors.Is(err, runErr) { + t.Fatalf("result=%#v err=%v", result, err) + } + if pending.Count() != 0 { + t.Fatalf("pending count = %d, want 0", pending.Count()) + } + var sawError, sawFailedTool bool + for _, ev := range events { + if ev.eventType == "error" { + sawError = ev.data["conversationId"] == "conv-1" && ev.data["source"] == "eino" + } + if ev.eventType == "tool_result" { + sawFailedTool = ev.data["toolCallId"] == "call-1" && ev.data["isError"] == true + } + } + if !sawError || !sawFailedTool { + t.Fatalf("events = %#v", events) + } +} + +func TestEinoRunCancellationHandlerInterruptContinueProgress(t *testing.T) { + ctx, cancel := context.WithCancelCause(context.Background()) + cancel(ErrInterruptContinue) + runErr := context.Canceled + var eventType string + var data map[string]interface{} + + _, err := newEinoRunCancellationHandler(einoRunCancellationHandlerConfig{ + Context: ctx, + ConversationID: "conv-1", + Progress: func(et, _ string, raw interface{}) { + eventType = et + data, _ = raw.(map[string]interface{}) + }, + TakePartial: func(got error) (*RunResult, error) { + return nil, got + }, + }).Handle(runErr) + + if !errors.Is(err, runErr) { + t.Fatalf("err = %v", err) + } + if eventType != "progress" || data["kind"] != "interrupt_continue" || data["conversationId"] != "conv-1" { + t.Fatalf("eventType=%q data=%#v", eventType, data) + } +} + +func TestEinoRunCancellationHandlerNilSafe(t *testing.T) { + runErr := errors.New("boom") + var h *einoRunCancellationHandler + result, err := h.Handle(runErr) + if result != nil || !errors.Is(err, runErr) { + t.Fatalf("result=%#v err=%v", result, err) + } + result, err = newEinoRunCancellationHandler(einoRunCancellationHandlerConfig{}).Handle(runErr) + if result != nil || !errors.Is(err, runErr) { + t.Fatalf("result=%#v err=%v", result, err) + } +} diff --git a/internal/multiagent/eino_run_completion_handler.go b/internal/multiagent/eino_run_completion_handler.go new file mode 100644 index 00000000..acda966d --- /dev/null +++ b/internal/multiagent/eino_run_completion_handler.go @@ -0,0 +1,81 @@ +package multiagent + +import ( + "errors" + "os" + + "go.uber.org/zap" +) + +type einoRunCompletionHandler struct { + conversationID string + orchMode string + progress func(eventType, message string, data interface{}) + logger *zap.Logger + + pending *einoPendingToolCalls + cpStore *fileCheckPointStore + checkPointID string +} + +type einoRunCompletionHandlerConfig struct { + ConversationID string + OrchMode string + Progress func(eventType, message string, data interface{}) + Logger *zap.Logger + Pending *einoPendingToolCalls + Checkpoint *fileCheckPointStore + CheckpointID string +} + +func newEinoRunCompletionHandler(cfg einoRunCompletionHandlerConfig) *einoRunCompletionHandler { + return &einoRunCompletionHandler{ + conversationID: cfg.ConversationID, + orchMode: cfg.OrchMode, + progress: cfg.Progress, + logger: cfg.Logger, + pending: cfg.Pending, + cpStore: cfg.Checkpoint, + checkPointID: cfg.CheckpointID, + } +} + +func (h *einoRunCompletionHandler) Complete() { + if h == nil { + return + } + h.flushOrphanedPending() + h.cleanupCheckpoint() +} + +func (h *einoRunCompletionHandler) flushOrphanedPending() { + if h.pending == nil { + return + } + orphanCount := h.pending.Count() + if orphanCount <= 0 { + return + } + h.pending.FlushAsFailed(errors.New("pending tool call missing result before run completion")) + if h.progress != nil { + h.progress("eino_pending_orphaned", "pending tool calls were force-closed at run end", map[string]interface{}{ + "conversationId": h.conversationID, + "source": "eino", + "orchestration": h.orchMode, + "pendingCount": orphanCount, + }) + } +} + +func (h *einoRunCompletionHandler) cleanupCheckpoint() { + if h.cpStore == nil || h.checkPointID == "" { + return + } + p, err := h.cpStore.path(h.checkPointID) + if err != nil { + return + } + if rmErr := os.Remove(p); rmErr != nil && !os.IsNotExist(rmErr) && h.logger != nil { + h.logger.Warn("eino checkpoint cleanup failed", zap.String("path", p), zap.Error(rmErr)) + } +} diff --git a/internal/multiagent/eino_run_completion_handler_test.go b/internal/multiagent/eino_run_completion_handler_test.go new file mode 100644 index 00000000..289205fa --- /dev/null +++ b/internal/multiagent/eino_run_completion_handler_test.go @@ -0,0 +1,72 @@ +package multiagent + +import ( + "context" + "os" + "testing" +) + +func TestEinoRunCompletionHandlerFlushesOrphansAndCleansCheckpoint(t *testing.T) { + var events []struct { + eventType string + data map[string]interface{} + } + progress := func(eventType, _ string, data interface{}) { + m, _ := data.(map[string]interface{}) + events = append(events, struct { + eventType string + data map[string]interface{} + }{eventType: eventType, data: m}) + } + pending := newEinoPendingToolCalls("conv-1", progress) + pending.Mark(toolCallPendingInfo{ToolCallID: "call-1", ToolName: "execute", EinoAgent: "lead", EinoRole: "orchestrator"}) + store, err := newFileCheckPointStore(t.TempDir()) + if err != nil { + t.Fatal(err) + } + if err := store.Set(context.Background(), "cp-1", []byte("checkpoint")); err != nil { + t.Fatal(err) + } + cpPath, err := store.path("cp-1") + if err != nil { + t.Fatal(err) + } + + newEinoRunCompletionHandler(einoRunCompletionHandlerConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + Progress: progress, + Pending: pending, + Checkpoint: store, + CheckpointID: "cp-1", + }).Complete() + + if pending.Count() != 0 { + t.Fatalf("pending count = %d, want 0", pending.Count()) + } + if _, err := os.Stat(cpPath); !os.IsNotExist(err) { + t.Fatalf("checkpoint should be removed, stat err=%v", err) + } + var orphanEvent map[string]interface{} + var failedToolResult map[string]interface{} + for _, ev := range events { + switch ev.eventType { + case "eino_pending_orphaned": + orphanEvent = ev.data + case "tool_result": + failedToolResult = ev.data + } + } + if orphanEvent == nil || orphanEvent["conversationId"] != "conv-1" || orphanEvent["orchestration"] != "deep" || orphanEvent["pendingCount"] != 1 { + t.Fatalf("orphan event = %#v", orphanEvent) + } + if failedToolResult == nil || failedToolResult["toolCallId"] != "call-1" || failedToolResult["isError"] != true { + t.Fatalf("failed tool result = %#v", failedToolResult) + } +} + +func TestEinoRunCompletionHandlerNoopWithoutState(t *testing.T) { + newEinoRunCompletionHandler(einoRunCompletionHandlerConfig{}).Complete() + var h *einoRunCompletionHandler + h.Complete() +} diff --git a/internal/multiagent/eino_run_error_handler.go b/internal/multiagent/eino_run_error_handler.go new file mode 100644 index 00000000..0e816c8b --- /dev/null +++ b/internal/multiagent/eino_run_error_handler.go @@ -0,0 +1,93 @@ +package multiagent + +import ( + "context" + "errors" + + "github.com/cloudwego/eino/adk" +) + +type einoRunErrorHandler struct { + conversationID string + orchMode string + progress func(eventType, message string, data interface{}) + pending *einoPendingToolCalls + nativeCancelFallback func() error +} + +type einoRunErrorHandlerConfig struct { + ConversationID string + OrchMode string + Progress func(eventType, message string, data interface{}) + Pending *einoPendingToolCalls + NativeCancelFallback func() error +} + +func newEinoRunErrorHandler(cfg einoRunErrorHandlerConfig) *einoRunErrorHandler { + return &einoRunErrorHandler{ + conversationID: cfg.ConversationID, + orchMode: cfg.OrchMode, + progress: cfg.Progress, + pending: cfg.Pending, + nativeCancelFallback: cfg.NativeCancelFallback, + } +} + +func (h *einoRunErrorHandler) Handle(runErr error) error { + if h == nil || runErr == nil { + return runErr + } + var cancelErr *adk.CancelError + if errors.As(runErr, &cancelErr) { + h.flushPending(runErr) + if h.nativeCancelFallback != nil { + return h.nativeCancelFallback() + } + return context.Canceled + } + if errors.Is(runErr, context.DeadlineExceeded) { + h.flushPending(runErr) + h.emitError(runErr, "timeout") + return runErr + } + if errors.Is(runErr, context.Canceled) { + h.flushPending(runErr) + h.emitError(runErr, "") + return runErr + } + if isEinoIterationLimitError(runErr) { + h.flushPending(runErr) + if h.progress != nil { + h.progress("iteration_limit_reached", runErr.Error(), map[string]interface{}{ + "conversationId": h.conversationID, + "source": "eino", + "orchestration": h.orchMode, + }) + } + h.emitError(runErr, "iteration_limit") + return runErr + } + h.flushPending(runErr) + h.emitError(runErr, "") + return runErr +} + +func (h *einoRunErrorHandler) flushPending(err error) { + if h != nil && h.pending != nil { + h.pending.FlushAsFailed(err) + } +} + +func (h *einoRunErrorHandler) emitError(err error, kind string) { + if h == nil || h.progress == nil || err == nil { + return + } + data := map[string]interface{}{ + "conversationId": h.conversationID, + "source": "eino", + } + if kind != "" { + data["errorKind"] = kind + } + h.progress("error", err.Error(), data) +} diff --git a/internal/multiagent/eino_run_error_handler_test.go b/internal/multiagent/eino_run_error_handler_test.go new file mode 100644 index 00000000..c310af64 --- /dev/null +++ b/internal/multiagent/eino_run_error_handler_test.go @@ -0,0 +1,101 @@ +package multiagent + +import ( + "context" + "errors" + "testing" + + "github.com/cloudwego/eino/adk" +) + +func TestEinoRunErrorHandlerCancelUsesNativeFallback(t *testing.T) { + pending := newEinoPendingToolCalls("conv-1", nil) + pending.Mark(toolCallPendingInfo{ToolCallID: "call-1", ToolName: "execute"}) + want := errors.New("native cancel") + + got := newEinoRunErrorHandler(einoRunErrorHandlerConfig{ + ConversationID: "conv-1", + Pending: pending, + NativeCancelFallback: func() error { + return want + }, + }).Handle(&adk.CancelError{Info: &adk.AgentCancelInfo{}}) + + if !errors.Is(got, want) { + t.Fatalf("err = %v, want native fallback", got) + } + if pending.Count() != 0 { + t.Fatalf("pending count = %d, want 0", pending.Count()) + } +} + +func TestEinoRunErrorHandlerTimeoutAndGeneralErrorProgress(t *testing.T) { + for _, tc := range []struct { + name string + err error + errorKind interface{} + }{ + {name: "timeout", err: context.DeadlineExceeded, errorKind: "timeout"}, + {name: "general", err: errors.New("boom"), errorKind: nil}, + } { + t.Run(tc.name, func(t *testing.T) { + var data map[string]interface{} + got := newEinoRunErrorHandler(einoRunErrorHandlerConfig{ + ConversationID: "conv-1", + Progress: func(eventType, _ string, raw interface{}) { + if eventType == "error" { + data, _ = raw.(map[string]interface{}) + } + }, + }).Handle(tc.err) + if !errors.Is(got, tc.err) { + t.Fatalf("err = %v", got) + } + if data["conversationId"] != "conv-1" || data["source"] != "eino" { + t.Fatalf("data = %#v", data) + } + if gotKind := data["errorKind"]; gotKind != tc.errorKind { + t.Fatalf("errorKind = %#v, want %#v", gotKind, tc.errorKind) + } + }) + } +} + +func TestEinoRunErrorHandlerIterationLimitProgress(t *testing.T) { + var events []string + var errorKind interface{} + err := errors.New("maximum iteration reached") + + got := newEinoRunErrorHandler(einoRunErrorHandlerConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + Progress: func(eventType, _ string, raw interface{}) { + events = append(events, eventType) + if eventType == "error" { + data, _ := raw.(map[string]interface{}) + errorKind = data["errorKind"] + } + }, + }).Handle(err) + + if !errors.Is(got, err) { + t.Fatalf("err = %v", got) + } + if len(events) != 2 || events[0] != "iteration_limit_reached" || events[1] != "error" { + t.Fatalf("events = %#v", events) + } + if errorKind != "iteration_limit" { + t.Fatalf("errorKind = %#v", errorKind) + } +} + +func TestEinoRunErrorHandlerNilSafe(t *testing.T) { + var h *einoRunErrorHandler + if h.Handle(nil) != nil { + t.Fatal("nil handler nil err should return nil") + } + err := errors.New("boom") + if got := h.Handle(err); !errors.Is(got, err) { + t.Fatalf("nil handler err = %v", got) + } +} diff --git a/internal/multiagent/eino_run_event_drain.go b/internal/multiagent/eino_run_event_drain.go new file mode 100644 index 00000000..5333caf7 --- /dev/null +++ b/internal/multiagent/eino_run_event_drain.go @@ -0,0 +1,244 @@ +package multiagent + +import ( + "context" + "fmt" + "sync/atomic" + + "cyberstrike-ai/internal/agent" + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/einomcp" + + "github.com/cloudwego/eino/adk" + "go.uber.org/zap" +) + +type einoRunEventDrainConfig struct { + Context context.Context + ConversationID string + OrchMode string + OrchestratorName string + Progress func(eventType, message string, data interface{}) + Logger *zap.Logger + BaseMessages []adk.Message + SnapshotMCPIDs func() []string + StreamsMainAssistant func(agent string) bool + EinoRoleTag func(agent string) string + MiddlewareConfig *config.MultiAgentEinoMiddlewareConfig + + FilesystemMonitorAgent *agent.Agent + FilesystemMonitorRecord einomcp.ExecutionRecorder + MCPExecutionBinder *MCPExecutionBinder +} + +type einoRunEventDrain struct { + cfg einoRunEventDrainConfig + + runMessages *einoRunMessageAccumulator + assistantOutput *einoAssistantOutputAccumulator + runProgress *einoRunProgressTracker + pendingToolCalls *einoPendingToolCalls + stdoutSuppressor *einoExecuteStdoutSuppressor + toolResultEmitter *einoToolResultProgressEmitter + usage *einoRunUsageAccumulator + + reasoningStreamSeq int64 + subReplyStreamSeq int64 + mainResponseStreamSeq int64 + + toolResultHandler *einoToolResultEventHandler + assistantStreamHandler *einoAssistantStreamEventHandler + materializedMessageHandler *einoMaterializedMessageEventHandler +} + +func newEinoRunEventDrain(cfg einoRunEventDrainConfig) *einoRunEventDrain { + if cfg.Context == nil { + cfg.Context = context.Background() + } + if cfg.SnapshotMCPIDs == nil { + cfg.SnapshotMCPIDs = func() []string { return nil } + } + if cfg.StreamsMainAssistant == nil { + cfg.StreamsMainAssistant = func(agentName string) bool { + return agentName == "" || agentName == cfg.OrchestratorName + } + } + if cfg.EinoRoleTag == nil { + cfg.EinoRoleTag = func(agentName string) string { + if cfg.StreamsMainAssistant(agentName) { + return "orchestrator" + } + return "sub" + } + } + + runMessages := newEinoRunMessageAccumulator(cfg.BaseMessages) + assistantOutput := newEinoAssistantOutputAccumulator(cfg.OrchMode) + runProgress := newEinoRunProgressTracker( + cfg.OrchMode, + cfg.OrchestratorName, + cfg.ConversationID, + cfg.Progress, + cfg.StreamsMainAssistant, + cfg.EinoRoleTag, + ) + pendingToolCalls := newEinoPendingToolCalls(cfg.ConversationID, cfg.Progress) + stdoutSuppressor := newEinoExecuteStdoutSuppressor() + usage := newEinoRunUsageAccumulator() + toolResultEmitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{ + ConversationID: cfg.ConversationID, + OrchestratorName: cfg.OrchestratorName, + Progress: cfg.Progress, + EinoRoleTag: cfg.EinoRoleTag, + Pending: pendingToolCalls, + ExecuteStdoutDup: stdoutSuppressor, + RunMessages: runMessages, + FilesystemMonitorAgent: cfg.FilesystemMonitorAgent, + FilesystemMonitorRecord: cfg.FilesystemMonitorRecord, + MCPExecutionBinder: cfg.MCPExecutionBinder, + }) + + return &einoRunEventDrain{ + cfg: cfg, + runMessages: runMessages, + assistantOutput: assistantOutput, + runProgress: runProgress, + pendingToolCalls: pendingToolCalls, + stdoutSuppressor: stdoutSuppressor, + toolResultEmitter: toolResultEmitter, + usage: usage, + } +} + +func (d *einoRunEventDrain) BindHandlers(confirmRecovery func()) { + if d == nil { + return + } + d.toolResultHandler = newEinoToolResultEventHandler(einoToolResultEventHandlerConfig{ + Context: d.cfg.Context, + Logger: d.cfg.Logger, + RunMessages: d.runMessages, + Emitter: d.toolResultEmitter, + ConfirmRecovery: confirmRecovery, + }) + streamToolCallCompletion := newEinoStreamToolCallCompletionHandler(einoStreamToolCallCompletionHandlerConfig{ + ConversationID: d.cfg.ConversationID, + OrchMode: d.cfg.OrchMode, + Progress: d.cfg.Progress, + RunProgress: d.runProgress, + RunMessages: d.runMessages, + MarkPending: d.markPendingWithMonitor, + }) + d.assistantStreamHandler = newEinoAssistantStreamEventHandler(einoAssistantStreamEventHandlerConfig{ + Context: d.cfg.Context, + ConversationID: d.cfg.ConversationID, + OrchMode: d.cfg.OrchMode, + Progress: d.cfg.Progress, + Logger: d.cfg.Logger, + SnapshotMCPIDs: d.cfg.SnapshotMCPIDs, + StreamsMainAssistant: d.cfg.StreamsMainAssistant, + EinoRoleTag: d.cfg.EinoRoleTag, + RunProgress: d.runProgress, + StdoutSuppressor: d.stdoutSuppressor, + AssistantOutput: d.assistantOutput, + RunMessages: d.runMessages, + Usage: d.usage, + ToolCallCompletion: streamToolCallCompletion, + NextMainStreamID: d.nextMainStreamID, + NextReasoningStreamID: d.nextReasoningStreamID, + NextSubAgentReplyStreamID: d.nextSubAgentReplyStreamID, + }) + d.materializedMessageHandler = newEinoMaterializedMessageEventHandler(einoMaterializedMessageEventHandlerConfig{ + ConversationID: d.cfg.ConversationID, + OrchMode: d.cfg.OrchMode, + Progress: d.cfg.Progress, + SnapshotMCPIDs: d.cfg.SnapshotMCPIDs, + StreamsMainAssistant: d.cfg.StreamsMainAssistant, + EinoRoleTag: d.cfg.EinoRoleTag, + RunProgress: d.runProgress, + StdoutSuppressor: d.stdoutSuppressor, + AssistantOutput: d.assistantOutput, + RunMessages: d.runMessages, + Usage: d.usage, + ToolResultHandler: d.toolResultHandler, + MarkPending: d.markPendingWithMonitor, + NextMainStreamID: d.nextMainStreamID, + }) +} + +func (d *einoRunEventDrain) RunMessages() *einoRunMessageAccumulator { + if d == nil { + return nil + } + return d.runMessages +} + +func (d *einoRunEventDrain) AssistantOutput() *einoAssistantOutputAccumulator { + if d == nil { + return nil + } + return d.assistantOutput +} + +func (d *einoRunEventDrain) PendingToolCalls() *einoPendingToolCalls { + if d == nil { + return nil + } + return d.pendingToolCalls +} + +func (d *einoRunEventDrain) Usage() *einoRunUsageAccumulator { + if d == nil { + return nil + } + return d.usage +} + +func (d *einoRunEventDrain) ObserveAgent(agentName string) { + if d == nil || d.runProgress == nil { + return + } + d.runProgress.ObserveAgent(agentName) +} + +func (d *einoRunEventDrain) HandleToolResultStreaming(mv *adk.MessageVariant, agentName string) bool { + return d != nil && d.toolResultHandler != nil && d.toolResultHandler.HandleStreaming(mv, agentName) +} + +func (d *einoRunEventDrain) HandleAssistantStream(mv *adk.MessageVariant, agentName string) (bool, error) { + if d == nil || d.assistantStreamHandler == nil { + return false, nil + } + return d.assistantStreamHandler.Handle(mv, agentName) +} + +func (d *einoRunEventDrain) HandleMaterialized(mv *adk.MessageVariant, msg adk.Message, agentName string) bool { + return d != nil && d.materializedMessageHandler != nil && d.materializedMessageHandler.Handle(mv, msg, agentName) +} + +func (d *einoRunEventDrain) markPendingWithMonitor(tc toolCallPendingInfo) { + if d == nil || d.pendingToolCalls == nil { + return + } + d.pendingToolCalls.Mark(tc) + beginEinoADKFilesystemToolMonitor( + d.cfg.Context, + d.cfg.FilesystemMonitorAgent, + d.cfg.FilesystemMonitorRecord, + d.cfg.MCPExecutionBinder, + tc.ToolCallID, + tc.ToolName, + ) +} + +func (d *einoRunEventDrain) nextMainStreamID() string { + return fmt.Sprintf("eino-main-%s-%d", d.cfg.ConversationID, atomic.AddInt64(&d.mainResponseStreamSeq, 1)) +} + +func (d *einoRunEventDrain) nextReasoningStreamID() string { + return fmt.Sprintf("eino-reasoning-%s-%d", d.cfg.ConversationID, atomic.AddInt64(&d.reasoningStreamSeq, 1)) +} + +func (d *einoRunEventDrain) nextSubAgentReplyStreamID() string { + return fmt.Sprintf("eino-sub-reply-%s-%d", d.cfg.ConversationID, atomic.AddInt64(&d.subReplyStreamSeq, 1)) +} diff --git a/internal/multiagent/eino_run_event_drain_test.go b/internal/multiagent/eino_run_event_drain_test.go new file mode 100644 index 00000000..293b1b97 --- /dev/null +++ b/internal/multiagent/eino_run_event_drain_test.go @@ -0,0 +1,79 @@ +package multiagent + +import ( + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestEinoRunEventDrainDefaultsAndStreamIDs(t *testing.T) { + drain := newEinoRunEventDrain(einoRunEventDrainConfig{ + ConversationID: "conv-1", + OrchestratorName: "lead", + BaseMessages: []adk.Message{schema.UserMessage("base")}, + }) + + if drain.RunMessages().BaseCount() != 1 { + t.Fatalf("base count = %d, want 1", drain.RunMessages().BaseCount()) + } + if !drain.cfg.StreamsMainAssistant("lead") || drain.cfg.StreamsMainAssistant("worker") { + t.Fatal("default main-assistant predicate should match only orchestrator") + } + if got := drain.cfg.EinoRoleTag("lead"); got != "orchestrator" { + t.Fatalf("lead role = %q, want orchestrator", got) + } + if got := drain.cfg.EinoRoleTag("worker"); got != "sub" { + t.Fatalf("worker role = %q, want sub", got) + } + if got := drain.nextMainStreamID(); got != "eino-main-conv-1-1" { + t.Fatalf("first main stream id = %q", got) + } + if got := drain.nextMainStreamID(); got != "eino-main-conv-1-2" { + t.Fatalf("second main stream id = %q", got) + } +} + +func TestEinoRunEventDrainBindsHandlersAndRecordsEvents(t *testing.T) { + var events []string + recovered := false + drain := newEinoRunEventDrain(einoRunEventDrainConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + Progress: func(eventType, _ string, _ interface{}) { + events = append(events, eventType) + }, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + }) + drain.BindHandlers(func() { recovered = true }) + + drain.ObserveAgent("lead") + if !drain.HandleMaterialized(&adk.MessageVariant{Role: schema.Assistant}, schema.AssistantMessage("done", nil), "lead") { + t.Fatal("materialized assistant should be handled") + } + if got := drain.AssistantOutput().LastAssistant(); got != "done" { + t.Fatalf("last assistant = %q, want done", got) + } + + stream := schema.StreamReaderFromArray([]*schema.Message{ + {Role: schema.Tool, Content: "ok", ToolCallID: "call-1"}, + }) + if !drain.HandleToolResultStreaming(&adk.MessageVariant{ + IsStreaming: true, + Role: schema.Tool, + ToolName: "execute", + MessageStream: stream, + }, "lead") { + t.Fatal("streaming tool result should be handled") + } + if !recovered { + t.Fatal("tool stream completion should confirm recovery") + } + if len(drain.RunMessages().Messages()) != 3 { + t.Fatalf("run messages = %#v, want base + assistant + tool", drain.RunMessages().Messages()) + } + if !containsString(events, "iteration") || !containsString(events, "response_start") || !containsString(events, "tool_result") { + t.Fatalf("events = %#v, want iteration, response and tool_result", events) + } +} diff --git a/internal/multiagent/eino_run_message_accumulator.go b/internal/multiagent/eino_run_message_accumulator.go new file mode 100644 index 00000000..2f4f7fbb --- /dev/null +++ b/internal/multiagent/eino_run_message_accumulator.go @@ -0,0 +1,69 @@ +package multiagent + +import ( + "strings" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +type einoRunMessageAccumulator struct { + baseCount int + msgs []adk.Message +} + +func newEinoRunMessageAccumulator(base []adk.Message) *einoRunMessageAccumulator { + msgs := append([]adk.Message(nil), base...) + return &einoRunMessageAccumulator{ + baseCount: len(msgs), + msgs: msgs, + } +} + +func (a *einoRunMessageAccumulator) Append(msg adk.Message) bool { + if a == nil || msg == nil { + return false + } + a.msgs = append(a.msgs, msg) + return true +} + +func (a *einoRunMessageAccumulator) AppendToolMessage(content, toolCallID string, opts ...schema.ToolMessageOption) bool { + if strings.TrimSpace(toolCallID) == "" { + return false + } + return a.Append(schema.ToolMessage(content, toolCallID, opts...)) +} + +func (a *einoRunMessageAccumulator) AppendAssistantText(content string) bool { + content = strings.TrimSpace(content) + if content == "" { + return false + } + return a.Append(schema.AssistantMessage(content, nil)) +} + +func (a *einoRunMessageAccumulator) AppendAssistantToolCalls(toolCalls []schema.ToolCall) bool { + if len(toolCalls) == 0 { + return false + } + return a.Append(schema.AssistantMessage("", toolCalls)) +} + +func (a *einoRunMessageAccumulator) Messages() []adk.Message { + if a == nil { + return nil + } + return a.msgs +} + +func (a *einoRunMessageAccumulator) BaseCount() int { + if a == nil { + return 0 + } + return a.baseCount +} + +func (a *einoRunMessageAccumulator) HasNewMessages() bool { + return a != nil && len(a.msgs) > a.baseCount +} diff --git a/internal/multiagent/eino_run_message_accumulator_test.go b/internal/multiagent/eino_run_message_accumulator_test.go new file mode 100644 index 00000000..68751ff8 --- /dev/null +++ b/internal/multiagent/eino_run_message_accumulator_test.go @@ -0,0 +1,64 @@ +package multiagent + +import ( + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestEinoRunMessageAccumulatorTracksBaseAndAppends(t *testing.T) { + acc := newEinoRunMessageAccumulator([]adk.Message{schema.UserMessage("hi")}) + + if acc.BaseCount() != 1 { + t.Fatalf("base count = %d, want 1", acc.BaseCount()) + } + if acc.HasNewMessages() { + t.Fatal("fresh accumulator should not have new messages") + } + + if !acc.AppendAssistantText(" hello ") { + t.Fatal("assistant text should append") + } + if !acc.HasNewMessages() { + t.Fatal("expected new messages after append") + } + msgs := acc.Messages() + if len(msgs) != 2 || msgs[1].Role != schema.Assistant || msgs[1].Content != "hello" { + t.Fatalf("messages = %#v", msgs) + } +} + +func TestEinoRunMessageAccumulatorToolMessage(t *testing.T) { + acc := newEinoRunMessageAccumulator(nil) + if acc.AppendToolMessage("ignored", "") { + t.Fatal("blank tool call id should not append") + } + if !acc.AppendToolMessage("result", "call-1", schema.WithToolName("execute")) { + t.Fatal("tool message should append") + } + msgs := acc.Messages() + if len(msgs) != 1 || msgs[0].Role != schema.Tool || msgs[0].Content != "result" || msgs[0].ToolCallID != "call-1" || msgs[0].ToolName != "execute" { + t.Fatalf("tool message = %#v", msgs) + } +} + +func TestEinoRunMessageAccumulatorAssistantToolCalls(t *testing.T) { + acc := newEinoRunMessageAccumulator(nil) + if acc.AppendAssistantToolCalls(nil) { + t.Fatal("empty tool calls should not append") + } + if !acc.AppendAssistantToolCalls([]schema.ToolCall{{ + ID: "call-1", + Function: schema.FunctionCall{ + Name: "execute", + Arguments: `{}`, + }, + }}) { + t.Fatal("assistant tool calls should append") + } + msgs := acc.Messages() + if len(msgs) != 1 || msgs[0].Role != schema.Assistant || len(msgs[0].ToolCalls) != 1 { + t.Fatalf("assistant tool call message = %#v", msgs) + } +} diff --git a/internal/multiagent/eino_run_progress_tracker.go b/internal/multiagent/eino_run_progress_tracker.go new file mode 100644 index 00000000..ae612ced --- /dev/null +++ b/internal/multiagent/eino_run_progress_tracker.go @@ -0,0 +1,144 @@ +package multiagent + +import ( + "fmt" + "strings" + + "github.com/cloudwego/eino/schema" +) + +type einoRunProgressTracker struct { + orchMode string + orchestratorName string + conversationID string + progress func(eventType, message string, data interface{}) + + streamsMainAssistant func(agent string) bool + einoRoleTag func(agent string) string + + mainRound int + lastAgent string + toolEmitSeen map[string]struct{} + subAgentToolStep map[string]int + mainAgentToolStep map[string]int +} + +func newEinoRunProgressTracker( + orchMode, orchestratorName, conversationID string, + progress func(eventType, message string, data interface{}), + streamsMainAssistant func(agent string) bool, + einoRoleTag func(agent string) string, +) *einoRunProgressTracker { + if streamsMainAssistant == nil { + streamsMainAssistant = func(agent string) bool { + return agent == "" || agent == orchestratorName + } + } + if einoRoleTag == nil { + einoRoleTag = func(agent string) string { + if streamsMainAssistant(agent) { + return "orchestrator" + } + return "sub" + } + } + return &einoRunProgressTracker{ + orchMode: orchMode, + orchestratorName: orchestratorName, + conversationID: conversationID, + progress: progress, + streamsMainAssistant: streamsMainAssistant, + einoRoleTag: einoRoleTag, + toolEmitSeen: make(map[string]struct{}), + subAgentToolStep: make(map[string]int), + mainAgentToolStep: make(map[string]int), + } +} + +func (t *einoRunProgressTracker) ObserveAgent(agentName string) { + if t == nil || strings.TrimSpace(agentName) == "" || t.progress == nil { + return + } + iterEinoAgent := t.orchestratorName + if t.orchMode == "plan_execute" { + if a := strings.TrimSpace(agentName); a != "" { + iterEinoAgent = a + } + } + if t.streamsMainAssistant(agentName) { + mainIterKey := einoMainIterationKey(iterEinoAgent, t.orchestratorName) + if t.mainRound == 0 { + t.mainRound = 1 + t.mainAgentToolStep[mainIterKey] = 1 + t.emitMainIteration(iterEinoAgent, t.mainRound) + } else if t.lastAgent != "" { + needBump := false + if !t.streamsMainAssistant(t.lastAgent) { + needBump = true + } else if t.lastAgent != agentName { + needBump = true + } + if needBump { + t.mainRound++ + t.mainAgentToolStep[mainIterKey] = t.mainRound + t.emitMainIteration(iterEinoAgent, t.mainRound) + } + } + } + if t.lastAgent != agentName { + t.progress("progress", fmt.Sprintf("[Eino] %s", agentName), map[string]interface{}{ + "conversationId": t.conversationID, + "einoAgent": agentName, + "einoRole": t.einoRoleTag(agentName), + "orchestration": t.orchMode, + }) + } + t.lastAgent = agentName +} + +func (t *einoRunProgressTracker) MainIteration(agentName string) int { + if t == nil { + return 0 + } + key := einoMainIterationKey(agentName, t.orchestratorName) + if n := t.mainAgentToolStep[key]; n > 0 { + return n + } + return t.mainRound +} + +func (t *einoRunProgressTracker) EmitToolCalls(msg *schema.Message, agentName string, markPending func(toolCallPendingInfo)) { + if t == nil { + return + } + before := t.MainIteration(agentName) + tryEmitToolCallsOnce( + msg, + agentName, + t.orchestratorName, + t.conversationID, + t.orchMode, + t.progress, + t.toolEmitSeen, + t.subAgentToolStep, + t.mainAgentToolStep, + markPending, + ) + if t.streamsMainAssistant(agentName) { + if after := t.MainIteration(agentName); after > before { + t.mainRound = after + } + } +} + +func (t *einoRunProgressTracker) emitMainIteration(agentName string, iteration int) { + t.progress("iteration", "", map[string]interface{}{ + "iteration": iteration, + "einoScope": "main", + "einoRole": "orchestrator", + "einoAgent": agentName, + "orchestration": t.orchMode, + "conversationId": t.conversationID, + "source": "eino", + }) +} diff --git a/internal/multiagent/eino_run_progress_tracker_test.go b/internal/multiagent/eino_run_progress_tracker_test.go new file mode 100644 index 00000000..a3e9c635 --- /dev/null +++ b/internal/multiagent/eino_run_progress_tracker_test.go @@ -0,0 +1,185 @@ +package multiagent + +import ( + "testing" + + "github.com/cloudwego/eino/schema" +) + +func TestEinoRunProgressTrackerMainToolCallAdvancesResponseIteration(t *testing.T) { + var events []string + var iterations []int + progress := func(eventType, _ string, raw interface{}) { + events = append(events, eventType) + data, _ := raw.(map[string]interface{}) + if eventType == "iteration" { + if n, ok := data["iteration"].(int); ok { + iterations = append(iterations, n) + } + } + } + tracker := newEinoRunProgressTracker( + "eino_single", "main", "conv-1", progress, + func(agent string) bool { return agent == "" || agent == "main" }, + nil, + ) + + tracker.ObserveAgent("main") + if got := tracker.MainIteration("main"); got != 1 { + t.Fatalf("initial main iteration = %d, want 1", got) + } + tracker.EmitToolCalls(&schema.Message{ToolCalls: []schema.ToolCall{{ + ID: "call-1", + Type: "function", + Function: schema.FunctionCall{ + Name: "execute", + Arguments: `{"command":"pwd"}`, + }, + }}}, "main", nil) + if got := tracker.MainIteration("main"); got != 2 { + t.Fatalf("post-tool main iteration = %d, want 2", got) + } + if len(iterations) != 2 || iterations[0] != 1 || iterations[1] != 2 { + t.Fatalf("iteration events = %#v, want [1 2]; events=%#v", iterations, events) + } +} + +func TestEinoRunProgressTrackerMainAgentSwitchAdvancesIteration(t *testing.T) { + var iterations []int + progress := func(eventType, _ string, raw interface{}) { + if eventType != "iteration" { + return + } + data, _ := raw.(map[string]interface{}) + if n, ok := data["iteration"].(int); ok { + iterations = append(iterations, n) + } + } + tracker := newEinoRunProgressTracker( + "supervisor", "lead", "conv-1", progress, + func(agent string) bool { return agent == "" || agent == "lead" }, + nil, + ) + + tracker.ObserveAgent("lead") + tracker.ObserveAgent("sub") + tracker.ObserveAgent("lead") + + if got := tracker.MainIteration("lead"); got != 2 { + t.Fatalf("main iteration after sub->main = %d, want 2", got) + } + if len(iterations) != 2 || iterations[0] != 1 || iterations[1] != 2 { + t.Fatalf("iteration events = %#v, want [1 2]", iterations) + } +} + +func TestEinoRunProgressTrackerDedupesToolCalls(t *testing.T) { + var toolCalls int + progress := func(eventType, _ string, _ interface{}) { + if eventType == "tool_call" { + toolCalls++ + } + } + tracker := newEinoRunProgressTracker("deep", "lead", "conv-1", progress, nil, nil) + msg := &schema.Message{ToolCalls: []schema.ToolCall{{ + ID: "call-1", + Type: "function", + Function: schema.FunctionCall{ + Name: "search", + Arguments: `{"q":"x"}`, + }, + }}} + + tracker.EmitToolCalls(msg, "lead", nil) + tracker.EmitToolCalls(msg, "lead", nil) + + if toolCalls != 1 { + t.Fatalf("tool call events = %d, want 1", toolCalls) + } +} + +func TestEinoRunProgressTrackerHidesModelOutputRecoveryToolCalls(t *testing.T) { + var eventTypes []string + var marked []toolCallPendingInfo + progress := func(eventType, _ string, _ interface{}) { + eventTypes = append(eventTypes, eventType) + } + tracker := newEinoRunProgressTracker("deep", "lead", "conv-1", progress, nil, nil) + msg := &schema.Message{ToolCalls: []schema.ToolCall{{ + ID: "call-recovery", + Type: "function", + Function: schema.FunctionCall{ + Name: "task", + Arguments: `{"_cyberstrike_model_output_recovery":{"reason":"invalid_tool_arguments_json","repair_attempt":1}}`, + }, + }}} + + tracker.EmitToolCalls(msg, "lead", func(info toolCallPendingInfo) { + marked = append(marked, info) + }) + + if containsString(eventTypes, "tool_calls_detected") || containsString(eventTypes, "tool_call") { + t.Fatalf("event types = %#v, want no visible recovery tool call events", eventTypes) + } + if len(marked) != 0 { + t.Fatalf("marked pending = %#v, want none", marked) + } +} + +func TestEinoRunProgressTrackerHidesAnonymousToolCallFragments(t *testing.T) { + var eventTypes []string + var marked []toolCallPendingInfo + progress := func(eventType, _ string, _ interface{}) { + eventTypes = append(eventTypes, eventType) + } + tracker := newEinoRunProgressTracker("eino_single", "lead", "conv-1", progress, nil, nil) + idx := 0 + msg := &schema.Message{ToolCalls: []schema.ToolCall{{ + Type: "function", + Index: &idx, + Function: schema.FunctionCall{ + Arguments: `"`, + }, + }}} + + tracker.EmitToolCalls(msg, "lead", func(info toolCallPendingInfo) { + marked = append(marked, info) + }) + + if containsString(eventTypes, "tool_calls_detected") || containsString(eventTypes, "tool_call") { + t.Fatalf("event types = %#v, want no visible anonymous fragment tool call events", eventTypes) + } + if len(marked) != 0 { + t.Fatalf("marked pending = %#v, want none", marked) + } +} + +func TestEinoRunProgressTrackerKeepsNamedInvalidToolCallsVisible(t *testing.T) { + var toolCalls int + var marked []toolCallPendingInfo + progress := func(eventType, _ string, _ interface{}) { + if eventType == "tool_call" { + toolCalls++ + } + } + tracker := newEinoRunProgressTracker("eino_single", "lead", "conv-1", progress, nil, nil) + msg := &schema.Message{ToolCalls: []schema.ToolCall{{ + ID: "call-bad-args", + Type: "function", + Function: schema.FunctionCall{ + Name: "exec", + Arguments: `command`, + }, + }}} + + tracker.EmitToolCalls(msg, "lead", func(info toolCallPendingInfo) { + marked = append(marked, info) + }) + + if toolCalls != 1 { + t.Fatalf("tool call events = %d, want 1", toolCalls) + } + if len(marked) != 1 || marked[0].ToolName != "exec" { + t.Fatalf("marked pending = %#v, want one exec call", marked) + } +} diff --git a/internal/multiagent/eino_run_recovery_handler.go b/internal/multiagent/eino_run_recovery_handler.go new file mode 100644 index 00000000..367287b2 --- /dev/null +++ b/internal/multiagent/eino_run_recovery_handler.go @@ -0,0 +1,70 @@ +package multiagent + +import ( + "github.com/cloudwego/eino/adk" + "go.uber.org/zap" +) + +type einoRunRecoveryHandlerConfig struct { + ConversationID string + OrchMode string + Args *einoADKRunLoopArgs + BaseMsgs []adk.Message + Progress func(eventType, message string, data interface{}) + Logger *zap.Logger + RunError *einoRunErrorHandler + ContextOverflow *einoContextOverflowRetryHandler + Transient *einoTransientRunRetryHandler +} + +type einoRunRecoveryResult struct { + Handled bool + Restarted bool + RestartMsgs []adk.Message + Fatal error +} + +type einoRunRecoveryHandler struct { + cfg einoRunRecoveryHandlerConfig +} + +func newEinoRunRecoveryHandler(cfg einoRunRecoveryHandlerConfig) *einoRunRecoveryHandler { + if cfg.Args == nil { + cfg.Args = &einoADKRunLoopArgs{} + } + return &einoRunRecoveryHandler{cfg: cfg} +} + +func (h *einoRunRecoveryHandler) Handle(runErr error, accumulated []adk.Message, baseCount int) einoRunRecoveryResult { + if h == nil || runErr == nil { + return einoRunRecoveryResult{} + } + if willRetry, ok := isEinoNativeWillRetry(runErr); ok { + emitEinoNativeModelRetryProgress(h.cfg.ConversationID, h.cfg.OrchMode, willRetry, h.cfg.Progress, h.cfg.Logger, runErr) + return einoRunRecoveryResult{Handled: true} + } + if h.cfg.ContextOverflow != nil { + if overflowRetry := h.cfg.ContextOverflow.Prepare(runErr, accumulated, baseCount); overflowRetry.Handled { + return einoRunRecoveryResult{Handled: true, Restarted: true, RestartMsgs: overflowRetry.RestartMsgs} + } + } + if h.cfg.Transient != nil { + if runRetry := h.cfg.Transient.Prepare(runErr, accumulated, baseCount); runRetry.Handled { + if runRetry.Fatal != nil { + return einoRunRecoveryResult{Handled: true, Fatal: runRetry.Fatal} + } + if !runRetry.Restarted { + return einoRunRecoveryResult{Handled: true} + } + return einoRunRecoveryResult{Handled: true, Restarted: true, RestartMsgs: runRetry.RestartMsgs} + } + } + return einoRunRecoveryResult{Handled: true, Fatal: h.handleFatal(runErr)} +} + +func (h *einoRunRecoveryHandler) handleFatal(runErr error) error { + if h != nil && h.cfg.RunError != nil { + return h.cfg.RunError.Handle(runErr) + } + return runErr +} diff --git a/internal/multiagent/eino_run_recovery_handler_test.go b/internal/multiagent/eino_run_recovery_handler_test.go new file mode 100644 index 00000000..1d121503 --- /dev/null +++ b/internal/multiagent/eino_run_recovery_handler_test.go @@ -0,0 +1,49 @@ +package multiagent + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestEinoRunRecoveryHandlerRoutesContextOverflowBeforeTransient(t *testing.T) { + baseMsgs := []adk.Message{schema.UserMessage("base")} + overflow := newEinoContextOverflowRetryHandler(einoContextOverflowRetryConfig{ + Context: context.Background(), + Args: &einoADKRunLoopArgs{}, + BaseMsgs: baseMsgs, + }) + transient := newEinoTransientRunRetryHandler(einoTransientRunRetryHandlerConfig{ + Args: &einoADKRunLoopArgs{}, + BaseMsgs: baseMsgs, + Policy: einoTransientRunRetryPolicy{maxAttempts: 1, maxBackoff: time.Nanosecond}, + }) + handler := newEinoRunRecoveryHandler(einoRunRecoveryHandlerConfig{ + ContextOverflow: overflow, + Transient: transient, + BaseMsgs: baseMsgs, + }) + + result := handler.Handle(errors.New("context length exceeded: upstream returned 503"), nil, 0) + if !result.Handled || !result.Restarted || result.Fatal != nil { + t.Fatalf("result = %+v, want context overflow restart", result) + } + second := handler.Handle(errors.New("upstream returned 503"), nil, 0) + if !second.Handled || !second.Restarted || second.Fatal != nil { + t.Fatalf("second result = %+v, want transient restart", second) + } +} + +func TestEinoRunRecoveryHandlerRoutesFatalFallback(t *testing.T) { + handler := newEinoRunRecoveryHandler(einoRunRecoveryHandlerConfig{ + RunError: newEinoRunErrorHandler(einoRunErrorHandlerConfig{}), + }) + result := handler.Handle(errors.New("invalid api key"), nil, 0) + if !result.Handled || result.Restarted || result.Fatal == nil { + t.Fatalf("result = %+v, want fatal fallback", result) + } +} diff --git a/internal/multiagent/eino_run_result_builder.go b/internal/multiagent/eino_run_result_builder.go new file mode 100644 index 00000000..958068e5 --- /dev/null +++ b/internal/multiagent/eino_run_result_builder.go @@ -0,0 +1,220 @@ +package multiagent + +import ( + "encoding/json" + "strings" + + "cyberstrike-ai/internal/agent" + "cyberstrike-ai/internal/einomcp" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +type einoRunResultBuilderConfig struct { + OrchMode string + EmptyHint string + RunMessages *einoRunMessageAccumulator + AssistantOutput *einoAssistantOutputAccumulator + SnapshotMCPIDs func() []string + ModelFacingTrace func() []adk.Message +} + +type einoRunResultBuilder struct { + cfg einoRunResultBuilderConfig +} + +func newEinoRunResultBuilder(cfg einoRunResultBuilderConfig) *einoRunResultBuilder { + return &einoRunResultBuilder{cfg: cfg} +} + +func (b *einoRunResultBuilder) BuildPartial(runErr error) (*RunResult, error) { + if b == nil || b.cfg.RunMessages == nil || !b.cfg.RunMessages.HasNewMessages() { + return nil, runErr + } + return b.build(true), runErr +} + +func (b *einoRunResultBuilder) BuildFinal() *RunResult { + if b == nil { + return &RunResult{} + } + return b.build(false) +} + +func (b *einoRunResultBuilder) build(partial bool) *RunResult { + var runMsgs []adk.Message + if b.cfg.RunMessages != nil { + runMsgs = b.cfg.RunMessages.Messages() + } + var lastAssistant string + var lastPlanExecuteExecutor string + if b.cfg.AssistantOutput != nil { + lastAssistant = b.cfg.AssistantOutput.LastAssistant() + lastPlanExecuteExecutor = b.cfg.AssistantOutput.LastPlanExecuteExecutor() + } + var modelFacing []adk.Message + if b.cfg.ModelFacingTrace != nil { + modelFacing = b.cfg.ModelFacingTrace() + } + var ids []string + if b.cfg.SnapshotMCPIDs != nil { + ids = b.cfg.SnapshotMCPIDs() + } + return buildEinoRunResultFromAccumulated( + b.cfg.OrchMode, + runMsgs, + modelFacing, + lastAssistant, + lastPlanExecuteExecutor, + b.cfg.EmptyHint, + ids, + partial, + ) +} + +func einoPartialRunLastOutputHint() string { + return "[执行未正常结束(用户停止、超时或异常)。续跑时请基于上文已产生的工具与结果继续,勿重复已完成步骤。]\n" + + "[Run ended abnormally; continue from the trace above without repeating completed steps.]" +} + +func buildEinoRunResultFromAccumulated( + orchMode string, + runAccumulatedMsgs []adk.Message, + persistMsgs []adk.Message, + lastAssistant string, + lastPlanExecuteExecutor string, + emptyHint string, + mcpIDs []string, + partial bool, +) *RunResult { + traceForJSON := persistMsgs + traceJSON := "" + if len(traceForJSON) > 0 { + traceForJSON = markModelFacingTraceForPersistence(traceForJSON) + if histJSON, err := json.Marshal(traceForJSON); err == nil { + traceJSON = string(histJSON) + } + } + cleaned := strings.TrimSpace(lastAssistant) + if orchMode == "plan_execute" { + if e := strings.TrimSpace(lastPlanExecuteExecutor); e != "" { + cleaned = e + } else { + cleaned = UnwrapPlanExecuteUserText(cleaned) + } + } + if cleaned == "" { + if fb := strings.TrimSpace(einoExtractFallbackAssistantFromMsgs(runAccumulatedMsgs)); fb != "" { + cleaned = fb + } + } + cleaned = dedupeRepeatedParagraphs(cleaned, 80) + cleaned = dedupeParagraphsByLineFingerprint(cleaned, 100) + const maxResponseRunes = 100000 + if rs := []rune(cleaned); len(rs) > maxResponseRunes { + cleaned = string(rs[:maxResponseRunes]) + "\n\n... (response truncated / 响应已截断)" + } + lastOut := cleaned + resp := cleaned + if partial && cleaned == "" { + lastOut = einoPartialRunLastOutputHint() + resp = emptyHint + } + out := &RunResult{ + Response: resp, + MCPExecutionIDs: mcpIDs, + LastAgentTraceInput: traceJSON, + LastAgentTraceOutput: lastOut, + } + if !partial && out.Response == "" { + out.Response = emptyHint + out.LastAgentTraceOutput = out.Response + } + return out +} + +func markModelFacingTraceForPersistence(msgs []adk.Message) []adk.Message { + out := cloneADKMessagesForTrace(msgs) + if len(out) == 0 || out[0] == nil { + return out + } + if out[0].Extra == nil { + out[0].Extra = make(map[string]any, 1) + } + out[0].Extra[agent.ModelFacingTraceVersionKey] = 1 + return out +} + +// einoExtractFallbackAssistantFromMsgs 在「主通道未产出助手正文」时,从 Eino ADK 轨迹中回填用户可见回复。 +// 典型场景:监督者仅调用 exit(final_result 落在 Tool 消息中),或工具结果已写入历史但 lastAssistant 未更新。 +// +// 优先级:最后一次 exit 工具输出 → 最后一条含 exit 的助手 tool_calls 参数中的 final_result。 +func einoExtractFallbackAssistantFromMsgs(msgs []adk.Message) string { + for i := len(msgs) - 1; i >= 0; i-- { + m := msgs[i] + if m == nil || m.Role != schema.Tool { + continue + } + if !strings.EqualFold(strings.TrimSpace(m.ToolName), adk.ToolInfoExit.Name) { + continue + } + content := strings.TrimSpace(m.Content) + if content == "" || strings.HasPrefix(content, einomcp.ToolErrorPrefix) { + continue + } + return content + } + for i := len(msgs) - 1; i >= 0; i-- { + m := msgs[i] + if m == nil || m.Role != schema.Assistant { + continue + } + if s := einoExtractExitFinalFromAssistantToolCalls(m); s != "" { + return s + } + } + return "" +} + +func einoExtractExitFinalFromAssistantToolCalls(msg *schema.Message) string { + if msg == nil || len(msg.ToolCalls) == 0 { + return "" + } + for i := len(msg.ToolCalls) - 1; i >= 0; i-- { + tc := msg.ToolCalls[i] + if !strings.EqualFold(strings.TrimSpace(tc.Function.Name), adk.ToolInfoExit.Name) { + continue + } + if s := einoParseExitFinalResultArguments(tc.Function.Arguments); s != "" { + return s + } + } + return "" +} + +func einoParseExitFinalResultArguments(arguments string) string { + arguments = strings.TrimSpace(arguments) + if arguments == "" { + return "" + } + var wrap struct { + FinalResult json.RawMessage `json:"final_result"` + } + if err := json.Unmarshal([]byte(arguments), &wrap); err != nil || len(wrap.FinalResult) == 0 { + return "" + } + var s string + if err := json.Unmarshal(wrap.FinalResult, &s); err == nil { + return strings.TrimSpace(s) + } + var anyVal interface{} + if err := json.Unmarshal(wrap.FinalResult, &anyVal); err != nil { + return "" + } + b, err := json.Marshal(anyVal) + if err != nil { + return "" + } + return strings.TrimSpace(string(b)) +} diff --git a/internal/multiagent/eino_run_result_builder_test.go b/internal/multiagent/eino_run_result_builder_test.go new file mode 100644 index 00000000..45d28da5 --- /dev/null +++ b/internal/multiagent/eino_run_result_builder_test.go @@ -0,0 +1,75 @@ +package multiagent + +import ( + "errors" + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestEinoRunResultBuilderPartialWithoutNewMessagesReturnsOriginalError(t *testing.T) { + runMessages := newEinoRunMessageAccumulator([]adk.Message{schema.UserMessage("base")}) + wantErr := errors.New("stream failed") + + got, err := newEinoRunResultBuilder(einoRunResultBuilderConfig{ + RunMessages: runMessages, + EmptyHint: "empty", + }).BuildPartial(wantErr) + + if got != nil { + t.Fatalf("partial result = %#v, want nil", got) + } + if !errors.Is(err, wantErr) { + t.Fatalf("err = %v, want %v", err, wantErr) + } +} + +func TestEinoRunResultBuilderFinalUsesSnapshots(t *testing.T) { + runMessages := newEinoRunMessageAccumulator([]adk.Message{schema.UserMessage("base")}) + runMessages.Append(schema.AssistantMessage("assistant done", nil)) + assistantOutput := newEinoAssistantOutputAccumulator("deep") + assistantOutput.RecordMainAssistant("orchestrator", "assistant done") + + got := newEinoRunResultBuilder(einoRunResultBuilderConfig{ + OrchMode: "deep", + EmptyHint: "empty", + RunMessages: runMessages, + AssistantOutput: assistantOutput, + SnapshotMCPIDs: func() []string { + return []string{"exec-1"} + }, + ModelFacingTrace: func() []adk.Message { + return []adk.Message{schema.UserMessage("model-facing")} + }, + }).BuildFinal() + + if got.Response != "assistant done" { + t.Fatalf("response = %q, want assistant done", got.Response) + } + if len(got.MCPExecutionIDs) != 1 || got.MCPExecutionIDs[0] != "exec-1" { + t.Fatalf("mcp ids = %#v", got.MCPExecutionIDs) + } + if got.LastAgentTraceInput == "" { + t.Fatal("model-facing trace should be persisted") + } +} + +func TestEinoRunResultBuilderPlanExecutePrefersExecutorOutput(t *testing.T) { + runMessages := newEinoRunMessageAccumulator(nil) + runMessages.Append(schema.AssistantMessage(`{"response":"planner text"}`, nil)) + assistantOutput := newEinoAssistantOutputAccumulator("plan_execute") + assistantOutput.RecordMainAssistant("planner", `{"response":"planner text"}`) + assistantOutput.RecordMainAssistant("executor", `{"response":"executor text"}`) + + got := newEinoRunResultBuilder(einoRunResultBuilderConfig{ + OrchMode: "plan_execute", + EmptyHint: "empty", + RunMessages: runMessages, + AssistantOutput: assistantOutput, + }).BuildFinal() + + if got.Response != "executor text" { + t.Fatalf("response = %q, want executor text", got.Response) + } +} diff --git a/internal/multiagent/eino_run_runtime_session.go b/internal/multiagent/eino_run_runtime_session.go new file mode 100644 index 00000000..e582156d --- /dev/null +++ b/internal/multiagent/eino_run_runtime_session.go @@ -0,0 +1,372 @@ +package multiagent + +import ( + "context" + "sync/atomic" + + "github.com/cloudwego/eino/adk" + "go.uber.org/zap" +) + +type einoRunRuntimeSessionConfig struct { + Context context.Context + Args *einoADKRunLoopArgs + Drain *einoRunEventDrain + BaseMessages []adk.Message + EmptyHint string + SnapshotMCPIDs func() []string + EinoRoleTag func(agent string) string +} + +type einoRunRuntimeErrorResult struct { + Restarted bool + Result *RunResult + Err error +} + +type einoRunRuntimeSession struct { + ctx context.Context + args *einoADKRunLoopArgs + orchMode string + conversationID string + progress func(eventType, message string, data interface{}) + logger *zap.Logger + baseMsgs []adk.Message + msgs []adk.Message + drain *einoRunEventDrain + runMessages *einoRunMessageAccumulator + usage *einoRunUsageAccumulator + + iter *adk.AsyncIterator[*adk.AgentEvent] + startFreshIter einoAgentEventIteratorStarter + + unregisterAgentCancel func() + unregisterTurnLoopInterrupt func() + nativeCancelCause atomic.Value + + transientRetry *einoTransientRunRetryHandler + runRecoveryHandler *einoRunRecoveryHandler + resultBuilder *einoRunResultBuilder + streamErrorHandler *einoStreamErrorHandler + completionHandler *einoRunCompletionHandler + cancellationHandler *einoRunCancellationHandler +} + +func newEinoRunRuntimeSession(cfg einoRunRuntimeSessionConfig) *einoRunRuntimeSession { + if cfg.Context == nil { + cfg.Context = context.Background() + } + if cfg.Args == nil { + cfg.Args = &einoADKRunLoopArgs{} + } + if cfg.SnapshotMCPIDs == nil { + cfg.SnapshotMCPIDs = func() []string { return nil } + } + s := &einoRunRuntimeSession{ + ctx: cfg.Context, + args: cfg.Args, + orchMode: cfg.Args.OrchMode, + conversationID: cfg.Args.ConversationID, + progress: cfg.Args.Progress, + logger: cfg.Args.Logger, + baseMsgs: cfg.BaseMessages, + msgs: append([]adk.Message(nil), cfg.BaseMessages...), + drain: cfg.Drain, + } + if s.drain != nil { + s.runMessages = s.drain.RunMessages() + s.usage = s.drain.Usage() + } + if s.runMessages == nil { + s.runMessages = newEinoRunMessageAccumulator(s.msgs) + } + s.initIteratorRuntime() + s.initRecoveryRuntime() + s.initResultRuntime(cfg.EmptyHint, cfg.SnapshotMCPIDs, cfg.EinoRoleTag) + return s +} + +func (s *einoRunRuntimeSession) Iterator() *adk.AsyncIterator[*adk.AgentEvent] { + if s == nil { + return nil + } + return s.iter +} + +func (s *einoRunRuntimeSession) Close() { + if s == nil { + return + } + callAndClearUnregister(&s.unregisterAgentCancel) + callAndClearUnregister(&s.unregisterTurnLoopInterrupt) +} + +func (s *einoRunRuntimeSession) HandleIteratorContextError(err error) (*RunResult, error) { + if s == nil || s.cancellationHandler == nil { + return nil, err + } + return s.cancellationHandler.Handle(err) +} + +func (s *einoRunRuntimeSession) HandleIteratorEnd() (completed bool, result *RunResult, err error) { + if s == nil { + return true, nil, nil + } + if ctxErr := s.ctx.Err(); ctxErr != nil { + result, err = s.HandleIteratorContextError(ctxErr) + return false, result, err + } + if s.completionHandler != nil { + s.completionHandler.Complete() + } + return true, nil, nil +} + +func (s *einoRunRuntimeSession) HandleRunError(runErr error) einoRunRuntimeErrorResult { + if s == nil || runErr == nil { + return einoRunRuntimeErrorResult{} + } + restarted, fatal := s.maybeRestart(runErr) + if fatal != nil { + result, err := s.takePartial(fatal) + return einoRunRuntimeErrorResult{Result: result, Err: err} + } + return einoRunRuntimeErrorResult{Restarted: restarted} +} + +func (s *einoRunRuntimeSession) HandleStreamError(streamErr error, agentName string) einoRunRuntimeErrorResult { + if s == nil || s.streamErrorHandler == nil || streamErr == nil { + return einoRunRuntimeErrorResult{} + } + handled := s.streamErrorHandler.Handle(streamErr, agentName) + return einoRunRuntimeErrorResult{ + Restarted: handled.Restarted, + Result: handled.Result, + Err: handled.Err, + } +} + +func (s *einoRunRuntimeSession) ConfirmRecovery() { + if s != nil && s.transientRetry != nil { + s.transientRetry.ConfirmRecovery() + } +} + +func (s *einoRunRuntimeSession) BuildFinalResult() *RunResult { + if s == nil || s.resultBuilder == nil { + return &RunResult{} + } + s.emitUsageSummary("final") + return s.resultBuilder.BuildFinal() +} + +func (s *einoRunRuntimeSession) takePartial(err error) (*RunResult, error) { + if s == nil || s.resultBuilder == nil { + return nil, err + } + s.emitUsageSummary("partial") + return s.resultBuilder.BuildPartial(err) +} + +func (s *einoRunRuntimeSession) maybeRestart(runErr error) (restarted bool, fatal error) { + if s == nil || s.runRecoveryHandler == nil { + return false, runErr + } + recovery := s.runRecoveryHandler.Handle(runErr, s.runMessages.Messages(), s.runMessages.BaseCount()) + if recovery.Fatal != nil { + return false, recovery.Fatal + } + if !recovery.Restarted { + return false, nil + } + s.msgs = recovery.RestartMsgs + s.iter = s.startFreshIter(s.msgs) + return true, nil +} + +func (s *einoRunRuntimeSession) initIteratorRuntime() { + if s == nil || s.args == nil { + return + } + runnerCfg := adk.RunnerConfig{ + Agent: s.args.DA, + // 启用 ADK 流式事件:plan_execute 也需要输出 reasoning/response 流, + // 与 deep/supervisor/eino_single 的前端体验保持一致。 + EnableStreaming: true, + } + var cpStore *fileCheckPointStore + var checkPointID string + if checkpoint := newEinoCheckpointRuntime(s.args.CheckpointDir, s.conversationID, s.orchMode, s.logger); checkpoint != nil { + cpStore = checkpoint.Store + checkPointID = checkpoint.CheckPointID + runnerCfg.CheckPointStore = checkpoint.Store + } + runner := adk.NewRunner(s.ctx, runnerCfg) + runtimeCancelRegistrar := agentRuntimeCancelRegistrarFromContext(s.ctx) + turnLoopInterruptRegistrar := agentTurnLoopInterruptRegistrarFromContext(s.ctx) + runnerStarter := newEinoRunnerIteratorStarter(einoRunnerIteratorStarterConfig{ + Context: s.ctx, + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Logger: s.logger, + Runner: runner, + CheckPointID: checkPointID, + NativeCancelCause: &s.nativeCancelCause, + UnregisterAgentCancel: &s.unregisterAgentCancel, + RuntimeCancelRegistrar: runtimeCancelRegistrar, + }) + turnLoopStarter := newEinoTurnLoopIteratorStarter(einoTurnLoopIteratorStarterConfig{ + Context: s.ctx, + Agent: s.args.DA, + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Progress: s.progress, + Logger: s.logger, + Store: cpStore, + CheckPointID: checkPointID, + InterruptTimeout: s.args.TurnLoopInterruptTimeout, + NativeCancelCause: &s.nativeCancelCause, + UnregisterAgentCancel: &s.unregisterAgentCancel, + UnregisterTurnLoopInterrupt: &s.unregisterTurnLoopInterrupt, + RuntimeCancelRegistrar: runtimeCancelRegistrar, + TurnLoopInterruptRegistrar: turnLoopInterruptRegistrar, + }) + useTurnLoop := turnLoopInterruptRegistrar != nil + s.startFreshIter = runnerStarter.Start + if useTurnLoop { + s.startFreshIter = turnLoopStarter.Start + } + if !useTurnLoop && cpStore != nil && checkPointID != "" { + s.iter = newEinoCheckpointResumeHandler(einoCheckpointResumeHandlerConfig{ + Context: s.ctx, + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Progress: s.progress, + Logger: s.logger, + Store: cpStore, + CheckPointID: checkPointID, + Resume: runnerStarter.Resume, + }).TryResume() + } + s.iter = newEinoInitialIteratorStartHandler(einoInitialIteratorStartHandlerConfig{ + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Progress: s.progress, + UseTurnLoop: useTurnLoop, + StartRunner: runnerStarter.Start, + StartTurnLoop: turnLoopStarter.Start, + }).StartIfNeeded(s.iter, s.msgs) + + pending := s.pending() + s.completionHandler = newEinoRunCompletionHandler(einoRunCompletionHandlerConfig{ + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Progress: s.progress, + Logger: s.logger, + Pending: pending, + Checkpoint: cpStore, + CheckpointID: checkPointID, + }) + s.cancellationHandler = newEinoRunCancellationHandler(einoRunCancellationHandlerConfig{ + Context: s.ctx, + ConversationID: s.conversationID, + Progress: s.progress, + Pending: pending, + TakePartial: s.takePartial, + }) +} + +func (s *einoRunRuntimeSession) initRecoveryRuntime() { + if s == nil || s.args == nil { + return + } + pending := s.pending() + contextOverflowRetry := newEinoContextOverflowRetryHandler(einoContextOverflowRetryConfig{ + Context: s.ctx, + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Args: s.args, + BaseMsgs: s.baseMsgs, + Progress: s.progress, + Logger: s.logger, + }) + s.transientRetry = newEinoTransientRunRetryHandler(einoTransientRunRetryHandlerConfig{ + Context: s.ctx, + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Args: s.args, + BaseMsgs: s.baseMsgs, + Progress: s.progress, + Logger: s.logger, + Pending: pending, + }) + runErrorHandler := newEinoRunErrorHandler(einoRunErrorHandlerConfig{ + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Progress: s.progress, + Pending: pending, + NativeCancelFallback: s.nativeCancelCauseOrCanceled, + }) + s.runRecoveryHandler = newEinoRunRecoveryHandler(einoRunRecoveryHandlerConfig{ + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Args: s.args, + BaseMsgs: s.baseMsgs, + Progress: s.progress, + Logger: s.logger, + RunError: runErrorHandler, + ContextOverflow: contextOverflowRetry, + Transient: s.transientRetry, + }) +} + +func (s *einoRunRuntimeSession) initResultRuntime(emptyHint string, snapshotMCPIDs func() []string, einoRoleTag func(agent string) string) { + if s == nil { + return + } + var assistantOutput *einoAssistantOutputAccumulator + if s.drain != nil { + assistantOutput = s.drain.AssistantOutput() + } + s.resultBuilder = newEinoRunResultBuilder(einoRunResultBuilderConfig{ + OrchMode: s.orchMode, + EmptyHint: emptyHint, + RunMessages: s.runMessages, + AssistantOutput: assistantOutput, + SnapshotMCPIDs: snapshotMCPIDs, + ModelFacingTrace: func() []adk.Message { return modelFacingTraceSnapshot(s.args) }, + }) + s.streamErrorHandler = newEinoStreamErrorHandler( + s.ctx, + s.conversationID, + s.progress, + einoRoleTag, + s.maybeRestart, + s.takePartial, + ) +} + +func (s *einoRunRuntimeSession) pending() *einoPendingToolCalls { + if s == nil || s.drain == nil { + return nil + } + return s.drain.PendingToolCalls() +} + +func (s *einoRunRuntimeSession) nativeCancelCauseOrCanceled() error { + if s != nil { + if v := s.nativeCancelCause.Load(); v != nil { + if err, ok := v.(error); ok && err != nil { + return err + } + } + } + return context.Canceled +} + +func (s *einoRunRuntimeSession) emitUsageSummary(reason string) bool { + if s == nil || s.usage == nil { + return false + } + return s.usage.EmitOnce(s.conversationID, s.orchMode, reason, s.progress, s.logger) +} diff --git a/internal/multiagent/eino_run_runtime_session_test.go b/internal/multiagent/eino_run_runtime_session_test.go new file mode 100644 index 00000000..9e14c573 --- /dev/null +++ b/internal/multiagent/eino_run_runtime_session_test.go @@ -0,0 +1,196 @@ +package multiagent + +import ( + "context" + "errors" + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +type fakeRuntimeSessionAgent struct { + runMessages []adk.Message + runOpts int +} + +func (a *fakeRuntimeSessionAgent) Name(context.Context) string { + return "lead" +} + +func (a *fakeRuntimeSessionAgent) Description(context.Context) string { + return "fake runtime session agent" +} + +func (a *fakeRuntimeSessionAgent) Run(_ context.Context, input *adk.AgentInput, opts ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] { + if input != nil { + a.runMessages = input.Messages + } + a.runOpts = len(opts) + iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + gen.Close() + return iter +} + +func TestEinoRunRuntimeSessionStartsRunner(t *testing.T) { + agent := &fakeRuntimeSessionAgent{} + drain := newEinoRunEventDrain(einoRunEventDrainConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + BaseMessages: []adk.Message{schema.UserMessage("base")}, + }) + session := newEinoRunRuntimeSession(einoRunRuntimeSessionConfig{ + Context: context.Background(), + Args: &einoADKRunLoopArgs{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + DA: agent, + }, + Drain: drain, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + EmptyHint: "empty", + }) + defer session.Close() + + if session.Iterator() == nil { + t.Fatal("session should start an iterator") + } + if len(agent.runMessages) != 1 || agent.runMessages[0].Content != "base" { + t.Fatalf("run messages = %#v", agent.runMessages) + } + if agent.runOpts != 1 { + t.Fatalf("run opts = %d, want native cancel option", agent.runOpts) + } +} + +func TestEinoRunRuntimeSessionCompletionFlushesPending(t *testing.T) { + agent := &fakeRuntimeSessionAgent{} + var events []string + drain := newEinoRunEventDrain(einoRunEventDrainConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + Progress: func(eventType, _ string, _ interface{}) { + events = append(events, eventType) + }, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + }) + session := newEinoRunRuntimeSession(einoRunRuntimeSessionConfig{ + Context: context.Background(), + Args: &einoADKRunLoopArgs{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + Progress: func(eventType, _ string, _ interface{}) { + events = append(events, eventType) + }, + DA: agent, + }, + Drain: drain, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + EmptyHint: "empty", + }) + defer session.Close() + + drain.PendingToolCalls().Mark(toolCallPendingInfo{ + ToolCallID: "call-1", + ToolName: "execute", + EinoAgent: "lead", + EinoRole: "orchestrator", + }) + completed, result, err := session.HandleIteratorEnd() + + if !completed || result != nil || err != nil { + t.Fatalf("completed=%v result=%#v err=%v", completed, result, err) + } + if !containsString(events, "tool_result") || !containsString(events, "eino_pending_orphaned") { + t.Fatalf("events = %#v, want orphan pending flush", events) + } +} + +func TestEinoRunRuntimeSessionCancellationReturnsPartialError(t *testing.T) { + agent := &fakeRuntimeSessionAgent{} + var events []string + drain := newEinoRunEventDrain(einoRunEventDrainConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + Progress: func(eventType, _ string, _ interface{}) { + events = append(events, eventType) + }, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + }) + session := newEinoRunRuntimeSession(einoRunRuntimeSessionConfig{ + Context: context.Background(), + Args: &einoADKRunLoopArgs{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + Progress: func(eventType, _ string, _ interface{}) { + events = append(events, eventType) + }, + DA: agent, + }, + Drain: drain, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + EmptyHint: "empty", + }) + defer session.Close() + + stopErr := errors.New("stop") + result, err := session.HandleIteratorContextError(stopErr) + + if result != nil { + t.Fatalf("result = %#v, want nil without new messages", result) + } + if !errors.Is(err, stopErr) { + t.Fatalf("err = %v, want %v", err, stopErr) + } + if !containsString(events, "error") { + t.Fatalf("events = %#v, want cancellation error event", events) + } +} + +func TestEinoRunRuntimeSessionBuildFinalEmitsUsageSummary(t *testing.T) { + agent := &fakeRuntimeSessionAgent{} + var usageEvent map[string]interface{} + progress := func(eventType, _ string, data interface{}) { + if eventType != "eino_usage_summary" { + return + } + usageEvent, _ = data.(map[string]interface{}) + } + drain := newEinoRunEventDrain(einoRunEventDrainConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + Progress: progress, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + }) + session := newEinoRunRuntimeSession(einoRunRuntimeSessionConfig{ + Context: context.Background(), + Args: &einoADKRunLoopArgs{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + Progress: progress, + DA: agent, + }, + Drain: drain, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + EmptyHint: "empty", + }) + defer session.Close() + + drain.Usage().AddUsage(&schema.TokenUsage{PromptTokens: 3, CompletionTokens: 4, TotalTokens: 7}) + _ = session.BuildFinalResult() + + if usageEvent == nil { + t.Fatal("usage summary event was not emitted") + } + if usageEvent["conversationId"] != "conv-1" || usageEvent["orchestration"] != "deep" || usageEvent["reason"] != "final" || usageEvent["totalTokens"] != 7 { + t.Fatalf("usage event = %#v", usageEvent) + } +} diff --git a/internal/multiagent/eino_run_trace.go b/internal/multiagent/eino_run_trace.go new file mode 100644 index 00000000..d94c742f --- /dev/null +++ b/internal/multiagent/eino_run_trace.go @@ -0,0 +1,41 @@ +package multiagent + +import ( + "fmt" + "strings" + + "github.com/google/uuid" +) + +func newEinoRunID() string { + return uuid.New().String() +} + +func withEinoRunIDProgress( + runID string, + progress func(eventType, message string, data interface{}), +) func(eventType, message string, data interface{}) { + runID = strings.TrimSpace(runID) + if progress == nil || runID == "" { + return progress + } + return func(eventType, message string, data interface{}) { + progress(eventType, message, addEinoRunIDToProgressData(runID, data)) + } +} + +func addEinoRunIDToProgressData(runID string, data interface{}) interface{} { + runID = strings.TrimSpace(runID) + if runID == "" { + return data + } + switch v := data.(type) { + case map[string]interface{}: + if existing, ok := v["runId"]; !ok || strings.TrimSpace(fmt.Sprint(existing)) == "" { + v["runId"] = runID + } + return v + default: + return data + } +} diff --git a/internal/multiagent/eino_run_trace_test.go b/internal/multiagent/eino_run_trace_test.go new file mode 100644 index 00000000..b1349e66 --- /dev/null +++ b/internal/multiagent/eino_run_trace_test.go @@ -0,0 +1,39 @@ +package multiagent + +import "testing" + +func TestWithEinoRunIDProgressAddsRunIDToMapData(t *testing.T) { + var gotType, gotMessage string + var gotData interface{} + progress := withEinoRunIDProgress("run-1", func(eventType, message string, data interface{}) { + gotType = eventType + gotMessage = message + gotData = data + }) + + progress("progress", "hello", map[string]interface{}{"source": "eino"}) + + if gotType != "progress" || gotMessage != "hello" { + t.Fatalf("event = (%q, %q)", gotType, gotMessage) + } + m, ok := gotData.(map[string]interface{}) + if !ok { + t.Fatalf("data type = %T", gotData) + } + if m["runId"] != "run-1" || m["source"] != "eino" { + t.Fatalf("data = %#v", m) + } +} + +func TestWithEinoRunIDProgressPreservesExistingRunID(t *testing.T) { + var got map[string]interface{} + progress := withEinoRunIDProgress("outer-run", func(_, _ string, data interface{}) { + got, _ = data.(map[string]interface{}) + }) + + progress("progress", "", map[string]interface{}{"runId": "inner-run"}) + + if got["runId"] != "inner-run" { + t.Fatalf("runId = %q, want inner-run", got["runId"]) + } +} diff --git a/internal/multiagent/eino_run_usage_accumulator.go b/internal/multiagent/eino_run_usage_accumulator.go new file mode 100644 index 00000000..8daa5b53 --- /dev/null +++ b/internal/multiagent/eino_run_usage_accumulator.go @@ -0,0 +1,150 @@ +package multiagent + +import ( + "sync" + + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" +) + +type einoRunUsageSummary struct { + ModelCalls int + PromptTokens int + CompletionTokens int + TotalTokens int + CachedTokens int + ReasoningTokens int +} + +type einoRunUsageAccumulator struct { + mu sync.Mutex + summary einoRunUsageSummary + emitted bool +} + +func newEinoRunUsageAccumulator() *einoRunUsageAccumulator { + return &einoRunUsageAccumulator{} +} + +func (a *einoRunUsageAccumulator) AddMessage(msg *schema.Message) bool { + if msg == nil || msg.ResponseMeta == nil || msg.ResponseMeta.Usage == nil { + return false + } + return a.AddUsage(msg.ResponseMeta.Usage) +} + +func (a *einoRunUsageAccumulator) AddUsage(usage *schema.TokenUsage) bool { + if a == nil || usage == nil || tokenUsageEmpty(usage) { + return false + } + a.mu.Lock() + defer a.mu.Unlock() + a.summary.ModelCalls++ + a.summary.PromptTokens += usage.PromptTokens + a.summary.CompletionTokens += usage.CompletionTokens + a.summary.TotalTokens += usage.TotalTokens + a.summary.CachedTokens += usage.PromptTokenDetails.CachedTokens + a.summary.ReasoningTokens += usage.CompletionTokensDetails.ReasoningTokens + return true +} + +func (a *einoRunUsageAccumulator) Summary() einoRunUsageSummary { + if a == nil { + return einoRunUsageSummary{} + } + a.mu.Lock() + defer a.mu.Unlock() + return a.summary +} + +func (a *einoRunUsageAccumulator) EmitOnce( + conversationID string, + orchestration string, + reason string, + progress func(eventType, message string, data interface{}), + logger *zap.Logger, +) bool { + if a == nil { + return false + } + a.mu.Lock() + if a.emitted || a.summary.ModelCalls == 0 { + a.mu.Unlock() + return false + } + a.emitted = true + s := a.summary + a.mu.Unlock() + + data := map[string]interface{}{ + "conversationId": conversationID, + "source": "eino", + "orchestration": orchestration, + "reason": reason, + "modelCalls": s.ModelCalls, + "promptTokens": s.PromptTokens, + "completionTokens": s.CompletionTokens, + "totalTokens": s.TotalTokens, + "cachedTokens": s.CachedTokens, + "reasoningTokens": s.ReasoningTokens, + } + if progress != nil { + progress("eino_usage_summary", "Eino token usage summary", data) + } + if logger != nil { + logger.Info("eino token usage summary", + zap.String("conversationId", conversationID), + zap.String("orchestration", orchestration), + zap.String("reason", reason), + zap.Int("modelCalls", s.ModelCalls), + zap.Int("promptTokens", s.PromptTokens), + zap.Int("completionTokens", s.CompletionTokens), + zap.Int("totalTokens", s.TotalTokens), + zap.Int("cachedTokens", s.CachedTokens), + zap.Int("reasoningTokens", s.ReasoningTokens), + ) + } + return true +} + +func maxEinoTokenUsage(dst *schema.TokenUsage, src *schema.TokenUsage) *schema.TokenUsage { + if src == nil { + return dst + } + if dst == nil { + return cloneEinoTokenUsage(src) + } + if src.PromptTokens > dst.PromptTokens { + dst.PromptTokens = src.PromptTokens + } + if src.CompletionTokens > dst.CompletionTokens { + dst.CompletionTokens = src.CompletionTokens + } + if src.TotalTokens > dst.TotalTokens { + dst.TotalTokens = src.TotalTokens + } + if src.PromptTokenDetails.CachedTokens > dst.PromptTokenDetails.CachedTokens { + dst.PromptTokenDetails.CachedTokens = src.PromptTokenDetails.CachedTokens + } + if src.CompletionTokensDetails.ReasoningTokens > dst.CompletionTokensDetails.ReasoningTokens { + dst.CompletionTokensDetails.ReasoningTokens = src.CompletionTokensDetails.ReasoningTokens + } + return dst +} + +func cloneEinoTokenUsage(src *schema.TokenUsage) *schema.TokenUsage { + if src == nil { + return nil + } + out := *src + return &out +} + +func tokenUsageEmpty(u *schema.TokenUsage) bool { + return u == nil || + (u.PromptTokens == 0 && + u.CompletionTokens == 0 && + u.TotalTokens == 0 && + u.PromptTokenDetails.CachedTokens == 0 && + u.CompletionTokensDetails.ReasoningTokens == 0) +} diff --git a/internal/multiagent/eino_run_usage_accumulator_test.go b/internal/multiagent/eino_run_usage_accumulator_test.go new file mode 100644 index 00000000..57c23779 --- /dev/null +++ b/internal/multiagent/eino_run_usage_accumulator_test.go @@ -0,0 +1,81 @@ +package multiagent + +import ( + "testing" + + "github.com/cloudwego/eino/schema" +) + +func TestEinoRunUsageAccumulatorSumsModelCalls(t *testing.T) { + acc := newEinoRunUsageAccumulator() + acc.AddUsage(&schema.TokenUsage{ + PromptTokens: 10, + CompletionTokens: 4, + TotalTokens: 14, + PromptTokenDetails: schema.PromptTokenDetails{ + CachedTokens: 3, + }, + CompletionTokensDetails: schema.CompletionTokensDetails{ + ReasoningTokens: 2, + }, + }) + msg := schema.AssistantMessage("ok", nil) + msg.ResponseMeta = &schema.ResponseMeta{Usage: &schema.TokenUsage{ + PromptTokens: 7, + CompletionTokens: 5, + TotalTokens: 12, + CompletionTokensDetails: schema.CompletionTokensDetails{ + ReasoningTokens: 1, + }, + }} + acc.AddMessage(msg) + + got := acc.Summary() + if got.ModelCalls != 2 || got.PromptTokens != 17 || got.CompletionTokens != 9 || got.TotalTokens != 26 || got.CachedTokens != 3 || got.ReasoningTokens != 3 { + t.Fatalf("summary = %#v", got) + } +} + +func TestEinoRunUsageAccumulatorEmitOnce(t *testing.T) { + acc := newEinoRunUsageAccumulator() + acc.AddUsage(&schema.TokenUsage{PromptTokens: 1, CompletionTokens: 2, TotalTokens: 3}) + var events []map[string]interface{} + progress := func(eventType, _ string, data interface{}) { + if eventType != "eino_usage_summary" { + return + } + if m, ok := data.(map[string]interface{}); ok { + events = append(events, m) + } + } + + if !acc.EmitOnce("conv-1", "deep", "final", progress, nil) { + t.Fatal("first emit should return true") + } + if acc.EmitOnce("conv-1", "deep", "partial", progress, nil) { + t.Fatal("second emit should return false") + } + if len(events) != 1 { + t.Fatalf("events = %#v, want one usage summary", events) + } + if events[0]["conversationId"] != "conv-1" || events[0]["orchestration"] != "deep" || events[0]["reason"] != "final" || events[0]["totalTokens"] != 3 { + t.Fatalf("event = %#v", events[0]) + } +} + +func TestMaxEinoTokenUsageUsesLargestStreamChunkValues(t *testing.T) { + var got *schema.TokenUsage + got = maxEinoTokenUsage(got, &schema.TokenUsage{PromptTokens: 10, CompletionTokens: 2, TotalTokens: 12}) + got = maxEinoTokenUsage(got, &schema.TokenUsage{ + PromptTokens: 9, + CompletionTokens: 5, + TotalTokens: 14, + CompletionTokensDetails: schema.CompletionTokensDetails{ + ReasoningTokens: 3, + }, + }) + + if got.PromptTokens != 10 || got.CompletionTokens != 5 || got.TotalTokens != 14 || got.CompletionTokensDetails.ReasoningTokens != 3 { + t.Fatalf("usage = %#v", got) + } +} diff --git a/internal/multiagent/eino_runner_iterator_starter.go b/internal/multiagent/eino_runner_iterator_starter.go new file mode 100644 index 00000000..cb1d9356 --- /dev/null +++ b/internal/multiagent/eino_runner_iterator_starter.go @@ -0,0 +1,94 @@ +package multiagent + +import ( + "context" + "sync/atomic" + + "github.com/cloudwego/eino/adk" + "go.uber.org/zap" +) + +type einoRunnerControl interface { + Run(context.Context, []adk.Message, ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] + Resume(context.Context, string, ...adk.AgentRunOption) (*adk.AsyncIterator[*adk.AgentEvent], error) +} + +type einoRunnerIteratorStarterConfig struct { + Context context.Context + ConversationID string + OrchMode string + Logger *zap.Logger + Runner einoRunnerControl + CheckPointID string + NativeCancelCause *atomic.Value + UnregisterAgentCancel *func() + RuntimeCancelRegistrar AgentRuntimeCancelRegistrar +} + +type einoRunnerIteratorStarter struct { + cfg einoRunnerIteratorStarterConfig +} + +func newEinoRunnerIteratorStarter(cfg einoRunnerIteratorStarterConfig) *einoRunnerIteratorStarter { + return &einoRunnerIteratorStarter{cfg: cfg} +} + +func (s *einoRunnerIteratorStarter) Start(runMsgs []adk.Message) *adk.AsyncIterator[*adk.AgentEvent] { + if s == nil || s.cfg.Runner == nil { + return nil + } + opts := s.newRunOptions() + if s.cfg.CheckPointID != "" { + opts = append(opts, adk.WithCheckPointID(s.cfg.CheckPointID)) + } + return s.cfg.Runner.Run(s.cfg.Context, runMsgs, opts...) +} + +func (s *einoRunnerIteratorStarter) Resume(checkPointID string) (*adk.AsyncIterator[*adk.AgentEvent], error) { + if s == nil || s.cfg.Runner == nil { + return nil, nil + } + return s.cfg.Runner.Resume(s.cfg.Context, checkPointID, s.newRunOptions()...) +} + +func (s *einoRunnerIteratorStarter) newRunOptions() []adk.AgentRunOption { + cancelOpt, cancelFn := adk.WithCancel() + callAndClearUnregister(s.cfg.UnregisterAgentCancel) + if s.cfg.RuntimeCancelRegistrar != nil && s.cfg.UnregisterAgentCancel != nil { + *s.cfg.UnregisterAgentCancel = s.cfg.RuntimeCancelRegistrar(func(cause error) bool { + s.storeNativeCancelCause(cause) + waitErr, submitted, handled := requestEinoNativeAgentCancel(cancelFn, cause) + s.logNativeCancelRequest(cause, waitErr, submitted, handled) + return handled + }) + } + return []adk.AgentRunOption{cancelOpt} +} + +func (s *einoRunnerIteratorStarter) storeNativeCancelCause(cause error) { + if s == nil || s.cfg.NativeCancelCause == nil || cause == nil { + return + } + s.cfg.NativeCancelCause.Store(cause) +} + +func (s *einoRunnerIteratorStarter) logNativeCancelRequest(cause error, waitErr error, submitted bool, handled bool) { + if s == nil || s.cfg.Logger == nil { + return + } + fields := []zap.Field{ + zap.String("conversation_id", s.cfg.ConversationID), + zap.String("orchestration", s.cfg.OrchMode), + zap.Bool("submitted", submitted), + zap.Bool("handled", handled), + } + if cause != nil { + fields = append(fields, zap.Error(cause)) + } + if waitErr != nil { + fields = append(fields, zap.NamedError("cancel_wait_error", waitErr)) + s.cfg.Logger.Debug("eino native cancel requested", fields...) + } else { + s.cfg.Logger.Info("eino native cancel requested", fields...) + } +} diff --git a/internal/multiagent/eino_runner_iterator_starter_test.go b/internal/multiagent/eino_runner_iterator_starter_test.go new file mode 100644 index 00000000..adcf7e6a --- /dev/null +++ b/internal/multiagent/eino_runner_iterator_starter_test.go @@ -0,0 +1,116 @@ +package multiagent + +import ( + "context" + "errors" + "sync/atomic" + "testing" + + "github.com/cloudwego/eino/adk" +) + +type fakeRunnerControl struct { + runMessages []adk.Message + runOpts int + resumeID string + resumeOpts int + resumeErr error +} + +func (f *fakeRunnerControl) Run(_ context.Context, messages []adk.Message, opts ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] { + f.runMessages = messages + f.runOpts = len(opts) + iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + gen.Close() + return iter +} + +func (f *fakeRunnerControl) Resume(_ context.Context, checkPointID string, opts ...adk.AgentRunOption) (*adk.AsyncIterator[*adk.AgentEvent], error) { + f.resumeID = checkPointID + f.resumeOpts = len(opts) + iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + gen.Close() + return iter, f.resumeErr +} + +func TestEinoRunnerIteratorStarterStartAddsCancelAndCheckpoint(t *testing.T) { + runner := &fakeRunnerControl{} + var cancelPush func(error) bool + var nativeCancelCause atomic.Value + oldUnregistered := false + newUnregistered := false + unregister := func() { oldUnregistered = true } + + iter := newEinoRunnerIteratorStarter(einoRunnerIteratorStarterConfig{ + Context: context.Background(), + Runner: runner, + CheckPointID: "cp-1", + NativeCancelCause: &nativeCancelCause, + UnregisterAgentCancel: &unregister, + RuntimeCancelRegistrar: func(push func(error) bool) func() { + cancelPush = push + return func() { newUnregistered = true } + }, + }).Start([]adk.Message{}) + + if iter == nil { + t.Fatal("iterator should be created") + } + if runner.runOpts != 2 { + t.Fatalf("run opts = %d, want cancel + checkpoint", runner.runOpts) + } + if !oldUnregistered { + t.Fatal("old unregister should be called before binding a new cancel hook") + } + if cancelPush == nil { + t.Fatal("cancel hook should be registered") + } + stopErr := errors.New("stop") + if cancelPush(stopErr) { + t.Fatal("unbound fake runner cancel should not report handled") + } + if got, _ := nativeCancelCause.Load().(error); !errors.Is(got, stopErr) { + t.Fatalf("native cancel cause = %v, want %v", got, stopErr) + } + unregister() + if !newUnregistered { + t.Fatal("new unregister should replace old unregister") + } +} + +func TestEinoRunnerIteratorStarterResumeUsesCancelOnly(t *testing.T) { + runner := &fakeRunnerControl{} + + iter, err := newEinoRunnerIteratorStarter(einoRunnerIteratorStarterConfig{ + Context: context.Background(), + Runner: runner, + CheckPointID: "fresh-run-checkpoint", + }).Resume("resume-cp") + + if err != nil { + t.Fatalf("resume err = %v", err) + } + if iter == nil { + t.Fatal("iterator should be created") + } + if runner.resumeID != "resume-cp" { + t.Fatalf("resume id = %q, want resume-cp", runner.resumeID) + } + if runner.resumeOpts != 1 { + t.Fatalf("resume opts = %d, want cancel only", runner.resumeOpts) + } +} + +func TestEinoRunnerIteratorStarterResumePropagatesError(t *testing.T) { + resumeErr := errors.New("resume failed") + runner := &fakeRunnerControl{resumeErr: resumeErr} + + _, err := newEinoRunnerIteratorStarter(einoRunnerIteratorStarterConfig{ + Context: context.Background(), + Runner: runner, + }).Resume("resume-cp") + + if !errors.Is(err, resumeErr) { + t.Fatalf("resume err = %v, want %v", err, resumeErr) + } +} diff --git a/internal/multiagent/eino_single_runner.go b/internal/multiagent/eino_single_runner.go new file mode 100644 index 00000000..6a1b494b --- /dev/null +++ b/internal/multiagent/eino_single_runner.go @@ -0,0 +1,233 @@ +package multiagent + +import ( + "context" + "fmt" + "sync" + + "cyberstrike-ai/internal/agent" + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/einomcp" + "cyberstrike-ai/internal/project" + "cyberstrike-ai/internal/reasoning" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/compose" + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" +) + +// einoSingleAgentName 与 ChatModelAgent.Name 一致,供流式事件映射主对话区。 +const einoSingleAgentName = "cyberstrike-eino-single" + +// RunEinoSingleChatModelAgent 使用 Eino TypedChatModelAgent[*schema.AgenticMessage] + adk.NewRunner.Run(官方 Quick Start 的 Query 同属 Runner API;此处用历史 + 用户消息切片等价于多轮 Query)。 +// 与 RunDeepAgent 共享 runEinoADKAgentLoop 的 SSE 映射与 MCP 桥。 +func RunEinoSingleChatModelAgent( + ctx context.Context, + appCfg *config.Config, + ma *config.MultiAgentConfig, + ag *agent.Agent, + db *database.DB, + logger *zap.Logger, + conversationID string, + projectID string, + userMessage string, + history []agent.ChatMessage, + roleTools []string, + progress func(eventType, message string, data interface{}), + reasoningClient *reasoning.ClientIntent, + systemPromptExtra string, +) (*RunResult, error) { + if appCfg == nil || ag == nil { + return nil, fmt.Errorf("eino single: 配置或 Agent 为空") + } + if ma == nil { + return nil, fmt.Errorf("eino single: multi_agent 配置为空") + } + runtimeUserMessage := prepareLatestUserMessageForModel(userMessage, appCfg, &ma.EinoMiddleware, conversationID, logger) + + einoLoc, einoSkillMW, einoFSTools, skillsRoot, einoErr := prepareEinoAgenticSkills(ctx, appCfg.SkillsDir, ma, logger) + if einoErr != nil { + return nil, einoErr + } + + holder := &einomcp.ConversationHolder{} + holder.Set(conversationID) + + var mcpIDsMu sync.Mutex + var mcpIDs []string + mcpExecBinder := NewMCPExecutionBinder() + recorder := func(id, toolCallID string) { + if id == "" { + return + } + mcpExecBinder.Bind(toolCallID, id) + mcpIDsMu.Lock() + mcpIDs = append(mcpIDs, id) + mcpIDsMu.Unlock() + } + + snapshotMCPIDs := func() []string { + mcpIDsMu.Lock() + defer mcpIDsMu.Unlock() + out := make([]string, len(mcpIDs)) + copy(out, mcpIDs) + return out + } + + toolInvokeNotify := einomcp.NewToolInvokeNotifyHolder() + einoExecBegin, einoExecAppendPartial, einoExecRegisterCancel, einoExecUnregisterCancel, einoExecFinish := newEinoExecuteMonitorCallbacks(ctx, ag, recorder) + mainDefs := ag.ToolsForRole(roleTools) + mainTools, err := einomcp.ToolsFromDefinitions(ag, holder, mainDefs, recorder, nil, toolInvokeNotify, einoSingleAgentName) + if err != nil { + return nil, err + } + + mainToolsForCfg, mainOrchestratorPre, singleToolSearchActive, err := prependEinoAgenticMiddlewares(ctx, &ma.EinoMiddleware, einoMWMain, mainTools, einoLoc, skillsRoot, conversationID, projectID, logger) + if err != nil { + return nil, fmt.Errorf("eino single eino 中间件: %w", err) + } + + baseHTTPClient := newEinoBaseHTTPClient() + agenticModelFactory := newEinoOpenAIAgenticChatModelFactory(baseHTTPClient, reasoningClient, logger) + mainModel, err := agenticModelFactory(ctx, appCfg.OpenAI, einoModelModeNormal) + if err != nil { + return nil, fmt.Errorf("eino single agentic 模型: %w", err) + } + modelRetryCfg := newEinoAgenticModelRetryConfig(&ma.EinoMiddleware, logger, "eino_single") + modelFailoverCfg, err := newEinoAgenticModelFailoverConfig(ctx, appCfg, &ma.EinoMiddleware, einoModelModeNormal, agenticModelFactory, logger, "eino_single", progress, "eino_single", conversationID) + if err != nil { + return nil, err + } + logEinoAgenticModelGate( + logger, + "eino_single", + "eino_single", + evaluateEinoAgenticModelGate(agenticModelGateFactory(agenticModelFactory, appCfg.OpenAI, einoModelModeNormal), einoAgenticRuntimeSupportV0914()), + ) + + mainSumMw, err := newEinoAgenticSummarizationMiddleware(ctx, mainModel, appCfg, &ma.EinoMiddleware, conversationID, db, projectID, logger) + if err != nil { + return nil, fmt.Errorf("eino single agentic summarization: %w", err) + } + + modelFacingTrace := newModelFacingTraceHolder() + + handlers := make([]adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], 0, 8) + if len(mainOrchestratorPre) > 0 { + handlers = append(handlers, mainOrchestratorPre...) + } + if einoSkillMW != nil { + if einoFSTools && einoLoc != nil { + fsMw, fsErr := subAgentAgenticFilesystemMiddleware(ctx, einoLoc, toolInvokeNotify, einoSingleAgentName, einoExecBegin, einoExecAppendPartial, einoExecRegisterCancel, einoExecUnregisterCancel, einoExecFinish, agentToolTimeoutMinutes(appCfg), agentToolWaitTimeoutSeconds(appCfg), agentShellNoOutputTimeoutSeconds(appCfg), nil) + if fsErr != nil { + return nil, fmt.Errorf("eino single filesystem 中间件: %w", fsErr) + } + handlers = append(handlers, fsMw) + } + handlers = append(handlers, einoSkillMW) + } + handlers = appendEinoAgenticChatModelTailMiddlewares(handlers, einoChatModelTailConfig{ + logger: logger, + phase: "eino_single", + agenticSummarization: mainSumMw, + modelName: appCfg.OpenAI.Model, + maxTotalTokens: appCfg.OpenAI.MaxTotalTokens, + toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware), + conversationID: conversationID, + trace: modelFacingTrace, + middlewareConfig: &ma.EinoMiddleware, + }) + + maxIter := agentMaxIterations(appCfg) + + mainToolsCfg := adk.ToolsConfig{ + ToolsNodeConfig: compose.ToolsNodeConfig{ + Tools: mainToolsForCfg, + UnknownToolsHandler: einomcp.UnknownToolReminderHandler(), + ToolCallMiddlewares: []compose.ToolMiddleware{ + modelOutputExecutionGuardMiddleware(), + localToolRBACMiddleware(), + hitlToolCallMiddleware(), + softRecoveryToolMiddleware(), + }, + }, + EmitInternalEvents: true, + } + ins := project.AppendSystemPromptBlock(ag.EinoSingleAgentSystemInstruction(), systemPromptExtra) + ins = project.AppendVisionImageAnalysisIfReady(ins, appCfg.Vision.Ready()) + ins = injectToolNamesOnlyInstruction(ctx, ins, mainTools, singleToolSearchActive) + if logger != nil { + names := collectToolNames(ctx, mainTools) + mountedNames := collectToolNames(ctx, mainToolsForCfg) + logger.Info("eino tool-name injection", + zap.String("scope", "eino_single"), + zap.Int("tool_names", len(names)), + zap.Int("mounted_tool_names", len(mountedNames)), + zap.Bool("tool_search_middleware", singleToolSearchActive), + ) + } + + chatCfg := einoAgenticChatModelAgentConfig{ + Name: einoSingleAgentName, + Description: "Eino ADK ChatModelAgent with MCP tools for authorized security testing.", + Instruction: ins, + GenModelInput: literalAgenticInstructionGenModelInput, + Model: mainModel, + ToolsConfig: mainToolsCfg, + MaxIterations: maxIter, + Handlers: handlers, + ModelRetryConfig: modelRetryCfg, + ModelFailoverConfig: modelFailoverCfg, + } + outKey, _ := deepExtrasFromConfig(ma) + if outKey != "" { + chatCfg.OutputKey = outKey + } + + chatAgent, err := newEinoAgenticChatModelAgentAdapter(ctx, chatCfg) + if err != nil { + return nil, fmt.Errorf("eino single Agentic ChatModelAgent: %w", err) + } + + baseMsgs := historyToMessages(history, appCfg, &ma.EinoMiddleware) + baseMsgs = appendUserMessageIfNeeded(baseMsgs, runtimeUserMessage) + + streamsMainAssistant := func(agent string) bool { + return agent == "" || agent == einoSingleAgentName + } + einoRoleTag := func(agent string) string { + _ = agent + return "orchestrator" + } + + return runEinoADKAgentLoop(ctx, &einoADKRunLoopArgs{ + OrchMode: "eino_single", + OrchestratorName: einoSingleAgentName, + ConversationID: conversationID, + Progress: progress, + Logger: logger, + SnapshotMCPIDs: snapshotMCPIDs, + StreamsMainAssistant: streamsMainAssistant, + EinoRoleTag: einoRoleTag, + CheckpointDir: ma.EinoMiddleware.CheckpointDir, + RunRetryMaxAttempts: RunRetryMaxAttemptsFromConfig(&ma.EinoMiddleware), + RunRetryMaxBackoffSec: int(einoRunRetryMaxBackoffFromConfig(&ma.EinoMiddleware).Seconds()), + McpIDsMu: &mcpIDsMu, + McpIDs: &mcpIDs, + FilesystemMonitorAgent: ag, + FilesystemMonitorRecord: recorder, + MCPExecutionBinder: mcpExecBinder, + ToolInvokeNotify: toolInvokeNotify, + DA: chatAgent, + ModelFacingTrace: modelFacingTrace, + EinoCallbacks: &ma.EinoCallbacks, + MaxTotalTokens: appCfg.OpenAI.MaxTotalTokens, + ToolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware), + ModelName: appCfg.OpenAI.Model, + MiddlewareConfig: &ma.EinoMiddleware, + EmptyResponseMessage: "(Eino ADK single-agent session completed but no assistant text was captured. Check process details or logs.) " + + "(Eino ADK 单代理会话已完成,但未捕获到助手文本输出。请查看过程详情或日志。)", + }, baseMsgs) +} diff --git a/internal/multiagent/eino_skills.go b/internal/multiagent/eino_skills.go new file mode 100644 index 00000000..9cbaec95 --- /dev/null +++ b/internal/multiagent/eino_skills.go @@ -0,0 +1,163 @@ +package multiagent + +import ( + "context" + "fmt" + "os" + "path/filepath" + "strings" + + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/einomcp" + "cyberstrike-ai/internal/security" + + localbk "github.com/cloudwego/eino-ext/adk/backend/local" + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/adk/middlewares/filesystem" + "github.com/cloudwego/eino/adk/middlewares/skill" + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" +) + +func prepareEinoAgenticSkills( + ctx context.Context, + skillsDir string, + ma *config.MultiAgentConfig, + logger *zap.Logger, +) (loc *localbk.Local, skillMW adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], fsTools bool, skillsRoot string, err error) { + if ma == nil { + return nil, nil, false, "", nil + } + needLocalBackend := ma.EinoMiddleware.ReductionEnable + newLocalBackend := func() (*localbk.Local, error) { + backend, backendErr := localbk.NewBackend(ctx, &localbk.Config{}) + if backendErr != nil { + return nil, fmt.Errorf("eino local backend: %w", backendErr) + } + return backend, nil + } + if ma.EinoSkills.Disable { + if !needLocalBackend { + return nil, nil, false, "", nil + } + loc, err = newLocalBackend() + return loc, nil, false, "", err + } + root := strings.TrimSpace(skillsDir) + if root == "" { + if logger != nil { + logger.Warn("eino agentic skills: skills_dir empty, skip") + } + if !needLocalBackend { + return nil, nil, false, "", nil + } + loc, err = newLocalBackend() + return loc, nil, false, "", err + } + abs, err := filepath.Abs(root) + if err != nil { + return nil, nil, false, "", fmt.Errorf("skills_dir abs: %w", err) + } + if st, err := os.Stat(abs); err != nil || !st.IsDir() { + if logger != nil { + logger.Warn("eino agentic skills: directory missing, skip", zap.String("dir", abs), zap.Error(err)) + } + if !needLocalBackend { + return nil, nil, false, "", nil + } + loc, err = newLocalBackend() + return loc, nil, false, "", err + } + + loc, err = newLocalBackend() + if err != nil { + return nil, nil, false, "", err + } + + skillBE, err := skill.NewBackendFromFilesystem(ctx, &skill.BackendFromFilesystemConfig{ + Backend: loc, + BaseDir: abs, + }) + if err != nil { + return nil, nil, false, "", fmt.Errorf("eino agentic skill filesystem backend: %w", err) + } + + sc := &skill.TypedConfig[*schema.AgenticMessage]{Backend: skillBE} + if name := strings.TrimSpace(ma.EinoSkills.SkillToolName); name != "" { + sc.SkillToolName = &name + } + skillMW, err = skill.NewTyped[*schema.AgenticMessage](ctx, sc) + if err != nil { + return nil, nil, false, "", fmt.Errorf("eino agentic skill middleware: %w", err) + } + + fsTools = ma.EinoSkills.EinoSkillFilesystemToolsEffective() + return loc, skillMW, fsTools, abs, nil +} + +func subAgentAgenticFilesystemMiddleware( + ctx context.Context, + loc *localbk.Local, + invokeNotify *einomcp.ToolInvokeNotifyHolder, + einoAgentName string, + beginMonitor func(toolCallID, command string) string, + appendPartialMonitor func(executionID, toolCallID, chunk string), + registerCancelMonitor func(executionID string, cancel context.CancelFunc), + unregisterCancelMonitor func(executionID string), + finishMonitor func(executionID, toolCallID, command, stdout string, success bool, invokeErr error), + toolTimeoutMinutes int, + toolWaitTimeoutSeconds int, + shellNoOutputTimeoutSec int, + outputChunk func(toolName, toolCallID, chunk string), +) (adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], error) { + if loc == nil { + return nil, nil + } + return filesystem.NewTyped[*schema.AgenticMessage](ctx, &filesystem.MiddlewareConfig{ + Backend: loc, + StreamingShell: &einoStreamingShellWrap{ + inner: security.NewEinoStreamingShell(), + invokeNotify: invokeNotify, + einoAgentName: strings.TrimSpace(einoAgentName), + outputChunk: outputChunk, + beginMonitor: beginMonitor, + appendPartialMonitor: appendPartialMonitor, + registerCancelMonitor: registerCancelMonitor, + unregisterCancelMonitor: unregisterCancelMonitor, + finishMonitor: finishMonitor, + toolTimeoutMinutes: toolTimeoutMinutes, + toolWaitTimeoutSeconds: toolWaitTimeoutSeconds, + shellNoOutputTimeoutSec: shellNoOutputTimeoutSec, + }, + }) +} + +// agentToolTimeoutMinutes 返回 agent.tool_timeout_minutes(与 executeToolViaMCP 一致);cfg 为 nil 时 0。 +func agentToolTimeoutMinutes(cfg *config.Config) int { + if cfg == nil { + return 0 + } + return cfg.Agent.ToolTimeoutMinutes +} + +func agentToolWaitTimeoutSeconds(cfg *config.Config) int { + if cfg == nil { + return 0 + } + return cfg.Agent.ToolWaitTimeoutSeconds +} + +// agentShellNoOutputTimeoutSeconds:0=默认 300s(5 分钟);-1=关闭;>0=自定义秒数。 +func agentShellNoOutputTimeoutSeconds(cfg *config.Config) int { + if cfg == nil { + return 300 + } + v := cfg.Agent.ShellNoOutputTimeoutSeconds + if v < 0 { + return 0 + } + if v == 0 { + return 300 + } + return v +} diff --git a/internal/multiagent/eino_skills_reduction_test.go b/internal/multiagent/eino_skills_reduction_test.go new file mode 100644 index 00000000..10345d0e --- /dev/null +++ b/internal/multiagent/eino_skills_reduction_test.go @@ -0,0 +1,27 @@ +package multiagent + +import ( + "context" + "testing" + + "cyberstrike-ai/internal/config" +) + +func TestPrepareEinoAgenticSkillsStillCreatesReductionBackendWhenSkillsDisabled(t *testing.T) { + ma := &config.MultiAgentConfig{ + EinoSkills: config.MultiAgentEinoSkillsConfig{Disable: true}, + EinoMiddleware: config.MultiAgentEinoMiddlewareConfig{ + ReductionEnable: true, + }, + } + loc, skillMW, fsTools, skillsRoot, err := prepareEinoAgenticSkills(context.Background(), "", ma, nil) + if err != nil { + t.Fatal(err) + } + if loc == nil { + t.Fatal("agentic reduction backend must exist even when Skills are disabled") + } + if skillMW != nil || fsTools || skillsRoot != "" { + t.Fatalf("Agentic Skills unexpectedly enabled: mw=%v fs=%v root=%q", skillMW, fsTools, skillsRoot) + } +} diff --git a/internal/multiagent/eino_stream_tool_call_completion_handler.go b/internal/multiagent/eino_stream_tool_call_completion_handler.go new file mode 100644 index 00000000..0bd4e5fe --- /dev/null +++ b/internal/multiagent/eino_stream_tool_call_completion_handler.go @@ -0,0 +1,49 @@ +package multiagent + +import "github.com/cloudwego/eino/schema" + +type einoStreamToolCallCompletionHandlerConfig struct { + ConversationID string + OrchMode string + Progress func(eventType, message string, data interface{}) + RunProgress *einoRunProgressTracker + RunMessages *einoRunMessageAccumulator + MarkPending func(toolCallPendingInfo) +} + +type einoStreamToolCallCompletionHandler struct { + conversationID string + orchMode string + progress func(eventType, message string, data interface{}) + runProgress *einoRunProgressTracker + runMessages *einoRunMessageAccumulator + markPending func(toolCallPendingInfo) +} + +func newEinoStreamToolCallCompletionHandler(cfg einoStreamToolCallCompletionHandlerConfig) *einoStreamToolCallCompletionHandler { + return &einoStreamToolCallCompletionHandler{ + conversationID: cfg.ConversationID, + orchMode: cfg.OrchMode, + progress: cfg.Progress, + runProgress: cfg.RunProgress, + runMessages: cfg.RunMessages, + markPending: cfg.MarkPending, + } +} + +func (h *einoStreamToolCallCompletionHandler) Complete(fragments []schema.ToolCall, agentName string) *schema.Message { + if h == nil { + return nil + } + var lastToolChunk *schema.Message + if merged := mergeStreamingToolCallFragments(fragments); len(merged) > 0 { + lastToolChunk = mergeMessageToolCalls(&schema.Message{ToolCalls: merged}) + } + if h.runProgress != nil { + h.runProgress.EmitToolCalls(lastToolChunk, agentName, h.markPending) + } + if lastToolChunk != nil && len(lastToolChunk.ToolCalls) > 0 && h.runMessages != nil { + h.runMessages.AppendAssistantToolCalls(lastToolChunk.ToolCalls) + } + return lastToolChunk +} diff --git a/internal/multiagent/eino_summarize_payload.go b/internal/multiagent/eino_summarize_payload.go new file mode 100644 index 00000000..254e7849 --- /dev/null +++ b/internal/multiagent/eino_summarize_payload.go @@ -0,0 +1,12 @@ +package multiagent + +import ( + copenai "cyberstrike-ai/internal/openai" +) + +// stripReasoningFromSummarizationPayload removes thinking / reasoning fields from a +// chat-completions JSON body. Applied only to summarization Generate calls via +// model.ModelOptions on the shared ChatModel — main-agent requests are unchanged. +func stripReasoningFromSummarizationPayload(rawBody []byte) ([]byte, error) { + return copenai.StripReasoningFromChatCompletionBody(rawBody) +}