diff --git a/agent/agent.go b/agent/agent.go index 8daf5756..e7d552d5 100644 --- a/agent/agent.go +++ b/agent/agent.go @@ -33,6 +33,13 @@ type ProviderConfig struct { // Middlewares wrap Run after agent history and context providers. Middlewares []Middleware + // ManagesToolExecution indicates that Run invokes function tools supplied through + // [WithTool], rather than only returning function call requests. When true, + // [New] applies [Config.FunctionMiddlewares] to those tools immediately before + // Run, after provider middleware. Run must use the tools in its options, not + // retained originals. The provider remains responsible for execution and approvals. + ManagesToolExecution bool + // Format creates a provider response format for a structured output value. Format func(v any) (ResponseFormat, error) @@ -176,10 +183,14 @@ func New(prov ProviderConfig, cfg Config) *Agent { providerDoesNotManageHistory: prov.ServiceDoesNotManageHistory, contextProviders: contextProviders, } + providerRun := prov.Run + if prov.ManagesToolExecution { + providerRun = wrapFuncTools(providerRun) + } if len(providerMiddlewares) == 0 { - a.providerPipeline = prov.Run + a.providerPipeline = providerRun } else { - a.providerPipeline = compileRunChain(prov.Run, providerMiddlewares) + a.providerPipeline = compileRunChain(providerRun, providerMiddlewares) } a.runPipeline = compileRunChain(a.invoke, agentMiddlewares) return a diff --git a/agent/harness/toolautocall/autocall.go b/agent/harness/toolautocall/autocall.go index c828a308..b38791ca 100644 --- a/agent/harness/toolautocall/autocall.go +++ b/agent/harness/toolautocall/autocall.go @@ -1164,7 +1164,7 @@ func (f *autocall) processFunctionCall(ctx context.Context, tools map[string]too } f.logger.Debug(ctx, "calling function", "funcName", funcCall.Name, slogx.SensitiveData("arguments", funcCall.Arguments)) start := time.Now() - ctx = toolmiddleware.WithCallID(ctx, funcCall.CallID) + ctx = agent.WithFuncCallID(ctx, funcCall.CallID) ctx, span := startToolSpan(ctx, funcCall, declaration) if span != nil { defer span.End() diff --git a/agent/middleware.go b/agent/middleware.go index 0566a2bd..1ffc54dc 100644 --- a/agent/middleware.go +++ b/agent/middleware.go @@ -75,8 +75,50 @@ type FunctionInvocationContext struct { // Multiple callbacks execute in registration order, with the first outermost. // Each call gets its own FunctionInvocationContext; callbacks must synchronize // shared application state if tools can execute concurrently. +// +// Providers that execute tools themselves set [ProviderConfig.ManagesToolExecution] +// and use [WithFuncCallID] to identify each invocation. type FunctionInvocationMiddleware func(next func(context.Context, *FunctionInvocationContext) (any, error), ctx context.Context, invocation *FunctionInvocationContext) (any, error) +// WithFuncCallID returns a context carrying callID for +// [FunctionInvocationContext.CallID]. Providers that execute tools themselves +// should pass the returned context to [tool.FuncTool.Call] so function invocation +// middleware receives the ID. Use an empty callID when the provider supplies no +// ID; this overrides any ID inherited from ctx. +func WithFuncCallID(ctx context.Context, callID string) context.Context { + return toolmiddleware.WithCallID(ctx, callID) +} + +// wrapFuncTools wraps option-provided function tools immediately before run, +// leaving the options seen by outer middleware unchanged. +func wrapFuncTools(run RunFunc) RunFunc { + return func(ctx context.Context, messages []*message.Message, options ...Option) iter.Seq2[*ResponseUpdate, error] { + var cloned bool + for _, opt := range options { + wrap, ok := opt.(toolmiddleware.Wrapper) + if !ok { + continue + } + for i, opt := range options { + t, ok := opt.(toolOpt) + if !ok { + continue + } + fn, ok := t.Tool.(tool.FuncTool) + if !ok { + continue + } + if !cloned { + options = slices.Clone(options) + cloned = true + } + options[i] = WithTool(wrap(fn)) + } + } + return run(ctx, messages, options...) + } +} + type functionInvocationTool struct { tool.FuncTool middlewares []FunctionInvocationMiddleware diff --git a/agent/middleware_test.go b/agent/middleware_test.go new file mode 100644 index 00000000..76b0f71e --- /dev/null +++ b/agent/middleware_test.go @@ -0,0 +1,301 @@ +// Copyright (c) Microsoft. All rights reserved. + +package agent_test + +import ( + "context" + "errors" + "iter" + "reflect" + "slices" + "testing" + + "github.com/microsoft/agent-framework-go/agent" + "github.com/microsoft/agent-framework-go/message" + "github.com/microsoft/agent-framework-go/tool" + "github.com/microsoft/agent-framework-go/tool/functool" +) + +func TestProviderConfig_ManagesToolExecution_ToolOptions(t *testing.T) { + for _, tc := range []struct { + name string + invokes bool + middleware bool + }{ + {name: "not opted in", middleware: true}, + {name: "no middleware", invokes: true}, + {name: "enabled", invokes: true, middleware: true}, + } { + t.Run(tc.name, func(t *testing.T) { + fn := functool.MustNew(functool.Config{Name: "lookup"}, func(context.Context, struct{}) (string, error) { + return "found", nil + }) + hosted := stubTool{name: "hosted"} + schemaOnly := struct{ tool.SchemaTool }{fn} + options := []agent.Option{ + agent.WithInstructions("first"), agent.WithTool(fn), agent.WithTool(hosted), + agent.WithTool(nil), agent.WithTool(schemaOnly), agent.WithInstructions("last"), + } + cfg := agent.Config{} + if tc.middleware { + cfg.FunctionMiddlewares = []agent.FunctionInvocationMiddleware{func(next func(context.Context, *agent.FunctionInvocationContext) (any, error), ctx context.Context, invocation *agent.FunctionInvocationContext) (any, error) { + result, err := next(ctx, invocation) + if err != nil { + return nil, err + } + return "wrapped " + result.(string), nil + }} + } + run := func(ctx context.Context, _ []*message.Message, opts ...agent.Option) iter.Seq2[*agent.ResponseUpdate, error] { + return func(yield func(*agent.ResponseUpdate, error) bool) { + tools := slices.Collect(agent.AllOptions(opts, agent.WithTool)) + if len(tools) != 3 || tools[1] != hosted || tools[2] != schemaOnly { + t.Fatalf("non-function tools or tool order changed: %v", tools) + } + if (!tc.invokes || !tc.middleware) && tools[0] != fn { + t.Error("tool changed without native invocation middleware") + } + if got := slices.Collect(agent.AllOptions(opts, agent.WithInstructions)); !slices.Equal(got, []string{"first", "last"}) { + t.Errorf("instructions = %v, want [first last]", got) + } + result, err := tools[0].(tool.FuncTool).Call(ctx, "{}") + if err != nil { + yield(nil, err) + return + } + yield(&agent.ResponseUpdate{Contents: []message.Content{&message.TextContent{Text: result.(string)}}}, nil) + } + } + a := agent.New(agent.ProviderConfig{Run: run, ManagesToolExecution: tc.invokes}, cfg) + response, err := a.RunText(t.Context(), "lookup", options...).Collect() + if err != nil { + t.Fatal(err) + } + want := "found" + if tc.invokes && tc.middleware { + want = "wrapped found" + } + if response.String() != want { + t.Errorf("response = %q, want %q", response.String(), want) + } + if got := slices.Collect(agent.AllOptions(options, agent.WithTool)); !slices.Equal(got, []tool.Tool{fn, hosted, schemaOnly}) { + t.Error("caller tools changed") + } + }) + } +} + +func TestProviderConfig_ManagesToolExecution(t *testing.T) { + toolFailure := errors.New("tool failed") + for _, source := range []string{"configured", "run option", "context provider", "provider middleware"} { + for _, tc := range []struct { + name string + block bool + approval bool + toolErr error + }{ + {name: "composition"}, + {name: "short circuit", block: true}, + {name: "error propagation", toolErr: toolFailure}, + {name: "approval required", approval: true}, + } { + t.Run(source+"/"+tc.name, func(t *testing.T) { + type traceKey struct{} + type middlewareKey struct{} + var order []string + var callID string + fn := functool.MustNew(functool.Config{Name: "lookup", Description: "Look up a value"}, func(ctx context.Context, args struct { + Value string `json:"value"` + }, + ) (string, error) { + order = append(order, "tool") + if args.Value != "changed" || ctx.Value(middlewareKey{}) != callID { + t.Error("tool did not receive the middleware's arguments and context") + } + return args.Value, tc.toolErr + }) + if tc.approval { + fn = tool.ApprovalRequiredFunc(fn) + } + first := agent.FunctionInvocationMiddleware(func(next func(context.Context, *agent.FunctionInvocationContext) (any, error), ctx context.Context, invocation *agent.FunctionInvocationContext) (any, error) { + order = append(order, "first before") + if invocation.Function != fn || invocation.CallID != callID || invocation.Arguments != `{"value":"original"}` { + t.Errorf("unexpected invocation: %#v", invocation) + } + if ctx.Value(traceKey{}) != "trace" { + t.Error("middleware lost the provider's context") + } + if tc.block { + return "blocked", nil + } + invocation.Arguments = `{"value":"changed"}` + result, err := next(context.WithValue(ctx, middlewareKey{}, invocation.CallID), invocation) + order = append(order, "first after") + if err != nil { + return nil, err + } + return "wrapped " + result.(string), nil + }) + second := agent.FunctionInvocationMiddleware(func(next func(context.Context, *agent.FunctionInvocationContext) (any, error), ctx context.Context, invocation *agent.FunctionInvocationContext) (any, error) { + order = append(order, "second before") + result, err := next(ctx, invocation) + order = append(order, "second after") + return result, err + }) + toolOptions := []agent.Option{agent.WithTool(fn)} + cfg := agent.Config{FunctionMiddlewares: []agent.FunctionInvocationMiddleware{first, nil, second}} + var runOptions []agent.Option + var providerMiddlewares []agent.Middleware + switch source { + case "configured": + cfg.Tools = []tool.Tool{fn} + case "run option": + runOptions = toolOptions + case "context provider": + cfg.ContextProviders = []agent.ContextProvider{agent.NewContextProvider(agent.ContextProviderConfig{ + SourceID: "tools", + Provide: func(context.Context, agent.InvokingContext) ([]*message.Message, []agent.Option, error) { + return nil, toolOptions, nil + }, + })} + case "provider middleware": + providerMiddlewares = []agent.Middleware{agent.MiddlewareFunc(func(next agent.RunFunc, ctx context.Context, messages []*message.Message, options ...agent.Option) iter.Seq2[*agent.ResponseUpdate, error] { + return next(ctx, messages, append(slices.Clone(options), toolOptions...)...) + })} + } + run := func(ctx context.Context, _ []*message.Message, options ...agent.Option) iter.Seq2[*agent.ResponseUpdate, error] { + return func(yield func(*agent.ResponseUpdate, error) bool) { + var function tool.FuncTool + for tl := range agent.AllOptions(options, agent.WithTool) { + function, _ = tl.(tool.FuncTool) + } + if function == nil { + t.Fatal("provider did not receive the function tool") + } + if function.Name() != fn.Name() || function.Description() != fn.Description() || + !reflect.DeepEqual(function.Schema(), fn.Schema()) || !reflect.DeepEqual(function.ReturnSchema(), fn.ReturnSchema()) { + t.Error("wrapped tool metadata changed") + } + approval, ok := function.(tool.ApprovalRequiredTool) + if got := ok && approval.ApprovalRequired(); got != tc.approval { + t.Errorf("ApprovalRequired() = %v, want %v", got, tc.approval) + } + if len(order) != 0 { + t.Fatal("middleware ran before the tool was invoked") + } + result, err := function.Call(agent.WithFuncCallID(ctx, callID), `{"value":"original"}`) + if err != nil { + yield(nil, err) + return + } + yield(&agent.ResponseUpdate{ + Role: message.RoleAssistant, + Contents: []message.Content{&message.TextContent{Text: result.(string)}}, + }, nil) + } + } + a := agent.New(agent.ProviderConfig{Run: run, Middlewares: providerMiddlewares, ManagesToolExecution: true}, cfg) + ctx := context.WithValue(t.Context(), traceKey{}, "trace") + ctx = agent.WithFuncCallID(ctx, "parent-call") + for _, id := range []string{"call-1", ""} { + callID = id + order = nil + response, err := a.RunText(ctx, "lookup", runOptions...).Collect() + if !errors.Is(err, tc.toolErr) { + t.Fatalf("RunText() error = %v, want %v", err, tc.toolErr) + } + wantOrder := []string{"first before", "second before", "tool", "second after", "first after"} + wantResult := "wrapped changed" + if tc.block { + wantOrder = []string{"first before"} + wantResult = "blocked" + } + if !slices.Equal(order, wantOrder) { + t.Errorf("callback order = %v, want %v", order, wantOrder) + } + if err == nil && response.String() != wantResult { + t.Errorf("response = %q, want %q", response.String(), wantResult) + } + } + if original, _ := agent.GetOption(toolOptions, agent.WithTool); original != fn { + t.Error("wrapping mutated the caller's options") + } + }) + } + } +} + +func TestProviderConfig_ManagesToolExecution_InsideProviderMiddleware(t *testing.T) { + fn := functool.MustNew(functool.Config{Name: "lookup"}, func(context.Context, struct{}) (string, error) { + return "found", nil + }) + functionMiddleware := agent.FunctionInvocationMiddleware(func(next func(context.Context, *agent.FunctionInvocationContext) (any, error), ctx context.Context, invocation *agent.FunctionInvocationContext) (any, error) { + if invocation.Function != fn { + t.Error("middleware did not receive the original tool") + } + result, err := next(ctx, invocation) + if err != nil { + return nil, err + } + return "wrapped " + result.(string), nil + }) + providerMiddleware := agent.MiddlewareFunc(func(next agent.RunFunc, ctx context.Context, messages []*message.Message, options ...agent.Option) iter.Seq2[*agent.ResponseUpdate, error] { + return func(yield func(*agent.ResponseUpdate, error) bool) { + for range 2 { + original, _ := agent.GetOption(options, agent.WithTool) + if original != fn { + t.Fatal("provider middleware did not receive the original tool") + } + for update, err := range next(ctx, messages, options...) { + if !yield(update, err) { + return + } + } + if original, _ := agent.GetOption(options, agent.WithTool); original != fn { + t.Fatal("wrapping mutated provider middleware's options") + } + } + } + }) + run := func(ctx context.Context, _ []*message.Message, options ...agent.Option) iter.Seq2[*agent.ResponseUpdate, error] { + return func(yield func(*agent.ResponseUpdate, error) bool) { + tl, _ := agent.GetOption(options, agent.WithTool) + result, err := tl.(tool.FuncTool).Call(ctx, "{}") + if err != nil { + yield(nil, err) + return + } + yield(&agent.ResponseUpdate{Contents: []message.Content{&message.TextContent{Text: result.(string)}}}, nil) + } + } + a := agent.New(agent.ProviderConfig{ + Run: run, Middlewares: []agent.Middleware{providerMiddleware}, ManagesToolExecution: true, + }, agent.Config{Tools: []tool.Tool{fn}, FunctionMiddlewares: []agent.FunctionInvocationMiddleware{functionMiddleware}}) + var results []string + for update, err := range a.RunText(t.Context(), "lookup") { + if err != nil { + t.Fatal(err) + } + results = append(results, update.String()) + } + if !slices.Equal(results, []string{"wrapped found", "wrapped found"}) { + t.Errorf("results = %v, want [wrapped found wrapped found]", results) + } +} + +func TestWithFuncCallID_PreservesContext(t *testing.T) { + type contextKey struct{} + parent, cancel := context.WithCancel(context.WithValue(t.Context(), contextKey{}, "value")) + defer cancel() + ctx := agent.WithFuncCallID(parent, "call-1") + if ctx.Value(contextKey{}) != "value" { + t.Error("WithFuncCallID lost a context value") + } + if ctx.Done() != parent.Done() { + t.Error("WithFuncCallID changed the cancellation channel") + } + cancel() + if ctx.Err() != context.Canceled { + t.Errorf("context error = %v, want context.Canceled", ctx.Err()) + } +} diff --git a/provider/copilotprovider/copilot.go b/provider/copilotprovider/copilot.go index 15944181..b0e2240e 100644 --- a/provider/copilotprovider/copilot.go +++ b/provider/copilotprovider/copilot.go @@ -19,7 +19,6 @@ import ( copilot "github.com/github/copilot-sdk/go" "github.com/microsoft/agent-framework-go/agent" - "github.com/microsoft/agent-framework-go/internal/toolmiddleware" "github.com/microsoft/agent-framework-go/message" "github.com/microsoft/agent-framework-go/tool" ) @@ -64,8 +63,9 @@ func NewAgent(cclient *copilot.Client, config AgentConfig) *agent.Agent { cfg: config, } return agent.New(agent.ProviderConfig{ - ProviderName: "copilot", - Run: p.run, + ProviderName: "copilot", + Run: p.run, + ManagesToolExecution: true, }, config.Config) } @@ -470,11 +470,6 @@ func copilotTools(options []agent.Option) []copilot.Tool { if !ok { continue } - for _, opt := range options { - if wrap, ok := opt.(toolmiddleware.Wrapper); ok { - funcTool = wrap(funcTool) - } - } converted, err := toCopilotTool(funcTool) if err != nil { converted = copilot.Tool{ @@ -505,7 +500,7 @@ func toCopilotTool(funcTool tool.FuncTool) (copilot.Tool, error) { if ctx == nil { ctx = context.Background() } - ctx = toolmiddleware.WithCallID(ctx, invocation.ToolCallID) + ctx = agent.WithFuncCallID(ctx, invocation.ToolCallID) result, err := funcTool.Call(ctx, arguments) if err != nil { return copilot.ToolResult{}, err diff --git a/provider/copilotprovider/copilot_internal_test.go b/provider/copilotprovider/copilot_internal_test.go index cf312e03..5d41b92a 100644 --- a/provider/copilotprovider/copilot_internal_test.go +++ b/provider/copilotprovider/copilot_internal_test.go @@ -79,7 +79,7 @@ func TestCopilotTool_FunctionInvocationIdentity(t *testing.T) { } else { cfg.Tools = []tool.Tool{fn} } - a := agent.New(agent.ProviderConfig{Run: run}, cfg) + a := agent.New(agent.ProviderConfig{Run: run, ManagesToolExecution: true}, cfg) if _, err := a.RunText(t.Context(), "lookup").Collect(); err != nil { t.Fatal(err) }