mirror of
https://github.com/Ed1s0nZ/CyberStrikeAI.git
synced 2026-09-30 05:02:06 +02:00
fix: make Eino exit and transfer work on AgenticMessage state
Official ExitTool still writes classic Message react state, which breaks supervisor delivery on the Agentic path. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -0,0 +1,107 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/cloudwego/eino/adk"
|
||||
"github.com/cloudwego/eino/components/tool"
|
||||
"github.com/cloudwego/eino/compose"
|
||||
"github.com/cloudwego/eino/schema"
|
||||
)
|
||||
|
||||
const agenticExitToolName = "exit"
|
||||
|
||||
type agenticCompatibleExitTool struct{}
|
||||
|
||||
func (agenticCompatibleExitTool) Info(context.Context) (*schema.ToolInfo, error) {
|
||||
return adk.ToolInfoExit, nil
|
||||
}
|
||||
|
||||
func (agenticCompatibleExitTool) InvokableRun(ctx context.Context, argumentsInJSON string, _ ...tool.Option) (string, error) {
|
||||
return invokeAgenticBuiltinActionTool(ctx, agenticExitToolName, argumentsInJSON)
|
||||
}
|
||||
|
||||
func replaceClassicExitTool(t tool.BaseTool) tool.BaseTool {
|
||||
switch t.(type) {
|
||||
case adk.ExitTool, *adk.ExitTool:
|
||||
return agenticCompatibleExitTool{}
|
||||
default:
|
||||
return t
|
||||
}
|
||||
}
|
||||
|
||||
func attachAgenticBuiltinActionToolMiddleware(cfg *adk.ToolsConfig) {
|
||||
if cfg == nil {
|
||||
return
|
||||
}
|
||||
cfg.ToolsNodeConfig.ToolCallMiddlewares = append(
|
||||
cfg.ToolsNodeConfig.ToolCallMiddlewares,
|
||||
agenticBuiltinActionToolMiddleware(),
|
||||
)
|
||||
}
|
||||
|
||||
func agenticBuiltinActionToolMiddleware() compose.ToolMiddleware {
|
||||
return compose.ToolMiddleware{
|
||||
Invokable: func(next compose.InvokableToolEndpoint) compose.InvokableToolEndpoint {
|
||||
return func(ctx context.Context, input *compose.ToolInput) (*compose.ToolOutput, error) {
|
||||
if input != nil && isAgenticBuiltinActionTool(input.Name) {
|
||||
result, err := invokeAgenticBuiltinActionTool(ctx, input.Name, input.Arguments)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &compose.ToolOutput{Result: result}, nil
|
||||
}
|
||||
return next(ctx, input)
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func isAgenticBuiltinActionTool(name string) bool {
|
||||
switch strings.TrimSpace(name) {
|
||||
case agenticExitToolName, adk.TransferToAgentToolName:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func invokeAgenticBuiltinActionTool(ctx context.Context, name, argumentsInJSON string) (string, error) {
|
||||
switch strings.TrimSpace(name) {
|
||||
case agenticExitToolName:
|
||||
var params struct {
|
||||
FinalResult string `json:"final_result"`
|
||||
}
|
||||
if err := unmarshalBuiltinActionArgs(argumentsInJSON, ¶ms); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := sendADKToolGenAction(ctx, agenticExitToolName, adk.NewExitAction()); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return params.FinalResult, nil
|
||||
case adk.TransferToAgentToolName:
|
||||
var params struct {
|
||||
AgentName string `json:"agent_name"`
|
||||
}
|
||||
if err := unmarshalBuiltinActionArgs(argumentsInJSON, ¶ms); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := sendADKToolGenAction(ctx, adk.TransferToAgentToolName, adk.NewTransferToAgentAction(params.AgentName)); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return fmt.Sprintf("successfully transferred to agent [%s]", params.AgentName), nil
|
||||
default:
|
||||
return "", fmt.Errorf("unsupported agentic builtin action tool %q", name)
|
||||
}
|
||||
}
|
||||
|
||||
func unmarshalBuiltinActionArgs(argumentsInJSON string, dest any) error {
|
||||
argumentsInJSON = strings.TrimSpace(argumentsInJSON)
|
||||
if argumentsInJSON == "" {
|
||||
argumentsInJSON = "{}"
|
||||
}
|
||||
return json.Unmarshal([]byte(argumentsInJSON), dest)
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/cloudwego/eino/adk"
|
||||
"github.com/cloudwego/eino/compose"
|
||||
"github.com/cloudwego/eino/schema"
|
||||
)
|
||||
|
||||
type agenticShapedReactState struct {
|
||||
ToolGenActions map[string]*adk.AgentAction
|
||||
}
|
||||
|
||||
func TestSendADKToolGenActionWritesAgenticShapedState(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
chain := compose.NewChain[string, string](compose.WithGenLocalState(func(context.Context) *agenticShapedReactState {
|
||||
return &agenticShapedReactState{}
|
||||
}))
|
||||
chain.AppendLambda(compose.InvokableLambda(func(ctx context.Context, in string) (string, error) {
|
||||
if err := sendADKToolGenAction(ctx, adk.TransferToAgentToolName, adk.NewTransferToAgentAction("expert")); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return in, compose.ProcessState(ctx, func(_ context.Context, st *agenticShapedReactState) error {
|
||||
action := st.ToolGenActions[adk.TransferToAgentToolName]
|
||||
if action == nil || action.TransferToAgent == nil || action.TransferToAgent.DestAgentName != "expert" {
|
||||
return errors.New("missing transfer tool gen action")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}))
|
||||
r, err := chain.Compile(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("compile: %v", err)
|
||||
}
|
||||
if _, err := r.Invoke(ctx, "ok"); err != nil {
|
||||
t.Fatalf("invoke: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClearADKReturnDirectlyZerosExportedFields(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
chain := compose.NewChain[string, string](compose.WithGenLocalState(func(context.Context) *adk.State {
|
||||
return &adk.State{HasReturnDirectly: true, ReturnDirectlyToolCallID: "call-1"}
|
||||
}))
|
||||
chain.AppendLambda(compose.InvokableLambda(func(ctx context.Context, in string) (string, error) {
|
||||
if err := clearADKReturnDirectly(ctx); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return in, compose.ProcessState(ctx, func(_ context.Context, st *adk.State) error {
|
||||
if st.HasReturnDirectly || st.ReturnDirectlyToolCallID != "" || st.ReturnDirectlyEvent != nil {
|
||||
return errors.New("return-directly fields were not cleared")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}))
|
||||
r, err := chain.Compile(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("compile: %v", err)
|
||||
}
|
||||
if _, err := r.Invoke(ctx, "ok"); err != nil {
|
||||
t.Fatalf("invoke: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgenticExitToolDoesNotFailOnAgenticChatModelAgent(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
fakeModel := &capturingAgenticChatModel{
|
||||
output: agenticAssistantToolCall("tool-call-1", "exit", `{"final_result":"This is the final result"}`),
|
||||
}
|
||||
agent, err := newEinoAgenticChatModelAgentAdapter(ctx, einoAgenticChatModelAgentConfig{
|
||||
Name: "agentic-exit",
|
||||
Description: "exit regression",
|
||||
Instruction: "finish with exit",
|
||||
Model: fakeModel,
|
||||
Exit: &adk.ExitTool{},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("newEinoAgenticChatModelAgentAdapter: %v", err)
|
||||
}
|
||||
|
||||
var sawExit bool
|
||||
var exitContent string
|
||||
iter := agent.Run(ctx, &adk.AgentInput{Messages: []*schema.Message{schema.UserMessage("please exit")}})
|
||||
for {
|
||||
ev, ok := iter.Next()
|
||||
if !ok {
|
||||
break
|
||||
}
|
||||
if ev.Err != nil {
|
||||
t.Fatalf("agent event error: %v", ev.Err)
|
||||
}
|
||||
if ev.Action != nil && ev.Action.Exit {
|
||||
sawExit = true
|
||||
if ev.Output != nil && ev.Output.MessageOutput != nil && ev.Output.MessageOutput.Message != nil {
|
||||
exitContent = ev.Output.MessageOutput.Message.Content
|
||||
}
|
||||
}
|
||||
if ev.Output != nil && ev.Output.MessageOutput != nil && ev.Output.MessageOutput.Message != nil {
|
||||
msg := ev.Output.MessageOutput.Message
|
||||
if strings.Contains(msg.Content, "cannot find state with type") {
|
||||
t.Fatalf("exit still failed with classic state mismatch: %q", msg.Content)
|
||||
}
|
||||
}
|
||||
}
|
||||
if !sawExit {
|
||||
t.Fatal("expected Exit action on agentic ChatModelAgent")
|
||||
}
|
||||
if exitContent != "This is the final result" {
|
||||
t.Fatalf("exit content = %q, want final result", exitContent)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgenticBuiltinActionMiddlewareInterceptsTransfer(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
nextCalled := false
|
||||
mw := agenticBuiltinActionToolMiddleware()
|
||||
endpoint := mw.Invokable(func(context.Context, *compose.ToolInput) (*compose.ToolOutput, error) {
|
||||
nextCalled = true
|
||||
return nil, errors.New("official transfer tool should not run")
|
||||
})
|
||||
chain := compose.NewChain[string, string](compose.WithGenLocalState(func(context.Context) *agenticShapedReactState {
|
||||
return &agenticShapedReactState{}
|
||||
}))
|
||||
chain.AppendLambda(compose.InvokableLambda(func(ctx context.Context, in string) (string, error) {
|
||||
out, err := endpoint(ctx, &compose.ToolInput{
|
||||
Name: adk.TransferToAgentToolName,
|
||||
Arguments: `{"agent_name":"expert"}`,
|
||||
})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if out == nil || !strings.Contains(out.Result, "expert") {
|
||||
return "", errors.New("missing transfer result")
|
||||
}
|
||||
if nextCalled {
|
||||
return "", errors.New("official transfer tool ran")
|
||||
}
|
||||
return in, nil
|
||||
}))
|
||||
r, err := chain.Compile(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("compile: %v", err)
|
||||
}
|
||||
if _, err := r.Invoke(ctx, "ok"); err != nil {
|
||||
t.Fatalf("invoke: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -30,6 +30,8 @@ func newEinoAgenticChatModelAgent(ctx context.Context, cfg einoAgenticChatModelA
|
||||
if cfg.Model == nil {
|
||||
return nil, fmt.Errorf("eino agentic ChatModelAgent: model is required")
|
||||
}
|
||||
attachAgenticBuiltinActionToolMiddleware(&cfg.ToolsConfig)
|
||||
cfg.Exit = replaceClassicExitTool(cfg.Exit)
|
||||
typedCfg := &adk.TypedChatModelAgentConfig[*schema.AgenticMessage]{
|
||||
Name: cfg.Name,
|
||||
Description: cfg.Description,
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
package multiagent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"reflect"
|
||||
|
||||
"github.com/cloudwego/eino/adk"
|
||||
"github.com/cloudwego/eino/compose"
|
||||
)
|
||||
|
||||
// mutateADKReactState updates ChatModelAgent react state regardless of whether
|
||||
// the live graph stores typedState[*schema.Message] or typedState[*schema.AgenticMessage].
|
||||
// Eino v0.9.14 only exports SendToolGenAction against the classic Message state.
|
||||
func mutateADKReactState(ctx context.Context, mutate func(st reflect.Value) error) error {
|
||||
if mutate == nil {
|
||||
return fmt.Errorf("adk react state mutate is nil")
|
||||
}
|
||||
return compose.ProcessState(ctx, func(_ context.Context, st any) error {
|
||||
if st == nil {
|
||||
return fmt.Errorf("adk react state is nil")
|
||||
}
|
||||
v := reflect.ValueOf(st)
|
||||
if v.Kind() != reflect.Pointer || v.IsNil() {
|
||||
return fmt.Errorf("unexpected adk react state type %T", st)
|
||||
}
|
||||
return mutate(v.Elem())
|
||||
})
|
||||
}
|
||||
|
||||
func sendADKToolGenAction(ctx context.Context, toolName string, action *adk.AgentAction) error {
|
||||
if action == nil {
|
||||
return fmt.Errorf("adk tool gen action is nil")
|
||||
}
|
||||
key := toolName
|
||||
if toolCallID := compose.GetToolCallID(ctx); toolCallID != "" {
|
||||
key = toolCallID
|
||||
}
|
||||
return mutateADKReactState(ctx, func(st reflect.Value) error {
|
||||
field := st.FieldByName("ToolGenActions")
|
||||
if !field.IsValid() || field.Kind() != reflect.Map {
|
||||
return fmt.Errorf("adk react state missing ToolGenActions")
|
||||
}
|
||||
if !field.CanSet() {
|
||||
return fmt.Errorf("cannot set ToolGenActions on adk react state")
|
||||
}
|
||||
if field.IsNil() {
|
||||
field.Set(reflect.MakeMap(field.Type()))
|
||||
}
|
||||
field.SetMapIndex(reflect.ValueOf(key), reflect.ValueOf(action))
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func clearADKReturnDirectly(ctx context.Context) error {
|
||||
return mutateADKReactState(ctx, func(st reflect.Value) error {
|
||||
zeroExportedField(st, "ReturnDirectlyToolCallID")
|
||||
zeroExportedField(st, "HasReturnDirectly")
|
||||
zeroExportedField(st, "ReturnDirectlyEvent")
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func zeroExportedField(st reflect.Value, name string) {
|
||||
field := st.FieldByName(name)
|
||||
if !field.IsValid() || !field.CanSet() {
|
||||
return
|
||||
}
|
||||
field.Set(reflect.Zero(field.Type()))
|
||||
}
|
||||
@@ -55,15 +55,7 @@ func hitlClearReturnDirectlyIfTransfer(ctx context.Context, toolName string) {
|
||||
if !strings.EqualFold(strings.TrimSpace(toolName), adk.TransferToAgentToolName) {
|
||||
return
|
||||
}
|
||||
_ = compose.ProcessState[*adk.State](ctx, func(_ context.Context, st *adk.State) error {
|
||||
if st == nil {
|
||||
return nil
|
||||
}
|
||||
st.ReturnDirectlyToolCallID = ""
|
||||
st.HasReturnDirectly = false
|
||||
st.ReturnDirectlyEvent = nil
|
||||
return nil
|
||||
})
|
||||
_ = clearADKReturnDirectly(ctx)
|
||||
}
|
||||
|
||||
func hitlEditedArgumentsNotice(original, edited string) string {
|
||||
|
||||
@@ -464,6 +464,7 @@ func RunDeepAgent(
|
||||
},
|
||||
EmitInternalEvents: true,
|
||||
}
|
||||
attachAgenticBuiltinActionToolMiddleware(&mainToolsCfg)
|
||||
|
||||
deepAgenticOutKey, agenticTaskGen := deepAgenticExtrasFromConfig(ma)
|
||||
|
||||
@@ -545,7 +546,7 @@ func RunDeepAgent(
|
||||
ToolsConfig: mainToolsCfg,
|
||||
MaxIterations: deepMaxIter,
|
||||
Handlers: supHandlers,
|
||||
Exit: &adk.ExitTool{},
|
||||
Exit: agenticCompatibleExitTool{},
|
||||
ModelRetryConfig: agenticModelRetryCfg,
|
||||
ModelFailoverConfig: agenticModelFailoverCfg,
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user