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_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 index a0ad6829..051984e2 100644 --- a/internal/multiagent/eino_orchestration.go +++ b/internal/multiagent/eino_orchestration.go @@ -19,7 +19,7 @@ import ( // PlanExecuteRootArgs 构建 Eino adk/prebuilt/planexecute 根 Agent 所需参数。 type PlanExecuteRootArgs struct { MainToolCallingModel model.ToolCallingChatModel - ExecModel model.ToolCallingChatModel + AgenticExecModel model.AgenticModel OrchInstruction string ToolsCfg adk.ToolsConfig ExecMaxIter int @@ -34,17 +34,16 @@ type PlanExecuteRootArgs struct { Logger *zap.Logger // ModelName is used for model input token estimation logs. ModelName string - // ExecPreMiddlewares 是由 prependEinoMiddlewares 构建的前置中间件(patchtoolcalls, reduction, toolsearch, plantask), - // 与 Deep/Supervisor 主代理的 mainOrchestratorPre 一致。 - ExecPreMiddlewares []adk.ChatModelAgentMiddleware - // SkillMiddleware 是 Eino 官方 skill 渐进式披露中间件(可选)。 - SkillMiddleware adk.ChatModelAgentMiddleware - // FilesystemMiddleware 是 Eino filesystem 中间件,当 eino_skills.filesystem_tools 启用时提供本机文件读写与 Shell 能力(可选)。 - FilesystemMiddleware adk.ChatModelAgentMiddleware + // 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 + ModelFacingTrace *modelFacingTraceHolder + AgenticModelRetryConfig *adk.TypedModelRetryConfig[*schema.AgenticMessage] + AgenticModelFailoverConfig *adk.ModelFailoverConfig[*schema.AgenticMessage] } // NewPlanExecuteRoot 返回 plan → execute → replan 预置编排根节点(与 Deep / Supervisor 并列)。 @@ -52,7 +51,7 @@ func NewPlanExecuteRoot(ctx context.Context, a *PlanExecuteRootArgs) (adk.Resuma if a == nil { return nil, fmt.Errorf("plan_execute: args 为空") } - if a.MainToolCallingModel == nil || a.ExecModel == nil { + if a.MainToolCallingModel == nil || a.AgenticExecModel == nil { return nil, fmt.Errorf("plan_execute: 模型为空") } tcm, ok := interface{}(a.MainToolCallingModel).(model.ToolCallingChatModel) @@ -79,16 +78,16 @@ func NewPlanExecuteRoot(ctx context.Context, a *PlanExecuteRootArgs) (adk.Resuma return nil, fmt.Errorf("plan_execute replanner: %w", err) } - execHandlers, err := buildPlanExecuteExecutorHandlers(ctx, a) - if err != nil { - return nil, err + var executor adk.Agent + agenticExecHandlers, herr := buildPlanExecuteAgenticExecutorHandlers(ctx, a) + if herr != nil { + return nil, herr } - executor, err := newPlanExecuteExecutor(ctx, &planexecute.ExecutorConfig{ - Model: a.ExecModel, + 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), - }, execHandlers) + }, a.AgenticExecModel, agenticExecHandlers, a.AgenticModelRetryConfig, a.AgenticModelFailoverConfig) if err != nil { return nil, fmt.Errorf("plan_execute executor: %w", err) } @@ -104,37 +103,35 @@ func NewPlanExecuteRoot(ctx context.Context, a *PlanExecuteRootArgs) (adk.Resuma }) } -// buildPlanExecuteExecutorHandlers 组装 Executor 中间件栈(outermost first),与 Deep/Supervisor 主代理对齐: -// ExecPreMiddlewares(patch / reduction / toolsearch / plantask)→ filesystem → skill → summarization tail。 -func buildPlanExecuteExecutorHandlers(ctx context.Context, a *PlanExecuteRootArgs) ([]adk.ChatModelAgentMiddleware, error) { +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.ChatModelAgentMiddleware - if len(a.ExecPreMiddlewares) > 0 { - execHandlers = append(execHandlers, a.ExecPreMiddlewares...) + var execHandlers []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] + if len(a.AgenticExecPreMiddlewares) > 0 { + execHandlers = append(execHandlers, a.AgenticExecPreMiddlewares...) } - if a.FilesystemMiddleware != nil { - execHandlers = append(execHandlers, a.FilesystemMiddleware) + if a.AgenticFilesystemMiddleware != nil { + execHandlers = append(execHandlers, a.AgenticFilesystemMiddleware) } - if a.SkillMiddleware != nil { - execHandlers = append(execHandlers, a.SkillMiddleware) + if a.AgenticSkillMiddleware != nil { + execHandlers = append(execHandlers, a.AgenticSkillMiddleware) } if a.AppCfg != nil { - sumMw, sumErr := newEinoSummarizationMiddleware(ctx, a.ExecModel, a.AppCfg, a.MwCfg, a.ConversationID, a.DB, a.ProjectID, a.Logger) + 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 executor summarization: %w", sumErr) + return nil, fmt.Errorf("plan_execute agentic executor summarization: %w", sumErr) } - execHandlers = appendEinoChatModelTailMiddlewares(execHandlers, einoChatModelTailConfig{ - logger: a.Logger, - phase: "plan_execute_executor", - summarization: sumMw, - modelName: a.ModelName, - maxTotalTokens: a.AppCfg.OpenAI.MaxTotalTokens, - toolMaxBytes: toolMaxBytesFromMW(a.MwCfg), - conversationID: a.ConversationID, - trace: a.ModelFacingTrace, - middlewareConfig: a.MwCfg, + 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 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_reasoning_stream_emitter_test.go b/internal/multiagent/eino_reasoning_stream_emitter_test.go new file mode 100644 index 00000000..f741aa07 --- /dev/null +++ b/internal/multiagent/eino_reasoning_stream_emitter_test.go @@ -0,0 +1,86 @@ +package multiagent + +import ( + "testing" + + "cyberstrike-ai/internal/openai" +) + +func TestEinoReasoningStreamEmitterStreamingLifecycle(t *testing.T) { + type progressEvent struct { + eventType string + message string + data map[string]interface{} + } + var events []progressEvent + progress := func(eventType, message string, data interface{}) { + m, _ := data.(map[string]interface{}) + events = append(events, progressEvent{eventType: eventType, message: message, data: m}) + } + emitter := newEinoReasoningStreamEmitter("conv-1", "deep", "lead", "orchestrator", progress, func() string { + return "reasoning-1" + }) + + if !emitter.EmitDelta("he") { + t.Fatal("first reasoning delta should emit") + } + if !emitter.EmitDelta("hello") { + t.Fatal("cumulative reasoning chunk should emit tail") + } + if got := emitter.Finish(); got != "hello" { + t.Fatalf("finish body = %q, want hello", got) + } + + if len(events) != 4 { + t.Fatalf("events = %#v, want start + 2 deltas + end", events) + } + if events[0].eventType != "reasoning_chain_stream_start" || events[0].message != " " { + t.Fatalf("event[0] = %#v", events[0]) + } + if events[1].eventType != "reasoning_chain_stream_delta" || events[1].message != "he" { + t.Fatalf("event[1] = %#v", events[1]) + } + if events[2].eventType != "reasoning_chain_stream_delta" || events[2].message != "llo" { + t.Fatalf("event[2] = %#v", events[2]) + } + if got := events[2].data[openai.SSEAccumulatedKey]; got != "hello" { + t.Fatalf("accumulated = %#v, want hello", got) + } + if events[3].eventType != "reasoning_chain_stream_end" || events[3].message != "hello" { + t.Fatalf("event[3] = %#v", events[3]) + } + if got := events[3].data["einoRole"]; got != "orchestrator" { + t.Fatalf("einoRole = %#v", got) + } +} + +func TestEinoReasoningStreamEmitterComplete(t *testing.T) { + var eventType, message string + var data map[string]interface{} + progress := func(et, msg string, raw interface{}) { + eventType = et + message = msg + data, _ = raw.(map[string]interface{}) + } + + ok := newEinoReasoningStreamEmitter("conv-1", "supervisor", "worker", "sub", progress, nil).EmitComplete(" thought ") + if !ok { + t.Fatal("complete reasoning should emit") + } + if eventType != "reasoning_chain" || message != "thought" { + t.Fatalf("event = %s %q", eventType, message) + } + if data["conversationId"] != "conv-1" || data["einoAgent"] != "worker" || data["einoRole"] != "sub" { + t.Fatalf("bad event data: %#v", data) + } +} + +func TestEinoReasoningStreamEmitterNoProgressStillBuffers(t *testing.T) { + emitter := newEinoReasoningStreamEmitter("conv", "deep", "lead", "orchestrator", nil, nil) + if emitter.EmitDelta("hello") { + t.Fatal("nil progress should not emit") + } + if got := emitter.Finish(); got != "hello" { + t.Fatalf("finish body = %q, want hello", got) + } +} 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_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_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_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) + } +}