From cf07dbc538e0927f396635ca93b6b71edfc60bc2 Mon Sep 17 00:00:00 2001 From: qmuntal Date: Tue, 22 Sep 2026 17:31:31 +0200 Subject: [PATCH] Support tool middleware in third-party providers Add ProviderConfig.ManagesToolExecution so providers receive function tools wrapped with middleware without importing internal packages. Apply wrapping immediately before the provider runs while preserving the options seen by outer middleware. Expose WithFuncCallID for provider-owned invocations and enable the capability in Copilot. Keep harness wrapping for additional tools. Test capability opt-in, dynamic tools, middleware behavior, option preservation, and repeated invocations. --- agent/agent.go | 15 +- agent/harness/toolautocall/autocall.go | 2 +- agent/middleware.go | 42 +++ agent/middleware_test.go | 301 ++++++++++++++++++ provider/copilotprovider/copilot.go | 13 +- .../copilotprovider/copilot_internal_test.go | 2 +- 6 files changed, 362 insertions(+), 13 deletions(-) create mode 100644 agent/middleware_test.go 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) }