mirror of
https://github.com/Ed1s0nZ/CyberStrikeAI.git
synced 2026-08-15 07:30:53 +02:00
Add files via upload
This commit is contained in:
@@ -0,0 +1,102 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/cloudwego/eino/adk"
|
||||
)
|
||||
|
||||
const (
|
||||
einoNativeCancelImmediateWait = 1200 * time.Millisecond
|
||||
einoNativeCancelSafePointWait = 3500 * time.Millisecond
|
||||
einoNativeCancelSafePointTTL = 3 * time.Second
|
||||
)
|
||||
|
||||
type agentRuntimeCancelRegistrarKey struct{}
|
||||
type agentTurnLoopInterruptRegistrarKey struct{}
|
||||
|
||||
// AgentRuntimeCancelRegistrar binds the currently active Eino ADK cancel hook
|
||||
// into the host task manager. The hook returns true when Eino accepted and
|
||||
// handled the cancel request, so the host can avoid canceling the parent context.
|
||||
type AgentRuntimeCancelRegistrar func(cancel func(error) bool) (unregister func())
|
||||
|
||||
// WithAgentRuntimeCancelRegistrar lets the HTTP/task layer trigger Eino's native
|
||||
// Agent Cancel before falling back to the existing context cancellation path.
|
||||
func WithAgentRuntimeCancelRegistrar(ctx context.Context, registrar AgentRuntimeCancelRegistrar) context.Context {
|
||||
if ctx == nil || registrar == nil {
|
||||
return ctx
|
||||
}
|
||||
return context.WithValue(ctx, agentRuntimeCancelRegistrarKey{}, registrar)
|
||||
}
|
||||
|
||||
func agentRuntimeCancelRegistrarFromContext(ctx context.Context) AgentRuntimeCancelRegistrar {
|
||||
if ctx == nil {
|
||||
return nil
|
||||
}
|
||||
if v, ok := ctx.Value(agentRuntimeCancelRegistrarKey{}).(AgentRuntimeCancelRegistrar); ok {
|
||||
return v
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// AgentTurnLoopInterruptRegistrar binds a conversation-level TurnLoop interrupt
|
||||
// pusher into the host task manager. The pusher receives the user supplied note
|
||||
// and returns true when the note was accepted by the loop.
|
||||
type AgentTurnLoopInterruptRegistrar func(push func(note string) bool) (unregister func())
|
||||
|
||||
// WithAgentTurnLoopInterruptRegistrar lets the HTTP/task layer enqueue a user
|
||||
// supplement into an active Eino TurnLoop before falling back to cancellation.
|
||||
func WithAgentTurnLoopInterruptRegistrar(ctx context.Context, registrar AgentTurnLoopInterruptRegistrar) context.Context {
|
||||
if ctx == nil || registrar == nil {
|
||||
return ctx
|
||||
}
|
||||
return context.WithValue(ctx, agentTurnLoopInterruptRegistrarKey{}, registrar)
|
||||
}
|
||||
|
||||
func agentTurnLoopInterruptRegistrarFromContext(ctx context.Context) AgentTurnLoopInterruptRegistrar {
|
||||
if ctx == nil {
|
||||
return nil
|
||||
}
|
||||
if v, ok := ctx.Value(agentTurnLoopInterruptRegistrarKey{}).(AgentTurnLoopInterruptRegistrar); ok {
|
||||
return v
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func requestEinoNativeAgentCancel(cancelFn adk.AgentCancelFunc, cause error) (waitErr error, submitted bool, handled bool) {
|
||||
if cancelFn == nil {
|
||||
return nil, false, false
|
||||
}
|
||||
opts, waitFor := einoNativeCancelOptions(cause)
|
||||
handle, submitted := cancelFn(opts...)
|
||||
if !submitted || handle == nil {
|
||||
return nil, submitted, false
|
||||
}
|
||||
waitCh := make(chan error, 1)
|
||||
go func() {
|
||||
waitCh <- handle.Wait()
|
||||
}()
|
||||
select {
|
||||
case err := <-waitCh:
|
||||
handled := err == nil || errors.Is(err, adk.ErrCancelTimeout) || errors.Is(err, adk.ErrExecutionEnded)
|
||||
return err, submitted, handled
|
||||
case <-time.After(waitFor):
|
||||
return context.DeadlineExceeded, submitted, false
|
||||
}
|
||||
}
|
||||
|
||||
func einoNativeCancelOptions(cause error) ([]adk.AgentCancelOption, time.Duration) {
|
||||
if errors.Is(cause, ErrInterruptContinue) {
|
||||
return []adk.AgentCancelOption{
|
||||
adk.WithAgentCancelMode(adk.CancelAfterChatModel | adk.CancelAfterToolCalls),
|
||||
adk.WithAgentCancelTimeout(einoNativeCancelSafePointTTL),
|
||||
adk.WithRecursive(),
|
||||
}, einoNativeCancelSafePointWait
|
||||
}
|
||||
return []adk.AgentCancelOption{
|
||||
adk.WithAgentCancelMode(adk.CancelImmediate),
|
||||
adk.WithRecursive(),
|
||||
}, einoNativeCancelImmediateWait
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/cloudwego/eino/adk"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
func emitEinoNativeModelRetryProgress(
|
||||
conversationID, orchMode string,
|
||||
willRetry *adk.WillRetryError,
|
||||
progress func(eventType, message string, data interface{}),
|
||||
logger *zap.Logger,
|
||||
runErr error,
|
||||
) bool {
|
||||
if willRetry == nil {
|
||||
return false
|
||||
}
|
||||
if progress != nil {
|
||||
reason := ""
|
||||
if willRetry.RejectReason() != nil {
|
||||
reason = fmt.Sprint(willRetry.RejectReason())
|
||||
}
|
||||
progress("eino_model_retry", "模型调用遇到临时问题,Eino 正在原生重试…", map[string]interface{}{
|
||||
"conversationId": conversationID,
|
||||
"source": "eino",
|
||||
"orchestration": orchMode,
|
||||
"attempt": willRetry.RetryAttempt,
|
||||
"reason": reason,
|
||||
"error": willRetry.Error(),
|
||||
})
|
||||
}
|
||||
if logger != nil {
|
||||
logger.Warn("eino native model retry event",
|
||||
zap.String("orchestration", orchMode),
|
||||
zap.Int("attempt", willRetry.RetryAttempt),
|
||||
zap.Error(runErr))
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/cloudwego/eino/adk"
|
||||
"go.uber.org/zap"
|
||||
"go.uber.org/zap/zaptest/observer"
|
||||
)
|
||||
|
||||
func TestEmitEinoNativeModelRetryProgress(t *testing.T) {
|
||||
willRetry := &adk.WillRetryError{
|
||||
ErrStr: "HTTP 429 Too Many Requests",
|
||||
RetryAttempt: 2,
|
||||
}
|
||||
var gotType, gotMessage string
|
||||
var gotData map[string]interface{}
|
||||
called := emitEinoNativeModelRetryProgress("conv-1", "deep_agent", willRetry, func(eventType, message string, data interface{}) {
|
||||
gotType = eventType
|
||||
gotMessage = message
|
||||
var ok bool
|
||||
gotData, ok = data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("progress data type = %T, want map[string]interface{}", data)
|
||||
}
|
||||
}, nil, willRetry)
|
||||
if !called {
|
||||
t.Fatal("called = false, want true")
|
||||
}
|
||||
if gotType != "eino_model_retry" {
|
||||
t.Fatalf("event type = %q, want eino_model_retry", gotType)
|
||||
}
|
||||
if gotMessage != "模型调用遇到临时问题,Eino 正在原生重试…" {
|
||||
t.Fatalf("message = %q", gotMessage)
|
||||
}
|
||||
assertNativeRetryMapValue(t, gotData, "conversationId", "conv-1")
|
||||
assertNativeRetryMapValue(t, gotData, "source", "eino")
|
||||
assertNativeRetryMapValue(t, gotData, "orchestration", "deep_agent")
|
||||
assertNativeRetryMapValue(t, gotData, "attempt", 2)
|
||||
assertNativeRetryMapValue(t, gotData, "reason", "")
|
||||
assertNativeRetryMapValue(t, gotData, "error", "HTTP 429 Too Many Requests")
|
||||
}
|
||||
|
||||
func TestEmitEinoNativeModelRetryProgressNilSafe(t *testing.T) {
|
||||
calledProgress := false
|
||||
called := emitEinoNativeModelRetryProgress("conv-1", "deep_agent", nil, func(string, string, interface{}) {
|
||||
calledProgress = true
|
||||
}, nil, nil)
|
||||
if called {
|
||||
t.Fatal("called = true, want false")
|
||||
}
|
||||
if calledProgress {
|
||||
t.Fatal("progress called for nil willRetry")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmitEinoNativeModelRetryProgressLogsEvent(t *testing.T) {
|
||||
core, logs := observer.New(zap.WarnLevel)
|
||||
logger := zap.New(core)
|
||||
willRetry := &adk.WillRetryError{
|
||||
ErrStr: "HTTP 500",
|
||||
RetryAttempt: 3,
|
||||
}
|
||||
|
||||
emitEinoNativeModelRetryProgress("conv-1", "single_agent", willRetry, nil, logger, willRetry)
|
||||
|
||||
entry := logs.FilterMessage("eino native model retry event").TakeAll()
|
||||
if len(entry) != 1 {
|
||||
t.Fatalf("log count = %d, want 1", len(entry))
|
||||
}
|
||||
fields := entry[0].ContextMap()
|
||||
if fields["orchestration"] != "single_agent" {
|
||||
t.Fatalf("orchestration field = %v", fields["orchestration"])
|
||||
}
|
||||
if fields["attempt"] != int64(3) {
|
||||
t.Fatalf("attempt field = %v", fields["attempt"])
|
||||
}
|
||||
}
|
||||
|
||||
func assertNativeRetryMapValue(t *testing.T, data map[string]interface{}, key string, want interface{}) {
|
||||
t.Helper()
|
||||
if got := data[key]; got != want {
|
||||
t.Fatalf("%s = %v, want %v", key, got, want)
|
||||
}
|
||||
}
|
||||
@@ -19,7 +19,7 @@ import (
|
||||
// PlanExecuteRootArgs 构建 Eino adk/prebuilt/planexecute 根 Agent 所需参数。
|
||||
type PlanExecuteRootArgs struct {
|
||||
MainToolCallingModel model.ToolCallingChatModel
|
||||
ExecModel model.ToolCallingChatModel
|
||||
AgenticExecModel model.AgenticModel
|
||||
OrchInstruction string
|
||||
ToolsCfg adk.ToolsConfig
|
||||
ExecMaxIter int
|
||||
@@ -34,17 +34,16 @@ type PlanExecuteRootArgs struct {
|
||||
Logger *zap.Logger
|
||||
// ModelName is used for model input token estimation logs.
|
||||
ModelName string
|
||||
// ExecPreMiddlewares 是由 prependEinoMiddlewares 构建的前置中间件(patchtoolcalls, reduction, toolsearch, plantask),
|
||||
// 与 Deep/Supervisor 主代理的 mainOrchestratorPre 一致。
|
||||
ExecPreMiddlewares []adk.ChatModelAgentMiddleware
|
||||
// SkillMiddleware 是 Eino 官方 skill 渐进式披露中间件(可选)。
|
||||
SkillMiddleware adk.ChatModelAgentMiddleware
|
||||
// FilesystemMiddleware 是 Eino filesystem 中间件,当 eino_skills.filesystem_tools 启用时提供本机文件读写与 Shell 能力(可选)。
|
||||
FilesystemMiddleware adk.ChatModelAgentMiddleware
|
||||
// AgenticExecPreMiddlewares 是由 prependEinoAgenticMiddlewares 构建的前置中间件。
|
||||
AgenticExecPreMiddlewares []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage]
|
||||
AgenticSkillMiddleware adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage]
|
||||
AgenticFilesystemMiddleware adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage]
|
||||
// PlannerReplannerRewriteHandlers applies BeforeModelRewriteState pipeline for planner/replanner input.
|
||||
PlannerReplannerRewriteHandlers []adk.ChatModelAgentMiddleware
|
||||
// ModelFacingTrace 可选:由 Executor Handlers 链末尾写入,供 last_react 与 summarization 后上下文对齐。
|
||||
ModelFacingTrace *modelFacingTraceHolder
|
||||
ModelFacingTrace *modelFacingTraceHolder
|
||||
AgenticModelRetryConfig *adk.TypedModelRetryConfig[*schema.AgenticMessage]
|
||||
AgenticModelFailoverConfig *adk.ModelFailoverConfig[*schema.AgenticMessage]
|
||||
}
|
||||
|
||||
// NewPlanExecuteRoot 返回 plan → execute → replan 预置编排根节点(与 Deep / Supervisor 并列)。
|
||||
@@ -52,7 +51,7 @@ func NewPlanExecuteRoot(ctx context.Context, a *PlanExecuteRootArgs) (adk.Resuma
|
||||
if a == nil {
|
||||
return nil, fmt.Errorf("plan_execute: args 为空")
|
||||
}
|
||||
if a.MainToolCallingModel == nil || a.ExecModel == nil {
|
||||
if a.MainToolCallingModel == nil || a.AgenticExecModel == nil {
|
||||
return nil, fmt.Errorf("plan_execute: 模型为空")
|
||||
}
|
||||
tcm, ok := interface{}(a.MainToolCallingModel).(model.ToolCallingChatModel)
|
||||
@@ -79,16 +78,16 @@ func NewPlanExecuteRoot(ctx context.Context, a *PlanExecuteRootArgs) (adk.Resuma
|
||||
return nil, fmt.Errorf("plan_execute replanner: %w", err)
|
||||
}
|
||||
|
||||
execHandlers, err := buildPlanExecuteExecutorHandlers(ctx, a)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
var executor adk.Agent
|
||||
agenticExecHandlers, herr := buildPlanExecuteAgenticExecutorHandlers(ctx, a)
|
||||
if herr != nil {
|
||||
return nil, herr
|
||||
}
|
||||
executor, err := newPlanExecuteExecutor(ctx, &planexecute.ExecutorConfig{
|
||||
Model: a.ExecModel,
|
||||
executor, err = newPlanExecuteAgenticExecutor(ctx, &planexecute.ExecutorConfig{
|
||||
ToolsConfig: a.ToolsCfg,
|
||||
MaxIterations: a.ExecMaxIter,
|
||||
GenInputFn: planExecuteExecutorGenInput(a.OrchInstruction, a.AppCfg, a.MwCfg, a.Logger, a.ModelName, a.ConversationID),
|
||||
}, execHandlers)
|
||||
}, a.AgenticExecModel, agenticExecHandlers, a.AgenticModelRetryConfig, a.AgenticModelFailoverConfig)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("plan_execute executor: %w", err)
|
||||
}
|
||||
@@ -104,37 +103,35 @@ func NewPlanExecuteRoot(ctx context.Context, a *PlanExecuteRootArgs) (adk.Resuma
|
||||
})
|
||||
}
|
||||
|
||||
// buildPlanExecuteExecutorHandlers 组装 Executor 中间件栈(outermost first),与 Deep/Supervisor 主代理对齐:
|
||||
// ExecPreMiddlewares(patch / reduction / toolsearch / plantask)→ filesystem → skill → summarization tail。
|
||||
func buildPlanExecuteExecutorHandlers(ctx context.Context, a *PlanExecuteRootArgs) ([]adk.ChatModelAgentMiddleware, error) {
|
||||
func buildPlanExecuteAgenticExecutorHandlers(ctx context.Context, a *PlanExecuteRootArgs) ([]adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], error) {
|
||||
if a == nil {
|
||||
return nil, fmt.Errorf("plan_execute: args 为空")
|
||||
}
|
||||
var execHandlers []adk.ChatModelAgentMiddleware
|
||||
if len(a.ExecPreMiddlewares) > 0 {
|
||||
execHandlers = append(execHandlers, a.ExecPreMiddlewares...)
|
||||
var execHandlers []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage]
|
||||
if len(a.AgenticExecPreMiddlewares) > 0 {
|
||||
execHandlers = append(execHandlers, a.AgenticExecPreMiddlewares...)
|
||||
}
|
||||
if a.FilesystemMiddleware != nil {
|
||||
execHandlers = append(execHandlers, a.FilesystemMiddleware)
|
||||
if a.AgenticFilesystemMiddleware != nil {
|
||||
execHandlers = append(execHandlers, a.AgenticFilesystemMiddleware)
|
||||
}
|
||||
if a.SkillMiddleware != nil {
|
||||
execHandlers = append(execHandlers, a.SkillMiddleware)
|
||||
if a.AgenticSkillMiddleware != nil {
|
||||
execHandlers = append(execHandlers, a.AgenticSkillMiddleware)
|
||||
}
|
||||
if a.AppCfg != nil {
|
||||
sumMw, sumErr := newEinoSummarizationMiddleware(ctx, a.ExecModel, a.AppCfg, a.MwCfg, a.ConversationID, a.DB, a.ProjectID, a.Logger)
|
||||
sumMw, sumErr := newEinoAgenticSummarizationMiddleware(ctx, a.AgenticExecModel, a.AppCfg, a.MwCfg, a.ConversationID, a.DB, a.ProjectID, a.Logger)
|
||||
if sumErr != nil {
|
||||
return nil, fmt.Errorf("plan_execute executor summarization: %w", sumErr)
|
||||
return nil, fmt.Errorf("plan_execute agentic executor summarization: %w", sumErr)
|
||||
}
|
||||
execHandlers = appendEinoChatModelTailMiddlewares(execHandlers, einoChatModelTailConfig{
|
||||
logger: a.Logger,
|
||||
phase: "plan_execute_executor",
|
||||
summarization: sumMw,
|
||||
modelName: a.ModelName,
|
||||
maxTotalTokens: a.AppCfg.OpenAI.MaxTotalTokens,
|
||||
toolMaxBytes: toolMaxBytesFromMW(a.MwCfg),
|
||||
conversationID: a.ConversationID,
|
||||
trace: a.ModelFacingTrace,
|
||||
middlewareConfig: a.MwCfg,
|
||||
execHandlers = appendEinoAgenticChatModelTailMiddlewares(execHandlers, einoChatModelTailConfig{
|
||||
logger: a.Logger,
|
||||
phase: "plan_execute_executor",
|
||||
agenticSummarization: sumMw,
|
||||
modelName: a.ModelName,
|
||||
maxTotalTokens: a.AppCfg.OpenAI.MaxTotalTokens,
|
||||
toolMaxBytes: toolMaxBytesFromMW(a.MwCfg),
|
||||
conversationID: a.ConversationID,
|
||||
trace: a.ModelFacingTrace,
|
||||
middlewareConfig: a.MwCfg,
|
||||
})
|
||||
}
|
||||
return execHandlers, nil
|
||||
|
||||
@@ -0,0 +1,125 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type einoPendingToolCalls struct {
|
||||
conversationID string
|
||||
progress func(eventType, message string, data interface{})
|
||||
|
||||
mu sync.Mutex
|
||||
byID map[string]toolCallPendingInfo
|
||||
queueByAgent map[string][]string
|
||||
}
|
||||
|
||||
func newEinoPendingToolCalls(conversationID string, progress func(eventType, message string, data interface{})) *einoPendingToolCalls {
|
||||
return &einoPendingToolCalls{
|
||||
conversationID: conversationID,
|
||||
progress: progress,
|
||||
byID: make(map[string]toolCallPendingInfo),
|
||||
queueByAgent: make(map[string][]string),
|
||||
}
|
||||
}
|
||||
|
||||
func (p *einoPendingToolCalls) Mark(tc toolCallPendingInfo) {
|
||||
if p == nil || strings.TrimSpace(tc.ToolCallID) == "" {
|
||||
return
|
||||
}
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
p.byID[tc.ToolCallID] = tc
|
||||
p.queueByAgent[tc.EinoAgent] = append(p.queueByAgent[tc.EinoAgent], tc.ToolCallID)
|
||||
}
|
||||
|
||||
func (p *einoPendingToolCalls) PopNextForAgent(agentName string) (toolCallPendingInfo, bool) {
|
||||
if p == nil {
|
||||
return toolCallPendingInfo{}, false
|
||||
}
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
q := p.queueByAgent[agentName]
|
||||
for len(q) > 0 {
|
||||
id := q[0]
|
||||
q = q[1:]
|
||||
p.queueByAgent[agentName] = q
|
||||
if tc, ok := p.byID[id]; ok {
|
||||
delete(p.byID, id)
|
||||
return tc, true
|
||||
}
|
||||
}
|
||||
return toolCallPendingInfo{}, false
|
||||
}
|
||||
|
||||
func (p *einoPendingToolCalls) RemoveByID(toolCallID string) {
|
||||
if p == nil || strings.TrimSpace(toolCallID) == "" {
|
||||
return
|
||||
}
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
delete(p.byID, toolCallID)
|
||||
}
|
||||
|
||||
func (p *einoPendingToolCalls) PopAny() (toolCallPendingInfo, bool) {
|
||||
if p == nil {
|
||||
return toolCallPendingInfo{}, false
|
||||
}
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
for id, tc := range p.byID {
|
||||
delete(p.byID, id)
|
||||
return tc, true
|
||||
}
|
||||
return toolCallPendingInfo{}, false
|
||||
}
|
||||
|
||||
func (p *einoPendingToolCalls) Count() int {
|
||||
if p == nil {
|
||||
return 0
|
||||
}
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
return len(p.byID)
|
||||
}
|
||||
|
||||
func (p *einoPendingToolCalls) FlushAsFailed(err error) {
|
||||
if p == nil {
|
||||
return
|
||||
}
|
||||
p.mu.Lock()
|
||||
pendingSnapshot := make([]toolCallPendingInfo, 0, len(p.byID))
|
||||
for _, tc := range p.byID {
|
||||
pendingSnapshot = append(pendingSnapshot, tc)
|
||||
}
|
||||
p.byID = make(map[string]toolCallPendingInfo)
|
||||
p.queueByAgent = make(map[string][]string)
|
||||
p.mu.Unlock()
|
||||
|
||||
if p.progress == nil {
|
||||
return
|
||||
}
|
||||
msg := ""
|
||||
if err != nil {
|
||||
msg = err.Error()
|
||||
}
|
||||
for _, tc := range pendingSnapshot {
|
||||
toolName := tc.ToolName
|
||||
if strings.TrimSpace(toolName) == "" {
|
||||
toolName = "unknown"
|
||||
}
|
||||
p.progress("tool_result", fmt.Sprintf("工具结果 (%s)", toolName), map[string]interface{}{
|
||||
"toolName": toolName,
|
||||
"success": false,
|
||||
"isError": true,
|
||||
"result": msg,
|
||||
"resultPreview": msg,
|
||||
"toolCallId": tc.ToolCallID,
|
||||
"conversationId": p.conversationID,
|
||||
"einoAgent": tc.EinoAgent,
|
||||
"einoRole": tc.EinoRole,
|
||||
"source": "eino",
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestEinoPendingToolCallsPopNextForAgentSkipsRemovedIDs(t *testing.T) {
|
||||
p := newEinoPendingToolCalls("conv", nil)
|
||||
p.Mark(toolCallPendingInfo{ToolCallID: "call-1", ToolName: "first", EinoAgent: "agent"})
|
||||
p.Mark(toolCallPendingInfo{ToolCallID: "call-2", ToolName: "second", EinoAgent: "agent"})
|
||||
p.RemoveByID("call-1")
|
||||
|
||||
got, ok := p.PopNextForAgent("agent")
|
||||
if !ok {
|
||||
t.Fatal("expected pending tool call")
|
||||
}
|
||||
if got.ToolCallID != "call-2" {
|
||||
t.Fatalf("toolCallID = %q, want call-2", got.ToolCallID)
|
||||
}
|
||||
if p.Count() != 0 {
|
||||
t.Fatalf("pending count = %d, want 0", p.Count())
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoPendingToolCallsPopAny(t *testing.T) {
|
||||
p := newEinoPendingToolCalls("conv", nil)
|
||||
p.Mark(toolCallPendingInfo{ToolCallID: "call-1", ToolName: "tool"})
|
||||
|
||||
got, ok := p.PopAny()
|
||||
if !ok || got.ToolCallID != "call-1" {
|
||||
t.Fatalf("PopAny = %#v ok=%v", got, ok)
|
||||
}
|
||||
if _, ok := p.PopAny(); ok {
|
||||
t.Fatal("PopAny should be empty after first pop")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoPendingToolCallsFlushAsFailedEmitsAndClears(t *testing.T) {
|
||||
var events []struct {
|
||||
eventType string
|
||||
message string
|
||||
data map[string]interface{}
|
||||
}
|
||||
p := newEinoPendingToolCalls("conv-1", func(eventType, message string, data interface{}) {
|
||||
m, _ := data.(map[string]interface{})
|
||||
events = append(events, struct {
|
||||
eventType string
|
||||
message string
|
||||
data map[string]interface{}
|
||||
}{eventType: eventType, message: message, data: m})
|
||||
})
|
||||
p.Mark(toolCallPendingInfo{
|
||||
ToolCallID: "call-err",
|
||||
ToolName: "",
|
||||
EinoAgent: "agent",
|
||||
EinoRole: "sub",
|
||||
})
|
||||
|
||||
p.FlushAsFailed(errors.New("boom"))
|
||||
|
||||
if p.Count() != 0 {
|
||||
t.Fatalf("pending count = %d, want 0", p.Count())
|
||||
}
|
||||
if len(events) != 1 {
|
||||
t.Fatalf("events = %#v, want one", events)
|
||||
}
|
||||
ev := events[0]
|
||||
if ev.eventType != "tool_result" || ev.message != "工具结果 (unknown)" {
|
||||
t.Fatalf("event = %#v", ev)
|
||||
}
|
||||
if ev.data["toolCallId"] != "call-err" ||
|
||||
ev.data["conversationId"] != "conv-1" ||
|
||||
ev.data["einoAgent"] != "agent" ||
|
||||
ev.data["einoRole"] != "sub" ||
|
||||
ev.data["isError"] != true ||
|
||||
ev.data["success"] != false ||
|
||||
ev.data["result"] != "boom" {
|
||||
t.Fatalf("payload = %#v", ev.data)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,111 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"cyberstrike-ai/internal/openai"
|
||||
)
|
||||
|
||||
type einoReasoningStreamEmitter struct {
|
||||
progress func(eventType, message string, data interface{})
|
||||
conversation string
|
||||
orchMode string
|
||||
agentName string
|
||||
einoRole string
|
||||
nextStreamID func() string
|
||||
|
||||
streamID string
|
||||
rawBuf string
|
||||
displayPrev string
|
||||
}
|
||||
|
||||
func newEinoReasoningStreamEmitter(
|
||||
conversationID, orchMode, agentName, einoRole string,
|
||||
progress func(eventType, message string, data interface{}),
|
||||
nextStreamID func() string,
|
||||
) *einoReasoningStreamEmitter {
|
||||
return &einoReasoningStreamEmitter{
|
||||
progress: progress,
|
||||
conversation: conversationID,
|
||||
orchMode: orchMode,
|
||||
agentName: agentName,
|
||||
einoRole: einoRole,
|
||||
nextStreamID: nextStreamID,
|
||||
}
|
||||
}
|
||||
|
||||
func (e *einoReasoningStreamEmitter) EmitDelta(reasoningContent string) bool {
|
||||
if e == nil || strings.TrimSpace(reasoningContent) == "" {
|
||||
return false
|
||||
}
|
||||
var rawDelta string
|
||||
e.rawBuf, rawDelta = normalizeStreamingDelta(e.rawBuf, reasoningContent)
|
||||
if rawDelta == "" || e.progress == nil {
|
||||
return false
|
||||
}
|
||||
fullDisplay := openai.DisplayReasoningContent(e.rawBuf)
|
||||
displayDelta := fullDisplay
|
||||
if strings.HasPrefix(fullDisplay, e.displayPrev) {
|
||||
displayDelta = fullDisplay[len(e.displayPrev):]
|
||||
}
|
||||
e.displayPrev = fullDisplay
|
||||
if displayDelta == "" {
|
||||
return false
|
||||
}
|
||||
if e.streamID == "" {
|
||||
if e.nextStreamID != nil {
|
||||
e.streamID = e.nextStreamID()
|
||||
}
|
||||
if e.streamID == "" {
|
||||
e.streamID = "eino-reasoning"
|
||||
}
|
||||
e.progress("reasoning_chain_stream_start", " ", map[string]interface{}{
|
||||
"streamId": e.streamID,
|
||||
"source": "eino",
|
||||
"einoAgent": e.agentName,
|
||||
"einoRole": e.einoRole,
|
||||
"orchestration": e.orchMode,
|
||||
})
|
||||
}
|
||||
e.progress("reasoning_chain_stream_delta", displayDelta, openai.WithSSEAccumulated(map[string]interface{}{
|
||||
"streamId": e.streamID,
|
||||
}, fullDisplay))
|
||||
return true
|
||||
}
|
||||
|
||||
func (e *einoReasoningStreamEmitter) Finish() string {
|
||||
if e == nil {
|
||||
return ""
|
||||
}
|
||||
display := openai.DisplayReasoningContent(strings.TrimSpace(e.rawBuf))
|
||||
if display == "" || e.streamID == "" || e.progress == nil {
|
||||
return display
|
||||
}
|
||||
e.progress("reasoning_chain_stream_end", display, map[string]interface{}{
|
||||
"streamId": e.streamID,
|
||||
"conversationId": e.conversation,
|
||||
"source": "eino",
|
||||
"einoAgent": e.agentName,
|
||||
"einoRole": e.einoRole,
|
||||
"orchestration": e.orchMode,
|
||||
})
|
||||
return display
|
||||
}
|
||||
|
||||
func (e *einoReasoningStreamEmitter) EmitComplete(reasoningContent string) bool {
|
||||
if e == nil || e.progress == nil {
|
||||
return false
|
||||
}
|
||||
display := openai.DisplayReasoningContent(strings.TrimSpace(reasoningContent))
|
||||
if display == "" {
|
||||
return false
|
||||
}
|
||||
e.progress("reasoning_chain", display, map[string]interface{}{
|
||||
"conversationId": e.conversation,
|
||||
"source": "eino",
|
||||
"einoAgent": e.agentName,
|
||||
"einoRole": e.einoRole,
|
||||
"orchestration": e.orchMode,
|
||||
})
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"cyberstrike-ai/internal/openai"
|
||||
)
|
||||
|
||||
func TestEinoReasoningStreamEmitterStreamingLifecycle(t *testing.T) {
|
||||
type progressEvent struct {
|
||||
eventType string
|
||||
message string
|
||||
data map[string]interface{}
|
||||
}
|
||||
var events []progressEvent
|
||||
progress := func(eventType, message string, data interface{}) {
|
||||
m, _ := data.(map[string]interface{})
|
||||
events = append(events, progressEvent{eventType: eventType, message: message, data: m})
|
||||
}
|
||||
emitter := newEinoReasoningStreamEmitter("conv-1", "deep", "lead", "orchestrator", progress, func() string {
|
||||
return "reasoning-1"
|
||||
})
|
||||
|
||||
if !emitter.EmitDelta("he") {
|
||||
t.Fatal("first reasoning delta should emit")
|
||||
}
|
||||
if !emitter.EmitDelta("hello") {
|
||||
t.Fatal("cumulative reasoning chunk should emit tail")
|
||||
}
|
||||
if got := emitter.Finish(); got != "hello" {
|
||||
t.Fatalf("finish body = %q, want hello", got)
|
||||
}
|
||||
|
||||
if len(events) != 4 {
|
||||
t.Fatalf("events = %#v, want start + 2 deltas + end", events)
|
||||
}
|
||||
if events[0].eventType != "reasoning_chain_stream_start" || events[0].message != " " {
|
||||
t.Fatalf("event[0] = %#v", events[0])
|
||||
}
|
||||
if events[1].eventType != "reasoning_chain_stream_delta" || events[1].message != "he" {
|
||||
t.Fatalf("event[1] = %#v", events[1])
|
||||
}
|
||||
if events[2].eventType != "reasoning_chain_stream_delta" || events[2].message != "llo" {
|
||||
t.Fatalf("event[2] = %#v", events[2])
|
||||
}
|
||||
if got := events[2].data[openai.SSEAccumulatedKey]; got != "hello" {
|
||||
t.Fatalf("accumulated = %#v, want hello", got)
|
||||
}
|
||||
if events[3].eventType != "reasoning_chain_stream_end" || events[3].message != "hello" {
|
||||
t.Fatalf("event[3] = %#v", events[3])
|
||||
}
|
||||
if got := events[3].data["einoRole"]; got != "orchestrator" {
|
||||
t.Fatalf("einoRole = %#v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoReasoningStreamEmitterComplete(t *testing.T) {
|
||||
var eventType, message string
|
||||
var data map[string]interface{}
|
||||
progress := func(et, msg string, raw interface{}) {
|
||||
eventType = et
|
||||
message = msg
|
||||
data, _ = raw.(map[string]interface{})
|
||||
}
|
||||
|
||||
ok := newEinoReasoningStreamEmitter("conv-1", "supervisor", "worker", "sub", progress, nil).EmitComplete(" thought ")
|
||||
if !ok {
|
||||
t.Fatal("complete reasoning should emit")
|
||||
}
|
||||
if eventType != "reasoning_chain" || message != "thought" {
|
||||
t.Fatalf("event = %s %q", eventType, message)
|
||||
}
|
||||
if data["conversationId"] != "conv-1" || data["einoAgent"] != "worker" || data["einoRole"] != "sub" {
|
||||
t.Fatalf("bad event data: %#v", data)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoReasoningStreamEmitterNoProgressStillBuffers(t *testing.T) {
|
||||
emitter := newEinoReasoningStreamEmitter("conv", "deep", "lead", "orchestrator", nil, nil)
|
||||
if emitter.EmitDelta("hello") {
|
||||
t.Fatal("nil progress should not emit")
|
||||
}
|
||||
if got := emitter.Finish(); got != "hello" {
|
||||
t.Fatalf("finish body = %q, want hello", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
package multiagent
|
||||
|
||||
import "context"
|
||||
|
||||
type einoRunCancellationHandler struct {
|
||||
ctx context.Context
|
||||
conversationID string
|
||||
progress func(eventType, message string, data interface{})
|
||||
pending *einoPendingToolCalls
|
||||
takePartial einoPartialResultFunc
|
||||
}
|
||||
|
||||
type einoRunCancellationHandlerConfig struct {
|
||||
Context context.Context
|
||||
ConversationID string
|
||||
Progress func(eventType, message string, data interface{})
|
||||
Pending *einoPendingToolCalls
|
||||
TakePartial einoPartialResultFunc
|
||||
}
|
||||
|
||||
func newEinoRunCancellationHandler(cfg einoRunCancellationHandlerConfig) *einoRunCancellationHandler {
|
||||
return &einoRunCancellationHandler{
|
||||
ctx: cfg.Context,
|
||||
conversationID: cfg.ConversationID,
|
||||
progress: cfg.Progress,
|
||||
pending: cfg.Pending,
|
||||
takePartial: cfg.TakePartial,
|
||||
}
|
||||
}
|
||||
|
||||
func (h *einoRunCancellationHandler) Handle(runErr error) (*RunResult, error) {
|
||||
if h == nil {
|
||||
return nil, runErr
|
||||
}
|
||||
if h.pending != nil {
|
||||
h.pending.FlushAsFailed(runErr)
|
||||
}
|
||||
if h.progress != nil {
|
||||
if isInterruptContinue(h.ctx) {
|
||||
h.progress("progress", "已暂停当前输出,正在合并用户补充并继续…", map[string]interface{}{
|
||||
"conversationId": h.conversationID,
|
||||
"source": "eino",
|
||||
"kind": "interrupt_continue",
|
||||
})
|
||||
} else if runErr != nil {
|
||||
h.progress("error", runErr.Error(), map[string]interface{}{
|
||||
"conversationId": h.conversationID,
|
||||
"source": "eino",
|
||||
})
|
||||
}
|
||||
}
|
||||
if h.takePartial == nil {
|
||||
return nil, runErr
|
||||
}
|
||||
return h.takePartial(runErr)
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestEinoRunCancellationHandlerFlushesPendingAndEmitsError(t *testing.T) {
|
||||
runErr := errors.New("context canceled")
|
||||
var events []struct {
|
||||
eventType string
|
||||
data map[string]interface{}
|
||||
}
|
||||
progress := func(eventType, _ string, data interface{}) {
|
||||
m, _ := data.(map[string]interface{})
|
||||
events = append(events, struct {
|
||||
eventType string
|
||||
data map[string]interface{}
|
||||
}{eventType: eventType, data: m})
|
||||
}
|
||||
pending := newEinoPendingToolCalls("conv-1", progress)
|
||||
pending.Mark(toolCallPendingInfo{ToolCallID: "call-1", ToolName: "execute", EinoAgent: "lead", EinoRole: "orchestrator"})
|
||||
want := &RunResult{Response: "partial"}
|
||||
|
||||
result, err := newEinoRunCancellationHandler(einoRunCancellationHandlerConfig{
|
||||
Context: context.Background(),
|
||||
ConversationID: "conv-1",
|
||||
Progress: progress,
|
||||
Pending: pending,
|
||||
TakePartial: func(got error) (*RunResult, error) {
|
||||
if !errors.Is(got, runErr) {
|
||||
t.Fatalf("partial err = %v", got)
|
||||
}
|
||||
return want, got
|
||||
},
|
||||
}).Handle(runErr)
|
||||
|
||||
if result != want || !errors.Is(err, runErr) {
|
||||
t.Fatalf("result=%#v err=%v", result, err)
|
||||
}
|
||||
if pending.Count() != 0 {
|
||||
t.Fatalf("pending count = %d, want 0", pending.Count())
|
||||
}
|
||||
var sawError, sawFailedTool bool
|
||||
for _, ev := range events {
|
||||
if ev.eventType == "error" {
|
||||
sawError = ev.data["conversationId"] == "conv-1" && ev.data["source"] == "eino"
|
||||
}
|
||||
if ev.eventType == "tool_result" {
|
||||
sawFailedTool = ev.data["toolCallId"] == "call-1" && ev.data["isError"] == true
|
||||
}
|
||||
}
|
||||
if !sawError || !sawFailedTool {
|
||||
t.Fatalf("events = %#v", events)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoRunCancellationHandlerInterruptContinueProgress(t *testing.T) {
|
||||
ctx, cancel := context.WithCancelCause(context.Background())
|
||||
cancel(ErrInterruptContinue)
|
||||
runErr := context.Canceled
|
||||
var eventType string
|
||||
var data map[string]interface{}
|
||||
|
||||
_, err := newEinoRunCancellationHandler(einoRunCancellationHandlerConfig{
|
||||
Context: ctx,
|
||||
ConversationID: "conv-1",
|
||||
Progress: func(et, _ string, raw interface{}) {
|
||||
eventType = et
|
||||
data, _ = raw.(map[string]interface{})
|
||||
},
|
||||
TakePartial: func(got error) (*RunResult, error) {
|
||||
return nil, got
|
||||
},
|
||||
}).Handle(runErr)
|
||||
|
||||
if !errors.Is(err, runErr) {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
if eventType != "progress" || data["kind"] != "interrupt_continue" || data["conversationId"] != "conv-1" {
|
||||
t.Fatalf("eventType=%q data=%#v", eventType, data)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoRunCancellationHandlerNilSafe(t *testing.T) {
|
||||
runErr := errors.New("boom")
|
||||
var h *einoRunCancellationHandler
|
||||
result, err := h.Handle(runErr)
|
||||
if result != nil || !errors.Is(err, runErr) {
|
||||
t.Fatalf("result=%#v err=%v", result, err)
|
||||
}
|
||||
result, err = newEinoRunCancellationHandler(einoRunCancellationHandlerConfig{}).Handle(runErr)
|
||||
if result != nil || !errors.Is(err, runErr) {
|
||||
t.Fatalf("result=%#v err=%v", result, err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
type einoRunCompletionHandler struct {
|
||||
conversationID string
|
||||
orchMode string
|
||||
progress func(eventType, message string, data interface{})
|
||||
logger *zap.Logger
|
||||
|
||||
pending *einoPendingToolCalls
|
||||
cpStore *fileCheckPointStore
|
||||
checkPointID string
|
||||
}
|
||||
|
||||
type einoRunCompletionHandlerConfig struct {
|
||||
ConversationID string
|
||||
OrchMode string
|
||||
Progress func(eventType, message string, data interface{})
|
||||
Logger *zap.Logger
|
||||
Pending *einoPendingToolCalls
|
||||
Checkpoint *fileCheckPointStore
|
||||
CheckpointID string
|
||||
}
|
||||
|
||||
func newEinoRunCompletionHandler(cfg einoRunCompletionHandlerConfig) *einoRunCompletionHandler {
|
||||
return &einoRunCompletionHandler{
|
||||
conversationID: cfg.ConversationID,
|
||||
orchMode: cfg.OrchMode,
|
||||
progress: cfg.Progress,
|
||||
logger: cfg.Logger,
|
||||
pending: cfg.Pending,
|
||||
cpStore: cfg.Checkpoint,
|
||||
checkPointID: cfg.CheckpointID,
|
||||
}
|
||||
}
|
||||
|
||||
func (h *einoRunCompletionHandler) Complete() {
|
||||
if h == nil {
|
||||
return
|
||||
}
|
||||
h.flushOrphanedPending()
|
||||
h.cleanupCheckpoint()
|
||||
}
|
||||
|
||||
func (h *einoRunCompletionHandler) flushOrphanedPending() {
|
||||
if h.pending == nil {
|
||||
return
|
||||
}
|
||||
orphanCount := h.pending.Count()
|
||||
if orphanCount <= 0 {
|
||||
return
|
||||
}
|
||||
h.pending.FlushAsFailed(errors.New("pending tool call missing result before run completion"))
|
||||
if h.progress != nil {
|
||||
h.progress("eino_pending_orphaned", "pending tool calls were force-closed at run end", map[string]interface{}{
|
||||
"conversationId": h.conversationID,
|
||||
"source": "eino",
|
||||
"orchestration": h.orchMode,
|
||||
"pendingCount": orphanCount,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func (h *einoRunCompletionHandler) cleanupCheckpoint() {
|
||||
if h.cpStore == nil || h.checkPointID == "" {
|
||||
return
|
||||
}
|
||||
p, err := h.cpStore.path(h.checkPointID)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if rmErr := os.Remove(p); rmErr != nil && !os.IsNotExist(rmErr) && h.logger != nil {
|
||||
h.logger.Warn("eino checkpoint cleanup failed", zap.String("path", p), zap.Error(rmErr))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestEinoRunCompletionHandlerFlushesOrphansAndCleansCheckpoint(t *testing.T) {
|
||||
var events []struct {
|
||||
eventType string
|
||||
data map[string]interface{}
|
||||
}
|
||||
progress := func(eventType, _ string, data interface{}) {
|
||||
m, _ := data.(map[string]interface{})
|
||||
events = append(events, struct {
|
||||
eventType string
|
||||
data map[string]interface{}
|
||||
}{eventType: eventType, data: m})
|
||||
}
|
||||
pending := newEinoPendingToolCalls("conv-1", progress)
|
||||
pending.Mark(toolCallPendingInfo{ToolCallID: "call-1", ToolName: "execute", EinoAgent: "lead", EinoRole: "orchestrator"})
|
||||
store, err := newFileCheckPointStore(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := store.Set(context.Background(), "cp-1", []byte("checkpoint")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cpPath, err := store.path("cp-1")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
newEinoRunCompletionHandler(einoRunCompletionHandlerConfig{
|
||||
ConversationID: "conv-1",
|
||||
OrchMode: "deep",
|
||||
Progress: progress,
|
||||
Pending: pending,
|
||||
Checkpoint: store,
|
||||
CheckpointID: "cp-1",
|
||||
}).Complete()
|
||||
|
||||
if pending.Count() != 0 {
|
||||
t.Fatalf("pending count = %d, want 0", pending.Count())
|
||||
}
|
||||
if _, err := os.Stat(cpPath); !os.IsNotExist(err) {
|
||||
t.Fatalf("checkpoint should be removed, stat err=%v", err)
|
||||
}
|
||||
var orphanEvent map[string]interface{}
|
||||
var failedToolResult map[string]interface{}
|
||||
for _, ev := range events {
|
||||
switch ev.eventType {
|
||||
case "eino_pending_orphaned":
|
||||
orphanEvent = ev.data
|
||||
case "tool_result":
|
||||
failedToolResult = ev.data
|
||||
}
|
||||
}
|
||||
if orphanEvent == nil || orphanEvent["conversationId"] != "conv-1" || orphanEvent["orchestration"] != "deep" || orphanEvent["pendingCount"] != 1 {
|
||||
t.Fatalf("orphan event = %#v", orphanEvent)
|
||||
}
|
||||
if failedToolResult == nil || failedToolResult["toolCallId"] != "call-1" || failedToolResult["isError"] != true {
|
||||
t.Fatalf("failed tool result = %#v", failedToolResult)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoRunCompletionHandlerNoopWithoutState(t *testing.T) {
|
||||
newEinoRunCompletionHandler(einoRunCompletionHandlerConfig{}).Complete()
|
||||
var h *einoRunCompletionHandler
|
||||
h.Complete()
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"github.com/cloudwego/eino/adk"
|
||||
)
|
||||
|
||||
type einoRunErrorHandler struct {
|
||||
conversationID string
|
||||
orchMode string
|
||||
progress func(eventType, message string, data interface{})
|
||||
pending *einoPendingToolCalls
|
||||
nativeCancelFallback func() error
|
||||
}
|
||||
|
||||
type einoRunErrorHandlerConfig struct {
|
||||
ConversationID string
|
||||
OrchMode string
|
||||
Progress func(eventType, message string, data interface{})
|
||||
Pending *einoPendingToolCalls
|
||||
NativeCancelFallback func() error
|
||||
}
|
||||
|
||||
func newEinoRunErrorHandler(cfg einoRunErrorHandlerConfig) *einoRunErrorHandler {
|
||||
return &einoRunErrorHandler{
|
||||
conversationID: cfg.ConversationID,
|
||||
orchMode: cfg.OrchMode,
|
||||
progress: cfg.Progress,
|
||||
pending: cfg.Pending,
|
||||
nativeCancelFallback: cfg.NativeCancelFallback,
|
||||
}
|
||||
}
|
||||
|
||||
func (h *einoRunErrorHandler) Handle(runErr error) error {
|
||||
if h == nil || runErr == nil {
|
||||
return runErr
|
||||
}
|
||||
var cancelErr *adk.CancelError
|
||||
if errors.As(runErr, &cancelErr) {
|
||||
h.flushPending(runErr)
|
||||
if h.nativeCancelFallback != nil {
|
||||
return h.nativeCancelFallback()
|
||||
}
|
||||
return context.Canceled
|
||||
}
|
||||
if errors.Is(runErr, context.DeadlineExceeded) {
|
||||
h.flushPending(runErr)
|
||||
h.emitError(runErr, "timeout")
|
||||
return runErr
|
||||
}
|
||||
if errors.Is(runErr, context.Canceled) {
|
||||
h.flushPending(runErr)
|
||||
h.emitError(runErr, "")
|
||||
return runErr
|
||||
}
|
||||
if isEinoIterationLimitError(runErr) {
|
||||
h.flushPending(runErr)
|
||||
if h.progress != nil {
|
||||
h.progress("iteration_limit_reached", runErr.Error(), map[string]interface{}{
|
||||
"conversationId": h.conversationID,
|
||||
"source": "eino",
|
||||
"orchestration": h.orchMode,
|
||||
})
|
||||
}
|
||||
h.emitError(runErr, "iteration_limit")
|
||||
return runErr
|
||||
}
|
||||
h.flushPending(runErr)
|
||||
h.emitError(runErr, "")
|
||||
return runErr
|
||||
}
|
||||
|
||||
func (h *einoRunErrorHandler) flushPending(err error) {
|
||||
if h != nil && h.pending != nil {
|
||||
h.pending.FlushAsFailed(err)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *einoRunErrorHandler) emitError(err error, kind string) {
|
||||
if h == nil || h.progress == nil || err == nil {
|
||||
return
|
||||
}
|
||||
data := map[string]interface{}{
|
||||
"conversationId": h.conversationID,
|
||||
"source": "eino",
|
||||
}
|
||||
if kind != "" {
|
||||
data["errorKind"] = kind
|
||||
}
|
||||
h.progress("error", err.Error(), data)
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/cloudwego/eino/adk"
|
||||
)
|
||||
|
||||
func TestEinoRunErrorHandlerCancelUsesNativeFallback(t *testing.T) {
|
||||
pending := newEinoPendingToolCalls("conv-1", nil)
|
||||
pending.Mark(toolCallPendingInfo{ToolCallID: "call-1", ToolName: "execute"})
|
||||
want := errors.New("native cancel")
|
||||
|
||||
got := newEinoRunErrorHandler(einoRunErrorHandlerConfig{
|
||||
ConversationID: "conv-1",
|
||||
Pending: pending,
|
||||
NativeCancelFallback: func() error {
|
||||
return want
|
||||
},
|
||||
}).Handle(&adk.CancelError{Info: &adk.AgentCancelInfo{}})
|
||||
|
||||
if !errors.Is(got, want) {
|
||||
t.Fatalf("err = %v, want native fallback", got)
|
||||
}
|
||||
if pending.Count() != 0 {
|
||||
t.Fatalf("pending count = %d, want 0", pending.Count())
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoRunErrorHandlerTimeoutAndGeneralErrorProgress(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
err error
|
||||
errorKind interface{}
|
||||
}{
|
||||
{name: "timeout", err: context.DeadlineExceeded, errorKind: "timeout"},
|
||||
{name: "general", err: errors.New("boom"), errorKind: nil},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
var data map[string]interface{}
|
||||
got := newEinoRunErrorHandler(einoRunErrorHandlerConfig{
|
||||
ConversationID: "conv-1",
|
||||
Progress: func(eventType, _ string, raw interface{}) {
|
||||
if eventType == "error" {
|
||||
data, _ = raw.(map[string]interface{})
|
||||
}
|
||||
},
|
||||
}).Handle(tc.err)
|
||||
if !errors.Is(got, tc.err) {
|
||||
t.Fatalf("err = %v", got)
|
||||
}
|
||||
if data["conversationId"] != "conv-1" || data["source"] != "eino" {
|
||||
t.Fatalf("data = %#v", data)
|
||||
}
|
||||
if gotKind := data["errorKind"]; gotKind != tc.errorKind {
|
||||
t.Fatalf("errorKind = %#v, want %#v", gotKind, tc.errorKind)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoRunErrorHandlerIterationLimitProgress(t *testing.T) {
|
||||
var events []string
|
||||
var errorKind interface{}
|
||||
err := errors.New("maximum iteration reached")
|
||||
|
||||
got := newEinoRunErrorHandler(einoRunErrorHandlerConfig{
|
||||
ConversationID: "conv-1",
|
||||
OrchMode: "deep",
|
||||
Progress: func(eventType, _ string, raw interface{}) {
|
||||
events = append(events, eventType)
|
||||
if eventType == "error" {
|
||||
data, _ := raw.(map[string]interface{})
|
||||
errorKind = data["errorKind"]
|
||||
}
|
||||
},
|
||||
}).Handle(err)
|
||||
|
||||
if !errors.Is(got, err) {
|
||||
t.Fatalf("err = %v", got)
|
||||
}
|
||||
if len(events) != 2 || events[0] != "iteration_limit_reached" || events[1] != "error" {
|
||||
t.Fatalf("events = %#v", events)
|
||||
}
|
||||
if errorKind != "iteration_limit" {
|
||||
t.Fatalf("errorKind = %#v", errorKind)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoRunErrorHandlerNilSafe(t *testing.T) {
|
||||
var h *einoRunErrorHandler
|
||||
if h.Handle(nil) != nil {
|
||||
t.Fatal("nil handler nil err should return nil")
|
||||
}
|
||||
err := errors.New("boom")
|
||||
if got := h.Handle(err); !errors.Is(got, err) {
|
||||
t.Fatalf("nil handler err = %v", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,244 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync/atomic"
|
||||
|
||||
"cyberstrike-ai/internal/agent"
|
||||
"cyberstrike-ai/internal/config"
|
||||
"cyberstrike-ai/internal/einomcp"
|
||||
|
||||
"github.com/cloudwego/eino/adk"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
type einoRunEventDrainConfig struct {
|
||||
Context context.Context
|
||||
ConversationID string
|
||||
OrchMode string
|
||||
OrchestratorName string
|
||||
Progress func(eventType, message string, data interface{})
|
||||
Logger *zap.Logger
|
||||
BaseMessages []adk.Message
|
||||
SnapshotMCPIDs func() []string
|
||||
StreamsMainAssistant func(agent string) bool
|
||||
EinoRoleTag func(agent string) string
|
||||
MiddlewareConfig *config.MultiAgentEinoMiddlewareConfig
|
||||
|
||||
FilesystemMonitorAgent *agent.Agent
|
||||
FilesystemMonitorRecord einomcp.ExecutionRecorder
|
||||
MCPExecutionBinder *MCPExecutionBinder
|
||||
}
|
||||
|
||||
type einoRunEventDrain struct {
|
||||
cfg einoRunEventDrainConfig
|
||||
|
||||
runMessages *einoRunMessageAccumulator
|
||||
assistantOutput *einoAssistantOutputAccumulator
|
||||
runProgress *einoRunProgressTracker
|
||||
pendingToolCalls *einoPendingToolCalls
|
||||
stdoutSuppressor *einoExecuteStdoutSuppressor
|
||||
toolResultEmitter *einoToolResultProgressEmitter
|
||||
usage *einoRunUsageAccumulator
|
||||
|
||||
reasoningStreamSeq int64
|
||||
subReplyStreamSeq int64
|
||||
mainResponseStreamSeq int64
|
||||
|
||||
toolResultHandler *einoToolResultEventHandler
|
||||
assistantStreamHandler *einoAssistantStreamEventHandler
|
||||
materializedMessageHandler *einoMaterializedMessageEventHandler
|
||||
}
|
||||
|
||||
func newEinoRunEventDrain(cfg einoRunEventDrainConfig) *einoRunEventDrain {
|
||||
if cfg.Context == nil {
|
||||
cfg.Context = context.Background()
|
||||
}
|
||||
if cfg.SnapshotMCPIDs == nil {
|
||||
cfg.SnapshotMCPIDs = func() []string { return nil }
|
||||
}
|
||||
if cfg.StreamsMainAssistant == nil {
|
||||
cfg.StreamsMainAssistant = func(agentName string) bool {
|
||||
return agentName == "" || agentName == cfg.OrchestratorName
|
||||
}
|
||||
}
|
||||
if cfg.EinoRoleTag == nil {
|
||||
cfg.EinoRoleTag = func(agentName string) string {
|
||||
if cfg.StreamsMainAssistant(agentName) {
|
||||
return "orchestrator"
|
||||
}
|
||||
return "sub"
|
||||
}
|
||||
}
|
||||
|
||||
runMessages := newEinoRunMessageAccumulator(cfg.BaseMessages)
|
||||
assistantOutput := newEinoAssistantOutputAccumulator(cfg.OrchMode)
|
||||
runProgress := newEinoRunProgressTracker(
|
||||
cfg.OrchMode,
|
||||
cfg.OrchestratorName,
|
||||
cfg.ConversationID,
|
||||
cfg.Progress,
|
||||
cfg.StreamsMainAssistant,
|
||||
cfg.EinoRoleTag,
|
||||
)
|
||||
pendingToolCalls := newEinoPendingToolCalls(cfg.ConversationID, cfg.Progress)
|
||||
stdoutSuppressor := newEinoExecuteStdoutSuppressor()
|
||||
usage := newEinoRunUsageAccumulator()
|
||||
toolResultEmitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{
|
||||
ConversationID: cfg.ConversationID,
|
||||
OrchestratorName: cfg.OrchestratorName,
|
||||
Progress: cfg.Progress,
|
||||
EinoRoleTag: cfg.EinoRoleTag,
|
||||
Pending: pendingToolCalls,
|
||||
ExecuteStdoutDup: stdoutSuppressor,
|
||||
RunMessages: runMessages,
|
||||
FilesystemMonitorAgent: cfg.FilesystemMonitorAgent,
|
||||
FilesystemMonitorRecord: cfg.FilesystemMonitorRecord,
|
||||
MCPExecutionBinder: cfg.MCPExecutionBinder,
|
||||
})
|
||||
|
||||
return &einoRunEventDrain{
|
||||
cfg: cfg,
|
||||
runMessages: runMessages,
|
||||
assistantOutput: assistantOutput,
|
||||
runProgress: runProgress,
|
||||
pendingToolCalls: pendingToolCalls,
|
||||
stdoutSuppressor: stdoutSuppressor,
|
||||
toolResultEmitter: toolResultEmitter,
|
||||
usage: usage,
|
||||
}
|
||||
}
|
||||
|
||||
func (d *einoRunEventDrain) BindHandlers(confirmRecovery func()) {
|
||||
if d == nil {
|
||||
return
|
||||
}
|
||||
d.toolResultHandler = newEinoToolResultEventHandler(einoToolResultEventHandlerConfig{
|
||||
Context: d.cfg.Context,
|
||||
Logger: d.cfg.Logger,
|
||||
RunMessages: d.runMessages,
|
||||
Emitter: d.toolResultEmitter,
|
||||
ConfirmRecovery: confirmRecovery,
|
||||
})
|
||||
streamToolCallCompletion := newEinoStreamToolCallCompletionHandler(einoStreamToolCallCompletionHandlerConfig{
|
||||
ConversationID: d.cfg.ConversationID,
|
||||
OrchMode: d.cfg.OrchMode,
|
||||
Progress: d.cfg.Progress,
|
||||
RunProgress: d.runProgress,
|
||||
RunMessages: d.runMessages,
|
||||
MarkPending: d.markPendingWithMonitor,
|
||||
})
|
||||
d.assistantStreamHandler = newEinoAssistantStreamEventHandler(einoAssistantStreamEventHandlerConfig{
|
||||
Context: d.cfg.Context,
|
||||
ConversationID: d.cfg.ConversationID,
|
||||
OrchMode: d.cfg.OrchMode,
|
||||
Progress: d.cfg.Progress,
|
||||
Logger: d.cfg.Logger,
|
||||
SnapshotMCPIDs: d.cfg.SnapshotMCPIDs,
|
||||
StreamsMainAssistant: d.cfg.StreamsMainAssistant,
|
||||
EinoRoleTag: d.cfg.EinoRoleTag,
|
||||
RunProgress: d.runProgress,
|
||||
StdoutSuppressor: d.stdoutSuppressor,
|
||||
AssistantOutput: d.assistantOutput,
|
||||
RunMessages: d.runMessages,
|
||||
Usage: d.usage,
|
||||
ToolCallCompletion: streamToolCallCompletion,
|
||||
NextMainStreamID: d.nextMainStreamID,
|
||||
NextReasoningStreamID: d.nextReasoningStreamID,
|
||||
NextSubAgentReplyStreamID: d.nextSubAgentReplyStreamID,
|
||||
})
|
||||
d.materializedMessageHandler = newEinoMaterializedMessageEventHandler(einoMaterializedMessageEventHandlerConfig{
|
||||
ConversationID: d.cfg.ConversationID,
|
||||
OrchMode: d.cfg.OrchMode,
|
||||
Progress: d.cfg.Progress,
|
||||
SnapshotMCPIDs: d.cfg.SnapshotMCPIDs,
|
||||
StreamsMainAssistant: d.cfg.StreamsMainAssistant,
|
||||
EinoRoleTag: d.cfg.EinoRoleTag,
|
||||
RunProgress: d.runProgress,
|
||||
StdoutSuppressor: d.stdoutSuppressor,
|
||||
AssistantOutput: d.assistantOutput,
|
||||
RunMessages: d.runMessages,
|
||||
Usage: d.usage,
|
||||
ToolResultHandler: d.toolResultHandler,
|
||||
MarkPending: d.markPendingWithMonitor,
|
||||
NextMainStreamID: d.nextMainStreamID,
|
||||
})
|
||||
}
|
||||
|
||||
func (d *einoRunEventDrain) RunMessages() *einoRunMessageAccumulator {
|
||||
if d == nil {
|
||||
return nil
|
||||
}
|
||||
return d.runMessages
|
||||
}
|
||||
|
||||
func (d *einoRunEventDrain) AssistantOutput() *einoAssistantOutputAccumulator {
|
||||
if d == nil {
|
||||
return nil
|
||||
}
|
||||
return d.assistantOutput
|
||||
}
|
||||
|
||||
func (d *einoRunEventDrain) PendingToolCalls() *einoPendingToolCalls {
|
||||
if d == nil {
|
||||
return nil
|
||||
}
|
||||
return d.pendingToolCalls
|
||||
}
|
||||
|
||||
func (d *einoRunEventDrain) Usage() *einoRunUsageAccumulator {
|
||||
if d == nil {
|
||||
return nil
|
||||
}
|
||||
return d.usage
|
||||
}
|
||||
|
||||
func (d *einoRunEventDrain) ObserveAgent(agentName string) {
|
||||
if d == nil || d.runProgress == nil {
|
||||
return
|
||||
}
|
||||
d.runProgress.ObserveAgent(agentName)
|
||||
}
|
||||
|
||||
func (d *einoRunEventDrain) HandleToolResultStreaming(mv *adk.MessageVariant, agentName string) bool {
|
||||
return d != nil && d.toolResultHandler != nil && d.toolResultHandler.HandleStreaming(mv, agentName)
|
||||
}
|
||||
|
||||
func (d *einoRunEventDrain) HandleAssistantStream(mv *adk.MessageVariant, agentName string) (bool, error) {
|
||||
if d == nil || d.assistantStreamHandler == nil {
|
||||
return false, nil
|
||||
}
|
||||
return d.assistantStreamHandler.Handle(mv, agentName)
|
||||
}
|
||||
|
||||
func (d *einoRunEventDrain) HandleMaterialized(mv *adk.MessageVariant, msg adk.Message, agentName string) bool {
|
||||
return d != nil && d.materializedMessageHandler != nil && d.materializedMessageHandler.Handle(mv, msg, agentName)
|
||||
}
|
||||
|
||||
func (d *einoRunEventDrain) markPendingWithMonitor(tc toolCallPendingInfo) {
|
||||
if d == nil || d.pendingToolCalls == nil {
|
||||
return
|
||||
}
|
||||
d.pendingToolCalls.Mark(tc)
|
||||
beginEinoADKFilesystemToolMonitor(
|
||||
d.cfg.Context,
|
||||
d.cfg.FilesystemMonitorAgent,
|
||||
d.cfg.FilesystemMonitorRecord,
|
||||
d.cfg.MCPExecutionBinder,
|
||||
tc.ToolCallID,
|
||||
tc.ToolName,
|
||||
)
|
||||
}
|
||||
|
||||
func (d *einoRunEventDrain) nextMainStreamID() string {
|
||||
return fmt.Sprintf("eino-main-%s-%d", d.cfg.ConversationID, atomic.AddInt64(&d.mainResponseStreamSeq, 1))
|
||||
}
|
||||
|
||||
func (d *einoRunEventDrain) nextReasoningStreamID() string {
|
||||
return fmt.Sprintf("eino-reasoning-%s-%d", d.cfg.ConversationID, atomic.AddInt64(&d.reasoningStreamSeq, 1))
|
||||
}
|
||||
|
||||
func (d *einoRunEventDrain) nextSubAgentReplyStreamID() string {
|
||||
return fmt.Sprintf("eino-sub-reply-%s-%d", d.cfg.ConversationID, atomic.AddInt64(&d.subReplyStreamSeq, 1))
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/cloudwego/eino/adk"
|
||||
"github.com/cloudwego/eino/schema"
|
||||
)
|
||||
|
||||
func TestEinoRunEventDrainDefaultsAndStreamIDs(t *testing.T) {
|
||||
drain := newEinoRunEventDrain(einoRunEventDrainConfig{
|
||||
ConversationID: "conv-1",
|
||||
OrchestratorName: "lead",
|
||||
BaseMessages: []adk.Message{schema.UserMessage("base")},
|
||||
})
|
||||
|
||||
if drain.RunMessages().BaseCount() != 1 {
|
||||
t.Fatalf("base count = %d, want 1", drain.RunMessages().BaseCount())
|
||||
}
|
||||
if !drain.cfg.StreamsMainAssistant("lead") || drain.cfg.StreamsMainAssistant("worker") {
|
||||
t.Fatal("default main-assistant predicate should match only orchestrator")
|
||||
}
|
||||
if got := drain.cfg.EinoRoleTag("lead"); got != "orchestrator" {
|
||||
t.Fatalf("lead role = %q, want orchestrator", got)
|
||||
}
|
||||
if got := drain.cfg.EinoRoleTag("worker"); got != "sub" {
|
||||
t.Fatalf("worker role = %q, want sub", got)
|
||||
}
|
||||
if got := drain.nextMainStreamID(); got != "eino-main-conv-1-1" {
|
||||
t.Fatalf("first main stream id = %q", got)
|
||||
}
|
||||
if got := drain.nextMainStreamID(); got != "eino-main-conv-1-2" {
|
||||
t.Fatalf("second main stream id = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoRunEventDrainBindsHandlersAndRecordsEvents(t *testing.T) {
|
||||
var events []string
|
||||
recovered := false
|
||||
drain := newEinoRunEventDrain(einoRunEventDrainConfig{
|
||||
ConversationID: "conv-1",
|
||||
OrchMode: "deep",
|
||||
OrchestratorName: "lead",
|
||||
Progress: func(eventType, _ string, _ interface{}) {
|
||||
events = append(events, eventType)
|
||||
},
|
||||
BaseMessages: []adk.Message{schema.UserMessage("base")},
|
||||
})
|
||||
drain.BindHandlers(func() { recovered = true })
|
||||
|
||||
drain.ObserveAgent("lead")
|
||||
if !drain.HandleMaterialized(&adk.MessageVariant{Role: schema.Assistant}, schema.AssistantMessage("done", nil), "lead") {
|
||||
t.Fatal("materialized assistant should be handled")
|
||||
}
|
||||
if got := drain.AssistantOutput().LastAssistant(); got != "done" {
|
||||
t.Fatalf("last assistant = %q, want done", got)
|
||||
}
|
||||
|
||||
stream := schema.StreamReaderFromArray([]*schema.Message{
|
||||
{Role: schema.Tool, Content: "ok", ToolCallID: "call-1"},
|
||||
})
|
||||
if !drain.HandleToolResultStreaming(&adk.MessageVariant{
|
||||
IsStreaming: true,
|
||||
Role: schema.Tool,
|
||||
ToolName: "execute",
|
||||
MessageStream: stream,
|
||||
}, "lead") {
|
||||
t.Fatal("streaming tool result should be handled")
|
||||
}
|
||||
if !recovered {
|
||||
t.Fatal("tool stream completion should confirm recovery")
|
||||
}
|
||||
if len(drain.RunMessages().Messages()) != 3 {
|
||||
t.Fatalf("run messages = %#v, want base + assistant + tool", drain.RunMessages().Messages())
|
||||
}
|
||||
if !containsString(events, "iteration") || !containsString(events, "response_start") || !containsString(events, "tool_result") {
|
||||
t.Fatalf("events = %#v, want iteration, response and tool_result", events)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,64 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/cloudwego/eino/adk"
|
||||
"github.com/cloudwego/eino/schema"
|
||||
)
|
||||
|
||||
func TestEinoRunMessageAccumulatorTracksBaseAndAppends(t *testing.T) {
|
||||
acc := newEinoRunMessageAccumulator([]adk.Message{schema.UserMessage("hi")})
|
||||
|
||||
if acc.BaseCount() != 1 {
|
||||
t.Fatalf("base count = %d, want 1", acc.BaseCount())
|
||||
}
|
||||
if acc.HasNewMessages() {
|
||||
t.Fatal("fresh accumulator should not have new messages")
|
||||
}
|
||||
|
||||
if !acc.AppendAssistantText(" hello ") {
|
||||
t.Fatal("assistant text should append")
|
||||
}
|
||||
if !acc.HasNewMessages() {
|
||||
t.Fatal("expected new messages after append")
|
||||
}
|
||||
msgs := acc.Messages()
|
||||
if len(msgs) != 2 || msgs[1].Role != schema.Assistant || msgs[1].Content != "hello" {
|
||||
t.Fatalf("messages = %#v", msgs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoRunMessageAccumulatorToolMessage(t *testing.T) {
|
||||
acc := newEinoRunMessageAccumulator(nil)
|
||||
if acc.AppendToolMessage("ignored", "") {
|
||||
t.Fatal("blank tool call id should not append")
|
||||
}
|
||||
if !acc.AppendToolMessage("result", "call-1", schema.WithToolName("execute")) {
|
||||
t.Fatal("tool message should append")
|
||||
}
|
||||
msgs := acc.Messages()
|
||||
if len(msgs) != 1 || msgs[0].Role != schema.Tool || msgs[0].Content != "result" || msgs[0].ToolCallID != "call-1" || msgs[0].ToolName != "execute" {
|
||||
t.Fatalf("tool message = %#v", msgs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoRunMessageAccumulatorAssistantToolCalls(t *testing.T) {
|
||||
acc := newEinoRunMessageAccumulator(nil)
|
||||
if acc.AppendAssistantToolCalls(nil) {
|
||||
t.Fatal("empty tool calls should not append")
|
||||
}
|
||||
if !acc.AppendAssistantToolCalls([]schema.ToolCall{{
|
||||
ID: "call-1",
|
||||
Function: schema.FunctionCall{
|
||||
Name: "execute",
|
||||
Arguments: `{}`,
|
||||
},
|
||||
}}) {
|
||||
t.Fatal("assistant tool calls should append")
|
||||
}
|
||||
msgs := acc.Messages()
|
||||
if len(msgs) != 1 || msgs[0].Role != schema.Assistant || len(msgs[0].ToolCalls) != 1 {
|
||||
t.Fatalf("assistant tool call message = %#v", msgs)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/cloudwego/eino/adk"
|
||||
"github.com/cloudwego/eino/schema"
|
||||
)
|
||||
|
||||
func TestEinoRunResultBuilderPartialWithoutNewMessagesReturnsOriginalError(t *testing.T) {
|
||||
runMessages := newEinoRunMessageAccumulator([]adk.Message{schema.UserMessage("base")})
|
||||
wantErr := errors.New("stream failed")
|
||||
|
||||
got, err := newEinoRunResultBuilder(einoRunResultBuilderConfig{
|
||||
RunMessages: runMessages,
|
||||
EmptyHint: "empty",
|
||||
}).BuildPartial(wantErr)
|
||||
|
||||
if got != nil {
|
||||
t.Fatalf("partial result = %#v, want nil", got)
|
||||
}
|
||||
if !errors.Is(err, wantErr) {
|
||||
t.Fatalf("err = %v, want %v", err, wantErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoRunResultBuilderFinalUsesSnapshots(t *testing.T) {
|
||||
runMessages := newEinoRunMessageAccumulator([]adk.Message{schema.UserMessage("base")})
|
||||
runMessages.Append(schema.AssistantMessage("assistant done", nil))
|
||||
assistantOutput := newEinoAssistantOutputAccumulator("deep")
|
||||
assistantOutput.RecordMainAssistant("orchestrator", "assistant done")
|
||||
|
||||
got := newEinoRunResultBuilder(einoRunResultBuilderConfig{
|
||||
OrchMode: "deep",
|
||||
EmptyHint: "empty",
|
||||
RunMessages: runMessages,
|
||||
AssistantOutput: assistantOutput,
|
||||
SnapshotMCPIDs: func() []string {
|
||||
return []string{"exec-1"}
|
||||
},
|
||||
ModelFacingTrace: func() []adk.Message {
|
||||
return []adk.Message{schema.UserMessage("model-facing")}
|
||||
},
|
||||
}).BuildFinal()
|
||||
|
||||
if got.Response != "assistant done" {
|
||||
t.Fatalf("response = %q, want assistant done", got.Response)
|
||||
}
|
||||
if len(got.MCPExecutionIDs) != 1 || got.MCPExecutionIDs[0] != "exec-1" {
|
||||
t.Fatalf("mcp ids = %#v", got.MCPExecutionIDs)
|
||||
}
|
||||
if got.LastAgentTraceInput == "" {
|
||||
t.Fatal("model-facing trace should be persisted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoRunResultBuilderPlanExecutePrefersExecutorOutput(t *testing.T) {
|
||||
runMessages := newEinoRunMessageAccumulator(nil)
|
||||
runMessages.Append(schema.AssistantMessage(`{"response":"planner text"}`, nil))
|
||||
assistantOutput := newEinoAssistantOutputAccumulator("plan_execute")
|
||||
assistantOutput.RecordMainAssistant("planner", `{"response":"planner text"}`)
|
||||
assistantOutput.RecordMainAssistant("executor", `{"response":"executor text"}`)
|
||||
|
||||
got := newEinoRunResultBuilder(einoRunResultBuilderConfig{
|
||||
OrchMode: "plan_execute",
|
||||
EmptyHint: "empty",
|
||||
RunMessages: runMessages,
|
||||
AssistantOutput: assistantOutput,
|
||||
}).BuildFinal()
|
||||
|
||||
if got.Response != "executor text" {
|
||||
t.Fatalf("response = %q, want executor text", got.Response)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/cloudwego/eino/adk"
|
||||
)
|
||||
|
||||
type fakeRunnerControl struct {
|
||||
runMessages []adk.Message
|
||||
runOpts int
|
||||
resumeID string
|
||||
resumeOpts int
|
||||
resumeErr error
|
||||
}
|
||||
|
||||
func (f *fakeRunnerControl) Run(_ context.Context, messages []adk.Message, opts ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] {
|
||||
f.runMessages = messages
|
||||
f.runOpts = len(opts)
|
||||
iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]()
|
||||
gen.Close()
|
||||
return iter
|
||||
}
|
||||
|
||||
func (f *fakeRunnerControl) Resume(_ context.Context, checkPointID string, opts ...adk.AgentRunOption) (*adk.AsyncIterator[*adk.AgentEvent], error) {
|
||||
f.resumeID = checkPointID
|
||||
f.resumeOpts = len(opts)
|
||||
iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]()
|
||||
gen.Close()
|
||||
return iter, f.resumeErr
|
||||
}
|
||||
|
||||
func TestEinoRunnerIteratorStarterStartAddsCancelAndCheckpoint(t *testing.T) {
|
||||
runner := &fakeRunnerControl{}
|
||||
var cancelPush func(error) bool
|
||||
var nativeCancelCause atomic.Value
|
||||
oldUnregistered := false
|
||||
newUnregistered := false
|
||||
unregister := func() { oldUnregistered = true }
|
||||
|
||||
iter := newEinoRunnerIteratorStarter(einoRunnerIteratorStarterConfig{
|
||||
Context: context.Background(),
|
||||
Runner: runner,
|
||||
CheckPointID: "cp-1",
|
||||
NativeCancelCause: &nativeCancelCause,
|
||||
UnregisterAgentCancel: &unregister,
|
||||
RuntimeCancelRegistrar: func(push func(error) bool) func() {
|
||||
cancelPush = push
|
||||
return func() { newUnregistered = true }
|
||||
},
|
||||
}).Start([]adk.Message{})
|
||||
|
||||
if iter == nil {
|
||||
t.Fatal("iterator should be created")
|
||||
}
|
||||
if runner.runOpts != 2 {
|
||||
t.Fatalf("run opts = %d, want cancel + checkpoint", runner.runOpts)
|
||||
}
|
||||
if !oldUnregistered {
|
||||
t.Fatal("old unregister should be called before binding a new cancel hook")
|
||||
}
|
||||
if cancelPush == nil {
|
||||
t.Fatal("cancel hook should be registered")
|
||||
}
|
||||
stopErr := errors.New("stop")
|
||||
if cancelPush(stopErr) {
|
||||
t.Fatal("unbound fake runner cancel should not report handled")
|
||||
}
|
||||
if got, _ := nativeCancelCause.Load().(error); !errors.Is(got, stopErr) {
|
||||
t.Fatalf("native cancel cause = %v, want %v", got, stopErr)
|
||||
}
|
||||
unregister()
|
||||
if !newUnregistered {
|
||||
t.Fatal("new unregister should replace old unregister")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoRunnerIteratorStarterResumeUsesCancelOnly(t *testing.T) {
|
||||
runner := &fakeRunnerControl{}
|
||||
|
||||
iter, err := newEinoRunnerIteratorStarter(einoRunnerIteratorStarterConfig{
|
||||
Context: context.Background(),
|
||||
Runner: runner,
|
||||
CheckPointID: "fresh-run-checkpoint",
|
||||
}).Resume("resume-cp")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("resume err = %v", err)
|
||||
}
|
||||
if iter == nil {
|
||||
t.Fatal("iterator should be created")
|
||||
}
|
||||
if runner.resumeID != "resume-cp" {
|
||||
t.Fatalf("resume id = %q, want resume-cp", runner.resumeID)
|
||||
}
|
||||
if runner.resumeOpts != 1 {
|
||||
t.Fatalf("resume opts = %d, want cancel only", runner.resumeOpts)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEinoRunnerIteratorStarterResumePropagatesError(t *testing.T) {
|
||||
resumeErr := errors.New("resume failed")
|
||||
runner := &fakeRunnerControl{resumeErr: resumeErr}
|
||||
|
||||
_, err := newEinoRunnerIteratorStarter(einoRunnerIteratorStarterConfig{
|
||||
Context: context.Background(),
|
||||
Runner: runner,
|
||||
}).Resume("resume-cp")
|
||||
|
||||
if !errors.Is(err, resumeErr) {
|
||||
t.Fatalf("resume err = %v, want %v", err, resumeErr)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user