diff --git a/pkg/tools/builtin/code_interpreter.go b/pkg/tools/builtin/code_interpreter.go index 6a92445..546ae7b 100644 --- a/pkg/tools/builtin/code_interpreter.go +++ b/pkg/tools/builtin/code_interpreter.go @@ -145,7 +145,7 @@ func (t *CodeInterpreterTool) Execute(ctx context.Context, argsJSON string) (too if timeout <= 0 { timeout = 30 * time.Second } - runCtx, cancel := context.WithTimeout(ctx, timeout) + runCtx, cancel := context.WithTimeoutCause(ctx, timeout, tools.DeadlineCause("code execution", timeout)) defer cancel() cmd := exec.CommandContext(runCtx, binary) @@ -162,7 +162,10 @@ func (t *CodeInterpreterTool) Execute(ctx context.Context, argsJSON string) (too runErr := cmd.Run() elapsed := time.Since(start) - timedOut := runCtx.Err() == context.DeadlineExceeded + // Our own deadline, not an enclosing one: a cancelled turn also leaves + // runCtx.Err() == DeadlineExceeded, and reporting that as the code + // having timed out tells the model the wrong thing. + timedOut := tools.TimedOut(runCtx) exitCode := 0 if runErr != nil { if ee, ok := runErr.(*exec.ExitError); ok { diff --git a/pkg/tools/builtin/generate_video.go b/pkg/tools/builtin/generate_video.go index 331834c..737ae6d 100644 --- a/pkg/tools/builtin/generate_video.go +++ b/pkg/tools/builtin/generate_video.go @@ -134,7 +134,7 @@ func (t *GenerateVideoTool) Execute(ctx context.Context, argsJSON string) (tools // Bound the polling loop with both the caller's context and our own // timeout — whichever fires first wins. - pollCtx, cancel := context.WithTimeout(ctx, defaultPollTimeout) + pollCtx, cancel := context.WithTimeoutCause(ctx, defaultPollTimeout, tools.DeadlineCause("video polling", defaultPollTimeout)) defer cancel() // Emit an initial progress tick immediately so the UI shows motion @@ -146,8 +146,10 @@ func (t *GenerateVideoTool) Execute(ctx context.Context, argsJSON string) (tools for !op.Done { select { case <-pollCtx.Done(): - if pollCtx.Err() == context.DeadlineExceeded { - return tools.Result{}, fmt.Errorf("tools: generate_video: timed out after %v", defaultPollTimeout) + // Only our own poll ceiling counts as a timeout here; a cancelled + // turn must not be reported as the generation having run long. + if tools.TimedOut(pollCtx) { + return tools.Result{}, fmt.Errorf("tools: generate_video: %w", context.Cause(pollCtx)) } return tools.Result{}, pollCtx.Err() case <-time.After(defaultPollInterval): diff --git a/pkg/tools/builtin/sql_agent.go b/pkg/tools/builtin/sql_agent.go index a35af6d..fd2fde6 100644 --- a/pkg/tools/builtin/sql_agent.go +++ b/pkg/tools/builtin/sql_agent.go @@ -832,7 +832,7 @@ func (t *executeSQLTool) executeRead(ctx context.Context, sqlStr string, emit fu queryCtx := ctx if t.queryTimeout > 0 { var cancel context.CancelFunc - queryCtx, cancel = context.WithTimeout(ctx, t.queryTimeout) + queryCtx, cancel = context.WithTimeoutCause(ctx, t.queryTimeout, tools.DeadlineCause("sql query", t.queryTimeout)) defer cancel() } @@ -842,7 +842,7 @@ func (t *executeSQLTool) executeRead(ctx context.Context, sqlStr string, emit fu return emit(SQLResult{ SQL: effectiveSQL, ExecutionMs: time.Since(start).Milliseconds(), - Error: err.Error(), + Error: queryErrText(queryCtx, err), }) } defer rows.Close() @@ -922,11 +922,23 @@ func (t *executeSQLTool) executeRead(ctx context.Context, sqlStr string, emit fu // skipped — auto-injecting LIMIT into UPDATE/DELETE has dialect-specific // semantics (MySQL accepts it, Postgres doesn't) and silently changes // the meaning of the statement. +// queryErrText renders a failed statement's error for the SQLResult the +// model reads. A driver error caused by our own query budget arrives as a +// bare "context deadline exceeded", which tells the model nothing it can +// act on; naming the budget tells it to narrow the query instead of +// retrying verbatim. An enclosing cancellation is reported as-is. +func queryErrText(queryCtx context.Context, err error) string { + if tools.TimedOut(queryCtx) { + return fmt.Sprintf("%v (%v) — narrow the query, add a LIMIT, or filter on an indexed column", context.Cause(queryCtx), err) + } + return err.Error() +} + func (t *executeSQLTool) executeMutation(ctx context.Context, sqlStr string, emit func(SQLResult) (tools.Result, error)) (tools.Result, error) { queryCtx := ctx if t.queryTimeout > 0 { var cancel context.CancelFunc - queryCtx, cancel = context.WithTimeout(ctx, t.queryTimeout) + queryCtx, cancel = context.WithTimeoutCause(ctx, t.queryTimeout, tools.DeadlineCause("sql statement", t.queryTimeout)) defer cancel() } @@ -934,7 +946,7 @@ func (t *executeSQLTool) executeMutation(ctx context.Context, sqlStr string, emi res, err := t.db.ExecContext(queryCtx, sqlStr) ms := time.Since(start).Milliseconds() if err != nil { - return emit(SQLResult{SQL: sqlStr, ExecutionMs: ms, Error: err.Error()}) + return emit(SQLResult{SQL: sqlStr, ExecutionMs: ms, Error: queryErrText(queryCtx, err)}) } // RowsAffected error is driver-specific (some return it when the // statement type has no meaningful affected count); treat as zero diff --git a/pkg/tools/deadline_test.go b/pkg/tools/deadline_test.go new file mode 100644 index 0000000..f4d1c98 --- /dev/null +++ b/pkg/tools/deadline_test.go @@ -0,0 +1,148 @@ +package tools + +import ( + "context" + "errors" + "strings" + "testing" + "time" +) + +func TestTimedOut_OwnDeadline(t *testing.T) { + ctx, cancel := context.WithTimeoutCause(context.Background(), time.Millisecond, DeadlineCause("tool call", time.Millisecond)) + defer cancel() + <-ctx.Done() + + if !TimedOut(ctx) { + t.Fatalf("TimedOut = false, want true (cause = %v)", context.Cause(ctx)) + } + if !errors.Is(context.Cause(ctx), ErrTimeout) { + t.Fatal("cause does not match ErrTimeout") + } +} + +// The distinction the sentinel exists for: an enclosing deadline expiring +// first must not read as this layer's own timeout, even though ctx.Err() is +// context.DeadlineExceeded either way. +func TestTimedOut_OuterDeadlineIsNotOurs(t *testing.T) { + outer, cancelOuter := context.WithTimeout(context.Background(), time.Millisecond) + defer cancelOuter() + inner, cancelInner := context.WithTimeoutCause(outer, time.Hour, DeadlineCause("tool call", time.Hour)) + defer cancelInner() + <-inner.Done() + + if !errors.Is(inner.Err(), context.DeadlineExceeded) { + t.Fatalf("inner.Err() = %v, want DeadlineExceeded (precondition)", inner.Err()) + } + if TimedOut(inner) { + t.Fatal("TimedOut = true for an outer deadline; the layers are indistinguishable") + } +} + +func TestTimedOut_CancellationIsNotATimeout(t *testing.T) { + outer, cancelOuter := context.WithCancel(context.Background()) + inner, cancelInner := context.WithTimeoutCause(outer, time.Hour, DeadlineCause("tool call", time.Hour)) + defer cancelInner() + cancelOuter() + <-inner.Done() + + if TimedOut(inner) { + t.Fatal("TimedOut = true for a cancelled parent, want false") + } +} + +func TestTimedOut_LiveContext(t *testing.T) { + ctx, cancel := context.WithTimeoutCause(context.Background(), time.Hour, DeadlineCause("tool call", time.Hour)) + defer cancel() + + if TimedOut(ctx) { + t.Fatal("TimedOut = true for a context that has not expired") + } +} + +func TestDeadlineCause_NamesTheBudget(t *testing.T) { + err := DeadlineCause("sql query", 250*time.Millisecond) + + if !errors.Is(err, ErrTimeout) { + t.Fatal("DeadlineCause does not match ErrTimeout") + } + for _, want := range []string{"sql query", "250ms"} { + if !strings.Contains(err.Error(), want) { + t.Fatalf("cause %q does not mention %q", err, want) + } + } +} + +// slowTool blocks until its context ends, then reports why. +type slowTool struct{} + +func (slowTool) Descriptor() ToolDescriptor { + return ToolDescriptor{Name: "slow", Description: "blocks", Display: DefaultDisplay("slow", "blocks")} +} + +func (slowTool) Execute(ctx context.Context, _ string) (Result, error) { + <-ctx.Done() + return Result{}, ctx.Err() +} + +func TestWithTimeout_AttributesOwnDeadline(t *testing.T) { + tool := Chain(slowTool{}, WithTimeout(10*time.Millisecond)) + + _, err := tool.Execute(context.Background(), "{}") + if err == nil { + t.Fatal("expected an error") + } + if !errors.Is(err, ErrTimeout) { + t.Fatalf("err = %v, want it to match ErrTimeout", err) + } + // Existing matches must keep working: the tool's own error is preserved. + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("err = %v, want it to still match context.DeadlineExceeded", err) + } + if !strings.Contains(err.Error(), "10ms") { + t.Fatalf("err %q does not name the elapsed budget", err) + } +} + +// A turn cancelled by the caller must not be reported as the tool having +// exceeded its own budget. +func TestWithTimeout_LeavesOuterCancellationUnattributed(t *testing.T) { + tool := Chain(slowTool{}, WithTimeout(time.Hour)) + ctx, cancel := context.WithCancel(context.Background()) + go func() { + time.Sleep(5 * time.Millisecond) + cancel() + }() + + _, err := tool.Execute(ctx, "{}") + if err == nil { + t.Fatal("expected an error") + } + if errors.Is(err, ErrTimeout) { + t.Fatalf("err = %v, want no ErrTimeout attribution for a cancelled parent", err) + } + if !errors.Is(err, context.Canceled) { + t.Fatalf("err = %v, want context.Canceled", err) + } +} + +// fastTool returns immediately, so the timeout path is never taken. +type fastTool struct{} + +func (fastTool) Descriptor() ToolDescriptor { + return ToolDescriptor{Name: "fast", Description: "returns", Display: DefaultDisplay("fast", "returns")} +} + +func (fastTool) Execute(context.Context, string) (Result, error) { return Text("ok"), nil } + +func TestWithTimeout_SuccessIsUntouched(t *testing.T) { + tool := Chain(fastTool{}, WithTimeout(time.Hour)) + + res, err := tool.Execute(context.Background(), `{"a":1}`) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if res.Text != "ok" { + t.Fatalf("Text = %q, want %q", res.Text, "ok") + } +} diff --git a/pkg/tools/errors.go b/pkg/tools/errors.go index 15cfe5e..714d836 100644 --- a/pkg/tools/errors.go +++ b/pkg/tools/errors.go @@ -1,6 +1,11 @@ package tools -import "errors" +import ( + "context" + "errors" + "fmt" + "time" +) // Error classification sentinels. Tool authors wrap their real error with // one of these to let the agent loop and clients distinguish the failure @@ -32,8 +37,41 @@ var ( // ErrPermanent flags non-recoverable downstream failures. ErrPermanent = errors.New("tools: permanent error") + + // ErrTimeout flags a deadline the framework itself imposed — a tool-call + // budget, a query timeout, a poll ceiling. It is deliberately distinct + // from context.DeadlineExceeded, which any enclosing context also + // produces: a layer that sets its own deadline needs to know whether the + // deadline that fired was *its own* before reporting one. Attach it with + // DeadlineCause and test with TimedOut. + ErrTimeout = errors.New("tools: deadline exceeded") ) +// DeadlineCause builds the cause value to hand to context.WithTimeoutCause, +// naming the budget that will have elapsed. what should read as the thing +// being bounded ("tool call", "sql query"), not the tool's name. +// +// ctx, cancel := context.WithTimeoutCause(ctx, d, tools.DeadlineCause("sql query", d)) +// +// The returned error satisfies errors.Is(err, ErrTimeout) and carries the +// duration in its message, so a model-facing report can name the budget +// without the reporting site holding it. +func DeadlineCause(what string, d time.Duration) error { + return fmt.Errorf("%w: %s exceeded its %s budget", ErrTimeout, what, d) +} + +// TimedOut reports whether ctx is done because a deadline set with +// DeadlineCause elapsed. It answers "was it *my* deadline?" — an enclosing +// context that expires or is cancelled first leaves its own cause in place, +// so this returns false and the caller correctly reports cancellation +// instead of claiming its own timeout fired. +// +// Prefer this over comparing ctx.Err() to context.DeadlineExceeded, which +// cannot tell the two apart. +func TimedOut(ctx context.Context) bool { + return errors.Is(context.Cause(ctx), ErrTimeout) +} + // ClassifyError returns the classification sentinel that err unwraps to, // or nil when err matches none of them. Consumers use this for concise // branching in error handlers: diff --git a/pkg/tools/middleware.go b/pkg/tools/middleware.go index 86cacb4..892dc5d 100644 --- a/pkg/tools/middleware.go +++ b/pkg/tools/middleware.go @@ -156,14 +156,25 @@ func WithTiming(onDone func(toolName string, d time.Duration, err error)) Middle // WithTimeout wraps each Execute call with a per-call deadline. // The parent context's deadline, if shorter, still takes precedence. +// +// A failure caused by this middleware's own deadline is wrapped so the +// error names the elapsed budget and satisfies errors.Is(err, ErrTimeout). +// An outer deadline or a cancellation is left untouched: the tool did not +// exceed *its* budget, and reporting otherwise would misattribute a turn +// the caller cut short. The wrapped error still unwraps to whatever the +// tool returned, so existing context.DeadlineExceeded matches keep working. func WithTimeout(d time.Duration) Middleware { return func(next Tool) Tool { return &wrappedTool{ Tool: next, executeFn: func(ctx context.Context, argsJSON string) (Result, error) { - tCtx, cancel := context.WithTimeout(ctx, d) + tCtx, cancel := context.WithTimeoutCause(ctx, d, DeadlineCause("tool call", d)) defer cancel() - return next.Execute(tCtx, argsJSON) + result, err := next.Execute(tCtx, argsJSON) + if err != nil && TimedOut(tCtx) { + return result, fmt.Errorf("%w: %w", context.Cause(tCtx), err) + } + return result, err }, } }