diff --git a/cmd/wasm/toolsets_test.go b/cmd/wasm/toolsets_test.go index 8db8a52723..85d5bf5d9a 100644 --- a/cmd/wasm/toolsets_test.go +++ b/cmd/wasm/toolsets_test.go @@ -327,7 +327,7 @@ agents: - type: echo ` echo := &echoToolSet{} - model := newScriptedModel("mock/root", toolTurn("run_tools_with_javascript", `{"script":"return echo({text: \"ping\"}) + \"!\""}`), textTurn("done")) + model := newScriptedModel("mock/root", toolTurn("run_tools_with_javascript", `{"script":"return (await echo({text: \"ping\"})) + \"!\""}`), textTurn("done")) s := openTestSession(t, testHost(echo, map[string]provider.Provider{"root": model}), sessionOptions{YAML: yaml, AutoApprove: true}) var c collectingEmitter diff --git a/docs/features/code-mode/index.md b/docs/features/code-mode/index.md index 0fa6a20634..b8ae8055ee 100644 --- a/docs/features/code-mode/index.md +++ b/docs/features/code-mode/index.md @@ -12,7 +12,7 @@ _Let an agent write JavaScript that orchestrates several tool calls in one turn By default, a model calls one tool at a time: it emits a tool call, waits for the result, then decides what to call next. For a task that chains many tool calls together — "list every open issue, then for each one fetch its comments, then summarize" — that means one model round-trip per step. -**Code Mode** replaces the agent's individual tools with a single tool, `run_tools_with_javascript`, that runs a JavaScript script. Every tool the agent would otherwise call directly is exposed to that script as a plain JavaScript function (synchronous — no `await`/`async` needed). The model writes a script that calls as many of them as it needs, combines and filters the results, and returns a single string — all in one tool call. +**Code Mode** replaces the agent's individual tools with a single tool, `run_tools_with_javascript`, that runs a JavaScript script. Every tool the agent would otherwise call directly is exposed to that script as a JavaScript function returning a Promise. Scripts support top-level `await`; use `await` for dependent calls and `Promise.all` to run independent calls in parallel. The model writes a script that calls as many of them as it needs, combines and filters the results, and returns a single string — all in one tool call. ## Enabling Code Mode @@ -41,6 +41,20 @@ To force Code Mode for every agent in a run regardless of their individual confi $ docker agent run agent.yaml --code-mode-tools ``` +## Parallel Tool Calls + +Each tool call starts immediately and returns a Promise. Await a call before using its result, or group independent calls with `Promise.all`: + +```javascript +const results = await Promise.all([ + SearchIssues({query: "repo:docker/docker-agent is:open is:issue"}), + SearchIssues({query: "repo:docker/docker-agent is:open is:pr"}), +]); +return results.join("\n"); +``` + +Use the function names and arguments listed in the tool description. Tool failures reject their Promises and can be handled with `try`/`catch` or `Promise.allSettled`. Unhandled rejections include tool-call history in the response. All started tool calls finish before the script response is returned, unless execution is cancelled; parallel calls are not rolled back if one fails. Existing scripts must await tool results before inspecting or combining them. + ## When It Helps Code Mode is worth enabling when an agent's task typically needs **many tool calls chained together**, especially with conditional logic or filtering in between — for example, paging through a large result set, cross-referencing several API calls, or reducing a large payload down to the few fields the model actually needs before it ever sees them. Each of those becomes one model turn instead of many, which cuts both latency and token spend on tool-call/response round-trips. diff --git a/pkg/tools/codemode/codemode.go b/pkg/tools/codemode/codemode.go index 0ba5d73fb7..c1906908fe 100644 --- a/pkg/tools/codemode/codemode.go +++ b/pkg/tools/codemode/codemode.go @@ -20,7 +20,9 @@ and manipulate the results before returning them. Instructions: - The script has access to all the tools as plain javascript functions. - - "await"/"async" are never needed. All the tool calls are synchronous. + - Every tool call returns a Promise. Use "await" to get its result; top-level "await" is supported. + - Run independent tool calls in parallel with "await Promise.all([ToolA(args), ToolB(args)])". + - Await dependent tool calls sequentially. Tool failures reject their Promises and can be caught with try/catch. - The script must return a string result. - "console.*" functions can be used to print debug information. - It's often encouraged to group multiple tool calls in a single script to reduce the number of LLM interactions. diff --git a/pkg/tools/codemode/codemode_test.go b/pkg/tools/codemode/codemode_test.go index a183c4430c..45a5a0b5f3 100644 --- a/pkg/tools/codemode/codemode_test.go +++ b/pkg/tools/codemode/codemode_test.go @@ -115,7 +115,7 @@ func TestCodeModeTool_TypeScriptDeclarationsInDescription(t *testing.T) { assert.Contains(t, allTools[0].Description, "interface FindItemInput") assert.Contains(t, allTools[0].Description, "id: string;") assert.Contains(t, allTools[0].Description, "type FindItemOutput = boolean;") - assert.Contains(t, allTools[0].Description, "declare function FindItem(args: FindItemInput): FindItemOutput;") + assert.Contains(t, allTools[0].Description, "declare function FindItem(args: FindItemInput): Promise;") assert.NotContains(t, allTools[0].Description, "Where Input follows the following JSON schema") } @@ -198,7 +198,7 @@ func TestCodeModeTool_CallToolWithNonIdentifierName(t *testing.T) { allTools, err := tool.Tools(t.Context()) require.NoError(t, err) require.Len(t, allTools, 1) - assert.Contains(t, allTools[0].Description, "declare function HelloWorld(args: HelloWorldInput): HelloWorldOutput;") + assert.Contains(t, allTools[0].Description, "declare function HelloWorld(args: HelloWorldInput): Promise;") result, err := allTools[0].Handler(t.Context(), tools.ToolCall{ Function: tools.FunctionCall{ @@ -860,7 +860,7 @@ func TestCodeModeTool_FailureIncludesToolCalls(t *testing.T) { // Script calls tools successfully but then throws a runtime error result, err := allTools[0].Handler(t.Context(), tools.ToolCall{ Function: tools.FunctionCall{ - Arguments: `{"script":"var a = first_tool(); var b = second_tool(); throw new Error('runtime error');"}`, + Arguments: `{"script":"var a = await first_tool(); var b = await second_tool(); throw new Error('runtime error');"}`, }, }, tools.NopRuntime{}) require.NoError(t, err) @@ -948,7 +948,7 @@ func TestCodeModeTool_FailureIncludesToolArguments(t *testing.T) { result, err := allTools[0].Handler(t.Context(), tools.ToolCall{ Function: tools.FunctionCall{ - Arguments: `{"script":"tool_with_args({'value': 'test123'}); throw new Error('forced error');"}`, + Arguments: `{"script":"await tool_with_args({'value': 'test123'}); throw new Error('forced error');"}`, }, }, tools.NopRuntime{}) require.NoError(t, err) diff --git a/pkg/tools/codemode/exec.go b/pkg/tools/codemode/exec.go index 899c60641e..705ceba60f 100644 --- a/pkg/tools/codemode/exec.go +++ b/pkg/tools/codemode/exec.go @@ -6,6 +6,7 @@ import ( "crypto/sha256" "encoding/hex" "encoding/json" + "errors" "fmt" "slices" @@ -37,13 +38,36 @@ type toolCallTracker struct { calls []ToolCallInfo } -func (t *toolCallTracker) record(info ToolCallInfo) { - t.calls = append(t.calls, info) +type toolCompletion struct { + index int + info ToolCallInfo + settle func() error +} + +type toolEventLoop struct { + vm *goja.Runtime + tracker *toolCallTracker + completions chan toolCompletion + pending int } func (c *codeModeTool) runJavascript(ctx context.Context, rt tools.Runtime, script string) (ScriptResult, error) { + ctx, cancel := context.WithCancel(ctx) + defer cancel() + vm := goja.New() + stopInterrupt := context.AfterFunc(ctx, func() { vm.Interrupt(ctx.Err()) }) + defer stopInterrupt() tracker := &toolCallTracker{} + loop := &toolEventLoop{vm: vm, tracker: tracker, completions: make(chan toolCompletion)} + unhandled := make(map[*goja.Promise]bool) + vm.SetPromiseRejectionTracker(func(p *goja.Promise, operation goja.PromiseRejectionOperation) { + if operation == goja.PromiseRejectionReject { + unhandled[p] = true + } else { + delete(unhandled, p) + } + }) // Always stamp a hash + length so dashboards can correlate // identical scripts ("model ran the same script 200 times this @@ -85,7 +109,7 @@ func (c *codeModeTool) runJavascript(ctx context.Context, rt tools.Runtime, scri } for _, tool := range allTools { - call := callTool(ctx, rt, tool, tracker) + call := loop.callTool(ctx, rt, tool) _ = vm.Set(tool.Name, call) if name := typeName(tool.Name); name != tool.Name { _ = vm.Set(name, call) @@ -93,11 +117,40 @@ func (c *codeModeTool) runJavascript(ctx context.Context, rt tools.Runtime, scri } } - // Wrap the user script in an IIFE to allow top-level returns. - script = "(() => {\n" + script + "\n})()" + // Wrap the script to support top-level await and return. + script = "(async () => {\n" + script + "\n})()" // Run the script. v, err := vm.RunString(script) + if err == nil { + promise := v.Export().(*goja.Promise) + for loop.pending > 0 && err == nil { + select { + case completion := <-loop.completions: + loop.pending-- + tracker.calls[completion.index] = completion.info + err = completion.settle() + case <-ctx.Done(): + err = ctx.Err() + } + } + if err == nil { + switch promise.State() { + case goja.PromiseStateFulfilled: + v = promise.Result() + case goja.PromiseStateRejected: + err = fmt.Errorf("%s", promise.Result().String()) + case goja.PromiseStatePending: + err = errors.New("script returned a Promise that cannot settle: no pending tool calls") + } + } + if err == nil { + for p := range unhandled { + err = fmt.Errorf("unhandled Promise rejection: %s", p.Result().String()) + break + } + } + } if err != nil { // Script execution failed - include tool call history to help LLM understand what went wrong return ScriptResult{ @@ -121,25 +174,29 @@ func (c *codeModeTool) runJavascript(ctx context.Context, rt tools.Runtime, scri }, nil } -// callTool wraps a tool as a goja-callable function. rt is forwarded to the -// inner handler so nested tools keep their runtime capabilities (streaming -// output, recall) when invoked from a script. -func callTool(ctx context.Context, rt tools.Runtime, tool tools.Tool, tracker *toolCallTracker) func(args map[string]any) (string, error) { - return func(args map[string]any) (string, error) { - output, filtered, err := invokeTool(ctx, rt, tool, args) - - info := ToolCallInfo{ - Name: tool.Name, - Arguments: filtered, - } - if err != nil { - info.Error = err.Error() - } else { - info.Result = output - } - tracker.record(info) +// Tool handlers run concurrently, but only the event loop touches the VM and tracker. +func (l *toolEventLoop) callTool(ctx context.Context, rt tools.Runtime, tool tools.Tool) func(args map[string]any) *goja.Promise { + return func(args map[string]any) *goja.Promise { + promise, resolve, reject := l.vm.NewPromise() + index := len(l.tracker.calls) + l.tracker.calls = append(l.tracker.calls, ToolCallInfo{Name: tool.Name, Arguments: args}) + l.pending++ + + go func() { + output, filtered, err := invokeTool(ctx, rt, tool, args) + info := ToolCallInfo{Name: tool.Name, Arguments: filtered, Result: output} + settle := func() error { return resolve(output) } + if err != nil { + info.Error = err.Error() + settle = func() error { return reject(l.vm.NewGoError(err)) } + } + select { + case l.completions <- toolCompletion{index: index, info: info, settle: settle}: + case <-ctx.Done(): + } + }() - return output, err + return promise } } diff --git a/pkg/tools/codemode/exec_test.go b/pkg/tools/codemode/exec_test.go index 0089f6ac3e..d8438b90f2 100644 --- a/pkg/tools/codemode/exec_test.go +++ b/pkg/tools/codemode/exec_test.go @@ -1,7 +1,9 @@ package codemode import ( + "context" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -56,3 +58,110 @@ func TestRunJavascript_no_result(t *testing.T) { assert.Empty(t, result.StdOut) assert.Empty(t, result.StdErr) } + +func TestRunJavascript_ParallelTools(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + firstStarted := make(chan struct{}) + secondFinished := make(chan struct{}) + tool := Wrap(&testToolSet{tools: []tools.Tool{ + {Name: "first", Handler: tools.NewHandler(func(ctx context.Context, args map[string]any) (*tools.ToolCallResult, error) { + close(firstStarted) + select { + case <-secondFinished: + return tools.ResultSuccess("first"), nil + case <-ctx.Done(): + return nil, ctx.Err() + } + })}, + {Name: "second", Handler: tools.NewHandler(func(ctx context.Context, args map[string]any) (*tools.ToolCallResult, error) { + select { + case <-firstStarted: + close(secondFinished) + return tools.ResultSuccess("second"), nil + case <-ctx.Done(): + return nil, ctx.Err() + } + })}, + }}).(*codeModeTool) + result, err := tool.runJavascript(ctx, tools.NopRuntime{}, ` + const a = first(); + const b = Second(); + if (!(a instanceof Promise) || !(b instanceof Promise)) throw new Error("not Promises"); + const results = await Promise.all([a, b]); + console.log(results.join(",")); + throw new Error("forced failure"); + `) + require.NoError(t, err) + assert.Contains(t, result.Value, "forced failure") + assert.Equal(t, "first,second\n", result.StdOut) + require.Len(t, result.ToolCalls, 2) + assert.Equal(t, "first", result.ToolCalls[0].Name) + assert.Equal(t, "first", result.ToolCalls[0].Result) + assert.Equal(t, "second", result.ToolCalls[1].Name) + assert.Equal(t, "second", result.ToolCalls[1].Result) +} + +func TestRunJavascript_Promises(t *testing.T) { + t.Parallel() + + tests := []struct { + name, script, want string + failed bool + }{ + {name: "sequential awaits", script: `const a = await Echo({message: "hello"}); return await echo({message: a + " world"});`, want: "hello world"}, + {name: "then callback", script: `return echo({message: "hello"}).then(value => value + " world");`, want: "hello world"}, + {name: "caught rejection", script: `try { await fail(); } catch (e) { return "caught: " + e.message; }`, want: "caught: assert.AnError"}, + {name: "all settled", script: `const results = await Promise.allSettled([fail(), echo({message: "ok"})]); return results.map(r => r.status).join(",");`, want: "rejected,fulfilled"}, + {name: "all rejection", script: `return await Promise.all([fail(), echo({message: "ok"})]);`, want: "assert.AnError", failed: true}, + {name: "unhandled rejection", script: `fail(); return "ignored";`, want: "unhandled Promise rejection", failed: true}, + {name: "unavailable handler", script: `return await unavailable();`, want: `tool "unavailable" is not available in code mode`, failed: true}, + {name: "unsettleable promise", script: `return new Promise(() => {});`, want: "no pending tool calls", failed: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + tool := Wrap(&testToolSet{tools: []tools.Tool{ + {Name: "echo", Handler: tools.NewHandler(func(ctx context.Context, args map[string]any) (*tools.ToolCallResult, error) { + return tools.ResultSuccess(args["message"].(string)), nil + })}, + {Name: "fail", Handler: tools.NewHandler(func(ctx context.Context, args map[string]any) (*tools.ToolCallResult, error) { + return nil, assert.AnError + })}, + {Name: "unavailable"}, + }}).(*codeModeTool) + result, err := tool.runJavascript(t.Context(), tools.NopRuntime{}, tt.script) + require.NoError(t, err) + assert.Contains(t, result.Value, tt.want) + if !tt.failed { + assert.Empty(t, result.ToolCalls) + } + if tt.name == "all rejection" { + require.Len(t, result.ToolCalls, 2) + assert.Equal(t, "ok", result.ToolCalls[1].Result) + } + }) + } +} + +func TestRunJavascript_Cancellation(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + finished := make(chan struct{}) + tool := Wrap(&testToolSet{tools: []tools.Tool{{Name: "wait", Handler: tools.NewHandler(func(ctx context.Context, args map[string]any) (*tools.ToolCallResult, error) { + cancel() + <-ctx.Done() + close(finished) + return nil, ctx.Err() + })}}}).(*codeModeTool) + result, err := tool.runJavascript(ctx, tools.NopRuntime{}, `return await wait();`) + require.NoError(t, err) + assert.Contains(t, result.Value, "context canceled") + select { + case <-finished: + case <-time.After(5 * time.Second): + t.Fatal("tool handler did not stop") + } +} diff --git a/pkg/tools/codemode/functions.go b/pkg/tools/codemode/functions.go index c69b53ece3..07a0b9f992 100644 --- a/pkg/tools/codemode/functions.go +++ b/pkg/tools/codemode/functions.go @@ -31,7 +31,7 @@ func toolToTypeScript(tool tools.Tool) string { fmt.Fprintf(&doc, "type %s = %s;\n\n", inputName, schemaType(input, input, 0)) } fmt.Fprintf(&doc, "type %s = %s;\n\n", outputName, schemaType(output, output, 0)) - fmt.Fprintf(&doc, "declare function %s(args: %s): %s;\n", baseName, inputName, outputName) + fmt.Fprintf(&doc, "declare function %s(args: %s): Promise<%s>;\n", baseName, inputName, outputName) return doc.String() } diff --git a/pkg/tools/codemode/functions_test.go b/pkg/tools/codemode/functions_test.go index 7c930fc7ae..cade660a7c 100644 --- a/pkg/tools/codemode/functions_test.go +++ b/pkg/tools/codemode/functions_test.go @@ -38,7 +38,7 @@ interface CreateTodoInput { type CreateTodoOutput = string; -declare function CreateTodo(args: CreateTodoInput): CreateTodoOutput; +declare function CreateTodo(args: CreateTodoInput): Promise; `, declaration) } @@ -62,7 +62,7 @@ type ExampleToolInput = string; type ExampleToolOutput = boolean; -declare function ExampleTool(args: ExampleToolInput): ExampleToolOutput; +declare function ExampleTool(args: ExampleToolInput): Promise; `, }, { @@ -98,7 +98,7 @@ interface ExampleToolInput { type ExampleToolOutput = number[] | null; -declare function ExampleTool(args: ExampleToolInput): ExampleToolOutput; +declare function ExampleTool(args: ExampleToolInput): Promise; `, }, { @@ -126,7 +126,7 @@ type ExampleToolOutput = { active?: boolean; }; -declare function ExampleTool(args: ExampleToolInput): ExampleToolOutput; +declare function ExampleTool(args: ExampleToolInput): Promise; `, }, { @@ -152,7 +152,7 @@ interface ExampleToolInput { type ExampleToolOutput = "ok"; -declare function ExampleTool(args: ExampleToolInput): ExampleToolOutput; +declare function ExampleTool(args: ExampleToolInput): Promise; `, }, } @@ -324,5 +324,5 @@ func TestToolToTypeScriptNestedAndNullableTypes(t *testing.T) { }; }`) assert.Contains(t, declaration, "type SearchItemsOutput = number[] | null;") - assert.Contains(t, declaration, "declare function SearchItems(args: SearchItemsInput): SearchItemsOutput;") + assert.Contains(t, declaration, "declare function SearchItems(args: SearchItemsInput): Promise;") }