diff --git a/internal/multiagent/eino_adk_run_loop.go b/internal/multiagent/eino_adk_run_loop.go index 825f3c95..2e321419 100644 --- a/internal/multiagent/eino_adk_run_loop.go +++ b/internal/multiagent/eino_adk_run_loop.go @@ -2,23 +2,18 @@ package multiagent import ( "context" - "encoding/json" "errors" "fmt" - "io" - "os" - "path/filepath" "regexp" "strings" "sync" - "sync/atomic" + "time" "unicode/utf8" "cyberstrike-ai/internal/agent" "cyberstrike-ai/internal/config" "cyberstrike-ai/internal/einomcp" "cyberstrike-ai/internal/einoobserve" - "cyberstrike-ai/internal/openai" "cyberstrike-ai/internal/security" "github.com/cloudwego/eino/adk" @@ -92,7 +87,7 @@ type einoADKRunLoopArgs struct { FilesystemMonitorRecord einomcp.ExecutionRecorder MCPExecutionBinder *MCPExecutionBinder - // ToolInvokeNotify 与 einomcp.ToolsFromDefinitions 共享:run loop 在迭代前 Set,execute/MCP 桥 Fire 时立即推送 tool_result(ADK 晚到经 toolResultSent 去重)。 + // ToolInvokeNotify 与 einomcp.ToolsFromDefinitions 共享:run loop 在迭代前 Set,execute/MCP 桥 Fire 时立即推送 tool_result(ADK 晚到经 toolResultEmitter 去重)。 ToolInvokeNotify *einomcp.ToolInvokeNotifyHolder DA adk.Agent @@ -108,9 +103,13 @@ type einoADKRunLoopArgs struct { EinoCallbacks *config.MultiAgentEinoCallbacksConfig // MaxTotalTokens / ToolMaxBytes / ModelName 用于 context overflow 时的激进压缩续跑。 - MaxTotalTokens int - ToolMaxBytes int - ModelName string + MaxTotalTokens int + ToolMaxBytes int + ModelName string + MiddlewareConfig *config.MultiAgentEinoMiddlewareConfig + + // TurnLoopInterruptTimeout 仅供测试/特殊运行时覆盖;0 使用 EinoTurnLoopRuntime 默认值。 + TurnLoopInterruptTimeout time.Duration } func runEinoADKAgentLoop(ctx context.Context, args *einoADKRunLoopArgs, baseMsgs []adk.Message) (*RunResult, error) { @@ -130,6 +129,17 @@ func runEinoADKAgentLoop(ctx context.Context, args *einoADKRunLoopArgs, baseMsgs conversationID := args.ConversationID progress := args.Progress logger := args.Logger + runID := newEinoRunID() + progress = withEinoRunIDProgress(runID, progress) + args.Progress = progress + if logger != nil { + logger.Info("eino run session started", + zap.String("runId", runID), + zap.String("conversationId", conversationID), + zap.String("orchestration", orchMode), + zap.String("orchestratorName", orchestratorName), + ) + } snapshotMCPIDs := args.SnapshotMCPIDs if snapshotMCPIDs == nil { snapshotMCPIDs = func() []string { return nil } @@ -149,10 +159,6 @@ func runEinoADKAgentLoop(ctx context.Context, args *einoADKRunLoopArgs, baseMsgs return "sub" } } - da := args.DA - mcpIDsMu := args.McpIDsMu - mcpIDs := args.McpIDs - // panic recovery:防止 Eino 框架内部 panic 导致整个 goroutine 崩溃、连接无法正常关闭。 defer func() { if r := recover(); r != nil { @@ -168,11 +174,7 @@ func runEinoADKAgentLoop(ctx context.Context, args *einoADKRunLoopArgs, baseMsgs } }() - var lastAssistant string - var lastPlanExecuteExecutor string msgs := append([]adk.Message(nil), baseMsgs...) - runAccumulatedMsgs := append([]adk.Message(nil), msgs...) - baseAccumulatedCount := len(runAccumulatedMsgs) emptyHint := strings.TrimSpace(args.EmptyResponseMessage) if emptyHint == "" { @@ -180,195 +182,6 @@ func runEinoADKAgentLoop(ctx context.Context, args *einoADKRunLoopArgs, baseMsgs "(Eino 会话已完成,但未捕获到助手文本输出。请查看过程详情或日志。)" } - lastAssistant = "" - lastPlanExecuteExecutor = "" - var reasoningStreamSeq int64 - var einoSubReplyStreamSeq int64 - var mainResponseStreamSeq int64 - toolEmitSeen := make(map[string]struct{}) - var einoMainRound int - var einoLastAgent string - subAgentToolStep := make(map[string]int) - // mainAgentToolStep:主代理每次工具调用批次递增,供 UI 显示「第 N 轮」(单代理无子代理切换时原先会一直停在第 1 轮)。 - mainAgentToolStep := make(map[string]int) - pendingByID := make(map[string]toolCallPendingInfo) - pendingQueueByAgent := make(map[string][]string) - var pendingMu sync.Mutex - markPending := func(tc toolCallPendingInfo) { - if tc.ToolCallID == "" { - return - } - pendingMu.Lock() - defer pendingMu.Unlock() - pendingByID[tc.ToolCallID] = tc - pendingQueueByAgent[tc.EinoAgent] = append(pendingQueueByAgent[tc.EinoAgent], tc.ToolCallID) - } - markPendingWithMonitor := func(tc toolCallPendingInfo) { - markPending(tc) - beginEinoADKFilesystemToolMonitor( - ctx, - args.FilesystemMonitorAgent, - args.FilesystemMonitorRecord, - args.MCPExecutionBinder, - tc.ToolCallID, - tc.ToolName, - ) - } - popNextPendingForAgent := func(agentName string) (toolCallPendingInfo, bool) { - pendingMu.Lock() - defer pendingMu.Unlock() - q := pendingQueueByAgent[agentName] - for len(q) > 0 { - id := q[0] - q = q[1:] - pendingQueueByAgent[agentName] = q - if tc, ok := pendingByID[id]; ok { - delete(pendingByID, id) - return tc, true - } - } - return toolCallPendingInfo{}, false - } - removePendingByID := func(toolCallID string) { - if toolCallID == "" { - return - } - pendingMu.Lock() - defer pendingMu.Unlock() - delete(pendingByID, toolCallID) - } - popAnyPending := func() (toolCallPendingInfo, bool) { - pendingMu.Lock() - defer pendingMu.Unlock() - for id, tc := range pendingByID { - delete(pendingByID, id) - return tc, true - } - return toolCallPendingInfo{}, false - } - pendingCount := func() int { - pendingMu.Lock() - defer pendingMu.Unlock() - return len(pendingByID) - } - flushAllPendingAsFailed := func(err error) { - pendingMu.Lock() - pendingSnapshot := make([]toolCallPendingInfo, 0, len(pendingByID)) - for _, tc := range pendingByID { - pendingSnapshot = append(pendingSnapshot, tc) - } - pendingByID = make(map[string]toolCallPendingInfo) - pendingQueueByAgent = make(map[string][]string) - pendingMu.Unlock() - - if progress == nil { - return - } - msg := "" - if err != nil { - msg = err.Error() - } - for _, tc := range pendingSnapshot { - toolName := tc.ToolName - if strings.TrimSpace(toolName) == "" { - toolName = "unknown" - } - progress("tool_result", fmt.Sprintf("工具结果 (%s)", toolName), map[string]interface{}{ - "toolName": toolName, - "success": false, - "isError": true, - "result": msg, - "resultPreview": msg, - "toolCallId": tc.ToolCallID, - "conversationId": conversationID, - "einoAgent": tc.EinoAgent, - "einoRole": tc.EinoRole, - "source": "eino", - }) - } - } - - // 最近一次成功的 Eino filesystem execute 的标准输出(trim):用于抑制模型紧接着复述同一字符串时的重复「助手输出」时间线。 - var executeStdoutDupMu sync.Mutex - var pendingExecuteStdoutDup string - recordPendingExecuteStdoutDup := func(toolName, stdout string, isErr bool) { - if isErr || !strings.EqualFold(strings.TrimSpace(toolName), "execute") { - return - } - t := strings.TrimSpace(stdout) - if t == "" { - return - } - executeStdoutDupMu.Lock() - pendingExecuteStdoutDup = t - executeStdoutDupMu.Unlock() - } - - var toolResultSent sync.Map // toolCallID -> struct{};ADK Tool 事件去重(权威正文来自 reduction 处理后的 agent 上下文) - tryEmitToolResultProgress := func(toolName, content, toolCallID string, isErr bool, agentName string) { - // 仅由 ADK schema.Tool 事件调用;MCP/execute 桥在 reduction 前的 ToolInvokeNotify 不得推送 tool_result, - // 否则全量输出会先占位并触发 toolResultSent 去重,导致 UI/监控展示与 agent 实际收到的截断正文不一致。 - toolName = strings.TrimSpace(toolName) - if toolName == "" { - toolName = "unknown" - } - preview := content - if len(preview) > 200 { - preview = preview[:200] + "..." - } - backgroundRunning := isErr && isMCPBackgroundWaitResult(content) - displayIsErr := isErr && !backgroundRunning - data := map[string]interface{}{ - "toolName": toolName, - "success": !displayIsErr, - "isError": displayIsErr, - "result": content, - "resultPreview": preview, - "agentFacing": true, // 与 reduction 后送入 ChatModel 的正文一致,供前端展示 - "conversationId": conversationID, - "einoAgent": agentName, - "einoRole": einoRoleTag(agentName), - "source": "eino", - } - if backgroundRunning { - data["status"] = "background_running" - data["modelFacingIsError"] = isErr - if execID := mcpExecutionIDFromWaitResult(content); execID != "" { - data["executionId"] = execID - } - } - tid := strings.TrimSpace(toolCallID) - if tid == "" { - if inferred, ok := popNextPendingForAgent(agentName); ok { - tid = inferred.ToolCallID - } else if inferred, ok := popNextPendingForAgent(orchestratorName); ok { - tid = inferred.ToolCallID - } else if inferred, ok := popNextPendingForAgent(""); ok { - tid = inferred.ToolCallID - } else if inferred, ok := popAnyPending(); ok { - tid = inferred.ToolCallID - } - } - if tid != "" { - removePendingByID(tid) - if _, loaded := toolResultSent.LoadOrStore(tid, struct{}{}); loaded { - return - } - data["toolCallId"] = tid - toolCallID = tid - } - recordPendingExecuteStdoutDup(toolName, content, displayIsErr) - recordEinoADKFilesystemToolMonitor(ctx, args.FilesystemMonitorAgent, args.FilesystemMonitorRecord, args.MCPExecutionBinder, toolName, toolCallID, runAccumulatedMsgs, content, displayIsErr) - if args.FilesystemMonitorAgent != nil && args.MCPExecutionBinder != nil { - if execID := args.MCPExecutionBinder.ExecutionID(toolCallID); execID != "" { - args.FilesystemMonitorAgent.UpdateMCPExecutionDisplayResult(execID, content) - } - } - if progress != nil { - progress("tool_result", fmt.Sprintf("工具结果 (%s)", toolName), data) - } - } - if args.EinoCallbacks != nil { ctx = einoobserve.AttachAgentRunCallbacks(ctx, args.EinoCallbacks, einoobserve.Params{ Logger: logger, @@ -376,715 +189,92 @@ func runEinoADKAgentLoop(ctx context.Context, args *einoADKRunLoopArgs, baseMsgs ConversationID: conversationID, OrchMode: orchMode, OrchestratorName: orchestratorName, + RunID: runID, }) } - runnerCfg := adk.RunnerConfig{ - Agent: da, - // 启用 ADK 流式事件:plan_execute 也需要输出 reasoning/response 流, - // 与 deep/supervisor/eino_single 的前端体验保持一致。 - EnableStreaming: true, - } - var cpStore *fileCheckPointStore - var checkPointID string - if cp := strings.TrimSpace(args.CheckpointDir); cp != "" { - cpDir := filepath.Join(cp, sanitizeEinoPathSegment(conversationID)) - st, stErr := newFileCheckPointStore(cpDir) - if stErr != nil { - if logger != nil { - logger.Warn("eino checkpoint store disabled", zap.String("dir", cpDir), zap.Error(stErr)) - } - } else { - cpStore = st - checkPointID = buildEinoCheckpointID(orchMode) - runnerCfg.CheckPointStore = st - if logger != nil { - logger.Info("eino runner: checkpoint store enabled", - zap.String("dir", cpDir), - zap.String("checkPointID", checkPointID)) - } - } - } - runner := adk.NewRunner(ctx, runnerCfg) - startRunnerIter := func(runMsgs []adk.Message) *adk.AsyncIterator[*adk.AgentEvent] { - if checkPointID != "" { - return runner.Run(ctx, runMsgs, adk.WithCheckPointID(checkPointID)) - } - return runner.Run(ctx, runMsgs) - } - var iter *adk.AsyncIterator[*adk.AgentEvent] - if cpStore != nil && checkPointID != "" { - if _, existed, getErr := cpStore.Get(ctx, checkPointID); getErr != nil { - if logger != nil { - logger.Warn("eino checkpoint preflight get failed", zap.String("checkPointID", checkPointID), zap.Error(getErr)) - } - } else if existed { - if progress != nil { - progress("progress", "检测到断点,正在从中断节点恢复执行...", map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "orchestration": orchMode, - "checkPointID": checkPointID, - }) - } - if logger != nil { - logger.Info("eino runner: resume from checkpoint", zap.String("checkPointID", checkPointID)) - } - resumeIter, resumeErr := runner.Resume(ctx, checkPointID) - if resumeErr == nil { - iter = resumeIter - } else { - if logger != nil { - logger.Warn("eino runner: resume failed, fallback to fresh run", - zap.String("checkPointID", checkPointID), - zap.Error(resumeErr)) - } - if progress != nil { - progress("progress", "断点恢复失败,已回退为全新执行。", map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "orchestration": orchMode, - "checkPointID": checkPointID, - }) - } - } - } - } - if iter == nil { - iter = startRunnerIter(msgs) - } - transientRetrier := newEinoTransientRunRetrier(einoTransientRunRetryPolicyFromArgs(args)) - var contextOverflowRetried bool - handleRunErr := func(runErr error) error { - if runErr == nil { - return nil - } - if errors.Is(runErr, context.DeadlineExceeded) { - flushAllPendingAsFailed(runErr) - if progress != nil { - progress("error", runErr.Error(), map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "errorKind": "timeout", - }) - } - return runErr - } - // context.Canceled 是唯一应当直接终止编排的错误(用户关闭页面、主动停止等)。 - if errors.Is(runErr, context.Canceled) { - flushAllPendingAsFailed(runErr) - if progress != nil { - progress("error", runErr.Error(), map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - }) - } - return runErr - } - if isEinoIterationLimitError(runErr) { - flushAllPendingAsFailed(runErr) - if progress != nil { - progress("iteration_limit_reached", runErr.Error(), map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "orchestration": orchMode, - }) - progress("error", runErr.Error(), map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "errorKind": "iteration_limit", - }) - } - return runErr - } - flushAllPendingAsFailed(runErr) - if progress != nil { - progress("error", runErr.Error(), map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - }) - } - return runErr - } - - maybeRetryTransientRun := func(runErr error) (restarted bool, fatal error) { - if runErr == nil { - return false, nil - } - var rejected *modelOutputRejectedError - if errors.As(runErr, &rejected) { - if progress != nil { - progress("model_output_rejected", "模型输出不完整或工具参数不安全,已阻止执行。", map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "orchestration": orchMode, - "reason": rejected.Reason, - "finishReason": rejected.FinishReason, - "toolName": rejected.ToolName, - "toolCallId": rejected.ToolCallID, - "argumentsBytes": rejected.ArgumentsBytes, - "completionTokens": rejected.CompletionTokens, - "reasoningTokens": rejected.ReasoningTokens, - "repairAttempt": rejected.RepairAttempt, - "repairable": rejected.Repairable, - }) - } - if !rejected.Repairable { - return false, handleRunErr(runErr) - } - restartMsgs, ctxSource := einoMessagesForRunRestart(args, baseMsgs, runAccumulatedMsgs, baseAccumulatedCount) - restartMsgs = append(restartMsgs, schema.UserMessage(modelOutputRepairInstruction)) - if logger != nil { - logger.Warn("eino model output rejected, retrying once with concise instruction", - zap.Error(runErr), zap.String("orchestration", orchMode), - zap.String("contextSource", string(ctxSource)), zap.Int("repairAttempt", rejected.RepairAttempt)) - } - msgs = restartMsgs - iter = startRunnerIter(msgs) - return true, nil - } - if isEinoContextOverflowError(runErr) && !contextOverflowRetried { - contextOverflowRetried = true - restartMsgs, ctxSource := einoMessagesForRunRestart(args, baseMsgs, runAccumulatedMsgs, baseAccumulatedCount) - restartMsgs = aggressiveCompactMessagesForOverflow( - ctx, restartMsgs, args.MaxTotalTokens, args.ModelName, args.ToolMaxBytes, orchMode, logger, - ) - if logger != nil { - logger.Warn("eino context overflow, retrying with aggressive compaction", - zap.Error(runErr), - zap.String("orchestration", orchMode), - zap.String("contextSource", string(ctxSource)), - ) - } - if progress != nil { - progress("eino_context_overflow_retry", "上下文超限,正在激进压缩后重试…", map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "orchestration": orchMode, - "contextSource": string(ctxSource), - }) - } - msgs = restartMsgs - iter = startRunnerIter(msgs) - return true, nil - } - if !isEinoTransientRunError(runErr) { - return false, handleRunErr(runErr) - } - restarted, restartMsgs, ctxSource, backoff, retErr := transientRetrier.tryRetry( - ctx, runErr, args, baseMsgs, runAccumulatedMsgs, baseAccumulatedCount, - ) - if retErr != nil { - flushAllPendingAsFailed(runErr) - if logger != nil { - logger.Warn("eino transient retry exhausted", - zap.Error(retErr), - zap.String("orchestration", orchMode), - zap.Int("maxAttempts", transientRetrier.maxAttempts())) - } - return false, retErr - } - if !restarted { - return false, nil - } - attemptNo := transientRetrier.attempt() - maxAttempts := transientRetrier.maxAttempts() - if logger != nil { - logger.Warn("eino transient error, retrying after backoff", - zap.Error(runErr), - zap.String("orchestration", orchMode), - zap.Int("attempt", attemptNo), - zap.Int("maxAttempts", maxAttempts), - zap.Duration("backoff", backoff)) - } - if progress != nil { - errorKind, errorSummary := einoTransientRunErrorUserDetail(runErr) - progress("eino_run_retry", fmt.Sprintf("遇到临时错误,%d 秒后第 %d/%d 次重试。原因:%s", int(backoff.Seconds()), attemptNo, maxAttempts, errorSummary), map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "orchestration": orchMode, - "error": runErr.Error(), - "errorKind": errorKind, - "errorSummary": errorSummary, - "attempt": attemptNo, - "maxAttempts": maxAttempts, - "backoffSec": int(backoff.Seconds()), - }) - progress("eino_run_retry", "已恢复上下文,正在重试…", map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "orchestration": orchMode, - "error": runErr.Error(), - "errorKind": errorKind, - "errorSummary": errorSummary, - "attempt": attemptNo, - "maxAttempts": maxAttempts, - "backoffSec": int(backoff.Seconds()), - "contextSource": string(ctxSource), - }) - } - msgs = restartMsgs - iter = startRunnerIter(msgs) - return true, nil - } + drain := newEinoRunEventDrain(einoRunEventDrainConfig{ + Context: ctx, + ConversationID: conversationID, + OrchMode: orchMode, + OrchestratorName: orchestratorName, + Progress: progress, + Logger: logger, + BaseMessages: msgs, + SnapshotMCPIDs: snapshotMCPIDs, + StreamsMainAssistant: streamsMainAssistant, + EinoRoleTag: einoRoleTag, + MiddlewareConfig: args.MiddlewareConfig, + FilesystemMonitorAgent: args.FilesystemMonitorAgent, + FilesystemMonitorRecord: args.FilesystemMonitorRecord, + MCPExecutionBinder: args.MCPExecutionBinder, + }) + session := newEinoRunRuntimeSession(einoRunRuntimeSessionConfig{ + Context: ctx, + Args: args, + Drain: drain, + BaseMessages: msgs, + EmptyHint: emptyHint, + SnapshotMCPIDs: snapshotMCPIDs, + EinoRoleTag: einoRoleTag, + }) + defer session.Close() // 仅在退避重试后真正收到数据/完成一步时清零,避免重启后首个无错 ADK 事件误把计数打回 0。 - confirmTransientRetryRecovery := func() { - if transientRetrier.attempt() > 0 { - transientRetrier.reset() - } - } - - takePartial := func(runErr error) (*RunResult, error) { - if len(runAccumulatedMsgs) <= baseAccumulatedCount { - return nil, runErr - } - ids := snapshotMCPIDs() - return buildEinoRunResultFromAccumulated( - orchMode, runAccumulatedMsgs, modelFacingTraceSnapshot(args), - lastAssistant, lastPlanExecuteExecutor, emptyHint, ids, true, - ), runErr - } + drain.BindHandlers(session.ConfirmRecovery) for { // iter.Next 可能长时间阻塞(工具执行、模型推理);须与 ctx 联动,否则取消/超时无法及时 flush pending。 - ev, ok, iterCtxErr := nextAgentEventWithContext(ctx, iter) + ev, ok, iterCtxErr := nextAgentEventWithContext(ctx, session.Iterator()) if iterCtxErr != nil { - flushAllPendingAsFailed(iterCtxErr) - if progress != nil { - if isInterruptContinue(ctx) { - progress("progress", "已暂停当前输出,正在合并用户补充并继续…", map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "kind": "interrupt_continue", - }) - } else { - progress("error", iterCtxErr.Error(), map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - }) - } - } - return takePartial(iterCtxErr) + return session.HandleIteratorContextError(iterCtxErr) } if !ok { // iter 结束并不总是“正常完成”: // 当取消/超时发生在 iter.Next() 阻塞期间时,可能直接返回 !ok。 // 此时必须保留 checkpoint,避免后续恢复时被误判为“无断点”而全量重跑。 - if ctxErr := ctx.Err(); ctxErr != nil { - flushAllPendingAsFailed(ctxErr) - if progress != nil { - if isInterruptContinue(ctx) { - progress("progress", "已暂停当前输出,正在合并用户补充并继续…", map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "kind": "interrupt_continue", - }) - } else { - progress("error", ctxErr.Error(), map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - }) - } - } - return takePartial(ctxErr) + completed, result, err := session.HandleIteratorEnd() + if result != nil || err != nil { + return result, err } - if orphanCount := pendingCount(); orphanCount > 0 { - flushAllPendingAsFailed(errors.New("pending tool call missing result before run completion")) - if progress != nil { - progress("eino_pending_orphaned", "pending tool calls were force-closed at run end", map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "orchestration": orchMode, - "pendingCount": orphanCount, - }) - } + if completed { + break } - if cpStore != nil && checkPointID != "" { - if p, pErr := cpStore.path(checkPointID); pErr == nil { - if rmErr := os.Remove(p); rmErr != nil && !os.IsNotExist(rmErr) && logger != nil { - logger.Warn("eino checkpoint cleanup failed", zap.String("path", p), zap.Error(rmErr)) - } - } - } - break + continue } if ev == nil { continue } if ev.Err != nil { - restarted, retErr := maybeRetryTransientRun(ev.Err) - if retErr != nil { - return takePartial(retErr) + handled := session.HandleRunError(ev.Err) + if handled.Result != nil || handled.Err != nil { + return handled.Result, handled.Err } - if restarted { + if handled.Restarted { continue } } - if ev.AgentName != "" && progress != nil { - iterEinoAgent := orchestratorName - if orchMode == "plan_execute" { - if a := strings.TrimSpace(ev.AgentName); a != "" { - iterEinoAgent = a - } - } - if streamsMainAssistant(ev.AgentName) { - mainIterKey := einoMainIterationKey(iterEinoAgent, orchestratorName) - if einoMainRound == 0 { - einoMainRound = 1 - mainAgentToolStep[mainIterKey] = 1 - progress("iteration", "", map[string]interface{}{ - "iteration": 1, - "einoScope": "main", - "einoRole": "orchestrator", - "einoAgent": iterEinoAgent, - "orchestration": orchMode, - "conversationId": conversationID, - "source": "eino", - }) - } else if einoLastAgent != "" { - needBump := false - if !streamsMainAssistant(einoLastAgent) { - needBump = true // 子代理 → 主代理 - } else if einoLastAgent != ev.AgentName { - needBump = true // plan_execute:planner ↔ executor 等主代理切换 - } - if needBump { - einoMainRound++ - mainAgentToolStep[mainIterKey] = einoMainRound - progress("iteration", "", map[string]interface{}{ - "iteration": einoMainRound, - "einoScope": "main", - "einoRole": "orchestrator", - "einoAgent": iterEinoAgent, - "orchestration": orchMode, - "conversationId": conversationID, - "source": "eino", - }) - } - } - } - // 仅在代理切换时更新进度标题;同一代理的每个 ADK 事件不再重复刷 progress。 - if einoLastAgent != ev.AgentName { - progress("progress", fmt.Sprintf("[Eino] %s", ev.AgentName), map[string]interface{}{ - "conversationId": conversationID, - "einoAgent": ev.AgentName, - "einoRole": einoRoleTag(ev.AgentName), - "orchestration": orchMode, - }) - } - einoLastAgent = ev.AgentName - } + drain.ObserveAgent(ev.AgentName) if ev.Output == nil || ev.Output.MessageOutput == nil { continue } mv := ev.Output.MessageOutput - if mv.IsStreaming && mv.MessageStream != nil && mv.Role == schema.Tool { - toolName := strings.TrimSpace(mv.ToolName) - content, streamToolCallID, toolStreamRecvErr := recvSchemaMessageStream(ctx, mv.MessageStream) - isErr := einoToolResultIsError(toolName, content) - content = einoToolResultBody(content) - if streamToolCallID != "" { - opts := []schema.ToolMessageOption{schema.WithToolName(toolName)} - runAccumulatedMsgs = append(runAccumulatedMsgs, schema.ToolMessage(content, streamToolCallID, opts...)) - } - tryEmitToolResultProgress(toolName, content, streamToolCallID, isErr, ev.AgentName) - if toolStreamRecvErr != nil && logger != nil { - logger.Warn("eino tool result stream recv error", - zap.Error(toolStreamRecvErr), - zap.String("agent", ev.AgentName), - zap.String("tool", toolName)) - } - if toolStreamRecvErr == nil { - confirmTransientRetryRecovery() - } + if drain.HandleToolResultStreaming(mv, ev.AgentName) { continue } - if mv.IsStreaming && mv.MessageStream != nil { - mainStreamID := fmt.Sprintf("eino-main-%s-%d", conversationID, atomic.AddInt64(&mainResponseStreamSeq, 1)) - streamHeaderSent := false - var reasoningStreamID string - var toolStreamFragments []schema.ToolCall - var subAssistantBuf string - var subReplyStreamID string - var mainAssistantBuf string - // 已通过 response_delta 推到前端的正文(与 monitor.js normalizeStreamingDeltaJs 累积一致) - var mainAssistWireAccum string - var mainAssistDupTarget string // 非空表示本段主助手流需缓冲至 EOF,与 execute 输出比对去重 - var reasoningBuf string - var prevReasoningDisplay string // UI 用:剥离 Claude 内部 signature 尾缀后的累计展示 - var streamRecvErr error - type streamMsg struct { - chunk *schema.Message - err error - } - recvCh := make(chan streamMsg, 8) - go func() { - defer close(recvCh) - for { - ch, rerr := mv.MessageStream.Recv() - recvCh <- streamMsg{chunk: ch, err: rerr} - if rerr != nil { - return - } - } - }() - streamRecvLoop: - for { - select { - case <-ctx.Done(): - streamRecvErr = ctx.Err() - break streamRecvLoop - case sm, ok := <-recvCh: - if !ok { - break streamRecvLoop - } - chunk, rerr := sm.chunk, sm.err - if rerr != nil { - if errors.Is(rerr, io.EOF) { - break streamRecvLoop - } - if logger != nil { - logger.Warn("eino stream recv error, flushing incomplete stream", - zap.Error(rerr), - zap.String("agent", ev.AgentName), - zap.Int("toolFragments", len(toolStreamFragments))) - } - streamRecvErr = rerr - break streamRecvLoop - } - if chunk == nil { - continue - } - if progress != nil && strings.TrimSpace(chunk.ReasoningContent) != "" { - var reasoningDelta string - reasoningBuf, reasoningDelta = normalizeStreamingDelta(reasoningBuf, chunk.ReasoningContent) - if reasoningDelta != "" { - fullDisplay := openai.DisplayReasoningContent(reasoningBuf) - var displayDelta string - if strings.HasPrefix(fullDisplay, prevReasoningDisplay) { - displayDelta = fullDisplay[len(prevReasoningDisplay):] - } else { - displayDelta = fullDisplay - } - prevReasoningDisplay = fullDisplay - if displayDelta != "" { - if reasoningStreamID == "" { - reasoningStreamID = fmt.Sprintf("eino-reasoning-%s-%d", conversationID, atomic.AddInt64(&reasoningStreamSeq, 1)) - progress("reasoning_chain_stream_start", " ", map[string]interface{}{ - "streamId": reasoningStreamID, - "source": "eino", - "einoAgent": ev.AgentName, - "einoRole": einoRoleTag(ev.AgentName), - "orchestration": orchMode, - }) - } - progress("reasoning_chain_stream_delta", displayDelta, openai.WithSSEAccumulated(map[string]interface{}{ - "streamId": reasoningStreamID, - }, fullDisplay)) - } - } - } - if chunk.Content != "" { - if progress != nil && streamsMainAssistant(ev.AgentName) { - var contentDelta string - mainAssistantBuf, contentDelta = normalizeStreamingDelta(mainAssistantBuf, chunk.Content) - if contentDelta != "" { - if mainAssistDupTarget == "" { - executeStdoutDupMu.Lock() - if pendingExecuteStdoutDup != "" { - mainAssistDupTarget = pendingExecuteStdoutDup - } - executeStdoutDupMu.Unlock() - } - if mainAssistDupTarget != "" { - // 已展示过 tool_result,缓冲全文;EOF 后与 execute 输出相同则不再发助手流 - } else { - if !streamHeaderSent { - progress("response_start", "", map[string]interface{}{ - "conversationId": conversationID, - "mcpExecutionIds": snapshotMCPIDs(), - "messageGeneratedBy": "eino:" + ev.AgentName, - "einoRole": "orchestrator", - "einoAgent": ev.AgentName, - "orchestration": orchMode, - "iteration": einoMainRound, - "streamId": mainStreamID, - }) - streamHeaderSent = true - } - progress("response_delta", contentDelta, openai.WithSSEAccumulated(map[string]interface{}{ - "conversationId": conversationID, - "mcpExecutionIds": snapshotMCPIDs(), - "einoRole": "orchestrator", - "einoAgent": ev.AgentName, - "orchestration": orchMode, - "iteration": einoMainRound, - "streamId": mainStreamID, - }, mainAssistantBuf)) - mainAssistWireAccum, _ = normalizeStreamingDelta(mainAssistWireAccum, contentDelta) - } - } - } else if !streamsMainAssistant(ev.AgentName) { - var subDelta string - subAssistantBuf, subDelta = normalizeStreamingDelta(subAssistantBuf, chunk.Content) - if subDelta != "" { - if progress != nil { - if subReplyStreamID == "" { - subReplyStreamID = fmt.Sprintf("eino-sub-reply-%s-%d", conversationID, atomic.AddInt64(&einoSubReplyStreamSeq, 1)) - progress("eino_agent_reply_stream_start", "", map[string]interface{}{ - "streamId": subReplyStreamID, - "einoAgent": ev.AgentName, - "einoRole": "sub", - "conversationId": conversationID, - "source": "eino", - }) - } - progress("eino_agent_reply_stream_delta", subDelta, openai.WithSSEAccumulated(map[string]interface{}{ - "streamId": subReplyStreamID, - "conversationId": conversationID, - }, subAssistantBuf)) - } - } - } - } - if len(chunk.ToolCalls) > 0 { - toolStreamFragments = append(toolStreamFragments, chunk.ToolCalls...) - } - } - } - if progress != nil && reasoningStreamID != "" && strings.TrimSpace(reasoningBuf) != "" { - progress("reasoning_chain_stream_end", openai.DisplayReasoningContent(strings.TrimSpace(reasoningBuf)), map[string]interface{}{ - "streamId": reasoningStreamID, - "conversationId": conversationID, - "source": "eino", - "einoAgent": ev.AgentName, - "einoRole": einoRoleTag(ev.AgentName), - "orchestration": orchMode, - }) - } - if streamsMainAssistant(ev.AgentName) { - s := strings.TrimSpace(mainAssistantBuf) - if mainAssistDupTarget != "" { - executeStdoutDupMu.Lock() - pendingExecuteStdoutDup = "" - executeStdoutDupMu.Unlock() - if s != "" && s == mainAssistDupTarget { - // 与刚展示的 execute 结果完全一致:不再发助手流式事件,仍写入轨迹与最终回复字段 - lastAssistant = s - runAccumulatedMsgs = append(runAccumulatedMsgs, schema.AssistantMessage(s, nil)) - if orchMode == "plan_execute" && strings.EqualFold(strings.TrimSpace(ev.AgentName), "executor") { - lastPlanExecuteExecutor = UnwrapPlanExecuteUserText(s) - } - } else if s != "" { - if progress != nil { - // 仅用 TrimSpace 与 execute 比对;推到 UI 的必须是 mainAssistantBuf, - // 否则尾部空白/换行与已流式前缀不一致时,前端 normalize 会走拼接路径造成叠字。 - _, eofTail := normalizeStreamingDelta(mainAssistWireAccum, mainAssistantBuf) - if eofTail != "" { - if !streamHeaderSent { - progress("response_start", "", map[string]interface{}{ - "conversationId": conversationID, - "mcpExecutionIds": snapshotMCPIDs(), - "messageGeneratedBy": "eino:" + ev.AgentName, - "einoRole": "orchestrator", - "einoAgent": ev.AgentName, - "orchestration": orchMode, - "iteration": einoMainRound, - "streamId": mainStreamID, - }) - } - progress("response_delta", eofTail, openai.WithSSEAccumulated(map[string]interface{}{ - "conversationId": conversationID, - "mcpExecutionIds": snapshotMCPIDs(), - "einoRole": "orchestrator", - "einoAgent": ev.AgentName, - "orchestration": orchMode, - "iteration": einoMainRound, - "streamId": mainStreamID, - }, mainAssistantBuf)) - mainAssistWireAccum, _ = normalizeStreamingDelta(mainAssistWireAccum, eofTail) - } - } - lastAssistant = s - runAccumulatedMsgs = append(runAccumulatedMsgs, schema.AssistantMessage(s, nil)) - if orchMode == "plan_execute" && strings.EqualFold(strings.TrimSpace(ev.AgentName), "executor") { - lastPlanExecuteExecutor = UnwrapPlanExecuteUserText(s) - } - } - } else if s != "" { - lastAssistant = s - runAccumulatedMsgs = append(runAccumulatedMsgs, schema.AssistantMessage(s, nil)) - if orchMode == "plan_execute" && strings.EqualFold(strings.TrimSpace(ev.AgentName), "executor") { - lastPlanExecuteExecutor = UnwrapPlanExecuteUserText(s) - } - } - } - if strings.TrimSpace(subAssistantBuf) != "" && progress != nil { - if s := strings.TrimSpace(subAssistantBuf); s != "" { - if subReplyStreamID != "" { - progress("eino_agent_reply_stream_end", s, map[string]interface{}{ - "streamId": subReplyStreamID, - "einoAgent": ev.AgentName, - "einoRole": "sub", - "conversationId": conversationID, - "source": "eino", - }) - } else { - progress("eino_agent_reply", s, map[string]interface{}{ - "conversationId": conversationID, - "einoAgent": ev.AgentName, - "einoRole": "sub", - "source": "eino", - }) - } - } - } - var lastToolChunk *schema.Message - if merged := mergeStreamingToolCallFragments(toolStreamFragments); len(merged) > 0 { - lastToolChunk = mergeMessageToolCalls(&schema.Message{ToolCalls: merged}) - } - if progress != nil && lastToolChunk != nil { - for _, tc := range lastToolChunk.ToolCalls { - if marker, ok := modelOutputRecoveryFromToolCall(tc); ok { - progress("model_output_rejected", "模型工具调用不完整或参数不安全,已阻止执行并要求重写。", map[string]interface{}{ - "conversationId": conversationID, "source": "eino", "orchestration": orchMode, - "reason": marker.Reason, "finishReason": marker.FinishReason, - "toolName": tc.Function.Name, "toolCallId": tc.ID, - "argumentsBytes": marker.ArgumentsBytes, "completionTokens": marker.CompletionTokens, - "reasoningTokens": marker.ReasoningTokens, "repairAttempt": marker.RepairAttempt, - }) - } - } - } - tryEmitToolCallsOnce(lastToolChunk, ev.AgentName, orchestratorName, conversationID, orchMode, progress, toolEmitSeen, subAgentToolStep, mainAgentToolStep, markPendingWithMonitor) - // 流式路径此前只把 tool_calls 推给进度 UI,未写入 runAccumulatedMsgs;落库后 loadHistory→RepairOrphan 会删掉全部 tool 结果,表现为「续跑/下轮失忆」。 - if lastToolChunk != nil && len(lastToolChunk.ToolCalls) > 0 { - runAccumulatedMsgs = append(runAccumulatedMsgs, schema.AssistantMessage("", lastToolChunk.ToolCalls)) - } + if handledStream, streamRecvErr := drain.HandleAssistantStream(mv, ev.AgentName); handledStream { if streamRecvErr != nil { - if isInterruptContinue(ctx) { - return takePartial(streamRecvErr) + handled := session.HandleStreamError(streamRecvErr, ev.AgentName) + if handled.Result != nil || handled.Err != nil { + return handled.Result, handled.Err } - if progress != nil { - progress("eino_stream_error", streamRecvErr.Error(), map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "einoAgent": ev.AgentName, - "einoRole": einoRoleTag(ev.AgentName), - }) - } - restarted, retErr := maybeRetryTransientRun(streamRecvErr) - if retErr != nil { - return takePartial(retErr) - } - if restarted { + if handled.Restarted { continue } } else { - confirmTransientRetryRecovery() + session.ConfirmRecovery() } continue } @@ -1093,113 +283,11 @@ func runEinoADKAgentLoop(ctx context.Context, args *einoADKRunLoopArgs, baseMsgs if gerr != nil || msg == nil { continue } - runAccumulatedMsgs = append(runAccumulatedMsgs, msg) - if progress != nil { - for _, tc := range msg.ToolCalls { - if marker, ok := modelOutputRecoveryFromToolCall(tc); ok { - progress("model_output_rejected", "模型工具调用不完整或参数不安全,已阻止执行并要求重写。", map[string]interface{}{ - "conversationId": conversationID, "source": "eino", "orchestration": orchMode, - "reason": marker.Reason, "finishReason": marker.FinishReason, - "toolName": tc.Function.Name, "toolCallId": tc.ID, - "argumentsBytes": marker.ArgumentsBytes, "completionTokens": marker.CompletionTokens, - "reasoningTokens": marker.ReasoningTokens, "repairAttempt": marker.RepairAttempt, - }) - } - } - } - tryEmitToolCallsOnce(mergeMessageToolCalls(msg), ev.AgentName, orchestratorName, conversationID, orchMode, progress, toolEmitSeen, subAgentToolStep, mainAgentToolStep, markPendingWithMonitor) - - if mv.Role == schema.Assistant { - if progress != nil && strings.TrimSpace(msg.ReasoningContent) != "" { - progress("reasoning_chain", openai.DisplayReasoningContent(strings.TrimSpace(msg.ReasoningContent)), map[string]interface{}{ - "conversationId": conversationID, - "source": "eino", - "einoAgent": ev.AgentName, - "einoRole": einoRoleTag(ev.AgentName), - "orchestration": orchMode, - }) - } - body := strings.TrimSpace(msg.Content) - if body != "" { - if streamsMainAssistant(ev.AgentName) { - executeStdoutDupMu.Lock() - dup := pendingExecuteStdoutDup - if dup != "" && body == dup { - pendingExecuteStdoutDup = "" - executeStdoutDupMu.Unlock() - lastAssistant = body - if orchMode == "plan_execute" && strings.EqualFold(strings.TrimSpace(ev.AgentName), "executor") { - lastPlanExecuteExecutor = UnwrapPlanExecuteUserText(body) - } - // 非流式:与 execute 输出相同则跳过助手通道展示(msg 已在上方写入 runAccumulatedMsgs) - } else { - if dup != "" { - pendingExecuteStdoutDup = "" - } - executeStdoutDupMu.Unlock() - if progress != nil { - nonStreamID := fmt.Sprintf("eino-main-%s-%d", conversationID, atomic.AddInt64(&mainResponseStreamSeq, 1)) - progress("response_start", "", map[string]interface{}{ - "conversationId": conversationID, - "mcpExecutionIds": snapshotMCPIDs(), - "messageGeneratedBy": "eino:" + ev.AgentName, - "einoRole": "orchestrator", - "einoAgent": ev.AgentName, - "orchestration": orchMode, - "iteration": einoMainRound, - "streamId": nonStreamID, - }) - progress("response_delta", body, openai.WithSSEAccumulated(map[string]interface{}{ - "conversationId": conversationID, - "mcpExecutionIds": snapshotMCPIDs(), - "einoRole": "orchestrator", - "einoAgent": ev.AgentName, - "orchestration": orchMode, - "iteration": einoMainRound, - "streamId": nonStreamID, - }, body)) - } - lastAssistant = body - if orchMode == "plan_execute" && strings.EqualFold(strings.TrimSpace(ev.AgentName), "executor") { - lastPlanExecuteExecutor = UnwrapPlanExecuteUserText(body) - } - } - } else if progress != nil { - progress("eino_agent_reply", body, map[string]interface{}{ - "conversationId": conversationID, - "einoAgent": ev.AgentName, - "einoRole": "sub", - "source": "eino", - }) - } - } - } - - if (mv.Role == schema.Tool || msg.Role == schema.Tool) && progress != nil { - toolName := msg.ToolName - if toolName == "" { - toolName = mv.ToolName - } - - content := msg.Content - isErr := einoToolResultIsError(toolName, content) - content = einoToolResultBody(content) - - toolCallID := strings.TrimSpace(msg.ToolCallID) - tryEmitToolResultProgress(toolName, content, toolCallID, isErr, ev.AgentName) - } - confirmTransientRetryRecovery() + drain.HandleMaterialized(mv, msg, ev.AgentName) + session.ConfirmRecovery() } - mcpIDsMu.Lock() - ids := append([]string(nil), *mcpIDs...) - mcpIDsMu.Unlock() - - out := buildEinoRunResultFromAccumulated( - orchMode, runAccumulatedMsgs, modelFacingTraceSnapshot(args), - lastAssistant, lastPlanExecuteExecutor, emptyHint, ids, false, - ) - return out, nil + return session.BuildFinalResult(), nil } // modelFacingTraceSnapshot returns only the state that actually reached the model boundary. @@ -1214,11 +302,6 @@ func modelFacingTraceSnapshot(args *einoADKRunLoopArgs) []adk.Message { return nil } -func einoPartialRunLastOutputHint() string { - return "[执行未正常结束(用户停止、超时或异常)。续跑时请基于上文已产生的工具与结果继续,勿重复已完成步骤。]\n" + - "[Run ended abnormally; continue from the trace above without repeating completed steps.]" -} - // friendlyEinoExecuteInvokeTail 将 Eino execute 超时/中断/流异常转为简短提示。 // 命令非零退出(ExecuteExitError)已有 exec 对齐的正文,不再追加「执行未正常结束」。 func friendlyEinoExecuteInvokeTail(invokeErr error) string { @@ -1317,195 +400,23 @@ func nextAgentEventWithContext(ctx context.Context, iter *adk.AsyncIterator[*adk } // recvSchemaMessageStream 消费 ADK Tool 流式结果;ctx 取消时立即返回,避免 amass 等无输出时永久阻塞。 -func recvSchemaMessageStream(ctx context.Context, stream *schema.StreamReader[*schema.Message]) (content, toolCallID string, recvErr error) { +func recvSchemaMessageStream(ctx context.Context, stream *schema.StreamReader[*schema.Message]) (content, toolCallID, toolName string, recvErr error) { if stream == nil { - return "", "", nil + return "", "", "", nil } - type streamMsg struct { - chunk *schema.Message - err error - } - recvCh := make(chan streamMsg, 8) - go func() { - defer close(recvCh) - for { - ch, rerr := stream.Recv() - recvCh <- streamMsg{chunk: ch, err: rerr} - if rerr != nil { - return - } - } - }() var buf strings.Builder - for { - select { - case <-ctx.Done(): - return buf.String(), toolCallID, ctx.Err() - case sm, open := <-recvCh: - if !open { - return buf.String(), toolCallID, nil - } - rerr := sm.err - if errors.Is(rerr, io.EOF) { - return buf.String(), toolCallID, nil - } - if rerr != nil { - return buf.String(), toolCallID, rerr - } - chunk := sm.chunk - if chunk == nil { - continue - } - if chunk.Content != "" { - buf.WriteString(chunk.Content) - } - if tid := strings.TrimSpace(chunk.ToolCallID); tid != "" { - toolCallID = tid - } + recvErr = recvEinoSchemaMessageStreamWithContext(ctx, stream, 8, func(chunk *schema.Message) { + if chunk.Content != "" { + buf.WriteString(chunk.Content) } - } -} - -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) + if tid := strings.TrimSpace(chunk.ToolCallID); tid != "" { + toolCallID = tid } - } - cleaned := strings.TrimSpace(lastAssistant) - if orchMode == "plan_execute" { - if e := strings.TrimSpace(lastPlanExecuteExecutor); e != "" { - cleaned = e - } else { - cleaned = UnwrapPlanExecuteUserText(cleaned) + if name := strings.TrimSpace(chunk.ToolName); name != "" { + toolName = name } - } - if cleaned == "" { - if fb := strings.TrimSpace(einoExtractFallbackAssistantFromMsgs(runAccumulatedMsgs)); fb != "" { - cleaned = fb - } - } - cleaned = dedupeRepeatedParagraphs(cleaned, 80) - cleaned = dedupeParagraphsByLineFingerprint(cleaned, 100) - // 防止超长响应导致 JSON 序列化慢或 OOM(多代理拼接大量工具输出时可能触发)。 - 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)) + }) + return buf.String(), toolCallID, toolName, recvErr } func buildEinoCheckpointID(orchMode string) string { @@ -1515,3 +426,11 @@ func buildEinoCheckpointID(orchMode string) string { } return "runner-" + mode } + +func buildEinoTurnLoopCheckpointID(orchMode string) string { + mode := sanitizeEinoPathSegment(strings.TrimSpace(orchMode)) + if mode == "" { + mode = "default" + } + return "turn-loop-" + mode +} diff --git a/internal/multiagent/eino_adk_run_loop_stream_test.go b/internal/multiagent/eino_adk_run_loop_stream_test.go index 4c216938..2ca8381a 100644 --- a/internal/multiagent/eino_adk_run_loop_stream_test.go +++ b/internal/multiagent/eino_adk_run_loop_stream_test.go @@ -15,7 +15,7 @@ func TestRecvSchemaMessageStream_EOF(t *testing.T) { _ = sw.Send(schema.ToolMessage("hello", "tc-1"), nil) sw.Close() - content, tid, err := recvSchemaMessageStream(context.Background(), sr) + content, tid, toolName, err := recvSchemaMessageStream(context.Background(), sr) if err != nil { t.Fatalf("unexpected err: %v", err) } @@ -25,6 +25,23 @@ func TestRecvSchemaMessageStream_EOF(t *testing.T) { if tid != "tc-1" { t.Fatalf("toolCallID=%q want tc-1", tid) } + if toolName != "" { + t.Fatalf("toolName=%q want empty", toolName) + } +} + +func TestRecvSchemaMessageStream_CapturesToolName(t *testing.T) { + sr, sw := schema.Pipe[*schema.Message](4) + _ = sw.Send(schema.ToolMessage("hello", "tc-1", schema.WithToolName("execute")), nil) + sw.Close() + + content, tid, toolName, err := recvSchemaMessageStream(context.Background(), sr) + if err != nil { + t.Fatalf("unexpected err: %v", err) + } + if content != "hello" || tid != "tc-1" || toolName != "execute" { + t.Fatalf("content=%q tid=%q toolName=%q", content, tid, toolName) + } } func TestRecvSchemaMessageStream_ContextCancel(t *testing.T) { @@ -37,7 +54,7 @@ func TestRecvSchemaMessageStream_ContextCancel(t *testing.T) { cancel() }() - content, _, err := recvSchemaMessageStream(ctx, sr) + content, _, _, err := recvSchemaMessageStream(ctx, sr) if !errors.Is(err, context.Canceled) { t.Fatalf("want context.Canceled, got %v content=%q", err, content) } @@ -49,16 +66,16 @@ func TestRecvSchemaMessageStream_RecvError(t *testing.T) { _ = sw.Send(nil, want) sw.Close() - _, _, err := recvSchemaMessageStream(context.Background(), sr) + _, _, _, err := recvSchemaMessageStream(context.Background(), sr) if !errors.Is(err, want) { t.Fatalf("want %v, got %v", want, err) } } func TestRecvSchemaMessageStream_NilStream(t *testing.T) { - content, tid, err := recvSchemaMessageStream(context.Background(), nil) - if err != nil || content != "" || tid != "" { - t.Fatalf("nil stream: content=%q tid=%q err=%v", content, tid, err) + content, tid, toolName, err := recvSchemaMessageStream(context.Background(), nil) + if err != nil || content != "" || tid != "" || toolName != "" { + t.Fatalf("nil stream: content=%q tid=%q toolName=%q err=%v", content, tid, toolName, err) } } @@ -67,8 +84,39 @@ func TestRecvSchemaMessageStream_EOFViaEmptyRead(t *testing.T) { _ = sw.Send(nil, io.EOF) sw.Close() - _, _, err := recvSchemaMessageStream(context.Background(), sr) + _, _, _, err := recvSchemaMessageStream(context.Background(), sr) if err != nil { t.Fatalf("EOF should not surface as error, got %v", err) } } + +func TestRecvEinoSchemaMessageStreamWithContext_SkipsNilChunks(t *testing.T) { + sr, sw := schema.Pipe[*schema.Message](4) + _ = sw.Send(nil, nil) + _ = sw.Send(schema.AssistantMessage("hello", nil), nil) + sw.Close() + + var got []string + err := recvEinoSchemaMessageStreamWithContext(context.Background(), sr, 1, func(chunk *schema.Message) { + got = append(got, chunk.Content) + }) + if err != nil { + t.Fatalf("unexpected err: %v", err) + } + if len(got) != 1 || got[0] != "hello" { + t.Fatalf("chunks = %#v, want [hello]", got) + } +} + +func TestRecvEinoSchemaMessageStreamWithContext_NilStream(t *testing.T) { + called := false + err := recvEinoSchemaMessageStreamWithContext(context.Background(), nil, 0, func(*schema.Message) { + called = true + }) + if err != nil { + t.Fatalf("nil stream should not error, got %v", err) + } + if called { + t.Fatal("nil stream should not call handler") + } +} diff --git a/internal/multiagent/eino_agentic_agent_adapter.go b/internal/multiagent/eino_agentic_agent_adapter.go new file mode 100644 index 00000000..cdc7b2d9 --- /dev/null +++ b/internal/multiagent/eino_agentic_agent_adapter.go @@ -0,0 +1,81 @@ +package multiagent + +import ( + "context" + "fmt" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +type einoAgenticMessageAgentAdapter struct { + inner adk.TypedAgent[*schema.AgenticMessage] +} + +func newEinoAgenticMessageAgentAdapter(inner adk.TypedAgent[*schema.AgenticMessage]) adk.Agent { + if inner == nil { + return nil + } + return &einoAgenticMessageAgentAdapter{inner: inner} +} + +func (a *einoAgenticMessageAgentAdapter) Name(ctx context.Context) string { + if a == nil || a.inner == nil { + return "" + } + return a.inner.Name(ctx) +} + +func (a *einoAgenticMessageAgentAdapter) Description(ctx context.Context) string { + if a == nil || a.inner == nil { + return "" + } + return a.inner.Description(ctx) +} + +func (a *einoAgenticMessageAgentAdapter) Run(ctx context.Context, input *adk.AgentInput, opts ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] { + return a.runTyped(ctx, input, nil, opts...) +} + +func (a *einoAgenticMessageAgentAdapter) Resume(ctx context.Context, info *adk.ResumeInfo, opts ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] { + return a.runTyped(ctx, nil, info, opts...) +} + +func (a *einoAgenticMessageAgentAdapter) runTyped(ctx context.Context, input *adk.AgentInput, resumeInfo *adk.ResumeInfo, opts ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] { + iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + go func() { + defer gen.Close() + if a == nil || a.inner == nil { + gen.Send(&adk.AgentEvent{Err: fmt.Errorf("agentic adapter: inner agent is nil")}) + return + } + var agenticIter *adk.AsyncIterator[*adk.TypedAgentEvent[*schema.AgenticMessage]] + if resumeInfo != nil { + resumable, ok := a.inner.(adk.TypedResumableAgent[*schema.AgenticMessage]) + if !ok { + gen.Send(&adk.AgentEvent{Err: fmt.Errorf("agentic adapter: inner agent does not support resume")}) + return + } + agenticIter = resumable.Resume(ctx, resumeInfo, opts...) + } else { + agenticInput := &adk.TypedAgentInput[*schema.AgenticMessage]{} + if input != nil { + agenticInput.EnableStreaming = input.EnableStreaming + agenticInput.Messages = EinoMessagesToAgentic(input.Messages) + } + agenticIter = a.inner.Run(ctx, agenticInput, opts...) + } + for { + ev, ok := agenticIter.Next() + if !ok { + return + } + for _, adapted := range adaptAgenticEventToEinoEvents(ev) { + if adapted != nil { + gen.Send(adapted) + } + } + } + }() + return iter +} diff --git a/internal/multiagent/eino_agentic_agent_adapter_test.go b/internal/multiagent/eino_agentic_agent_adapter_test.go new file mode 100644 index 00000000..b28aecc0 --- /dev/null +++ b/internal/multiagent/eino_agentic_agent_adapter_test.go @@ -0,0 +1,145 @@ +package multiagent + +import ( + "context" + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +type fakeAgenticMessageAgent struct { + name string + description string + captured *adk.TypedAgentInput[*schema.AgenticMessage] + resumeInfo *adk.ResumeInfo + events []*adk.TypedAgentEvent[*schema.AgenticMessage] +} + +func (f *fakeAgenticMessageAgent) Name(context.Context) string { + return f.name +} + +func (f *fakeAgenticMessageAgent) Description(context.Context) string { + return f.description +} + +func (f *fakeAgenticMessageAgent) Run(_ context.Context, input *adk.TypedAgentInput[*schema.AgenticMessage], _ ...adk.AgentRunOption) *adk.AsyncIterator[*adk.TypedAgentEvent[*schema.AgenticMessage]] { + f.captured = input + iter, gen := adk.NewAsyncIteratorPair[*adk.TypedAgentEvent[*schema.AgenticMessage]]() + go func() { + defer gen.Close() + for _, ev := range f.events { + gen.Send(ev) + } + }() + return iter +} + +func (f *fakeAgenticMessageAgent) Resume(_ context.Context, info *adk.ResumeInfo, _ ...adk.AgentRunOption) *adk.AsyncIterator[*adk.TypedAgentEvent[*schema.AgenticMessage]] { + f.resumeInfo = info + iter, gen := adk.NewAsyncIteratorPair[*adk.TypedAgentEvent[*schema.AgenticMessage]]() + go func() { + defer gen.Close() + for _, ev := range f.events { + gen.Send(ev) + } + }() + return iter +} + +func TestEinoAgenticMessageAgentAdapterConvertsInputAndEvents(t *testing.T) { + inner := &fakeAgenticMessageAgent{ + name: "agentic", + description: "typed agent", + events: []*adk.TypedAgentEvent[*schema.AgenticMessage]{ + { + AgentName: "agentic", + Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{ + MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{ + Message: &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{ + schema.NewContentBlock(&schema.AssistantGenText{Text: "hello"}), + }, + }, + }, + }, + }, + }, + } + agent := newEinoAgenticMessageAgentAdapter(inner) + + if agent.Name(context.Background()) != "agentic" || agent.Description(context.Background()) != "typed agent" { + t.Fatalf("adapter metadata name=%q desc=%q", agent.Name(context.Background()), agent.Description(context.Background())) + } + iter := agent.Run(context.Background(), &adk.AgentInput{ + EnableStreaming: true, + Messages: []*schema.Message{ + schema.UserMessage("hi"), + }, + }) + + ev, ok := iter.Next() + if !ok { + t.Fatal("expected adapted event") + } + if inner.captured == nil || !inner.captured.EnableStreaming || len(inner.captured.Messages) != 1 { + t.Fatalf("captured input = %#v", inner.captured) + } + if inner.captured.Messages[0].Role != schema.AgenticRoleTypeUser || inner.captured.Messages[0].ContentBlocks[0].UserInputText.Text != "hi" { + t.Fatalf("captured message = %#v", inner.captured.Messages[0]) + } + if ev.AgentName != "agentic" || ev.Output == nil || ev.Output.MessageOutput == nil { + t.Fatalf("event = %#v", ev) + } + if ev.Output.MessageOutput.Role != schema.Assistant || ev.Output.MessageOutput.Message.Content != "hello" { + t.Fatalf("message output = %#v", ev.Output.MessageOutput) + } + if _, ok := iter.Next(); ok { + t.Fatal("expected iterator to close") + } +} + +func TestEinoAgenticMessageAgentAdapterNilInnerReturnsNil(t *testing.T) { + if got := newEinoAgenticMessageAgentAdapter(nil); got != nil { + t.Fatalf("adapter = %#v, want nil", got) + } +} + +func TestEinoAgenticMessageAgentAdapterResumeConvertsEvents(t *testing.T) { + inner := &fakeAgenticMessageAgent{ + name: "agentic", + events: []*adk.TypedAgentEvent[*schema.AgenticMessage]{ + { + AgentName: "agentic", + Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{ + MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{ + Message: &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{ + schema.NewContentBlock(&schema.AssistantGenText{Text: "resumed"}), + }, + }, + }, + }, + }, + }, + } + agent, ok := newEinoAgenticMessageAgentAdapter(inner).(adk.ResumableAgent) + if !ok { + t.Fatal("adapter must implement adk.ResumableAgent") + } + info := &adk.ResumeInfo{WasInterrupted: true} + iter := agent.Resume(context.Background(), info) + ev, ok := iter.Next() + if !ok { + t.Fatal("expected adapted resume event") + } + if inner.resumeInfo != info { + t.Fatalf("resume info = %#v, want original pointer", inner.resumeInfo) + } + if ev.Output == nil || ev.Output.MessageOutput == nil || ev.Output.MessageOutput.Message.Content != "resumed" { + t.Fatalf("resume event = %#v", ev) + } +} diff --git a/internal/multiagent/eino_agentic_chat_model_agent.go b/internal/multiagent/eino_agentic_chat_model_agent.go new file mode 100644 index 00000000..fb767154 --- /dev/null +++ b/internal/multiagent/eino_agentic_chat_model_agent.go @@ -0,0 +1,64 @@ +package multiagent + +import ( + "context" + "fmt" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/components/tool" + "github.com/cloudwego/eino/schema" +) + +type einoAgenticChatModelAgentConfig struct { + Name string + Description string + Instruction string + Model model.AgenticModel + ToolsConfig adk.ToolsConfig + MaxIterations int + Exit tool.BaseTool + + GenModelInput adk.TypedGenModelInput[*schema.AgenticMessage] + Handlers []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] + ModelRetryConfig *adk.TypedModelRetryConfig[*schema.AgenticMessage] + ModelFailoverConfig *adk.ModelFailoverConfig[*schema.AgenticMessage] + OutputKey string +} + +func newEinoAgenticChatModelAgent(ctx context.Context, cfg einoAgenticChatModelAgentConfig) (adk.TypedResumableAgent[*schema.AgenticMessage], error) { + if cfg.Model == nil { + return nil, fmt.Errorf("eino agentic ChatModelAgent: model is required") + } + typedCfg := &adk.TypedChatModelAgentConfig[*schema.AgenticMessage]{ + Name: cfg.Name, + Description: cfg.Description, + Instruction: cfg.Instruction, + Model: cfg.Model, + ToolsConfig: cfg.ToolsConfig, + MaxIterations: cfg.MaxIterations, + Exit: cfg.Exit, + GenModelInput: cfg.GenModelInput, + Handlers: cfg.Handlers, + ModelRetryConfig: cfg.ModelRetryConfig, + ModelFailoverConfig: cfg.ModelFailoverConfig, + OutputKey: cfg.OutputKey, + } + typedAgent, err := adk.NewTypedChatModelAgent(ctx, typedCfg) + if err != nil { + return nil, fmt.Errorf("eino agentic NewTypedChatModelAgent: %w", err) + } + return typedAgent, nil +} + +func newEinoAgenticChatModelAgentAdapter(ctx context.Context, cfg einoAgenticChatModelAgentConfig) (adk.Agent, error) { + typedAgent, err := newEinoAgenticChatModelAgent(ctx, cfg) + if err != nil { + return nil, err + } + agent := newEinoAgenticMessageAgentAdapter(typedAgent) + if agent == nil { + return nil, fmt.Errorf("eino agentic ChatModelAgent: adapter is nil") + } + return agent, nil +} diff --git a/internal/multiagent/eino_agentic_chat_model_agent_test.go b/internal/multiagent/eino_agentic_chat_model_agent_test.go new file mode 100644 index 00000000..63761c71 --- /dev/null +++ b/internal/multiagent/eino_agentic_chat_model_agent_test.go @@ -0,0 +1,163 @@ +package multiagent + +import ( + "context" + "strings" + "sync" + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/schema" +) + +type capturingAgenticChatModel struct { + mu sync.Mutex + inputs [][]*schema.AgenticMessage + output *schema.AgenticMessage +} + +func (m *capturingAgenticChatModel) Generate(_ context.Context, input []*schema.AgenticMessage, _ ...model.Option) (*schema.AgenticMessage, error) { + m.mu.Lock() + m.inputs = append(m.inputs, input) + m.mu.Unlock() + if m.output != nil { + return m.output, nil + } + return &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.AssistantGenText{Text: "agentic answer"})}, + }, nil +} + +func (m *capturingAgenticChatModel) Stream(_ context.Context, input []*schema.AgenticMessage, _ ...model.Option) (*schema.StreamReader[*schema.AgenticMessage], error) { + msg, err := m.Generate(context.Background(), input) + if err != nil { + return nil, err + } + return schema.StreamReaderFromArray([]*schema.AgenticMessage{msg}), nil +} + +func (m *capturingAgenticChatModel) snapshotInputs() [][]*schema.AgenticMessage { + m.mu.Lock() + defer m.mu.Unlock() + out := make([][]*schema.AgenticMessage, len(m.inputs)) + copy(out, m.inputs) + return out +} + +func TestNewEinoAgenticChatModelAgentAdapterRunsThroughClassicAgentBoundary(t *testing.T) { + t.Parallel() + ctx := context.Background() + trace := newModelFacingTraceHolder() + fakeModel := &capturingAgenticChatModel{} + agent, err := newEinoAgenticChatModelAgentAdapter(ctx, einoAgenticChatModelAgentConfig{ + Name: "agentic", + Description: "agentic adapter test", + Instruction: "system instruction", + Model: fakeModel, + Handlers: appendEinoAgenticChatModelTailMiddlewares(nil, einoChatModelTailConfig{ + phase: "agentic", + trace: trace, + }), + }) + if err != nil { + t.Fatalf("newEinoAgenticChatModelAgentAdapter: %v", err) + } + + iter := agent.Run(ctx, &adk.AgentInput{ + Messages: []*schema.Message{schema.UserMessage("classic input")}, + }) + var last *adk.AgentEvent + for { + ev, ok := iter.Next() + if !ok { + break + } + if ev.Err != nil { + t.Fatalf("agent event error: %v", ev.Err) + } + last = ev + } + if last == nil || last.Output == nil || last.Output.MessageOutput == nil { + t.Fatalf("last event = %#v, want message output", last) + } + if got := last.Output.MessageOutput.Message.Content; got != "agentic answer" { + t.Fatalf("classic output content = %q, want agentic answer", got) + } + + inputs := fakeModel.snapshotInputs() + if len(inputs) != 1 { + t.Fatalf("model calls = %d, want 1", len(inputs)) + } + if len(inputs[0]) != 2 { + t.Fatalf("model input messages = %d, want instruction + user", len(inputs[0])) + } + if inputs[0][0].Role != schema.AgenticRoleTypeSystem || agenticMessageText(inputs[0][0]) != "system instruction" { + t.Fatalf("first agentic input = %#v", inputs[0][0]) + } + if inputs[0][1].Role != schema.AgenticRoleTypeUser || agenticMessageText(inputs[0][1]) != "classic input" { + t.Fatalf("second agentic input = %#v", inputs[0][1]) + } + + snapshot := trace.Snapshot() + if len(snapshot) != 2 || snapshot[0].Role != schema.System || snapshot[1].Role != schema.User { + t.Fatalf("trace snapshot = %#v, want classic system + user trace", snapshot) + } +} + +func TestNewEinoAgenticChatModelAgentAdapterPreservesTypedToolCallsForToolLayerRecovery(t *testing.T) { + t.Parallel() + ctx := context.Background() + fakeModel := &capturingAgenticChatModel{ + output: &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.FunctionToolCall{ + CallID: "call-1", + Name: "exec", + Arguments: `{"command":"` + strings.Repeat("x", 20000) + `"}`, + })}, + }, + } + agent, err := newEinoAgenticChatModelAgentAdapter(ctx, einoAgenticChatModelAgentConfig{ + Name: "agentic", + Description: "agentic adapter test", + Model: fakeModel, + Handlers: appendEinoAgenticChatModelTailMiddlewares(nil, einoChatModelTailConfig{ + phase: "agentic", + }), + }) + if err != nil { + t.Fatalf("newEinoAgenticChatModelAgentAdapter: %v", err) + } + iter := agent.Run(ctx, &adk.AgentInput{Messages: []*schema.Message{schema.UserMessage("run")}}) + var last *adk.AgentEvent + for { + ev, ok := iter.Next() + if !ok { + break + } + if ev.Err != nil { + t.Fatalf("agent event error: %v", ev.Err) + } + last = ev + } + if last == nil || last.Output == nil || last.Output.MessageOutput == nil { + t.Fatalf("last event = %#v, want message output", last) + } + msg := last.Output.MessageOutput.Message + if len(msg.ToolCalls) != 1 { + t.Fatalf("tool calls = %#v, want one tool call", msg.ToolCalls) + } + args := msg.ToolCalls[0].Function.Arguments + if !strings.Contains(args, strings.Repeat("x", 32)) || strings.Contains(args, modelOutputRecoveryKey) { + t.Fatalf("agentic tool args were unexpectedly rewritten: %q", args) + } +} + +func TestNewEinoAgenticChatModelAgentAdapterRequiresModel(t *testing.T) { + t.Parallel() + if _, err := newEinoAgenticChatModelAgentAdapter(context.Background(), einoAgenticChatModelAgentConfig{}); err == nil { + t.Fatal("expected missing model error") + } +} diff --git a/internal/multiagent/eino_agentic_chat_model_tail_middleware.go b/internal/multiagent/eino_agentic_chat_model_tail_middleware.go new file mode 100644 index 00000000..a82cd1ae --- /dev/null +++ b/internal/multiagent/eino_agentic_chat_model_tail_middleware.go @@ -0,0 +1,209 @@ +package multiagent + +import ( + "context" + "strings" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" +) + +// appendEinoAgenticChatModelTailMiddlewares appends protocol-neutral handlers for +// TypedChatModelAgent[*schema.AgenticMessage]. Classic ReAct history repair +// handlers stay on the schema.Message path because AgenticMessage has native +// content blocks for function calls/results. +func appendEinoAgenticChatModelTailMiddlewares( + handlers []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], + cfg einoChatModelTailConfig, +) []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] { + handlers = append(handlers, newAgenticSystemMessageNormalizerMiddleware(cfg.logger, cfg.phase)) + handlers = append(handlers, newAgenticContinuationUserDedupMiddleware(cfg.logger, cfg.phase)) + if cfg.agenticSummarization != nil { + handlers = append(handlers, cfg.agenticSummarization) + } + if !cfg.skipTrace && cfg.trace != nil { + if capMw := newAgenticModelFacingTraceMiddleware(cfg.trace); capMw != nil { + handlers = append(handlers, capMw) + } + } + return handlers +} + +type agenticSystemMessageNormalizerMiddleware struct { + *adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage] + logger *zap.Logger + phase string +} + +func newAgenticSystemMessageNormalizerMiddleware(logger *zap.Logger, phase string) adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] { + return &agenticSystemMessageNormalizerMiddleware{ + TypedBaseChatModelAgentMiddleware: &adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage]{}, + logger: logger, + phase: phase, + } +} + +func (m *agenticSystemMessageNormalizerMiddleware) BeforeModelRewriteState( + ctx context.Context, + state *adk.TypedChatModelAgentState[*schema.AgenticMessage], + mc *adk.TypedModelContext[*schema.AgenticMessage], +) (context.Context, *adk.TypedChatModelAgentState[*schema.AgenticMessage], error) { + _ = mc + if m == nil || state == nil || len(state.Messages) == 0 { + return ctx, state, nil + } + before := countAgenticSystemMessages(state.Messages) + if before <= 1 { + return ctx, state, nil + } + normalized := normalizeSingleLeadingAgenticSystemMessage(state.Messages) + if len(normalized) == len(state.Messages) && countAgenticSystemMessages(normalized) >= before { + return ctx, state, nil + } + if m.logger != nil { + m.logger.Info("eino agentic system messages merged", + zap.String("phase", m.phase), + zap.Int("system_before", before), + zap.Int("system_after", countAgenticSystemMessages(normalized)), + zap.Int("messages_before", len(state.Messages)), + zap.Int("messages_after", len(normalized)), + ) + } + out := *state + out.Messages = normalized + return ctx, &out, nil +} + +type agenticContinuationUserDedupMiddleware struct { + *adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage] + logger *zap.Logger + phase string +} + +func newAgenticContinuationUserDedupMiddleware(logger *zap.Logger, phase string) adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] { + return &agenticContinuationUserDedupMiddleware{ + TypedBaseChatModelAgentMiddleware: &adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage]{}, + logger: logger, + phase: phase, + } +} + +func (m *agenticContinuationUserDedupMiddleware) BeforeModelRewriteState( + ctx context.Context, + state *adk.TypedChatModelAgentState[*schema.AgenticMessage], + mc *adk.TypedModelContext[*schema.AgenticMessage], +) (context.Context, *adk.TypedChatModelAgentState[*schema.AgenticMessage], error) { + _ = mc + if m == nil || state == nil || len(state.Messages) == 0 { + return ctx, state, nil + } + deduped, dropped := dedupAgenticContinuationUserMessages(state.Messages) + if dropped == 0 { + return ctx, state, nil + } + if m.logger != nil { + m.logger.Info("eino agentic continuation user messages deduplicated", + zap.String("phase", m.phase), + zap.Int("dropped", dropped), + zap.Int("messages_before", len(state.Messages)), + zap.Int("messages_after", len(deduped)), + ) + } + out := *state + out.Messages = deduped + return ctx, &out, nil +} + +func countAgenticSystemMessages(msgs []*schema.AgenticMessage) int { + n := 0 + for _, msg := range msgs { + if msg != nil && msg.Role == schema.AgenticRoleTypeSystem { + n++ + } + } + return n +} + +func normalizeSingleLeadingAgenticSystemMessage(msgs []*schema.AgenticMessage) []*schema.AgenticMessage { + var systemParts []string + out := make([]*schema.AgenticMessage, 0, len(msgs)) + for _, msg := range msgs { + if msg == nil { + continue + } + if msg.Role == schema.AgenticRoleTypeSystem { + if text := strings.TrimSpace(agenticMessageText(msg)); text != "" { + systemParts = append(systemParts, text) + } + continue + } + out = append(out, msg) + } + if len(systemParts) == 0 { + return out + } + merged := schema.SystemAgenticMessage(strings.Join(systemParts, "\n\n")) + return append([]*schema.AgenticMessage{merged}, out...) +} + +func dedupAgenticContinuationUserMessages(msgs []*schema.AgenticMessage) ([]*schema.AgenticMessage, int) { + lastIdx := -1 + contCount := 0 + for i, msg := range msgs { + if !isAgenticContinuationUserMessage(msg) { + continue + } + contCount++ + lastIdx = i + } + if contCount <= 1 { + return msgs, 0 + } + out := make([]*schema.AgenticMessage, 0, len(msgs)-(contCount-1)) + dropped := 0 + for i, msg := range msgs { + if isAgenticContinuationUserMessage(msg) && i != lastIdx { + dropped++ + continue + } + out = append(out, msg) + } + return out, dropped +} + +func isAgenticContinuationUserMessage(msg *schema.AgenticMessage) bool { + if msg == nil || msg.Role != schema.AgenticRoleTypeUser { + return false + } + return strings.Contains(agenticMessageText(msg), continuationSessionMarker) +} + +func agenticMessageText(msg *schema.AgenticMessage) string { + if msg == nil { + return "" + } + var b strings.Builder + for _, block := range msg.ContentBlocks { + if block == nil { + continue + } + switch { + case block.UserInputText != nil: + if s := strings.TrimSpace(block.UserInputText.Text); s != "" { + if b.Len() > 0 { + b.WriteByte('\n') + } + b.WriteString(s) + } + case block.AssistantGenText != nil: + if s := strings.TrimSpace(block.AssistantGenText.Text); s != "" { + if b.Len() > 0 { + b.WriteByte('\n') + } + b.WriteString(s) + } + } + } + return b.String() +} diff --git a/internal/multiagent/eino_agentic_chat_model_tail_middleware_test.go b/internal/multiagent/eino_agentic_chat_model_tail_middleware_test.go new file mode 100644 index 00000000..d020eb7d --- /dev/null +++ b/internal/multiagent/eino_agentic_chat_model_tail_middleware_test.go @@ -0,0 +1,112 @@ +package multiagent + +import ( + "context" + "strings" + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestAgenticSystemMessageNormalizerMiddlewareMergesDuplicates(t *testing.T) { + t.Parallel() + mw := newAgenticSystemMessageNormalizerMiddleware(nil, "test") + state := &adk.TypedChatModelAgentState[*schema.AgenticMessage]{ + Messages: []*schema.AgenticMessage{ + schema.SystemAgenticMessage("first"), + schema.UserAgenticMessage("hello"), + schema.SystemAgenticMessage("second"), + }, + } + _, out, err := mw.BeforeModelRewriteState(context.Background(), state, nil) + if err != nil { + t.Fatalf("BeforeModelRewriteState: %v", err) + } + if out == state { + t.Fatal("expected rewritten state") + } + if got := countAgenticSystemMessages(out.Messages); got != 1 { + t.Fatalf("system messages = %d, want 1", got) + } + if out.Messages[0].Role != schema.AgenticRoleTypeSystem { + t.Fatalf("first role = %s, want system", out.Messages[0].Role) + } + text := agenticMessageText(out.Messages[0]) + if !strings.Contains(text, "first") || !strings.Contains(text, "second") { + t.Fatalf("merged system text = %q", text) + } + if len(out.Messages) != 2 || agenticMessageText(out.Messages[1]) != "hello" { + t.Fatalf("normalized messages = %#v", out.Messages) + } +} + +func TestAgenticContinuationUserDedupMiddlewareKeepsLatest(t *testing.T) { + t.Parallel() + mw := newAgenticContinuationUserDedupMiddleware(nil, "test") + state := &adk.TypedChatModelAgentState[*schema.AgenticMessage]{ + Messages: []*schema.AgenticMessage{ + schema.UserAgenticMessage(continuationSessionMarker + "\nold"), + schema.UserAgenticMessage("real user request"), + schema.UserAgenticMessage(continuationSessionMarker + "\nnew"), + }, + } + _, out, err := mw.BeforeModelRewriteState(context.Background(), state, nil) + if err != nil { + t.Fatalf("BeforeModelRewriteState: %v", err) + } + if out == state { + t.Fatal("expected rewritten state") + } + if len(out.Messages) != 2 { + t.Fatalf("messages = %d, want 2", len(out.Messages)) + } + if strings.Contains(agenticMessageText(out.Messages[0]), continuationSessionMarker) { + t.Fatalf("old continuation was not dropped: %#v", out.Messages) + } + if !strings.Contains(agenticMessageText(out.Messages[1]), "new") { + t.Fatalf("latest continuation not retained: %#v", out.Messages) + } +} + +func TestAgenticModelFacingTraceMiddlewareStoresClassicTrace(t *testing.T) { + t.Parallel() + holder := newModelFacingTraceHolder() + mw := newAgenticModelFacingTraceMiddleware(holder) + state := &adk.TypedChatModelAgentState[*schema.AgenticMessage]{ + Messages: []*schema.AgenticMessage{ + schema.SystemAgenticMessage("instruction"), + { + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{ + schema.NewContentBlock(&schema.AssistantGenText{Text: "answer"}), + }, + }, + }, + } + if _, _, err := mw.BeforeModelRewriteState(context.Background(), state, nil); err != nil { + t.Fatalf("BeforeModelRewriteState: %v", err) + } + got := holder.Snapshot() + if len(got) != 2 { + t.Fatalf("trace len = %d, want 2", len(got)) + } + if got[0].Role != schema.System || got[0].Content != "instruction" { + t.Fatalf("system trace = %#v", got[0]) + } + if got[1].Role != schema.Assistant || got[1].Content != "answer" { + t.Fatalf("assistant trace = %#v", got[1]) + } +} + +func TestAppendEinoAgenticChatModelTailMiddlewares(t *testing.T) { + t.Parallel() + holder := newModelFacingTraceHolder() + handlers := appendEinoAgenticChatModelTailMiddlewares(nil, einoChatModelTailConfig{ + phase: "agentic", + trace: holder, + }) + if len(handlers) != 3 { + t.Fatalf("handlers = %d, want system + continuation + trace", len(handlers)) + } +} diff --git a/internal/multiagent/eino_agentic_event_adapter.go b/internal/multiagent/eino_agentic_event_adapter.go new file mode 100644 index 00000000..0b91c49d --- /dev/null +++ b/internal/multiagent/eino_agentic_event_adapter.go @@ -0,0 +1,106 @@ +package multiagent + +import ( + "io" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +// adaptAgenticEventToEinoEvents converts typed AgenticMessage ADK events into +// the classic schema.Message events consumed by the existing SSE/MCP drain. +func adaptAgenticEventToEinoEvents(ev *adk.TypedAgentEvent[*schema.AgenticMessage]) []*adk.AgentEvent { + if ev == nil { + return nil + } + base := func(output *adk.AgentOutput) *adk.AgentEvent { + return &adk.AgentEvent{ + AgentName: ev.AgentName, + RunPath: append([]adk.RunStep(nil), ev.RunPath...), + Output: output, + Action: ev.Action, + Err: ev.Err, + } + } + if ev.Output == nil { + return []*adk.AgentEvent{base(nil)} + } + customized := ev.Output.CustomizedOutput + mv := ev.Output.MessageOutput + if mv == nil { + return []*adk.AgentEvent{base(&adk.AgentOutput{CustomizedOutput: customized})} + } + if mv.IsStreaming { + return []*adk.AgentEvent{base(&adk.AgentOutput{ + MessageOutput: &adk.MessageVariant{ + IsStreaming: true, + MessageStream: agenticStreamToEinoStream(mv.MessageStream), + Role: agenticVariantRole(mv), + }, + CustomizedOutput: customized, + })} + } + + msgs := AgenticMessageToEino(mv.Message) + if len(msgs) == 0 { + return []*adk.AgentEvent{base(&adk.AgentOutput{CustomizedOutput: customized})} + } + out := make([]*adk.AgentEvent, 0, len(msgs)) + for i, msg := range msgs { + eventCustomized := any(nil) + if i == 0 { + eventCustomized = customized + } + out = append(out, base(&adk.AgentOutput{ + MessageOutput: &adk.MessageVariant{ + Message: msg, + Role: msg.Role, + ToolName: msg.ToolName, + }, + CustomizedOutput: eventCustomized, + })) + } + return out +} + +func agenticStreamToEinoStream(sr *schema.StreamReader[*schema.AgenticMessage]) *schema.StreamReader[*schema.Message] { + out, writer := schema.Pipe[*schema.Message](8) + go func() { + defer writer.Close() + if sr == nil { + return + } + defer sr.Close() + for { + chunk, err := sr.Recv() + if err != nil { + if err != io.EOF { + writer.Send(nil, err) + } + return + } + for _, msg := range AgenticMessageToEino(chunk) { + if msg != nil && writer.Send(msg, nil) { + return + } + } + } + }() + return out +} + +func agenticVariantRole(mv *adk.TypedMessageVariant[*schema.AgenticMessage]) schema.RoleType { + if mv == nil { + return schema.Assistant + } + switch mv.AgenticRole { + case schema.AgenticRoleTypeSystem: + return schema.System + case schema.AgenticRoleTypeUser: + // In Agentic ReAct output, user-role events from the graph are local + // FunctionToolResult messages emitted by AgenticToolsNode. + return schema.Tool + default: + return schema.Assistant + } +} diff --git a/internal/multiagent/eino_agentic_event_adapter_test.go b/internal/multiagent/eino_agentic_event_adapter_test.go new file mode 100644 index 00000000..e60bbbea --- /dev/null +++ b/internal/multiagent/eino_agentic_event_adapter_test.go @@ -0,0 +1,249 @@ +package multiagent + +import ( + "errors" + "io" + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestAdaptAgenticEventToEinoEventsAssistantMessage(t *testing.T) { + usage := &schema.TokenUsage{PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15} + ev := &adk.TypedAgentEvent[*schema.AgenticMessage]{ + AgentName: "agentic", + Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{ + MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{ + Message: &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeAssistant, + ResponseMeta: &schema.AgenticResponseMeta{TokenUsage: usage}, + ContentBlocks: []*schema.ContentBlock{ + schema.NewContentBlock(&schema.Reasoning{Text: "think"}), + schema.NewContentBlock(&schema.AssistantGenText{Text: "calling"}), + schema.NewContentBlock(&schema.FunctionToolCall{CallID: "call-1", Name: "scan", Arguments: `{"host":"127.0.0.1"}`}), + }, + }, + }, + CustomizedOutput: "custom", + }, + } + + got := adaptAgenticEventToEinoEvents(ev) + if len(got) != 1 { + t.Fatalf("events = %d, want 1", len(got)) + } + mv := got[0].Output.MessageOutput + if got[0].AgentName != "agentic" || got[0].Output.CustomizedOutput != "custom" { + t.Fatalf("event metadata = %#v", got[0]) + } + if mv.Role != schema.Assistant || mv.Message.Role != schema.Assistant { + t.Fatalf("role = %q/%q, want assistant", mv.Role, mv.Message.Role) + } + if mv.Message.Content != "calling" || mv.Message.ReasoningContent != "think" { + t.Fatalf("message text = %#v", mv.Message) + } + if len(mv.Message.ToolCalls) != 1 || mv.Message.ToolCalls[0].ID != "call-1" || mv.Message.ToolCalls[0].Function.Name != "scan" { + t.Fatalf("tool calls = %#v", mv.Message.ToolCalls) + } + if mv.Message.ResponseMeta == nil || mv.Message.ResponseMeta.Usage != usage { + t.Fatalf("usage = %#v, want original usage", mv.Message.ResponseMeta) + } +} + +func TestAdaptAgenticEventToEinoEventsPureToolResult(t *testing.T) { + ev := &adk.TypedAgentEvent[*schema.AgenticMessage]{ + Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{ + MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{ + Message: &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeUser, + ContentBlocks: []*schema.ContentBlock{ + schema.NewContentBlock(&schema.FunctionToolResult{ + CallID: "call-2", + Name: "execute", + Content: []*schema.FunctionToolResultContentBlock{{ + Type: schema.FunctionToolResultContentBlockTypeText, + Text: &schema.UserInputText{Text: "done"}, + }}, + }), + }, + }, + }, + }, + } + + got := adaptAgenticEventToEinoEvents(ev) + if len(got) != 1 { + t.Fatalf("events = %d, want 1", len(got)) + } + msg := got[0].Output.MessageOutput.Message + if got[0].Output.MessageOutput.Role != schema.Tool || msg.Role != schema.Tool || msg.ToolName != "execute" || msg.ToolCallID != "call-2" || msg.Content != "done" { + t.Fatalf("tool event = %#v message=%#v", got[0].Output.MessageOutput, msg) + } +} + +func TestAdaptAgenticEventToEinoEventsSplitsMixedToolResult(t *testing.T) { + ev := &adk.TypedAgentEvent[*schema.AgenticMessage]{ + Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{ + MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{ + Message: &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{ + schema.NewContentBlock(&schema.AssistantGenText{Text: "text"}), + schema.NewContentBlock(&schema.FunctionToolResult{ + CallID: "call-3", + Name: "grep", + Content: []*schema.FunctionToolResultContentBlock{{ + Type: schema.FunctionToolResultContentBlockTypeText, + Text: &schema.UserInputText{Text: "match"}, + }}, + }), + }, + }, + }, + }, + } + + got := adaptAgenticEventToEinoEvents(ev) + if len(got) != 2 { + t.Fatalf("events = %d, want assistant + tool", len(got)) + } + if got[0].Output.MessageOutput.Role != schema.Assistant || got[0].Output.MessageOutput.Message.Content != "text" { + t.Fatalf("assistant event = %#v", got[0].Output.MessageOutput) + } + if got[1].Output.MessageOutput.Role != schema.Tool || got[1].Output.MessageOutput.Message.ToolName != "grep" { + t.Fatalf("tool event = %#v", got[1].Output.MessageOutput) + } +} + +func TestAdaptAgenticEventToEinoEventsStreamingAssistant(t *testing.T) { + stream := schema.StreamReaderFromArray([]*schema.AgenticMessage{ + { + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{ + schema.NewContentBlock(&schema.AssistantGenText{Text: "hel"}), + }, + }, + { + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{ + schema.NewContentBlock(&schema.AssistantGenText{Text: "lo"}), + }, + }, + }) + ev := &adk.TypedAgentEvent[*schema.AgenticMessage]{ + Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{ + MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{ + IsStreaming: true, + MessageStream: stream, + AgenticRole: schema.AgenticRoleTypeAssistant, + }, + }, + } + + got := adaptAgenticEventToEinoEvents(ev) + if len(got) != 1 { + t.Fatalf("events = %d, want 1", len(got)) + } + mv := got[0].Output.MessageOutput + if !mv.IsStreaming || mv.Role != schema.Assistant { + t.Fatalf("stream variant = %#v", mv) + } + first, err := mv.MessageStream.Recv() + if err != nil || first.Content != "hel" { + t.Fatalf("first = %#v err=%v", first, err) + } + second, err := mv.MessageStream.Recv() + if err != nil || second.Content != "lo" { + t.Fatalf("second = %#v err=%v", second, err) + } + _, err = mv.MessageStream.Recv() + if !errors.Is(err, io.EOF) { + t.Fatalf("final err = %v, want EOF", err) + } +} + +func TestAdaptAgenticStreamingToolResultFeedsClassicToolResultHandler(t *testing.T) { + stream := schema.StreamReaderFromArray([]*schema.AgenticMessage{ + { + Role: schema.AgenticRoleTypeUser, + ContentBlocks: []*schema.ContentBlock{ + schema.NewContentBlock(&schema.FunctionToolResult{ + CallID: "call-agentic-stream", + Name: "execute", + Content: []*schema.FunctionToolResultContentBlock{{ + Type: schema.FunctionToolResultContentBlockTypeText, + Text: &schema.UserInputText{Text: "partial "}, + }}, + }), + }, + }, + { + Role: schema.AgenticRoleTypeUser, + ContentBlocks: []*schema.ContentBlock{ + schema.NewContentBlock(&schema.FunctionToolResult{ + CallID: "call-agentic-stream", + Name: "execute", + Content: []*schema.FunctionToolResultContentBlock{{ + Type: schema.FunctionToolResultContentBlockTypeText, + Text: &schema.UserInputText{Text: "done"}, + }}, + }), + }, + }, + }) + ev := &adk.TypedAgentEvent[*schema.AgenticMessage]{ + AgentName: "agentic", + Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{ + MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{ + IsStreaming: true, + MessageStream: stream, + AgenticRole: schema.AgenticRoleTypeUser, + }, + }, + } + + got := adaptAgenticEventToEinoEvents(ev) + if len(got) != 1 || got[0].Output == nil || got[0].Output.MessageOutput == nil { + t.Fatalf("events = %#v", got) + } + mv := got[0].Output.MessageOutput + if !mv.IsStreaming || mv.Role != schema.Tool { + t.Fatalf("streaming variant = %#v, want tool stream", mv) + } + + var event map[string]interface{} + runMessages := newEinoRunMessageAccumulator(nil) + emitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{ + ConversationID: "conv-agentic", + Progress: func(eventType, _ string, data interface{}) { + if eventType == "tool_result" { + event, _ = data.(map[string]interface{}) + } + }, + }) + handler := newEinoToolResultEventHandler(einoToolResultEventHandlerConfig{ + RunMessages: runMessages, + Emitter: emitter, + }) + if !handler.HandleStreaming(mv, "agentic") { + t.Fatal("agentic streaming tool result was not handled") + } + if event["toolName"] != "execute" || event["toolCallId"] != "call-agentic-stream" || event["result"] != "partial done" { + t.Fatalf("tool result event = %#v", event) + } + msgs := runMessages.Messages() + if len(msgs) != 1 || msgs[0].ToolName != "execute" || msgs[0].ToolCallID != "call-agentic-stream" || msgs[0].Content != "partial done" { + t.Fatalf("run messages = %#v", msgs) + } +} + +func TestAdaptAgenticEventToEinoEventsPreservesErrorOnlyEvent(t *testing.T) { + wantErr := errors.New("boom") + ev := &adk.TypedAgentEvent[*schema.AgenticMessage]{AgentName: "agentic", Err: wantErr} + + got := adaptAgenticEventToEinoEvents(ev) + if len(got) != 1 || got[0].AgentName != "agentic" || !errors.Is(got[0].Err, wantErr) { + t.Fatalf("events = %#v", got) + } +} diff --git a/internal/multiagent/eino_agentic_message.go b/internal/multiagent/eino_agentic_message.go new file mode 100644 index 00000000..b62e62ab --- /dev/null +++ b/internal/multiagent/eino_agentic_message.go @@ -0,0 +1,184 @@ +package multiagent + +import ( + "strings" + + "github.com/cloudwego/eino/schema" +) + +// EinoMessagesToAgentic converts the project's current ADK message history to +// Eino's native AgenticMessage shape. It intentionally covers the text, +// reasoning, function tool-call, and function tool-result channels used by the +// agent runtime today; unsupported multimodal/provider-specific fields stay in +// schema.Message until a real AgenticModel backend is wired. +func EinoMessagesToAgentic(msgs []*schema.Message) []*schema.AgenticMessage { + if len(msgs) == 0 { + return nil + } + out := make([]*schema.AgenticMessage, 0, len(msgs)) + for _, msg := range msgs { + if msg == nil { + continue + } + out = append(out, EinoMessageToAgentic(msg)) + } + return out +} + +func EinoMessageToAgentic(msg *schema.Message) *schema.AgenticMessage { + if msg == nil { + return nil + } + out := &schema.AgenticMessage{ + Role: messageRoleToAgentic(msg.Role), + Extra: cloneAnyMap(msg.Extra), + } + if msg.ResponseMeta != nil { + out.ResponseMeta = &schema.AgenticResponseMeta{TokenUsage: msg.ResponseMeta.Usage} + } + if text := strings.TrimSpace(msg.ReasoningContent); text != "" { + out.ContentBlocks = append(out.ContentBlocks, schema.NewContentBlock(&schema.Reasoning{Text: msg.ReasoningContent})) + } + switch msg.Role { + case schema.Assistant: + if msg.Content != "" { + out.ContentBlocks = append(out.ContentBlocks, schema.NewContentBlock(&schema.AssistantGenText{Text: msg.Content})) + } + for _, tc := range msg.ToolCalls { + out.ContentBlocks = append(out.ContentBlocks, schema.NewContentBlock(&schema.FunctionToolCall{ + CallID: tc.ID, + Name: tc.Function.Name, + Arguments: tc.Function.Arguments, + })) + } + case schema.Tool: + out.Role = schema.AgenticRoleTypeUser + out.ContentBlocks = append(out.ContentBlocks, schema.NewContentBlock(&schema.FunctionToolResult{ + CallID: msg.ToolCallID, + Name: msg.ToolName, + Content: []*schema.FunctionToolResultContentBlock{{ + Type: schema.FunctionToolResultContentBlockTypeText, + Text: &schema.UserInputText{Text: msg.Content}, + }}, + })) + default: + if msg.Content != "" { + out.ContentBlocks = append(out.ContentBlocks, schema.NewContentBlock(&schema.UserInputText{Text: msg.Content})) + } + } + return out +} + +// AgenticMessagesToEino converts AgenticMessage values back into the classic +// schema.Message form used by the existing ADK event drain and persistence code. +func AgenticMessagesToEino(msgs []*schema.AgenticMessage) []*schema.Message { + if len(msgs) == 0 { + return nil + } + out := make([]*schema.Message, 0, len(msgs)) + for _, msg := range msgs { + if msg == nil { + continue + } + out = append(out, AgenticMessageToEino(msg)...) + } + return out +} + +func AgenticMessageToEino(msg *schema.AgenticMessage) []*schema.Message { + if msg == nil { + return nil + } + base := &schema.Message{ + Role: agenticRoleToMessage(msg.Role), + Extra: cloneAnyMap(msg.Extra), + } + if msg.ResponseMeta != nil { + base.ResponseMeta = &schema.ResponseMeta{Usage: msg.ResponseMeta.TokenUsage} + } + var toolResults []*schema.Message + for _, block := range msg.ContentBlocks { + if block == nil { + continue + } + switch { + case block.Reasoning != nil: + base.ReasoningContent += block.Reasoning.Text + case block.UserInputText != nil: + base.Content += block.UserInputText.Text + case block.AssistantGenText != nil: + base.Role = schema.Assistant + base.Content += block.AssistantGenText.Text + case block.FunctionToolCall != nil: + base.Role = schema.Assistant + base.ToolCalls = append(base.ToolCalls, schema.ToolCall{ + ID: block.FunctionToolCall.CallID, + Type: "function", + Function: schema.FunctionCall{ + Name: block.FunctionToolCall.Name, + Arguments: block.FunctionToolCall.Arguments, + }, + }) + case block.FunctionToolResult != nil: + toolResults = append(toolResults, functionToolResultToMessage(block.FunctionToolResult)) + } + } + if len(toolResults) > 0 && base.Content == "" && base.ReasoningContent == "" && len(base.ToolCalls) == 0 { + return toolResults + } + out := []*schema.Message{base} + out = append(out, toolResults...) + return out +} + +func messageRoleToAgentic(role schema.RoleType) schema.AgenticRoleType { + switch role { + case schema.System: + return schema.AgenticRoleTypeSystem + case schema.Assistant: + return schema.AgenticRoleTypeAssistant + default: + return schema.AgenticRoleTypeUser + } +} + +func agenticRoleToMessage(role schema.AgenticRoleType) schema.RoleType { + switch role { + case schema.AgenticRoleTypeSystem: + return schema.System + case schema.AgenticRoleTypeAssistant: + return schema.Assistant + default: + return schema.User + } +} + +func functionToolResultToMessage(result *schema.FunctionToolResult) *schema.Message { + if result == nil { + return nil + } + parts := make([]string, 0, len(result.Content)) + for _, block := range result.Content { + if block == nil || block.Text == nil { + continue + } + parts = append(parts, block.Text.Text) + } + return &schema.Message{ + Role: schema.Tool, + Content: strings.Join(parts, ""), + ToolCallID: result.CallID, + ToolName: result.Name, + } +} + +func cloneAnyMap(in map[string]any) map[string]any { + if len(in) == 0 { + return nil + } + out := make(map[string]any, len(in)) + for k, v := range in { + out[k] = v + } + return out +} diff --git a/internal/multiagent/eino_agentic_message_test.go b/internal/multiagent/eino_agentic_message_test.go new file mode 100644 index 00000000..9fa89ab6 --- /dev/null +++ b/internal/multiagent/eino_agentic_message_test.go @@ -0,0 +1,154 @@ +package multiagent + +import ( + "testing" + + "github.com/cloudwego/eino/schema" +) + +func TestEinoMessageToAgenticPreservesAssistantToolCalls(t *testing.T) { + msg := &schema.Message{ + Role: schema.Assistant, + Content: "I will scan it.", + ReasoningContent: "Need enumerate first.", + ToolCalls: []schema.ToolCall{{ + ID: "call-1", + Type: "function", + Function: schema.FunctionCall{ + Name: "nmap", + Arguments: `{"target":"127.0.0.1"}`, + }, + }}, + Extra: map[string]any{"trace": "kept"}, + } + + got := EinoMessageToAgentic(msg) + if got.Role != schema.AgenticRoleTypeAssistant { + t.Fatalf("role = %q, want assistant", got.Role) + } + if len(got.ContentBlocks) != 3 { + t.Fatalf("blocks = %d, want 3", len(got.ContentBlocks)) + } + if got.ContentBlocks[0].Reasoning == nil || got.ContentBlocks[0].Reasoning.Text != msg.ReasoningContent { + t.Fatalf("reasoning block = %#v", got.ContentBlocks[0]) + } + if got.ContentBlocks[1].AssistantGenText == nil || got.ContentBlocks[1].AssistantGenText.Text != msg.Content { + t.Fatalf("assistant text block = %#v", got.ContentBlocks[1]) + } + call := got.ContentBlocks[2].FunctionToolCall + if call == nil || call.CallID != "call-1" || call.Name != "nmap" || call.Arguments != `{"target":"127.0.0.1"}` { + t.Fatalf("tool call block = %#v", got.ContentBlocks[2]) + } + if got.Extra["trace"] != "kept" { + t.Fatalf("extra = %#v", got.Extra) + } +} + +func TestEinoMessageToAgenticMapsToolResultAsUserFunctionResult(t *testing.T) { + msg := &schema.Message{ + Role: schema.Tool, + Content: "22/tcp open ssh", + ToolCallID: "call-ssh", + ToolName: "nmap", + } + + got := EinoMessageToAgentic(msg) + if got.Role != schema.AgenticRoleTypeUser { + t.Fatalf("role = %q, want user", got.Role) + } + if len(got.ContentBlocks) != 1 || got.ContentBlocks[0].FunctionToolResult == nil { + t.Fatalf("blocks = %#v", got.ContentBlocks) + } + result := got.ContentBlocks[0].FunctionToolResult + if result.CallID != "call-ssh" || result.Name != "nmap" { + t.Fatalf("tool result metadata = %#v", result) + } + if len(result.Content) != 1 || result.Content[0].Text == nil || result.Content[0].Text.Text != "22/tcp open ssh" { + t.Fatalf("tool result content = %#v", result.Content) + } +} + +func TestAgenticMessageToEinoPreservesAssistantBlocks(t *testing.T) { + msg := &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{ + schema.NewContentBlock(&schema.Reasoning{Text: "Think first."}), + schema.NewContentBlock(&schema.AssistantGenText{Text: "Calling scanner."}), + schema.NewContentBlock(&schema.FunctionToolCall{ + CallID: "call-2", + Name: "scan", + Arguments: `{"host":"example.com"}`, + }), + }, + } + + got := AgenticMessageToEino(msg) + if len(got) != 1 { + t.Fatalf("messages = %d, want 1", len(got)) + } + if got[0].Role != schema.Assistant || got[0].Content != "Calling scanner." || got[0].ReasoningContent != "Think first." { + t.Fatalf("assistant message = %#v", got[0]) + } + if len(got[0].ToolCalls) != 1 || got[0].ToolCalls[0].ID != "call-2" || got[0].ToolCalls[0].Function.Name != "scan" { + t.Fatalf("tool calls = %#v", got[0].ToolCalls) + } +} + +func TestAgenticMessageToEinoSplitsPureToolResult(t *testing.T) { + msg := &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeUser, + ContentBlocks: []*schema.ContentBlock{ + schema.NewContentBlock(&schema.FunctionToolResult{ + CallID: "call-3", + Name: "execute", + Content: []*schema.FunctionToolResultContentBlock{{ + Type: schema.FunctionToolResultContentBlockTypeText, + Text: &schema.UserInputText{Text: "done"}, + }}, + }), + }, + } + + got := AgenticMessageToEino(msg) + if len(got) != 1 { + t.Fatalf("messages = %d, want 1", len(got)) + } + if got[0].Role != schema.Tool || got[0].ToolCallID != "call-3" || got[0].ToolName != "execute" || got[0].Content != "done" { + t.Fatalf("tool message = %#v", got[0]) + } +} + +func TestEinoAgenticRoundTripForSupportedFields(t *testing.T) { + msgs := []*schema.Message{ + schema.SystemMessage("system"), + schema.UserMessage("user"), + { + Role: schema.Assistant, + Content: "assistant", + ToolCalls: []schema.ToolCall{{ + ID: "call-4", + Type: "function", + Function: schema.FunctionCall{Name: "grep", Arguments: `{"q":"token"}`}, + }}, + }, + { + Role: schema.Tool, + Content: "match", + ToolCallID: "call-4", + ToolName: "grep", + }, + } + + got := AgenticMessagesToEino(EinoMessagesToAgentic(msgs)) + if len(got) != len(msgs) { + t.Fatalf("round trip messages = %d, want %d: %#v", len(got), len(msgs), got) + } + for i := range msgs { + if got[i].Role != msgs[i].Role || got[i].Content != msgs[i].Content || got[i].ToolCallID != msgs[i].ToolCallID || got[i].ToolName != msgs[i].ToolName { + t.Fatalf("message[%d] = %#v, want %#v", i, got[i], msgs[i]) + } + if len(got[i].ToolCalls) != len(msgs[i].ToolCalls) { + t.Fatalf("message[%d] tool calls = %#v, want %#v", i, got[i].ToolCalls, msgs[i].ToolCalls) + } + } +} diff --git a/internal/multiagent/eino_agentic_model_gate.go b/internal/multiagent/eino_agentic_model_gate.go new file mode 100644 index 00000000..1b4e8db1 --- /dev/null +++ b/internal/multiagent/eino_agentic_model_gate.go @@ -0,0 +1,109 @@ +package multiagent + +import ( + "context" + "strings" + + "github.com/cloudwego/eino/components/model" + "go.uber.org/zap" +) + +type einoAgenticModelFactory func(context.Context) (model.AgenticModel, error) + +type einoAgenticRuntimeSupport struct { + TypedRunner bool + Streaming bool + CancelMonitoring bool + ModelRetry bool + ModelFailover bool + ToolResultObservation bool + MCPExecutionAudit bool +} + +type einoAgenticModelGate struct { + Ready bool + Reason string + Missing []string +} + +// Eino v0.9.14 wires AgenticMessage through the same generic TypedRunner, +// stream cancel monitoring, model retry, and model failover wrappers used by +// schema.Message. Keep this matrix explicit so future upgrades are audited +// deliberately instead of flipping the AgenticModel path by accident. +func einoAgenticRuntimeSupportV0914() einoAgenticRuntimeSupport { + return einoAgenticRuntimeSupport{ + TypedRunner: true, + Streaming: true, + CancelMonitoring: true, + ModelRetry: true, + ModelFailover: true, + ToolResultObservation: true, + MCPExecutionAudit: true, + } +} + +func evaluateEinoAgenticModelGate(factory einoAgenticModelFactory, support einoAgenticRuntimeSupport) einoAgenticModelGate { + missing := make([]string, 0, 8) + if factory == nil { + missing = append(missing, "model.AgenticModel backend") + } else { + if m, err := factory(context.Background()); err != nil || m == nil { + missing = append(missing, "model.AgenticModel backend") + } + } + if !support.TypedRunner { + missing = append(missing, "adk.TypedRunner[*schema.AgenticMessage]") + } + if !support.Streaming { + missing = append(missing, "AgenticMessage streaming") + } + if !support.CancelMonitoring { + missing = append(missing, "AgenticMessage model-stream cancel monitoring") + } + if !support.ModelRetry { + missing = append(missing, "AgenticMessage ModelRetry") + } + if !support.ModelFailover { + missing = append(missing, "AgenticMessage ModelFailover") + } + if !support.ToolResultObservation { + missing = append(missing, "AgenticMessage tool-result observation") + } + if !support.MCPExecutionAudit { + missing = append(missing, "AgenticMessage MCP execution audit") + } + if len(missing) == 0 { + return einoAgenticModelGate{Ready: true, Reason: "ready"} + } + return einoAgenticModelGate{ + Reason: "agentic_model_not_ready: " + strings.Join(missing, ", "), + Missing: missing, + } +} + +func logEinoAgenticModelGate(logger *zap.Logger, scope, orchestration string, gate einoAgenticModelGate) { + if logger == nil { + return + } + fields := []zap.Field{ + zap.String("scope", scope), + zap.String("orchestration", orchestration), + zap.Bool("ready", gate.Ready), + zap.String("reason", gate.Reason), + zap.Strings("missing", gate.Missing), + } + if gate.Ready { + logger.Info("eino agentic model gate ready", fields...) + return + } + logger.Info("eino agentic model gate disabled", fields...) +} + +func agenticTextModelFactory(m model.AgenticModel) einoAgenticModelFactory { + if m == nil { + return nil + } + return func(context.Context) (model.AgenticModel, error) { + return m, nil + } +} diff --git a/internal/multiagent/eino_agentic_model_gate_test.go b/internal/multiagent/eino_agentic_model_gate_test.go new file mode 100644 index 00000000..78ec18c1 --- /dev/null +++ b/internal/multiagent/eino_agentic_model_gate_test.go @@ -0,0 +1,93 @@ +package multiagent + +import ( + "context" + "errors" + "testing" + + "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/schema" +) + +type fakeAgenticGateModel struct{} + +func (m *fakeAgenticGateModel) Generate(context.Context, []*schema.AgenticMessage, ...model.Option) (*schema.AgenticMessage, error) { + return &schema.AgenticMessage{Role: schema.AgenticRoleTypeAssistant}, nil +} + +func (m *fakeAgenticGateModel) Stream(context.Context, []*schema.AgenticMessage, ...model.Option) (*schema.StreamReader[*schema.AgenticMessage], error) { + return schema.StreamReaderFromArray([]*schema.AgenticMessage{{Role: schema.AgenticRoleTypeAssistant}}), nil +} + +func TestEinoAgenticModelGateV0914WaitsOnlyForBackend(t *testing.T) { + gate := evaluateEinoAgenticModelGate(nil, einoAgenticRuntimeSupportV0914()) + + if gate.Ready { + t.Fatal("v0.9.14 gate should stay disabled without an AgenticModel backend") + } + if !containsString(gate.Missing, "model.AgenticModel backend") { + t.Fatalf("missing = %#v, want backend reason", gate.Missing) + } + for _, unexpected := range []string{ + "AgenticMessage model-stream cancel monitoring", + "AgenticMessage ModelRetry", + "AgenticMessage ModelFailover", + "AgenticMessage tool-result observation", + "AgenticMessage MCP execution audit", + } { + if containsString(gate.Missing, unexpected) { + t.Fatalf("missing = %#v, should not include %q for v0.9.14 runtime support", gate.Missing, unexpected) + } + } +} + +func TestEinoAgenticModelGateV0914ReadyWithBackend(t *testing.T) { + gate := evaluateEinoAgenticModelGate(agenticTextModelFactory(&fakeAgenticGateModel{}), einoAgenticRuntimeSupportV0914()) + + if !gate.Ready { + t.Fatalf("gate = %#v, want ready when v0.9.14 runtime support has a backend", gate) + } + if gate.Reason != "ready" || len(gate.Missing) != 0 { + t.Fatalf("gate details = %#v", gate) + } +} + +func TestEinoAgenticModelGateReadyWhenBackendAndRuntimeParityExist(t *testing.T) { + gate := evaluateEinoAgenticModelGate(agenticTextModelFactory(&fakeAgenticGateModel{}), einoAgenticRuntimeSupport{ + TypedRunner: true, + Streaming: true, + CancelMonitoring: true, + ModelRetry: true, + ModelFailover: true, + ToolResultObservation: true, + MCPExecutionAudit: true, + }) + + if !gate.Ready { + t.Fatalf("gate = %#v, want ready", gate) + } + if gate.Reason != "ready" || len(gate.Missing) != 0 { + t.Fatalf("gate details = %#v", gate) + } +} + +func TestEinoAgenticModelGateTreatsFactoryErrorAsMissingBackend(t *testing.T) { + gate := evaluateEinoAgenticModelGate(func(context.Context) (model.AgenticModel, error) { + return nil, errors.New("not implemented") + }, einoAgenticRuntimeSupport{ + TypedRunner: true, + Streaming: true, + CancelMonitoring: true, + ModelRetry: true, + ModelFailover: true, + ToolResultObservation: true, + MCPExecutionAudit: true, + }) + + if gate.Ready { + t.Fatal("factory error should disable gate") + } + if !containsString(gate.Missing, "model.AgenticModel backend") { + t.Fatalf("missing = %#v, want backend reason", gate.Missing) + } +} diff --git a/internal/multiagent/eino_agentic_summarize.go b/internal/multiagent/eino_agentic_summarize.go new file mode 100644 index 00000000..d4f1e45b --- /dev/null +++ b/internal/multiagent/eino_agentic_summarize.go @@ -0,0 +1,278 @@ +package multiagent + +import ( + "context" + "fmt" + "os" + "path/filepath" + "strings" + + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/database" + + einoopenai "github.com/cloudwego/eino-ext/components/model/openai" + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/adk/middlewares/summarization" + "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" +) + +// newEinoAgenticSummarizationMiddleware wires the project's domain-specific +// compaction policy into Eino's native typed AgenticMessage summarization. +func newEinoAgenticSummarizationMiddleware( + ctx context.Context, + summaryModel model.BaseModel[*schema.AgenticMessage], + appCfg *config.Config, + mwCfg *config.MultiAgentEinoMiddlewareConfig, + conversationID string, + db *database.DB, + projectID string, + logger *zap.Logger, +) (adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], error) { + if summaryModel == nil || appCfg == nil { + return nil, fmt.Errorf("multiagent: agentic summarization 需要 model 与配置") + } + maxTotal := appCfg.OpenAI.MaxTotalTokens + if maxTotal <= 0 { + maxTotal = 120000 + } + triggerRatio := 0.8 + emitInternalEvents := true + outputReserve := config.DefaultSummarizationOutputReserveTokens + userLedgerMaxRunes := config.DefaultSummarizationUserIntentLedgerMaxRunes + userLedgerEntryMaxRunes := config.DefaultSummarizationUserIntentLedgerEntryMaxRunes + toolMaxBytes := config.MultiAgentEinoMiddlewareConfig{}.ReductionMaxLengthForTruncEffective() + if mwCfg != nil { + triggerRatio = mwCfg.SummarizationTriggerRatioEffective() + emitInternalEvents = mwCfg.SummarizationEmitInternalEventsEffective() + outputReserve = mwCfg.SummarizationOutputReserveTokensEffective() + userLedgerMaxRunes = mwCfg.SummarizationUserIntentLedgerMaxRunesEffective() + userLedgerEntryMaxRunes = mwCfg.SummarizationUserIntentLedgerEntryMaxRunesEffective() + toolMaxBytes = mwCfg.ReductionMaxLengthForTruncEffective() + } + + ledgerWindowCap := modelFacingRuneBudget(maxTotal, 0.20) + userLedgerMaxRunes = minPositiveInt(userLedgerMaxRunes, ledgerWindowCap) + userLedgerEntryMaxRunes = minPositiveInt(userLedgerEntryMaxRunes, userLedgerMaxRunes) + trigger := int(float64(maxTotal) * triggerRatio) + if trigger < 4096 { + trigger = maxTotal + if trigger < 4096 { + trigger = 4096 + } + } + modelName := strings.TrimSpace(appCfg.OpenAI.Model) + if modelName == "" { + modelName = "gpt-4o" + } + classicTokenCounter := einoSummarizationTokenCounter(modelName) + agenticTokenCounter := func(ctx context.Context, input *summarization.TypedTokenCounterInput[*schema.AgenticMessage]) (int, error) { + if input == nil { + return 0, nil + } + return classicTokenCounter(ctx, &summarization.TokenCounterInput{ + Messages: AgenticMessagesToEino(input.Messages), + Tools: input.Tools, + }) + } + recentTrailMax := trigger / 4 + if recentTrailMax < 2048 { + recentTrailMax = 2048 + } + if recentTrailMax > trigger/2 { + recentTrailMax = trigger / 2 + } + summaryInputMax := trigger - outputReserve + if summaryInputMax < 4096 { + summaryInputMax = trigger * 80 / 100 + } + if summaryInputMax < 4096 { + summaryInputMax = 4096 + } + + transcriptPath := "" + if conv := strings.TrimSpace(conversationID); conv != "" { + baseRoot := filepath.Join(os.TempDir(), "cyberstrike-summarization") + if dbPath := strings.TrimSpace(appCfg.Database.Path); dbPath != "" { + baseRoot = filepath.Join(filepath.Dir(dbPath), "conversation_artifacts", sanitizeEinoPathSegment(conv), "summarization") + } + base := baseRoot + if abs, err := filepath.Abs(base); err == nil { + base = abs + } + if mkErr := os.MkdirAll(base, 0o755); mkErr == nil { + transcriptPath = filepath.Join(base, "transcript.txt") + } + } + + retryPolicy := einoTransientRunRetryPolicyFromMW(mwCfg) + retryMax := retryPolicy.maxAttempts + var summaryOverflowRetries int + summaryModelOpts := []model.Option{ + einoopenai.WithMaxCompletionTokens(outputReserve), + } + + mw, err := summarization.NewTyped[*schema.AgenticMessage](ctx, &summarization.TypedConfig[*schema.AgenticMessage]{ + Model: summaryModel, + ModelOptions: summaryModelOpts, + GenModelInput: func(ctx context.Context, sysInstruction, userInstruction *schema.AgenticMessage, originalMsgs []*schema.AgenticMessage) ([]*schema.AgenticMessage, error) { + classicOriginal := AgenticMessagesToEino(originalMsgs) + if transcriptPath != "" && len(classicOriginal) > 0 { + if werr := writeSummarizationTranscript(transcriptPath, classicOriginal); werr != nil && logger != nil { + logger.Warn("eino agentic summarization transcript preflight 写入失败", + zap.String("path", transcriptPath), zap.Error(werr)) + } + } + budget := summaryInputMax + aggressive := summaryOverflowRetries > 0 + if aggressive { + budget = summaryInputMax * 70 / 100 + if budget < 4096 { + budget = 4096 + } + } + input, dropped, berr := buildBudgetedSummarizationModelInput( + ctx, + agenticInstructionToClassic(sysInstruction, schema.System), + agenticInstructionToClassic(userInstruction, schema.User), + classicOriginal, + classicTokenCounter, + budget, + summarizationInputBudgetOpts{ + toolMaxBytes: toolMaxBytes, + spillRef: transcriptPath, + aggressive: aggressive, + }, + ) + if logger != nil && (berr != nil || dropped > 0 || aggressive) { + fields := []zap.Field{ + zap.Int("max_input_tokens", budget), + zap.Int("trigger_context_tokens", trigger), + zap.Int("output_reserve_tokens", outputReserve), + zap.Int("dropped_rounds", dropped), + zap.Bool("aggressive", aggressive), + } + if berr != nil { + fields = append(fields, zap.Error(berr)) + logger.Warn("eino agentic summarization input budget failed", fields...) + } else { + logger.Info("eino agentic summarization input bounded", fields...) + } + } + return EinoMessagesToAgentic(input), berr + }, + Trigger: &summarization.TriggerCondition{ + ContextTokens: trigger, + }, + TokenCounter: agenticTokenCounter, + UserInstruction: einoSummarizeUserInstruction, + EmitInternalEvents: emitInternalEvents, + TranscriptFilePath: transcriptPath, + Retry: &summarization.TypedRetryConfig[*schema.AgenticMessage]{ + MaxRetries: &retryMax, + ShouldRetry: func(_ context.Context, _ *schema.AgenticMessage, err error) bool { + if isEinoContextOverflowError(err) && summaryOverflowRetries < 1 { + summaryOverflowRetries++ + if logger != nil { + logger.Warn("eino agentic summarization context overflow, retrying with aggressive compaction", + zap.Error(err), + ) + } + return true + } + retry := isEinoTransientRunError(err) + if retry && logger != nil { + logger.Warn("eino agentic summarization generate transient error, will retry if attempts remain", + zap.Error(err), + zap.Int("max_retries", retryMax), + ) + } + return retry + }, + }, + Finalize: func(ctx context.Context, originalMessages []*schema.AgenticMessage, summary *schema.AgenticMessage) ([]*schema.AgenticMessage, error) { + classicOriginal := AgenticMessagesToEino(originalMessages) + classicSummary := agenticSummaryToClassicMessage(summary) + if classicSummary == nil { + return nil, fmt.Errorf("agentic summarization returned empty summary") + } + compactionMessages := stripOriginalUserIntentLedgerFromMessages(classicOriginal) + defaultFinalized, derr := summarization.DefaultFinalize(ctx, compactionMessages, classicSummary) + if derr != nil { + return nil, derr + } + if len(defaultFinalized) == 0 { + return nil, fmt.Errorf("agentic summarization default finalize returned no messages") + } + summaryMsg := appendTranscriptPathToSummarizationMessage(defaultFinalized[len(defaultFinalized)-1], transcriptPath) + summaryMsg = stripAnalysisFromSummarizationMessage(summaryMsg) + userLedger := buildOriginalUserIntentLedgerMessage(classicOriginal, userLedgerMaxRunes, userLedgerEntryMaxRunes) + out, ferr := summarizeFinalizeWithRecentAssistantToolTrail(ctx, compactionMessages, summaryMsg, classicTokenCounter, recentTrailMax) + if ferr != nil { + return nil, ferr + } + out = mergeMessageIntoLeadingSystem(out, userLedger) + if appCfg != nil { + out = refreshFactIndexInMessages(out, db, projectID, appCfg.Project, logger) + } + return EinoMessagesToAgentic(out), nil + }, + Callback: func(ctx context.Context, before, after adk.TypedChatModelAgentState[*schema.AgenticMessage]) error { + classicBefore := AgenticMessagesToEino(before.Messages) + classicAfter := AgenticMessagesToEino(after.Messages) + if transcriptPath != "" && len(classicBefore) > 0 { + if werr := writeSummarizationTranscript(transcriptPath, classicBefore); werr != nil && logger != nil { + logger.Warn("eino agentic summarization transcript 写入失败", + zap.String("path", transcriptPath), + zap.Error(werr), + ) + } + } + if logger != nil { + beforeTokens, _ := classicTokenCounter(ctx, &summarization.TokenCounterInput{Messages: classicBefore}) + afterTokens, _ := classicTokenCounter(ctx, &summarization.TokenCounterInput{Messages: classicAfter}) + logger.Info("eino agentic summarization 已压缩上下文", + zap.Int("messages_before", len(before.Messages)), + zap.Int("messages_after", len(after.Messages)), + zap.Int("tokens_before_estimated", beforeTokens), + zap.Int("tokens_after_estimated", afterTokens), + zap.Int("max_total_tokens", maxTotal), + zap.Int("trigger_context_tokens", trigger), + zap.String("transcript_file", transcriptPath), + ) + } + return nil + }, + }) + if err != nil { + return nil, fmt.Errorf("summarization.NewTyped[AgenticMessage]: %w", err) + } + return mw, nil +} + +func agenticInstructionToClassic(msg *schema.AgenticMessage, fallbackRole schema.RoleType) *schema.Message { + msgs := AgenticMessageToEino(msg) + if len(msgs) > 0 && msgs[0] != nil { + return msgs[0] + } + return &schema.Message{Role: fallbackRole} +} + +func agenticSummaryToClassicMessage(msg *schema.AgenticMessage) *schema.Message { + msgs := AgenticMessageToEino(msg) + for _, m := range msgs { + if m == nil { + continue + } + if m.Role == schema.Assistant || strings.TrimSpace(m.Content) != "" || m.ReasoningContent != "" { + if m.Role != schema.Assistant { + cp := *m + cp.Role = schema.Assistant + return &cp + } + return m + } + } + return nil +} diff --git a/internal/multiagent/eino_agentic_summarize_test.go b/internal/multiagent/eino_agentic_summarize_test.go new file mode 100644 index 00000000..4472582c --- /dev/null +++ b/internal/multiagent/eino_agentic_summarize_test.go @@ -0,0 +1,210 @@ +package multiagent + +import ( + "context" + "path/filepath" + "strings" + "testing" + + "cyberstrike-ai/internal/config" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestNewEinoAgenticSummarizationMiddlewareCompactsWithNativeTypedMiddleware(t *testing.T) { + t.Parallel() + ctx := context.Background() + emit := false + summaryModel := &capturingAgenticChatModel{ + output: &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.AssistantGenText{Text: `检查历史 + +## 1. 授权范围与约束 +- 仅测试 example.com + +## 7. 当前进度、策略决策与下一步 +- 继续验证 SQL 注入路径 +`})}, + }, + } + appCfg := &config.Config{} + appCfg.OpenAI.Model = "gpt-4o" + appCfg.OpenAI.MaxTotalTokens = 5000 + appCfg.Database.Path = filepath.Join(t.TempDir(), "cyberstrike.db") + mwCfg := &config.MultiAgentEinoMiddlewareConfig{ + SummarizationEmitInternalEvents: &emit, + SummarizationOutputReserveTokens: 1024, + } + + mw, err := newEinoAgenticSummarizationMiddleware(ctx, summaryModel, appCfg, mwCfg, "conv-agentic", nil, "", nil) + if err != nil { + t.Fatalf("newEinoAgenticSummarizationMiddleware: %v", err) + } + state := &adk.TypedChatModelAgentState[*schema.AgenticMessage]{ + Messages: []*schema.AgenticMessage{ + schema.SystemAgenticMessage("system root"), + schema.UserAgenticMessage("授权范围 example.com\n" + strings.Repeat("历史扫描输出 ", 12000)), + agenticAssistantTextMessage("已记录范围"), + schema.UserAgenticMessage("继续验证 SQL 注入路径"), + }, + } + + _, after, err := mw.BeforeModelRewriteState(ctx, state, nil) + if err != nil { + t.Fatalf("BeforeModelRewriteState: %v", err) + } + inputs := summaryModel.snapshotInputs() + if len(inputs) != 1 || len(inputs[0]) == 0 { + t.Fatalf("summary model inputs = %#v, want one typed AgenticMessage call", inputs) + } + if after == nil { + t.Fatal("after state is nil") + } + classicAfter := AgenticMessagesToEino(after.Messages) + joined := joinClassicMessageContent(classicAfter) + if strings.Contains(joined, "") { + t.Fatalf("analysis block leaked into compacted context: %s", joined) + } + for _, want := range []string{"继续验证 SQL 注入路径", "原始用户输入与约束账本", "完整的对话记录位于"} { + if !strings.Contains(joined, want) { + t.Fatalf("compacted context missing %q:\n%s", want, joined) + } + } +} + +func TestEinoAgenticChatModelAgentCompactsContextBeforeBusinessModel(t *testing.T) { + t.Parallel() + ctx := context.Background() + emit := false + summaryModel := &capturingAgenticChatModel{ + output: agenticAssistantTextMessage(`internal scratchpad + +## 1. 授权范围与约束 +- 仅测试 example.com + +## 7. 当前进度、策略决策与下一步 +- 继续验证 SQL 注入路径 +`), + } + businessModel := &capturingAgenticChatModel{ + output: agenticAssistantTextMessage("business answer after compaction"), + } + appCfg := &config.Config{} + appCfg.OpenAI.Model = "gpt-4o" + appCfg.OpenAI.MaxTotalTokens = 5000 + appCfg.Database.Path = filepath.Join(t.TempDir(), "cyberstrike.db") + mwCfg := &config.MultiAgentEinoMiddlewareConfig{ + SummarizationEmitInternalEvents: &emit, + SummarizationOutputReserveTokens: 1024, + } + sumMw, err := newEinoAgenticSummarizationMiddleware(ctx, summaryModel, appCfg, mwCfg, "conv-agentic-e2e", nil, "", nil) + if err != nil { + t.Fatalf("newEinoAgenticSummarizationMiddleware: %v", err) + } + trace := newModelFacingTraceHolder() + agent, err := newEinoAgenticChatModelAgentAdapter(ctx, einoAgenticChatModelAgentConfig{ + Name: "agentic", + Description: "agentic compaction e2e test", + Instruction: "system root", + Model: businessModel, + Handlers: appendEinoAgenticChatModelTailMiddlewares(nil, einoChatModelTailConfig{ + phase: "agentic", + agenticSummarization: sumMw, + trace: trace, + }), + }) + if err != nil { + t.Fatalf("newEinoAgenticChatModelAgentAdapter: %v", err) + } + + rawHistory := "授权范围 example.com\n" + strings.Repeat("原始扫描输出SHOULD_NOT_REACH_BUSINESS_MODEL ", 12000) + iter := agent.Run(ctx, &adk.AgentInput{ + Messages: []*schema.Message{ + schema.UserMessage(rawHistory), + schema.AssistantMessage("已记录范围", nil), + schema.UserMessage("继续验证 SQL 注入路径"), + }, + }) + var last *adk.AgentEvent + for { + ev, ok := iter.Next() + if !ok { + break + } + if ev.Err != nil { + t.Fatalf("agent event error: %v", ev.Err) + } + last = ev + } + if last == nil || last.Output == nil || last.Output.MessageOutput == nil { + t.Fatalf("last event = %#v, want message output", last) + } + if got := last.Output.MessageOutput.Message.Content; got != "business answer after compaction" { + t.Fatalf("business output = %q", got) + } + + if inputs := summaryModel.snapshotInputs(); len(inputs) != 1 { + t.Fatalf("summary model calls = %d, want 1", len(inputs)) + } + businessInputs := businessModel.snapshotInputs() + if len(businessInputs) != 1 { + t.Fatalf("business model calls = %d, want 1", len(businessInputs)) + } + finalClassicInput := AgenticMessagesToEino(businessInputs[0]) + joined := joinClassicMessageContent(finalClassicInput) + for _, want := range []string{"继续验证 SQL 注入路径", "原始用户输入与约束账本", "完整的对话记录位于"} { + if !strings.Contains(joined, want) { + t.Fatalf("business model input missing %q:\n%s", want, joined) + } + } + if strings.Contains(joined, "") { + t.Fatalf("analysis leaked to business model input:\n%s", joined) + } + if strings.Count(joined, "原始扫描输出SHOULD_NOT_REACH_BUSINESS_MODEL") > 3 { + t.Fatalf("raw oversized history leaked to business model input, count=%d", strings.Count(joined, "原始扫描输出SHOULD_NOT_REACH_BUSINESS_MODEL")) + } + traceJoined := joinClassicMessageContent(trace.Snapshot()) + if !strings.Contains(traceJoined, "继续验证 SQL 注入路径") || strings.Count(traceJoined, "原始扫描输出SHOULD_NOT_REACH_BUSINESS_MODEL") > 3 { + t.Fatalf("model-facing trace not compacted:\n%s", traceJoined) + } +} + +func TestAppendEinoAgenticChatModelTailMiddlewaresIncludesTypedSummarization(t *testing.T) { + t.Parallel() + mw := newAgenticSystemMessageNormalizerMiddleware(nil, "summary") + handlers := appendEinoAgenticChatModelTailMiddlewares(nil, einoChatModelTailConfig{ + agenticSummarization: mw, + skipTrace: true, + }) + found := false + for _, h := range handlers { + if h == mw { + found = true + break + } + } + if !found { + t.Fatal("agentic summarization middleware was not appended") + } +} + +func agenticAssistantTextMessage(text string) *schema.AgenticMessage { + return &schema.AgenticMessage{ + Role: schema.AgenticRoleTypeAssistant, + ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.AssistantGenText{Text: text})}, + } +} + +func joinClassicMessageContent(msgs []*schema.Message) string { + var b strings.Builder + for _, msg := range msgs { + if msg == nil { + continue + } + b.WriteString(msg.Content) + b.WriteByte('\n') + } + return b.String() +} diff --git a/internal/multiagent/eino_assistant_output_accumulator.go b/internal/multiagent/eino_assistant_output_accumulator.go new file mode 100644 index 00000000..1ca53128 --- /dev/null +++ b/internal/multiagent/eino_assistant_output_accumulator.go @@ -0,0 +1,42 @@ +package multiagent + +import "strings" + +type einoAssistantOutputAccumulator struct { + orchMode string + lastAssistant string + lastPlanExecuteExecutor string +} + +func newEinoAssistantOutputAccumulator(orchMode string) *einoAssistantOutputAccumulator { + return &einoAssistantOutputAccumulator{orchMode: orchMode} +} + +func (a *einoAssistantOutputAccumulator) RecordMainAssistant(agentName, content string) bool { + if a == nil { + return false + } + content = strings.TrimSpace(content) + if content == "" { + return false + } + a.lastAssistant = content + if a.orchMode == "plan_execute" && strings.EqualFold(strings.TrimSpace(agentName), "executor") { + a.lastPlanExecuteExecutor = UnwrapPlanExecuteUserText(content) + } + return true +} + +func (a *einoAssistantOutputAccumulator) LastAssistant() string { + if a == nil { + return "" + } + return a.lastAssistant +} + +func (a *einoAssistantOutputAccumulator) LastPlanExecuteExecutor() string { + if a == nil { + return "" + } + return a.lastPlanExecuteExecutor +} diff --git a/internal/multiagent/eino_assistant_output_accumulator_test.go b/internal/multiagent/eino_assistant_output_accumulator_test.go new file mode 100644 index 00000000..27102b06 --- /dev/null +++ b/internal/multiagent/eino_assistant_output_accumulator_test.go @@ -0,0 +1,52 @@ +package multiagent + +import "testing" + +func TestEinoAssistantOutputAccumulatorRecordsMainAssistant(t *testing.T) { + acc := newEinoAssistantOutputAccumulator("deep") + if acc.RecordMainAssistant("lead", " hello ") != true { + t.Fatal("expected record") + } + if got := acc.LastAssistant(); got != "hello" { + t.Fatalf("last assistant = %q, want hello", got) + } + if got := acc.LastPlanExecuteExecutor(); got != "" { + t.Fatalf("plan execute executor = %q, want empty", got) + } + if acc.RecordMainAssistant("lead", " ") { + t.Fatal("blank content should not record") + } + if got := acc.LastAssistant(); got != "hello" { + t.Fatalf("blank content changed last assistant to %q", got) + } +} + +func TestEinoAssistantOutputAccumulatorPlanExecuteExecutor(t *testing.T) { + acc := newEinoAssistantOutputAccumulator("plan_execute") + raw := `{"response":"给用户看的正文","scratchpad":"internal"}` + acc.RecordMainAssistant("executor", raw) + + if got := acc.LastAssistant(); got != raw { + t.Fatalf("last assistant = %q, want raw", got) + } + if got := acc.LastPlanExecuteExecutor(); got != "给用户看的正文" { + t.Fatalf("executor output = %q", got) + } + acc.RecordMainAssistant("planner", "planner note") + if got := acc.LastAssistant(); got != "planner note" { + t.Fatalf("last assistant after planner = %q", got) + } + if got := acc.LastPlanExecuteExecutor(); got != "给用户看的正文" { + t.Fatalf("planner should not overwrite executor output, got %q", got) + } +} + +func TestEinoAssistantOutputAccumulatorNilSafe(t *testing.T) { + var acc *einoAssistantOutputAccumulator + if acc.RecordMainAssistant("agent", "hello") { + t.Fatal("nil accumulator should not record") + } + if acc.LastAssistant() != "" || acc.LastPlanExecuteExecutor() != "" { + t.Fatal("nil accumulator should return empty values") + } +} diff --git a/internal/multiagent/eino_assistant_stream_event_handler.go b/internal/multiagent/eino_assistant_stream_event_handler.go new file mode 100644 index 00000000..bc109efe --- /dev/null +++ b/internal/multiagent/eino_assistant_stream_event_handler.go @@ -0,0 +1,167 @@ +package multiagent + +import ( + "context" + "errors" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" +) + +type einoAssistantStreamEventHandlerConfig struct { + Context context.Context + ConversationID string + OrchMode string + Progress func(eventType, message string, data interface{}) + Logger *zap.Logger + SnapshotMCPIDs func() []string + StreamsMainAssistant func(agent string) bool + EinoRoleTag func(agent string) string + RunProgress *einoRunProgressTracker + StdoutSuppressor *einoExecuteStdoutSuppressor + AssistantOutput *einoAssistantOutputAccumulator + RunMessages *einoRunMessageAccumulator + Usage *einoRunUsageAccumulator + ToolCallCompletion *einoStreamToolCallCompletionHandler + NextMainStreamID func() string + NextReasoningStreamID func() string + NextSubAgentReplyStreamID func() string +} + +type einoAssistantStreamEventHandler struct { + ctx context.Context + conversationID string + orchMode string + progress func(eventType, message string, data interface{}) + logger *zap.Logger + snapshotMCPIDs func() []string + streamsMainAssistant func(agent string) bool + einoRoleTag func(agent string) string + runProgress *einoRunProgressTracker + stdoutSuppressor *einoExecuteStdoutSuppressor + assistantOutput *einoAssistantOutputAccumulator + runMessages *einoRunMessageAccumulator + usage *einoRunUsageAccumulator + toolCallCompletion *einoStreamToolCallCompletionHandler + nextMainStreamID func() string + nextReasoningStreamID func() string + nextSubAgentReplyStreamID func() string +} + +func newEinoAssistantStreamEventHandler(cfg einoAssistantStreamEventHandlerConfig) *einoAssistantStreamEventHandler { + 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(string) bool { return true } + } + if cfg.EinoRoleTag == nil { + cfg.EinoRoleTag = func(string) string { return "" } + } + if cfg.NextMainStreamID == nil { + cfg.NextMainStreamID = func() string { return "eino-main" } + } + if cfg.NextReasoningStreamID == nil { + cfg.NextReasoningStreamID = func() string { return "eino-reasoning" } + } + if cfg.NextSubAgentReplyStreamID == nil { + cfg.NextSubAgentReplyStreamID = func() string { return "eino-sub-reply" } + } + return &einoAssistantStreamEventHandler{ + ctx: cfg.Context, + conversationID: cfg.ConversationID, + orchMode: cfg.OrchMode, + progress: cfg.Progress, + logger: cfg.Logger, + snapshotMCPIDs: cfg.SnapshotMCPIDs, + streamsMainAssistant: cfg.StreamsMainAssistant, + einoRoleTag: cfg.EinoRoleTag, + runProgress: cfg.RunProgress, + stdoutSuppressor: cfg.StdoutSuppressor, + assistantOutput: cfg.AssistantOutput, + runMessages: cfg.RunMessages, + usage: cfg.Usage, + toolCallCompletion: cfg.ToolCallCompletion, + nextMainStreamID: cfg.NextMainStreamID, + nextReasoningStreamID: cfg.NextReasoningStreamID, + nextSubAgentReplyStreamID: cfg.NextSubAgentReplyStreamID, + } +} + +func (h *einoAssistantStreamEventHandler) Handle(mv *adk.MessageVariant, agentName string) (handled bool, recvErr error) { + if h == nil || mv == nil || !mv.IsStreaming || mv.MessageStream == nil || mv.Role == schema.Tool { + return false, nil + } + mainStreamID := h.nextMainStreamID() + mainEmitter := newEinoMainResponseStreamEmitter( + h.conversationID, h.orchMode, agentName, mainStreamID, h.mainIteration(agentName), h.progress, h.snapshotMCPIDs, + ) + reasoningEmitter := newEinoReasoningStreamEmitter( + h.conversationID, + h.orchMode, + agentName, + h.einoRoleTag(agentName), + h.progress, + h.nextReasoningStreamID, + ) + var toolStreamFragments []schema.ToolCall + var streamUsage *schema.TokenUsage + subReplyEmitter := newEinoSubAgentReplyEmitter( + h.conversationID, + agentName, + h.progress, + h.nextSubAgentReplyStreamID, + ) + mainAssistantStream := newEinoMainAssistantStreamHandler(einoMainAssistantStreamHandlerConfig{ + AgentName: agentName, + Emitter: mainEmitter, + StdoutSuppressor: h.stdoutSuppressor, + AssistantOutput: h.assistantOutput, + RunMessages: h.runMessages, + }) + recvErr = recvEinoSchemaMessageStreamWithContext(h.ctx, mv.MessageStream, 8, func(chunk *schema.Message) { + reasoningEmitter.EmitDelta(chunk.ReasoningContent) + if chunk.Content != "" { + if h.streamsMainAssistant(agentName) { + mainAssistantStream.EmitDelta(chunk.Content) + } else if !h.streamsMainAssistant(agentName) { + subReplyEmitter.EmitDelta(chunk.Content) + } + } + if len(chunk.ToolCalls) > 0 { + toolStreamFragments = append(toolStreamFragments, chunk.ToolCalls...) + } + if chunk.ResponseMeta != nil && chunk.ResponseMeta.Usage != nil { + streamUsage = maxEinoTokenUsage(streamUsage, chunk.ResponseMeta.Usage) + } + }) + if recvErr != nil && !errors.Is(recvErr, context.Canceled) && h.logger != nil { + h.logger.Warn("eino stream recv error, flushing incomplete stream", + zap.Error(recvErr), + zap.String("agent", agentName), + zap.Int("toolFragments", len(toolStreamFragments))) + } + reasoningEmitter.Finish() + if h.streamsMainAssistant(agentName) { + mainAssistantStream.Finish() + } + subReplyEmitter.Finish() + if h.toolCallCompletion != nil { + h.toolCallCompletion.Complete(toolStreamFragments, agentName) + } + if h.usage != nil { + h.usage.AddUsage(streamUsage) + } + return true, recvErr +} + +func (h *einoAssistantStreamEventHandler) mainIteration(agentName string) int { + if h == nil || h.runProgress == nil { + return 0 + } + return h.runProgress.MainIteration(agentName) +} diff --git a/internal/multiagent/eino_assistant_stream_event_handler_test.go b/internal/multiagent/eino_assistant_stream_event_handler_test.go new file mode 100644 index 00000000..05fecf1c --- /dev/null +++ b/internal/multiagent/eino_assistant_stream_event_handler_test.go @@ -0,0 +1,148 @@ +package multiagent + +import ( + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestEinoAssistantStreamEventHandlerHandlesMainAssistantStream(t *testing.T) { + var events []string + runMessages := newEinoRunMessageAccumulator(nil) + assistantOutput := newEinoAssistantOutputAccumulator("deep") + usage := newEinoRunUsageAccumulator() + handler := newEinoAssistantStreamEventHandler(einoAssistantStreamEventHandlerConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + RunMessages: runMessages, + Usage: usage, + AssistantOutput: assistantOutput, + StreamsMainAssistant: func(agent string) bool { return agent == "lead" }, + EinoRoleTag: func(string) string { return "orchestrator" }, + NextMainStreamID: func() string { return "main-stream-1" }, + Progress: func(eventType, _ string, _ interface{}) { + events = append(events, eventType) + }, + }) + mv := &adk.MessageVariant{ + IsStreaming: true, + Role: schema.Assistant, + MessageStream: schema.StreamReaderFromArray([]*schema.Message{ + {Role: schema.Assistant, Content: "he", ResponseMeta: &schema.ResponseMeta{Usage: &schema.TokenUsage{PromptTokens: 10, CompletionTokens: 1, TotalTokens: 11}}}, + {Role: schema.Assistant, Content: "hello", ResponseMeta: &schema.ResponseMeta{Usage: &schema.TokenUsage{PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15}}}, + }), + } + + handled, err := handler.Handle(mv, "lead") + if !handled || err != nil { + t.Fatalf("handled=%v err=%v", handled, err) + } + if assistantOutput.LastAssistant() != "hello" { + t.Fatalf("last assistant = %q", assistantOutput.LastAssistant()) + } + if msgs := runMessages.Messages(); len(msgs) != 1 || msgs[0].Content != "hello" { + t.Fatalf("run messages = %#v", msgs) + } + if got := usage.Summary(); got.ModelCalls != 1 || got.PromptTokens != 10 || got.CompletionTokens != 5 || got.TotalTokens != 15 { + t.Fatalf("usage = %#v, want one stream model call", got) + } + if !containsString(events, "response_start") || !containsString(events, "response_delta") { + t.Fatalf("events = %#v, want response stream events", events) + } +} + +func TestEinoAssistantStreamEventHandlerHandlesSubAgentStream(t *testing.T) { + var events []string + runMessages := newEinoRunMessageAccumulator(nil) + assistantOutput := newEinoAssistantOutputAccumulator("deep") + handler := newEinoAssistantStreamEventHandler(einoAssistantStreamEventHandlerConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + RunMessages: runMessages, + AssistantOutput: assistantOutput, + StreamsMainAssistant: func(agent string) bool { return agent == "lead" }, + EinoRoleTag: func(string) string { return "sub" }, + NextSubAgentReplyStreamID: func() string { + return "sub-stream-1" + }, + Progress: func(eventType, _ string, _ interface{}) { + events = append(events, eventType) + }, + }) + mv := &adk.MessageVariant{ + IsStreaming: true, + Role: schema.Assistant, + MessageStream: schema.StreamReaderFromArray([]*schema.Message{{Role: schema.Assistant, Content: "sub reply"}}), + } + + handled, err := handler.Handle(mv, "worker") + if !handled || err != nil { + t.Fatalf("handled=%v err=%v", handled, err) + } + if len(runMessages.Messages()) != 0 { + t.Fatalf("sub stream should not append main run text, got %#v", runMessages.Messages()) + } + if assistantOutput.LastAssistant() != "" { + t.Fatalf("sub stream should not record main assistant, got %q", assistantOutput.LastAssistant()) + } + if !containsString(events, "eino_agent_reply_stream_start") || + !containsString(events, "eino_agent_reply_stream_delta") || + !containsString(events, "eino_agent_reply_stream_end") { + t.Fatalf("events = %#v, want sub reply stream events", events) + } +} + +func TestEinoAssistantStreamEventHandlerCompletesToolFragments(t *testing.T) { + idx := 0 + var events []string + runMessages := newEinoRunMessageAccumulator(nil) + runProgress := newEinoRunProgressTracker( + "deep", "lead", "conv-1", + func(eventType, _ string, _ interface{}) { events = append(events, eventType) }, + func(agent string) bool { return agent == "lead" }, + nil, + ) + completion := newEinoStreamToolCallCompletionHandler(einoStreamToolCallCompletionHandlerConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + Progress: func(eventType, _ string, _ interface{}) { events = append(events, eventType) }, + RunProgress: runProgress, + RunMessages: runMessages, + }) + handler := newEinoAssistantStreamEventHandler(einoAssistantStreamEventHandlerConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + RunMessages: runMessages, + StreamsMainAssistant: func(string) bool { return true }, + ToolCallCompletion: completion, + }) + mv := &adk.MessageVariant{ + IsStreaming: true, + Role: schema.Assistant, + MessageStream: schema.StreamReaderFromArray([]*schema.Message{ + {Role: schema.Assistant, ToolCalls: []schema.ToolCall{{ID: "call-1", Index: &idx, Type: "function", Function: schema.FunctionCall{Name: "execute", Arguments: `{"command":`}}}}, + {Role: schema.Assistant, ToolCalls: []schema.ToolCall{{Index: &idx, Function: schema.FunctionCall{Arguments: `"pwd"}`}}}}, + }), + } + + handled, err := handler.Handle(mv, "lead") + if !handled || err != nil { + t.Fatalf("handled=%v err=%v", handled, err) + } + msgs := runMessages.Messages() + if len(msgs) != 1 || len(msgs[0].ToolCalls) != 1 || msgs[0].ToolCalls[0].Function.Arguments != `{"command":"pwd"}` { + t.Fatalf("run messages = %#v, want merged tool call", msgs) + } + if !containsString(events, "tool_call") { + t.Fatalf("events = %#v, want tool_call", events) + } +} + +func TestEinoAssistantStreamEventHandlerIgnoresToolStream(t *testing.T) { + handler := newEinoAssistantStreamEventHandler(einoAssistantStreamEventHandlerConfig{}) + handled, err := handler.Handle(&adk.MessageVariant{IsStreaming: true, Role: schema.Tool, MessageStream: schema.StreamReaderFromArray([]*schema.Message{})}, "lead") + if handled || err != nil { + t.Fatalf("handled=%v err=%v, want ignored", handled, err) + } +} diff --git a/internal/multiagent/eino_chat_model_tail_middleware.go b/internal/multiagent/eino_chat_model_tail_middleware.go index d1cd8f2d..66352887 100644 --- a/internal/multiagent/eino_chat_model_tail_middleware.go +++ b/internal/multiagent/eino_chat_model_tail_middleware.go @@ -4,6 +4,7 @@ import ( "cyberstrike-ai/internal/config" "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" "go.uber.org/zap" ) @@ -24,18 +25,19 @@ import ( // 11. telemetry // 12. model-facing trace snapshot type einoChatModelTailConfig struct { - logger *zap.Logger - phase string - summarization adk.ChatModelAgentMiddleware - modelName string - maxTotalTokens int - toolMaxBytes int - conversationID string - trace *modelFacingTraceHolder - middlewareConfig *config.MultiAgentEinoMiddlewareConfig - skipOrphanPruner bool - skipTelemetry bool - skipTrace bool + logger *zap.Logger + phase string + summarization adk.ChatModelAgentMiddleware + agenticSummarization adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] + modelName string + maxTotalTokens int + toolMaxBytes int + conversationID string + trace *modelFacingTraceHolder + middlewareConfig *config.MultiAgentEinoMiddlewareConfig + skipOrphanPruner bool + skipTelemetry bool + skipTrace bool } func appendEinoChatModelTailMiddlewares(handlers []adk.ChatModelAgentMiddleware, cfg einoChatModelTailConfig) []adk.ChatModelAgentMiddleware { @@ -65,7 +67,6 @@ func appendEinoChatModelTailMiddlewares(handlers []adk.ChatModelAgentMiddleware, handlers = append(handlers, capMw) } } - handlers = append(handlers, newModelOutputGuardMiddleware(cfg.middlewareConfig, cfg.logger, cfg.phase)) return handlers } diff --git a/internal/multiagent/eino_checkpoint_resume_handler.go b/internal/multiagent/eino_checkpoint_resume_handler.go new file mode 100644 index 00000000..cb2e1f8b --- /dev/null +++ b/internal/multiagent/eino_checkpoint_resume_handler.go @@ -0,0 +1,71 @@ +package multiagent + +import ( + "context" + + "github.com/cloudwego/eino/adk" + "go.uber.org/zap" +) + +type einoCheckpointResumeHandlerConfig struct { + Context context.Context + ConversationID string + OrchMode string + Progress func(eventType, message string, data interface{}) + Logger *zap.Logger + Store *fileCheckPointStore + CheckPointID string + Resume func(checkPointID string) (*adk.AsyncIterator[*adk.AgentEvent], error) +} + +type einoCheckpointResumeHandler struct { + cfg einoCheckpointResumeHandlerConfig +} + +func newEinoCheckpointResumeHandler(cfg einoCheckpointResumeHandlerConfig) *einoCheckpointResumeHandler { + if cfg.Context == nil { + cfg.Context = context.Background() + } + return &einoCheckpointResumeHandler{cfg: cfg} +} + +func (h *einoCheckpointResumeHandler) TryResume() *adk.AsyncIterator[*adk.AgentEvent] { + if h == nil || h.cfg.Store == nil || h.cfg.CheckPointID == "" || h.cfg.Resume == nil { + return nil + } + if _, existed, err := h.cfg.Store.Get(h.cfg.Context, h.cfg.CheckPointID); err != nil { + if h.cfg.Logger != nil { + h.cfg.Logger.Warn("eino checkpoint preflight get failed", zap.String("checkPointID", h.cfg.CheckPointID), zap.Error(err)) + } + return nil + } else if !existed { + return nil + } + h.emitProgress("检测到断点,正在从中断节点恢复执行...") + if h.cfg.Logger != nil { + h.cfg.Logger.Info("eino runner: resume from checkpoint", zap.String("checkPointID", h.cfg.CheckPointID)) + } + iter, err := h.cfg.Resume(h.cfg.CheckPointID) + if err == nil { + return iter + } + if h.cfg.Logger != nil { + h.cfg.Logger.Warn("eino runner: resume failed, fallback to fresh run", + zap.String("checkPointID", h.cfg.CheckPointID), + zap.Error(err)) + } + h.emitProgress("断点恢复失败,已回退为全新执行。") + return nil +} + +func (h *einoCheckpointResumeHandler) emitProgress(message string) { + if h == nil || h.cfg.Progress == nil { + return + } + h.cfg.Progress("progress", message, map[string]interface{}{ + "conversationId": h.cfg.ConversationID, + "source": "eino", + "orchestration": h.cfg.OrchMode, + "checkPointID": h.cfg.CheckPointID, + }) +} diff --git a/internal/multiagent/eino_checkpoint_resume_handler_test.go b/internal/multiagent/eino_checkpoint_resume_handler_test.go new file mode 100644 index 00000000..33b1372f --- /dev/null +++ b/internal/multiagent/eino_checkpoint_resume_handler_test.go @@ -0,0 +1,138 @@ +package multiagent + +import ( + "context" + "errors" + "testing" + + "github.com/cloudwego/eino/adk" + "go.uber.org/zap" + "go.uber.org/zap/zaptest/observer" +) + +func TestEinoCheckpointResumeHandlerSkipsWithoutCheckpoint(t *testing.T) { + called := false + handler := newEinoCheckpointResumeHandler(einoCheckpointResumeHandlerConfig{ + Resume: func(string) (*adk.AsyncIterator[*adk.AgentEvent], error) { + called = true + return nil, nil + }, + }) + if iter := handler.TryResume(); iter != nil { + t.Fatalf("iter = %#v, want nil", iter) + } + if called { + t.Fatal("resume should not be called without checkpoint state") + } +} + +func TestEinoCheckpointResumeHandlerResumesExistingCheckpoint(t *testing.T) { + 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) + } + var progressMessages []string + var resumedID string + core, logs := observer.New(zap.InfoLevel) + wantIter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + defer gen.Close() + handler := newEinoCheckpointResumeHandler(einoCheckpointResumeHandlerConfig{ + Context: context.Background(), + ConversationID: "conv-1", + OrchMode: "deep", + Store: store, + CheckPointID: "cp-1", + Logger: zap.New(core), + Progress: func(eventType, message string, data interface{}) { + if eventType != "progress" { + return + } + progressMessages = append(progressMessages, message) + m, _ := data.(map[string]interface{}) + if m["conversationId"] != "conv-1" || m["orchestration"] != "deep" || m["checkPointID"] != "cp-1" { + t.Fatalf("progress data = %#v", m) + } + }, + Resume: func(checkPointID string) (*adk.AsyncIterator[*adk.AgentEvent], error) { + resumedID = checkPointID + return wantIter, nil + }, + }) + + got := handler.TryResume() + if got != wantIter { + t.Fatalf("iter = %#v, want resume iterator", got) + } + if resumedID != "cp-1" { + t.Fatalf("resumed id = %q", resumedID) + } + if len(progressMessages) != 1 || progressMessages[0] != "检测到断点,正在从中断节点恢复执行..." { + t.Fatalf("progress messages = %#v", progressMessages) + } + if logs.FilterMessage("eino runner: resume from checkpoint").Len() != 1 { + t.Fatalf("expected resume log, got %d", logs.Len()) + } +} + +func TestEinoCheckpointResumeHandlerFallsBackOnResumeError(t *testing.T) { + 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) + } + var progressMessages []string + core, logs := observer.New(zap.WarnLevel) + handler := newEinoCheckpointResumeHandler(einoCheckpointResumeHandlerConfig{ + Context: context.Background(), + Store: store, + CheckPointID: "cp-1", + Logger: zap.New(core), + Progress: func(eventType, message string, _ interface{}) { + if eventType == "progress" { + progressMessages = append(progressMessages, message) + } + }, + Resume: func(string) (*adk.AsyncIterator[*adk.AgentEvent], error) { + return nil, errors.New("resume failed") + }, + }) + + if iter := handler.TryResume(); iter != nil { + t.Fatalf("iter = %#v, want nil fallback", iter) + } + if len(progressMessages) != 2 || progressMessages[1] != "断点恢复失败,已回退为全新执行。" { + t.Fatalf("progress messages = %#v", progressMessages) + } + if logs.FilterMessage("eino runner: resume failed, fallback to fresh run").Len() != 1 { + t.Fatalf("expected fallback log, got %d", logs.Len()) + } +} + +func TestEinoCheckpointResumeHandlerLogsPreflightError(t *testing.T) { + store, err := newFileCheckPointStore(t.TempDir()) + if err != nil { + t.Fatal(err) + } + core, logs := observer.New(zap.WarnLevel) + handler := newEinoCheckpointResumeHandler(einoCheckpointResumeHandlerConfig{ + Context: context.Background(), + Store: store, + CheckPointID: "bad/id", + Logger: zap.New(core), + Resume: func(string) (*adk.AsyncIterator[*adk.AgentEvent], error) { + t.Fatal("resume should not be called after preflight error") + return nil, nil + }, + }) + if iter := handler.TryResume(); iter != nil { + t.Fatalf("iter = %#v, want nil", iter) + } + if logs.FilterMessage("eino checkpoint preflight get failed").Len() != 1 { + t.Fatalf("expected preflight warning, got %d", logs.Len()) + } +} diff --git a/internal/multiagent/eino_checkpoint_runtime.go b/internal/multiagent/eino_checkpoint_runtime.go new file mode 100644 index 00000000..130f408d --- /dev/null +++ b/internal/multiagent/eino_checkpoint_runtime.go @@ -0,0 +1,38 @@ +package multiagent + +import ( + "path/filepath" + "strings" + + "go.uber.org/zap" +) + +type einoCheckpointRuntime struct { + Store *fileCheckPointStore + CheckPointID string +} + +func newEinoCheckpointRuntime(checkpointDir, conversationID, orchMode string, logger *zap.Logger) *einoCheckpointRuntime { + checkpointDir = strings.TrimSpace(checkpointDir) + if checkpointDir == "" { + return nil + } + cpDir := filepath.Join(checkpointDir, sanitizeEinoPathSegment(conversationID)) + store, err := newFileCheckPointStore(cpDir) + if err != nil { + if logger != nil { + logger.Warn("eino checkpoint store disabled", zap.String("dir", cpDir), zap.Error(err)) + } + return nil + } + checkPointID := buildEinoCheckpointID(orchMode) + if logger != nil { + logger.Info("eino runner: checkpoint store enabled", + zap.String("dir", cpDir), + zap.String("checkPointID", checkPointID)) + } + return &einoCheckpointRuntime{ + Store: store, + CheckPointID: checkPointID, + } +} diff --git a/internal/multiagent/eino_checkpoint_runtime_test.go b/internal/multiagent/eino_checkpoint_runtime_test.go new file mode 100644 index 00000000..23196ee2 --- /dev/null +++ b/internal/multiagent/eino_checkpoint_runtime_test.go @@ -0,0 +1,48 @@ +package multiagent + +import ( + "os" + "strings" + "testing" + + "go.uber.org/zap" + "go.uber.org/zap/zaptest/observer" +) + +func TestNewEinoCheckpointRuntimeDisabledWithoutDir(t *testing.T) { + if got := newEinoCheckpointRuntime(" ", "conv-1", "deep", nil); got != nil { + t.Fatalf("runtime = %#v, want nil", got) + } +} + +func TestNewEinoCheckpointRuntimeCreatesStore(t *testing.T) { + core, logs := observer.New(zap.InfoLevel) + runtime := newEinoCheckpointRuntime(t.TempDir(), "conv/1", "deep", zap.New(core)) + if runtime == nil || runtime.Store == nil { + t.Fatal("expected checkpoint runtime with store") + } + if runtime.CheckPointID != buildEinoCheckpointID("deep") { + t.Fatalf("checkpoint id = %q", runtime.CheckPointID) + } + if !strings.Contains(runtime.Store.dir, sanitizeEinoPathSegment("conv/1")) { + t.Fatalf("store dir = %q, want sanitized conversation segment", runtime.Store.dir) + } + if logs.FilterMessage("eino runner: checkpoint store enabled").Len() != 1 { + t.Fatalf("expected enabled log, got %d", logs.Len()) + } +} + +func TestNewEinoCheckpointRuntimeLogsCreateFailure(t *testing.T) { + filePath := t.TempDir() + "/not-a-dir" + if err := os.WriteFile(filePath, []byte("x"), 0o600); err != nil { + t.Fatal(err) + } + core, logs := observer.New(zap.WarnLevel) + runtime := newEinoCheckpointRuntime(filePath, "conv-1", "deep", zap.New(core)) + if runtime != nil { + t.Fatalf("runtime = %#v, want nil", runtime) + } + if logs.FilterMessage("eino checkpoint store disabled").Len() != 1 { + t.Fatalf("expected disabled log, got %d", logs.Len()) + } +} diff --git a/internal/multiagent/eino_context_overflow_retry.go b/internal/multiagent/eino_context_overflow_retry.go new file mode 100644 index 00000000..b242b440 --- /dev/null +++ b/internal/multiagent/eino_context_overflow_retry.go @@ -0,0 +1,90 @@ +package multiagent + +import ( + "context" + + "github.com/cloudwego/eino/adk" + "go.uber.org/zap" +) + +type einoContextOverflowRetryConfig struct { + Context context.Context + ConversationID string + OrchMode string + Args *einoADKRunLoopArgs + BaseMsgs []adk.Message + Progress func(eventType, message string, data interface{}) + Logger *zap.Logger +} + +type einoContextOverflowRetryResult struct { + Handled bool + RestartMsgs []adk.Message + ContextSrc einoRunRestartContextSource +} + +type einoContextOverflowRetryHandler struct { + cfg einoContextOverflowRetryConfig + retried bool +} + +func newEinoContextOverflowRetryHandler(cfg einoContextOverflowRetryConfig) *einoContextOverflowRetryHandler { + if cfg.Context == nil { + cfg.Context = context.Background() + } + if cfg.Args == nil { + cfg.Args = &einoADKRunLoopArgs{} + } + return &einoContextOverflowRetryHandler{cfg: cfg} +} + +func (h *einoContextOverflowRetryHandler) Prepare( + runErr error, + accumulated []adk.Message, + baseCount int, +) einoContextOverflowRetryResult { + if h == nil || !isEinoContextOverflowError(runErr) || h.retried { + return einoContextOverflowRetryResult{} + } + h.retried = true + restartMsgs, ctxSource := einoMessagesForRunRestart(h.cfg.Args, h.cfg.BaseMsgs, accumulated, baseCount) + restartMsgs = aggressiveCompactMessagesForOverflow( + h.cfg.Context, + restartMsgs, + h.cfg.Args.MaxTotalTokens, + h.cfg.Args.ModelName, + h.cfg.Args.ToolMaxBytes, + h.cfg.OrchMode, + h.cfg.Logger, + ) + if h.cfg.Logger != nil { + h.cfg.Logger.Warn("eino context overflow, retrying with aggressive compaction", + zap.Error(runErr), + zap.String("orchestration", h.cfg.OrchMode), + zap.String("contextSource", string(ctxSource)), + ) + } + emitEinoContextOverflowRetryProgress(h.cfg.Progress, h.cfg.ConversationID, h.cfg.OrchMode, ctxSource) + return einoContextOverflowRetryResult{ + Handled: true, + RestartMsgs: restartMsgs, + ContextSrc: ctxSource, + } +} + +func emitEinoContextOverflowRetryProgress( + progress func(eventType, message string, data interface{}), + conversationID, orchMode string, + ctxSource einoRunRestartContextSource, +) bool { + if progress == nil { + return false + } + progress("eino_context_overflow_retry", "上下文超限,正在激进压缩后重试…", map[string]interface{}{ + "conversationId": conversationID, + "source": "eino", + "orchestration": orchMode, + "contextSource": string(ctxSource), + }) + return true +} diff --git a/internal/multiagent/eino_context_overflow_retry_test.go b/internal/multiagent/eino_context_overflow_retry_test.go new file mode 100644 index 00000000..ca2ef603 --- /dev/null +++ b/internal/multiagent/eino_context_overflow_retry_test.go @@ -0,0 +1,90 @@ +package multiagent + +import ( + "context" + "errors" + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" + "go.uber.org/zap/zaptest/observer" +) + +func TestEinoContextOverflowRetryHandlerPreparesOnce(t *testing.T) { + baseMsgs := []adk.Message{ + schema.UserMessage("base"), + } + accumulated := []adk.Message{ + schema.UserMessage("base"), + schema.AssistantMessage("partial", nil), + } + var gotType, gotMessage string + var gotData map[string]interface{} + core, logs := observer.New(zap.WarnLevel) + handler := newEinoContextOverflowRetryHandler(einoContextOverflowRetryConfig{ + Context: context.Background(), + ConversationID: "conv-1", + OrchMode: "deep_agent", + Args: &einoADKRunLoopArgs{}, + BaseMsgs: baseMsgs, + Progress: 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) + } + }, + Logger: zap.New(core), + }) + + result := handler.Prepare(errors.New("context length exceeded"), accumulated, len(baseMsgs)) + if !result.Handled { + t.Fatal("handled = false, want true") + } + if result.ContextSrc != einoRestartContextAccumulated { + t.Fatalf("context source = %q, want %q", result.ContextSrc, einoRestartContextAccumulated) + } + if len(result.RestartMsgs) != len(accumulated) { + t.Fatalf("restart message count = %d, want %d", len(result.RestartMsgs), len(accumulated)) + } + if gotType != "eino_context_overflow_retry" { + t.Fatalf("event type = %q, want eino_context_overflow_retry", gotType) + } + if gotMessage != "上下文超限,正在激进压缩后重试…" { + t.Fatalf("message = %q", gotMessage) + } + assertContextOverflowMapValue(t, gotData, "conversationId", "conv-1") + assertContextOverflowMapValue(t, gotData, "source", "eino") + assertContextOverflowMapValue(t, gotData, "orchestration", "deep_agent") + assertContextOverflowMapValue(t, gotData, "contextSource", string(einoRestartContextAccumulated)) + if logs.FilterMessage("eino context overflow, retrying with aggressive compaction").Len() != 1 { + t.Fatalf("expected one context overflow retry log, got %d", logs.Len()) + } + + second := handler.Prepare(errors.New("maximum context length"), accumulated, len(baseMsgs)) + if second.Handled { + t.Fatalf("second result = %+v, want unhandled after first retry", second) + } +} + +func TestEinoContextOverflowRetryHandlerIgnoresOtherErrors(t *testing.T) { + handler := newEinoContextOverflowRetryHandler(einoContextOverflowRetryConfig{ + Context: context.Background(), + Args: &einoADKRunLoopArgs{}, + BaseMsgs: []adk.Message{schema.UserMessage("base")}, + }) + result := handler.Prepare(errors.New("HTTP 429 Too Many Requests"), nil, 0) + if result.Handled { + t.Fatalf("result = %+v, want unhandled", result) + } +} + +func assertContextOverflowMapValue(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_execute_stdout_suppressor.go b/internal/multiagent/eino_execute_stdout_suppressor.go new file mode 100644 index 00000000..455ce409 --- /dev/null +++ b/internal/multiagent/eino_execute_stdout_suppressor.go @@ -0,0 +1,57 @@ +package multiagent + +import ( + "strings" + "sync" +) + +type einoExecuteStdoutSuppressor struct { + mu sync.Mutex + pending string +} + +func newEinoExecuteStdoutSuppressor() *einoExecuteStdoutSuppressor { + return &einoExecuteStdoutSuppressor{} +} + +func (s *einoExecuteStdoutSuppressor) Record(toolName, stdout string, isErr bool) { + if s == nil || isErr || !strings.EqualFold(strings.TrimSpace(toolName), "execute") { + return + } + t := strings.TrimSpace(stdout) + if t == "" { + return + } + s.mu.Lock() + s.pending = t + s.mu.Unlock() +} + +func (s *einoExecuteStdoutSuppressor) Peek() string { + if s == nil { + return "" + } + s.mu.Lock() + defer s.mu.Unlock() + return s.pending +} + +func (s *einoExecuteStdoutSuppressor) Consume() string { + if s == nil { + return "" + } + s.mu.Lock() + defer s.mu.Unlock() + out := s.pending + s.pending = "" + return out +} + +func (s *einoExecuteStdoutSuppressor) Clear() { + if s == nil { + return + } + s.mu.Lock() + s.pending = "" + s.mu.Unlock() +} diff --git a/internal/multiagent/eino_execute_stdout_suppressor_test.go b/internal/multiagent/eino_execute_stdout_suppressor_test.go new file mode 100644 index 00000000..f4e99e01 --- /dev/null +++ b/internal/multiagent/eino_execute_stdout_suppressor_test.go @@ -0,0 +1,42 @@ +package multiagent + +import "testing" + +func TestEinoExecuteStdoutSuppressorRecordsOnlySuccessfulExecute(t *testing.T) { + s := newEinoExecuteStdoutSuppressor() + s.Record("read_file", "file body", false) + if got := s.Peek(); got != "" { + t.Fatalf("non-execute should not be recorded, got %q", got) + } + s.Record("execute", "failed", true) + if got := s.Peek(); got != "" { + t.Fatalf("failed execute should not be recorded, got %q", got) + } + s.Record(" execute ", " hello\n", false) + if got := s.Peek(); got != "hello" { + t.Fatalf("Peek = %q, want hello", got) + } +} + +func TestEinoExecuteStdoutSuppressorConsumeAndClear(t *testing.T) { + s := newEinoExecuteStdoutSuppressor() + s.Record("execute", "stdout", false) + if got := s.Peek(); got != "stdout" { + t.Fatalf("Peek = %q, want stdout", got) + } + if got := s.Peek(); got != "stdout" { + t.Fatalf("Peek should not clear, got %q", got) + } + if got := s.Consume(); got != "stdout" { + t.Fatalf("Consume = %q, want stdout", got) + } + if got := s.Peek(); got != "" { + t.Fatalf("Consume should clear, got %q", got) + } + + s.Record("execute", "again", false) + s.Clear() + if got := s.Consume(); got != "" { + t.Fatalf("Clear should remove pending value, got %q", got) + } +} diff --git a/internal/multiagent/eino_filesystem_tool_monitor_test.go b/internal/multiagent/eino_filesystem_tool_monitor_test.go new file mode 100644 index 00000000..332e81f4 --- /dev/null +++ b/internal/multiagent/eino_filesystem_tool_monitor_test.go @@ -0,0 +1,82 @@ +package multiagent + +import ( + "context" + "testing" + + "cyberstrike-ai/internal/agent" + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/einomcp" + "cyberstrike-ai/internal/mcp" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" +) + +func TestEinoADKFilesystemToolMonitorBindsFinishesAndUpdatesDisplayResult(t *testing.T) { + t.Parallel() + ctx := context.Background() + logger := zap.NewNop() + server := mcp.NewServer(logger) + ag := agent.NewAgent(&config.OpenAIConfig{}, &config.AgentConfig{}, server, nil, logger, 1) + binder := NewMCPExecutionBinder() + var recorded []string + rec := einomcp.ExecutionRecorder(func(executionID, toolCallID string) { + recorded = append(recorded, executionID+"|"+toolCallID) + }) + + beginEinoADKFilesystemToolMonitor(ctx, ag, rec, binder, "call-read", "read_file") + execID := binder.ExecutionID("call-read") + if execID == "" { + t.Fatal("expected begin to bind execution id") + } + exec, ok := server.GetExecution(execID) + if !ok || exec == nil || exec.Status != "running" || exec.ToolName != "eino_fs::read_file" { + t.Fatalf("begin execution = %#v ok=%v", exec, ok) + } + if len(recorded) != 1 || recorded[0] != execID+"|call-read" { + t.Fatalf("recorded begin ids = %#v", recorded) + } + + runMessages := newEinoRunMessageAccumulator([]adk.Message{ + &schema.Message{ + Role: schema.Assistant, + ToolCalls: []schema.ToolCall{{ + ID: "call-read", + Type: "function", + Function: schema.FunctionCall{ + Name: "read_file", + Arguments: `{"path":"/tmp/secret.txt"}`, + }, + }}, + }, + }) + emitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{ + ConversationID: "conv-1", + RunMessages: runMessages, + FilesystemMonitorAgent: ag, + FilesystemMonitorRecord: rec, + MCPExecutionBinder: binder, + }) + + if !emitter.Emit(ctx, "read_file", "model-facing truncated body", "call-read", false, "lead") { + t.Fatal("expected tool_result emit") + } + exec, ok = server.GetExecution(execID) + if !ok || exec == nil { + t.Fatalf("finished execution missing: ok=%v exec=%#v", ok, exec) + } + if exec.Status != "completed" || exec.ToolName != "eino_fs::read_file" { + t.Fatalf("finished execution status/name = %#v", exec) + } + if got, _ := exec.Arguments["path"].(string); got != "/tmp/secret.txt" { + t.Fatalf("execution args = %#v", exec.Arguments) + } + if exec.Result == nil || len(exec.Result.Content) != 1 || exec.Result.Content[0].Text != "model-facing truncated body" { + t.Fatalf("execution display result = %#v", exec.Result) + } + if len(recorded) != 1 { + t.Fatalf("finish should reuse existing execution without recording a second id, got %#v", recorded) + } +} diff --git a/internal/multiagent/eino_initial_iterator_start_handler.go b/internal/multiagent/eino_initial_iterator_start_handler.go new file mode 100644 index 00000000..630b6a14 --- /dev/null +++ b/internal/multiagent/eino_initial_iterator_start_handler.go @@ -0,0 +1,54 @@ +package multiagent + +import "github.com/cloudwego/eino/adk" + +type einoAgentEventIteratorStarter func([]adk.Message) *adk.AsyncIterator[*adk.AgentEvent] + +type einoInitialIteratorStartHandlerConfig struct { + ConversationID string + OrchMode string + Progress func(eventType, message string, data interface{}) + UseTurnLoop bool + StartRunner einoAgentEventIteratorStarter + StartTurnLoop einoAgentEventIteratorStarter +} + +type einoInitialIteratorStartHandler struct { + cfg einoInitialIteratorStartHandlerConfig +} + +func newEinoInitialIteratorStartHandler(cfg einoInitialIteratorStartHandlerConfig) *einoInitialIteratorStartHandler { + return &einoInitialIteratorStartHandler{cfg: cfg} +} + +func (h *einoInitialIteratorStartHandler) StartIfNeeded(existing *adk.AsyncIterator[*adk.AgentEvent], msgs []adk.Message) *adk.AsyncIterator[*adk.AgentEvent] { + if existing != nil { + return existing + } + if h == nil { + return nil + } + if h.cfg.UseTurnLoop { + h.emitTurnLoopTakeover() + if h.cfg.StartTurnLoop == nil { + return nil + } + return h.cfg.StartTurnLoop(msgs) + } + if h.cfg.StartRunner == nil { + return nil + } + return h.cfg.StartRunner(msgs) +} + +func (h *einoInitialIteratorStartHandler) emitTurnLoopTakeover() { + if h == nil || h.cfg.Progress == nil { + return + } + h.cfg.Progress("progress", "Eino TurnLoop 常驻多轮 runtime 已接管本轮会话。", map[string]interface{}{ + "conversationId": h.cfg.ConversationID, + "source": "eino", + "orchestration": h.cfg.OrchMode, + "kind": "turn_loop_takeover", + }) +} diff --git a/internal/multiagent/eino_initial_iterator_start_handler_test.go b/internal/multiagent/eino_initial_iterator_start_handler_test.go new file mode 100644 index 00000000..3564e3e1 --- /dev/null +++ b/internal/multiagent/eino_initial_iterator_start_handler_test.go @@ -0,0 +1,111 @@ +package multiagent + +import ( + "testing" + + "github.com/cloudwego/eino/adk" +) + +func TestEinoInitialIteratorStartHandlerKeepsExistingIterator(t *testing.T) { + existing, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + defer gen.Close() + + var started bool + got := newEinoInitialIteratorStartHandler(einoInitialIteratorStartHandlerConfig{ + UseTurnLoop: true, + StartTurnLoop: func([]adk.Message) *adk.AsyncIterator[*adk.AgentEvent] { + started = true + iter, iterGen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + iterGen.Close() + return iter + }, + Progress: func(string, string, interface{}) { + t.Fatal("progress should not be emitted when an iterator already exists") + }, + }).StartIfNeeded(existing, nil) + + if got != existing { + t.Fatal("existing iterator should be preserved") + } + if started { + t.Fatal("start function should not be called when an iterator already exists") + } +} + +func TestEinoInitialIteratorStartHandlerStartsRunner(t *testing.T) { + wantIter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + defer gen.Close() + + var runnerStarted bool + got := newEinoInitialIteratorStartHandler(einoInitialIteratorStartHandlerConfig{ + StartRunner: func(msgs []adk.Message) *adk.AsyncIterator[*adk.AgentEvent] { + runnerStarted = true + if msgs == nil { + t.Fatal("msgs should be forwarded") + } + return wantIter + }, + StartTurnLoop: func([]adk.Message) *adk.AsyncIterator[*adk.AgentEvent] { + t.Fatal("turn loop should not start when UseTurnLoop is false") + return nil + }, + Progress: func(string, string, interface{}) { + t.Fatal("runner start should not emit TurnLoop takeover progress") + }, + }).StartIfNeeded(nil, []adk.Message{}) + + if !runnerStarted { + t.Fatal("runner start was not called") + } + if got != wantIter { + t.Fatal("runner iterator should be returned") + } +} + +func TestEinoInitialIteratorStartHandlerStartsTurnLoopWithTakeoverProgress(t *testing.T) { + wantIter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + defer gen.Close() + + var turnLoopStarted bool + var gotType, gotMessage string + var gotData map[string]interface{} + got := newEinoInitialIteratorStartHandler(einoInitialIteratorStartHandlerConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + UseTurnLoop: true, + StartRunner: func([]adk.Message) *adk.AsyncIterator[*adk.AgentEvent] { + t.Fatal("runner should not start when UseTurnLoop is true") + return nil + }, + StartTurnLoop: func(msgs []adk.Message) *adk.AsyncIterator[*adk.AgentEvent] { + turnLoopStarted = true + if msgs == nil { + t.Fatal("msgs should be forwarded") + } + return wantIter + }, + Progress: func(eventType, message string, data interface{}) { + gotType = eventType + gotMessage = message + if m, ok := data.(map[string]interface{}); ok { + gotData = m + } + }, + }).StartIfNeeded(nil, []adk.Message{}) + + if !turnLoopStarted { + t.Fatal("turn loop start was not called") + } + if got != wantIter { + t.Fatal("turn loop iterator should be returned") + } + if gotType != "progress" { + t.Fatalf("progress type = %q, want progress", gotType) + } + if gotMessage != "Eino TurnLoop 常驻多轮 runtime 已接管本轮会话。" { + t.Fatalf("progress message = %q", gotMessage) + } + if gotData["conversationId"] != "conv-1" || gotData["source"] != "eino" || gotData["orchestration"] != "deep" { + t.Fatalf("progress data = %#v", gotData) + } +} diff --git a/internal/multiagent/eino_main_assistant_complete_handler.go b/internal/multiagent/eino_main_assistant_complete_handler.go new file mode 100644 index 00000000..427ee854 --- /dev/null +++ b/internal/multiagent/eino_main_assistant_complete_handler.go @@ -0,0 +1,49 @@ +package multiagent + +import "strings" + +type einoMainAssistantCompleteHandler struct { + agentName string + emitter *einoMainResponseStreamEmitter + stdoutSuppressor *einoExecuteStdoutSuppressor + assistantOutput *einoAssistantOutputAccumulator +} + +type einoMainAssistantCompleteHandlerConfig struct { + AgentName string + Emitter *einoMainResponseStreamEmitter + StdoutSuppressor *einoExecuteStdoutSuppressor + AssistantOutput *einoAssistantOutputAccumulator +} + +func newEinoMainAssistantCompleteHandler(cfg einoMainAssistantCompleteHandlerConfig) *einoMainAssistantCompleteHandler { + return &einoMainAssistantCompleteHandler{ + agentName: cfg.AgentName, + emitter: cfg.Emitter, + stdoutSuppressor: cfg.StdoutSuppressor, + assistantOutput: cfg.AssistantOutput, + } +} + +func (h *einoMainAssistantCompleteHandler) EmitComplete(content string) bool { + if h == nil { + return false + } + body := strings.TrimSpace(content) + if body == "" { + return false + } + if h.stdoutSuppressor != nil { + if dup := h.stdoutSuppressor.Consume(); dup != "" && body == dup { + if h.assistantOutput != nil { + h.assistantOutput.RecordMainAssistant(h.agentName, body) + } + return false + } + } + emitted := h.emitter.EmitDelta(body, body) + if h.assistantOutput != nil { + h.assistantOutput.RecordMainAssistant(h.agentName, body) + } + return emitted +} diff --git a/internal/multiagent/eino_main_assistant_complete_handler_test.go b/internal/multiagent/eino_main_assistant_complete_handler_test.go new file mode 100644 index 00000000..577604f5 --- /dev/null +++ b/internal/multiagent/eino_main_assistant_complete_handler_test.go @@ -0,0 +1,76 @@ +package multiagent + +import "testing" + +func TestEinoMainAssistantCompleteHandlerEmitsAndRecords(t *testing.T) { + var eventTypes []string + var messages []string + progress := func(eventType, message string, _ interface{}) { + eventTypes = append(eventTypes, eventType) + messages = append(messages, message) + } + out := newEinoAssistantOutputAccumulator("deep") + handler := newEinoMainAssistantCompleteHandler(einoMainAssistantCompleteHandlerConfig{ + AgentName: "lead", + Emitter: newEinoMainResponseStreamEmitter("conv-1", "deep", "lead", "stream-1", 2, progress, nil), + AssistantOutput: out, + }) + + if !handler.EmitComplete(" hello ") { + t.Fatal("complete assistant should emit") + } + if len(eventTypes) != 2 || eventTypes[0] != "response_start" || eventTypes[1] != "response_delta" { + t.Fatalf("events = %#v", eventTypes) + } + if messages[1] != "hello" { + t.Fatalf("delta message = %q", messages[1]) + } + if out.LastAssistant() != "hello" { + t.Fatalf("last assistant = %q", out.LastAssistant()) + } +} + +func TestEinoMainAssistantCompleteHandlerSuppressesDuplicateExecuteStdout(t *testing.T) { + var eventTypes []string + progress := func(eventType, _ string, _ interface{}) { + eventTypes = append(eventTypes, eventType) + } + stdoutDup := newEinoExecuteStdoutSuppressor() + stdoutDup.Record("execute", "hello", false) + out := newEinoAssistantOutputAccumulator("deep") + handler := newEinoMainAssistantCompleteHandler(einoMainAssistantCompleteHandlerConfig{ + AgentName: "lead", + Emitter: newEinoMainResponseStreamEmitter("conv-1", "deep", "lead", "stream-1", 1, progress, nil), + StdoutSuppressor: stdoutDup, + AssistantOutput: out, + }) + + if handler.EmitComplete("hello") { + t.Fatal("duplicate execute stdout should not emit") + } + if len(eventTypes) != 0 { + t.Fatalf("events = %#v, want none", eventTypes) + } + if out.LastAssistant() != "hello" { + t.Fatalf("last assistant = %q", out.LastAssistant()) + } + if stdoutDup.Peek() != "" { + t.Fatal("duplicate target should be consumed") + } +} + +func TestEinoMainAssistantCompleteHandlerRecordsWithoutProgress(t *testing.T) { + out := newEinoAssistantOutputAccumulator("plan_execute") + handler := newEinoMainAssistantCompleteHandler(einoMainAssistantCompleteHandlerConfig{ + AgentName: "executor", + Emitter: newEinoMainResponseStreamEmitter("conv-1", "plan_execute", "executor", "stream-1", 1, nil, nil), + AssistantOutput: out, + }) + + if handler.EmitComplete(`{"response":"done"}`) { + t.Fatal("nil progress should not emit") + } + if out.LastPlanExecuteExecutor() != "done" { + t.Fatalf("executor output = %q", out.LastPlanExecuteExecutor()) + } +} diff --git a/internal/multiagent/eino_main_assistant_stream_handler.go b/internal/multiagent/eino_main_assistant_stream_handler.go new file mode 100644 index 00000000..8e6d9207 --- /dev/null +++ b/internal/multiagent/eino_main_assistant_stream_handler.go @@ -0,0 +1,77 @@ +package multiagent + +import "strings" + +type einoMainAssistantStreamHandler struct { + agentName string + emitter *einoMainResponseStreamEmitter + stdoutSuppressor *einoExecuteStdoutSuppressor + assistantOutput *einoAssistantOutputAccumulator + runMessages *einoRunMessageAccumulator + + buf string + dupTarget string +} + +type einoMainAssistantStreamHandlerConfig struct { + AgentName string + Emitter *einoMainResponseStreamEmitter + StdoutSuppressor *einoExecuteStdoutSuppressor + AssistantOutput *einoAssistantOutputAccumulator + RunMessages *einoRunMessageAccumulator +} + +func newEinoMainAssistantStreamHandler(cfg einoMainAssistantStreamHandlerConfig) *einoMainAssistantStreamHandler { + return &einoMainAssistantStreamHandler{ + agentName: cfg.AgentName, + emitter: cfg.Emitter, + stdoutSuppressor: cfg.StdoutSuppressor, + assistantOutput: cfg.AssistantOutput, + runMessages: cfg.RunMessages, + } +} + +func (h *einoMainAssistantStreamHandler) EmitDelta(content string) bool { + if h == nil || content == "" { + return false + } + var delta string + h.buf, delta = normalizeStreamingDelta(h.buf, content) + if delta == "" { + return false + } + if h.dupTarget == "" && h.stdoutSuppressor != nil { + h.dupTarget = h.stdoutSuppressor.Peek() + } + if h.dupTarget != "" { + return false + } + return h.emitter.EmitDelta(delta, h.buf) +} + +func (h *einoMainAssistantStreamHandler) Finish() string { + if h == nil { + return "" + } + body := strings.TrimSpace(h.buf) + if body == "" { + return "" + } + if h.dupTarget != "" { + if h.stdoutSuppressor != nil { + h.stdoutSuppressor.Clear() + } + if body != h.dupTarget { + h.emitter.EmitTailFromFull(h.buf) + } + } else { + h.emitter.EmitTailFromFull(h.buf) + } + if h.assistantOutput != nil { + h.assistantOutput.RecordMainAssistant(h.agentName, body) + } + if h.runMessages != nil { + h.runMessages.AppendAssistantText(body) + } + return body +} diff --git a/internal/multiagent/eino_main_assistant_stream_handler_test.go b/internal/multiagent/eino_main_assistant_stream_handler_test.go new file mode 100644 index 00000000..5feb7adf --- /dev/null +++ b/internal/multiagent/eino_main_assistant_stream_handler_test.go @@ -0,0 +1,103 @@ +package multiagent + +import "testing" + +func TestEinoMainAssistantStreamHandlerEmitsAndRecords(t *testing.T) { + var eventTypes []string + var messages []string + progress := func(eventType, message string, _ interface{}) { + eventTypes = append(eventTypes, eventType) + messages = append(messages, message) + } + out := newEinoAssistantOutputAccumulator("deep") + runMsgs := newEinoRunMessageAccumulator(nil) + emitter := newEinoMainResponseStreamEmitter("conv-1", "deep", "lead", "stream-1", 2, progress, nil) + handler := newEinoMainAssistantStreamHandler(einoMainAssistantStreamHandlerConfig{ + AgentName: "lead", + Emitter: emitter, + AssistantOutput: out, + RunMessages: runMsgs, + }) + + if !handler.EmitDelta("he") { + t.Fatal("first delta should emit") + } + if !handler.EmitDelta("hello") { + t.Fatal("cumulative chunk should emit tail") + } + if got := handler.Finish(); got != "hello" { + t.Fatalf("finish = %q, want hello", got) + } + + if len(eventTypes) != 3 || eventTypes[0] != "response_start" || eventTypes[1] != "response_delta" || eventTypes[2] != "response_delta" { + t.Fatalf("events = %#v", eventTypes) + } + if messages[1] != "he" || messages[2] != "llo" { + t.Fatalf("delta messages = %#v", messages) + } + if out.LastAssistant() != "hello" { + t.Fatalf("last assistant = %q", out.LastAssistant()) + } + if msgs := runMsgs.Messages(); len(msgs) != 1 || msgs[0].Content != "hello" { + t.Fatalf("run messages = %#v", msgs) + } +} + +func TestEinoMainAssistantStreamHandlerSuppressesDuplicateExecuteStdout(t *testing.T) { + var eventTypes []string + progress := func(eventType, _ string, _ interface{}) { + eventTypes = append(eventTypes, eventType) + } + stdoutDup := newEinoExecuteStdoutSuppressor() + stdoutDup.Record("execute", "hello", false) + out := newEinoAssistantOutputAccumulator("deep") + runMsgs := newEinoRunMessageAccumulator(nil) + handler := newEinoMainAssistantStreamHandler(einoMainAssistantStreamHandlerConfig{ + AgentName: "lead", + Emitter: newEinoMainResponseStreamEmitter("conv-1", "deep", "lead", "stream-1", 1, progress, nil), + StdoutSuppressor: stdoutDup, + AssistantOutput: out, + RunMessages: runMsgs, + }) + + if handler.EmitDelta("hello") { + t.Fatal("duplicate execute stdout should not emit delta") + } + if got := handler.Finish(); got != "hello" { + t.Fatalf("finish = %q, want hello", got) + } + if len(eventTypes) != 0 { + t.Fatalf("events = %#v, want none", eventTypes) + } + if stdoutDup.Peek() != "" { + t.Fatal("duplicate target should be cleared on finish") + } + if out.LastAssistant() != "hello" { + t.Fatalf("last assistant = %q", out.LastAssistant()) + } + if msgs := runMsgs.Messages(); len(msgs) != 1 || msgs[0].Content != "hello" { + t.Fatalf("run messages = %#v", msgs) + } +} + +func TestEinoMainAssistantStreamHandlerRecordsWithoutProgress(t *testing.T) { + out := newEinoAssistantOutputAccumulator("plan_execute") + runMsgs := newEinoRunMessageAccumulator(nil) + handler := newEinoMainAssistantStreamHandler(einoMainAssistantStreamHandlerConfig{ + AgentName: "executor", + Emitter: newEinoMainResponseStreamEmitter("conv-1", "plan_execute", "executor", "stream-1", 1, nil, nil), + AssistantOutput: out, + RunMessages: runMsgs, + }) + + handler.EmitDelta(`{"response":"done"}`) + if got := handler.Finish(); got != `{"response":"done"}` { + t.Fatalf("finish = %q", got) + } + if out.LastPlanExecuteExecutor() != "done" { + t.Fatalf("executor output = %q", out.LastPlanExecuteExecutor()) + } + if msgs := runMsgs.Messages(); len(msgs) != 1 || msgs[0].Content != `{"response":"done"}` { + t.Fatalf("run messages = %#v", msgs) + } +} diff --git a/internal/multiagent/eino_main_response_stream_emitter.go b/internal/multiagent/eino_main_response_stream_emitter.go new file mode 100644 index 00000000..19dcc69d --- /dev/null +++ b/internal/multiagent/eino_main_response_stream_emitter.go @@ -0,0 +1,85 @@ +package multiagent + +import "cyberstrike-ai/internal/openai" + +type einoMainResponseStreamEmitter struct { + progress func(eventType, message string, data interface{}) + snapshotMCPIDs func() []string + conversationID string + orchMode string + agentName string + streamID string + iteration int + headerSent bool + wireAccum string +} + +func newEinoMainResponseStreamEmitter( + conversationID, orchMode, agentName, streamID string, + iteration int, + progress func(eventType, message string, data interface{}), + snapshotMCPIDs func() []string, +) *einoMainResponseStreamEmitter { + if snapshotMCPIDs == nil { + snapshotMCPIDs = func() []string { return nil } + } + return &einoMainResponseStreamEmitter{ + progress: progress, + snapshotMCPIDs: snapshotMCPIDs, + conversationID: conversationID, + orchMode: orchMode, + agentName: agentName, + streamID: streamID, + iteration: iteration, + } +} + +func (e *einoMainResponseStreamEmitter) EmitDelta(delta, accumulated string) bool { + if e == nil || e.progress == nil || delta == "" { + return false + } + e.emitStart() + e.progress("response_delta", delta, openai.WithSSEAccumulated(e.responseData(), accumulated)) + e.wireAccum, _ = normalizeStreamingDelta(e.wireAccum, delta) + return true +} + +func (e *einoMainResponseStreamEmitter) EmitTailFromFull(full string) bool { + if e == nil || full == "" { + return false + } + _, tail := normalizeStreamingDelta(e.wireAccum, full) + if tail == "" { + return false + } + return e.EmitDelta(tail, full) +} + +func (e *einoMainResponseStreamEmitter) emitStart() { + if e.headerSent || e.progress == nil { + return + } + e.progress("response_start", "", map[string]interface{}{ + "conversationId": e.conversationID, + "mcpExecutionIds": e.snapshotMCPIDs(), + "messageGeneratedBy": "eino:" + e.agentName, + "einoRole": "orchestrator", + "einoAgent": e.agentName, + "orchestration": e.orchMode, + "iteration": e.iteration, + "streamId": e.streamID, + }) + e.headerSent = true +} + +func (e *einoMainResponseStreamEmitter) responseData() map[string]interface{} { + return map[string]interface{}{ + "conversationId": e.conversationID, + "mcpExecutionIds": e.snapshotMCPIDs(), + "einoRole": "orchestrator", + "einoAgent": e.agentName, + "orchestration": e.orchMode, + "iteration": e.iteration, + "streamId": e.streamID, + } +} diff --git a/internal/multiagent/eino_main_response_stream_emitter_test.go b/internal/multiagent/eino_main_response_stream_emitter_test.go new file mode 100644 index 00000000..36188bd0 --- /dev/null +++ b/internal/multiagent/eino_main_response_stream_emitter_test.go @@ -0,0 +1,65 @@ +package multiagent + +import ( + "testing" + + "cyberstrike-ai/internal/openai" +) + +func TestEinoMainResponseStreamEmitterEmitsStartOnceAndTail(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 := newEinoMainResponseStreamEmitter( + "conv-1", "supervisor", "lead", "stream-1", 3, progress, func() []string { return []string{"mcp-1"} }, + ) + if !emitter.EmitDelta("he", "he") { + t.Fatal("first delta should be emitted") + } + if !emitter.EmitTailFromFull("hello") { + t.Fatal("tail should be emitted") + } + if emitter.EmitTailFromFull("hello") { + t.Fatal("duplicate tail should not be emitted") + } + + if len(events) != 3 { + t.Fatalf("events = %#v, want start + 2 deltas", events) + } + if events[0].eventType != "response_start" { + t.Fatalf("event[0] = %s, want response_start", events[0].eventType) + } + if events[1].eventType != "response_delta" || events[1].message != "he" { + t.Fatalf("event[1] = %#v, want first delta", events[1]) + } + if events[2].eventType != "response_delta" || events[2].message != "llo" { + t.Fatalf("event[2] = %#v, want tail delta", events[2]) + } + if got := events[2].data[openai.SSEAccumulatedKey]; got != "hello" { + t.Fatalf("accumulated = %#v, want hello", got) + } + if got := events[0].data["messageGeneratedBy"]; got != "eino:lead" { + t.Fatalf("messageGeneratedBy = %#v", got) + } + if got := events[0].data["iteration"]; got != 3 { + t.Fatalf("iteration = %#v", got) + } +} + +func TestEinoMainResponseStreamEmitterNoProgress(t *testing.T) { + emitter := newEinoMainResponseStreamEmitter("conv", "deep", "agent", "stream", 1, nil, nil) + if emitter.EmitDelta("hello", "hello") { + t.Fatal("nil progress should not emit") + } + if emitter.EmitTailFromFull("hello") { + t.Fatal("nil progress should not emit tail") + } +} diff --git a/internal/multiagent/eino_materialized_message_event_handler.go b/internal/multiagent/eino_materialized_message_event_handler.go new file mode 100644 index 00000000..810de05d --- /dev/null +++ b/internal/multiagent/eino_materialized_message_event_handler.go @@ -0,0 +1,115 @@ +package multiagent + +import ( + "strings" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +type einoMaterializedMessageEventHandlerConfig struct { + ConversationID string + OrchMode string + Progress func(eventType, message string, data interface{}) + SnapshotMCPIDs func() []string + StreamsMainAssistant func(agent string) bool + EinoRoleTag func(agent string) string + RunProgress *einoRunProgressTracker + StdoutSuppressor *einoExecuteStdoutSuppressor + AssistantOutput *einoAssistantOutputAccumulator + RunMessages *einoRunMessageAccumulator + Usage *einoRunUsageAccumulator + ToolResultHandler *einoToolResultEventHandler + MarkPending func(toolCallPendingInfo) + NextMainStreamID func() string +} + +type einoMaterializedMessageEventHandler struct { + conversationID string + orchMode string + progress func(eventType, message string, data interface{}) + snapshotMCPIDs func() []string + streamsMainAssistant func(agent string) bool + einoRoleTag func(agent string) string + runProgress *einoRunProgressTracker + stdoutSuppressor *einoExecuteStdoutSuppressor + assistantOutput *einoAssistantOutputAccumulator + runMessages *einoRunMessageAccumulator + usage *einoRunUsageAccumulator + toolResultHandler *einoToolResultEventHandler + markPending func(toolCallPendingInfo) + nextMainStreamID func() string +} + +func newEinoMaterializedMessageEventHandler(cfg einoMaterializedMessageEventHandlerConfig) *einoMaterializedMessageEventHandler { + if cfg.SnapshotMCPIDs == nil { + cfg.SnapshotMCPIDs = func() []string { return nil } + } + if cfg.StreamsMainAssistant == nil { + cfg.StreamsMainAssistant = func(string) bool { return true } + } + if cfg.EinoRoleTag == nil { + cfg.EinoRoleTag = func(string) string { return "" } + } + if cfg.NextMainStreamID == nil { + cfg.NextMainStreamID = func() string { return "eino-main" } + } + return &einoMaterializedMessageEventHandler{ + conversationID: cfg.ConversationID, + orchMode: cfg.OrchMode, + progress: cfg.Progress, + snapshotMCPIDs: cfg.SnapshotMCPIDs, + streamsMainAssistant: cfg.StreamsMainAssistant, + einoRoleTag: cfg.EinoRoleTag, + runProgress: cfg.RunProgress, + stdoutSuppressor: cfg.StdoutSuppressor, + assistantOutput: cfg.AssistantOutput, + runMessages: cfg.RunMessages, + usage: cfg.Usage, + toolResultHandler: cfg.ToolResultHandler, + markPending: cfg.MarkPending, + nextMainStreamID: cfg.NextMainStreamID, + } +} + +func (h *einoMaterializedMessageEventHandler) Handle(mv *adk.MessageVariant, msg adk.Message, agentName string) bool { + if h == nil || mv == nil || msg == nil { + return false + } + if h.runMessages != nil { + h.runMessages.Append(msg) + } + if msg.Role == schema.Assistant && h.usage != nil { + h.usage.AddMessage(msg) + } + if h.runProgress != nil { + h.runProgress.EmitToolCalls(mergeMessageToolCalls(msg), agentName, h.markPending) + } + if mv.Role == schema.Assistant { + newEinoReasoningStreamEmitter(h.conversationID, h.orchMode, agentName, h.einoRoleTag(agentName), h.progress, nil).EmitComplete(msg.ReasoningContent) + body := strings.TrimSpace(msg.Content) + if body != "" { + if h.streamsMainAssistant(agentName) { + newEinoMainAssistantCompleteHandler(einoMainAssistantCompleteHandlerConfig{ + AgentName: agentName, + Emitter: newEinoMainResponseStreamEmitter(h.conversationID, h.orchMode, agentName, h.nextMainStreamID(), h.mainIteration(agentName), h.progress, h.snapshotMCPIDs), + StdoutSuppressor: h.stdoutSuppressor, + AssistantOutput: h.assistantOutput, + }).EmitComplete(body) + } else { + newEinoSubAgentReplyEmitter(h.conversationID, agentName, h.progress, nil).EmitComplete(body) + } + } + } + if h.toolResultHandler != nil { + h.toolResultHandler.HandleMaterialized(mv, msg, agentName) + } + return true +} + +func (h *einoMaterializedMessageEventHandler) mainIteration(agentName string) int { + if h == nil || h.runProgress == nil { + return 0 + } + return h.runProgress.MainIteration(agentName) +} diff --git a/internal/multiagent/eino_materialized_message_event_handler_test.go b/internal/multiagent/eino_materialized_message_event_handler_test.go new file mode 100644 index 00000000..a67478c6 --- /dev/null +++ b/internal/multiagent/eino_materialized_message_event_handler_test.go @@ -0,0 +1,151 @@ +package multiagent + +import ( + "testing" + + "cyberstrike-ai/internal/einomcp" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestEinoMaterializedMessageEventHandlerHandlesMainAssistant(t *testing.T) { + var events []string + runMessages := newEinoRunMessageAccumulator(nil) + assistantOutput := newEinoAssistantOutputAccumulator("deep") + usage := newEinoRunUsageAccumulator() + runProgress := newEinoRunProgressTracker( + "deep", "lead", "conv-1", + func(eventType, _ string, _ interface{}) { events = append(events, eventType) }, + func(agent string) bool { return agent == "lead" }, + nil, + ) + handler := newEinoMaterializedMessageEventHandler(einoMaterializedMessageEventHandlerConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + Progress: func(eventType, _ string, _ interface{}) { events = append(events, eventType) }, + RunMessages: runMessages, + Usage: usage, + AssistantOutput: assistantOutput, + RunProgress: runProgress, + StreamsMainAssistant: func(agent string) bool { return agent == "lead" }, + EinoRoleTag: func(string) string { return "orchestrator" }, + NextMainStreamID: func() string { return "main-complete-1" }, + }) + msg := schema.AssistantMessage(" done ", nil) + msg.ReasoningContent = "thought" + msg.ResponseMeta = &schema.ResponseMeta{Usage: &schema.TokenUsage{ + PromptTokens: 11, + CompletionTokens: 7, + TotalTokens: 18, + }} + mv := &adk.MessageVariant{Role: schema.Assistant} + + if !handler.Handle(mv, msg, "lead") { + t.Fatal("main assistant message was not handled") + } + if assistantOutput.LastAssistant() != "done" { + t.Fatalf("last assistant = %q", assistantOutput.LastAssistant()) + } + if msgs := runMessages.Messages(); len(msgs) != 1 || msgs[0].Content != " done " { + t.Fatalf("run messages = %#v", msgs) + } + if got := usage.Summary(); got.ModelCalls != 1 || got.TotalTokens != 18 { + t.Fatalf("usage = %#v, want one assistant model call", got) + } + if !containsString(events, "reasoning_chain") || !containsString(events, "response_start") || !containsString(events, "response_delta") { + t.Fatalf("events = %#v, want reasoning and response events", events) + } +} + +func TestEinoMaterializedMessageEventHandlerHandlesSubAssistant(t *testing.T) { + var events []string + runMessages := newEinoRunMessageAccumulator(nil) + assistantOutput := newEinoAssistantOutputAccumulator("deep") + handler := newEinoMaterializedMessageEventHandler(einoMaterializedMessageEventHandlerConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + Progress: func(eventType, _ string, _ interface{}) { events = append(events, eventType) }, + RunMessages: runMessages, + AssistantOutput: assistantOutput, + StreamsMainAssistant: func(agent string) bool { return agent == "lead" }, + EinoRoleTag: func(string) string { return "sub" }, + }) + + if !handler.Handle(&adk.MessageVariant{Role: schema.Assistant}, schema.AssistantMessage("sub done", nil), "worker") { + t.Fatal("sub assistant message was not handled") + } + if assistantOutput.LastAssistant() != "" { + t.Fatalf("sub assistant should not update main output, got %q", assistantOutput.LastAssistant()) + } + if len(runMessages.Messages()) != 1 { + t.Fatalf("run messages = %#v, want appended original message", runMessages.Messages()) + } + if !containsString(events, "eino_agent_reply") { + t.Fatalf("events = %#v, want sub reply event", events) + } +} + +func TestEinoMaterializedMessageEventHandlerHandlesToolCallsAndToolResult(t *testing.T) { + var events []string + var marked []toolCallPendingInfo + runMessages := newEinoRunMessageAccumulator(nil) + runProgress := newEinoRunProgressTracker( + "deep", "lead", "conv-1", + func(eventType, _ string, _ interface{}) { events = append(events, eventType) }, + func(agent string) bool { return agent == "lead" }, + nil, + ) + toolResultEmitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{ + ConversationID: "conv-1", + Progress: func(eventType, _ string, _ interface{}) { events = append(events, eventType) }, + }) + toolResultHandler := newEinoToolResultEventHandler(einoToolResultEventHandlerConfig{Emitter: toolResultEmitter}) + handler := newEinoMaterializedMessageEventHandler(einoMaterializedMessageEventHandlerConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + Progress: func(eventType, _ string, _ interface{}) { events = append(events, eventType) }, + RunMessages: runMessages, + RunProgress: runProgress, + ToolResultHandler: toolResultHandler, + MarkPending: func(info toolCallPendingInfo) { + marked = append(marked, info) + }, + }) + + toolCallMsg := &schema.Message{ + Role: schema.Assistant, + ToolCalls: []schema.ToolCall{{ + ID: "call-1", + Type: "function", + Function: schema.FunctionCall{ + Name: "execute", + Arguments: `{"command":`, + }, + }}, + } + if !handler.Handle(&adk.MessageVariant{Role: schema.Assistant}, toolCallMsg, "lead") { + t.Fatal("tool call message was not handled") + } + toolMsg := schema.ToolMessage(einomcp.ToolErrorPrefix+"bad command", "call-1", schema.WithToolName("execute")) + if !handler.Handle(&adk.MessageVariant{Role: schema.Tool}, toolMsg, "lead") { + t.Fatal("tool message was not handled") + } + + if !containsString(events, "tool_call") || !containsString(events, "tool_result") || containsString(events, "model_output_rejected") { + t.Fatalf("events = %#v, want real tool_call and tool_result without model-output recovery", events) + } + if len(marked) != 1 || marked[0].ToolCallID != "call-1" || marked[0].ToolName != "execute" { + t.Fatalf("marked pending = %#v", marked) + } + if len(runMessages.Messages()) != 2 { + t.Fatalf("run messages = %#v, want assistant and tool messages", runMessages.Messages()) + } +} + +func TestEinoMaterializedMessageEventHandlerIgnoresNil(t *testing.T) { + handler := newEinoMaterializedMessageEventHandler(einoMaterializedMessageEventHandlerConfig{}) + if handler.Handle(nil, nil, "lead") { + t.Fatal("nil message should be ignored") + } +} diff --git a/internal/multiagent/eino_message_stream_receiver.go b/internal/multiagent/eino_message_stream_receiver.go new file mode 100644 index 00000000..1cddff2d --- /dev/null +++ b/internal/multiagent/eino_message_stream_receiver.go @@ -0,0 +1,61 @@ +package multiagent + +import ( + "context" + "errors" + "io" + + "github.com/cloudwego/eino/schema" +) + +// recvEinoSchemaMessageStreamWithContext consumes an Eino schema.Message stream +// and stops promptly when ctx is canceled. EOF and nil chunks are treated as a +// normal stream boundary. +func recvEinoSchemaMessageStreamWithContext( + ctx context.Context, + stream *schema.StreamReader[*schema.Message], + buffer int, + onChunk func(*schema.Message), +) error { + if stream == nil { + return nil + } + if buffer <= 0 { + buffer = 1 + } + type streamMsg struct { + chunk *schema.Message + err error + } + recvCh := make(chan streamMsg, buffer) + go func() { + defer close(recvCh) + for { + ch, rerr := stream.Recv() + recvCh <- streamMsg{chunk: ch, err: rerr} + if rerr != nil { + return + } + } + }() + for { + select { + case <-ctx.Done(): + return ctx.Err() + case sm, ok := <-recvCh: + if !ok { + return nil + } + if errors.Is(sm.err, io.EOF) { + return nil + } + if sm.err != nil { + return sm.err + } + if sm.chunk == nil || onChunk == nil { + continue + } + onChunk(sm.chunk) + } + } +} diff --git a/internal/multiagent/eino_middleware.go b/internal/multiagent/eino_middleware.go index 4e90bc02..1f2c057a 100644 --- a/internal/multiagent/eino_middleware.go +++ b/internal/multiagent/eino_middleware.go @@ -17,6 +17,7 @@ import ( "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" ) @@ -149,6 +150,43 @@ func buildReductionMiddleware(ctx context.Context, mw config.MultiAgentEinoMiddl 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. @@ -243,6 +281,97 @@ func prependEinoMiddlewares( 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 @@ -273,3 +402,34 @@ func deepExtrasFromConfig(ma *config.MultiAgentConfig) (outputKey string, taskDe } 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 index a3a0a4fd..45842d86 100644 --- a/internal/multiagent/eino_middleware_test.go +++ b/internal/multiagent/eino_middleware_test.go @@ -7,6 +7,10 @@ import ( "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" ) @@ -28,6 +32,169 @@ func TestReductionCacheRootDir(t *testing.T) { } } +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) { diff --git a/internal/multiagent/eino_model_facing_trace.go b/internal/multiagent/eino_model_facing_trace.go index e18f3307..33d8d011 100644 --- a/internal/multiagent/eino_model_facing_trace.go +++ b/internal/multiagent/eino_model_facing_trace.go @@ -6,6 +6,7 @@ import ( "sync" "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" ) // modelFacingTraceHolder 保存「即将送入 ChatModel」的消息快照(已走 summarization / reduction / orphan 修剪等), @@ -43,6 +44,19 @@ func (h *modelFacingTraceHolder) storeFromState(state *adk.ChatModelAgentState) 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 @@ -82,3 +96,29 @@ func (m *modelFacingTraceMiddleware) BeforeModelRewriteState( } 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_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) + } +}