Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 5 additions & 2 deletions pkg/tools/builtin/code_interpreter.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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 {
Expand Down
8 changes: 5 additions & 3 deletions pkg/tools/builtin/generate_video.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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):
Expand Down
20 changes: 16 additions & 4 deletions pkg/tools/builtin/sql_agent.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
}

Expand All @@ -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()
Expand Down Expand Up @@ -922,19 +922,31 @@ 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()
}

start := time.Now()
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
Expand Down
148 changes: 148 additions & 0 deletions pkg/tools/deadline_test.go
Original file line number Diff line number Diff line change
@@ -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")
}
}
40 changes: 39 additions & 1 deletion pkg/tools/errors.go
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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:
Expand Down
15 changes: 13 additions & 2 deletions pkg/tools/middleware.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
},
}
}
Expand Down
Loading