diff --git a/coderd/x/chatd/chatloop/chatloop.go b/coderd/x/chatd/chatloop/chatloop.go index 3b768cc41e..7211ea29c6 100644 --- a/coderd/x/chatd/chatloop/chatloop.go +++ b/coderd/x/chatd/chatloop/chatloop.go @@ -24,6 +24,7 @@ import ( "github.com/coder/coder/v2/coderd/x/chatd/chatprompt" "github.com/coder/coder/v2/coderd/x/chatd/chatretry" "github.com/coder/coder/v2/codersdk" + "github.com/coder/quartz" ) const ( @@ -100,6 +101,10 @@ type RunOptions struct { // first stream part before the attempt is canceled and // retried. Zero uses the production default. StartupTimeout time.Duration + // Clock creates startup guard timers. In production use a + // real clock; tests can inject quartz.NewMock(t) to make + // startup timeout behavior deterministic. + Clock quartz.Clock ActiveTools []string ContextLimitFallback int64 @@ -289,6 +294,9 @@ func Run(ctx context.Context, opts RunOptions) error { if opts.StartupTimeout <= 0 { opts.StartupTimeout = defaultStartupTimeout } + if opts.Clock == nil { + opts.Clock = quartz.NewReal() + } publishMessagePart := func(role codersdk.ChatMessageRole, part codersdk.ChatMessagePart) { if opts.PublishMessagePart == nil { @@ -364,6 +372,7 @@ func Run(ctx context.Context, opts RunOptions) error { attempt, streamErr := guardedStream( retryCtx, opts.Model.Provider(), + opts.Clock, opts.StartupTimeout, func(attemptCtx context.Context) (fantasy.StreamResponse, error) { return opts.Model.Stream(attemptCtx, call) @@ -660,17 +669,18 @@ type guardedAttempt struct { // stream startup. Exactly one outcome wins: the timer cancels // the attempt, or the first-part path disarms the timer. type startupGuard struct { - timer *time.Timer + timer *quartz.Timer cancel context.CancelCauseFunc once sync.Once } func newStartupGuard( + clock quartz.Clock, timeout time.Duration, cancel context.CancelCauseFunc, ) *startupGuard { guard := &startupGuard{cancel: cancel} - guard.timer = time.AfterFunc(timeout, guard.onTimeout) + guard.timer = clock.AfterFunc(timeout, guard.onTimeout, "startupGuard") return guard } @@ -707,11 +717,12 @@ func classifyStartupTimeout( func guardedStream( parent context.Context, provider string, + clock quartz.Clock, timeout time.Duration, openStream func(context.Context) (fantasy.StreamResponse, error), ) (guardedAttempt, error) { attemptCtx, cancelAttempt := context.WithCancelCause(parent) - guard := newStartupGuard(timeout, cancelAttempt) + guard := newStartupGuard(clock, timeout, cancelAttempt) var releaseOnce sync.Once release := func() { releaseOnce.Do(func() { diff --git a/coderd/x/chatd/chatloop/chatloop_test.go b/coderd/x/chatd/chatloop/chatloop_test.go index a3bd55decf..b96292fd65 100644 --- a/coderd/x/chatd/chatloop/chatloop_test.go +++ b/coderd/x/chatd/chatloop/chatloop_test.go @@ -19,10 +19,24 @@ import ( "github.com/coder/coder/v2/coderd/x/chatd/chaterror" "github.com/coder/coder/v2/coderd/x/chatd/chatretry" "github.com/coder/coder/v2/codersdk" + "github.com/coder/coder/v2/testutil" + "github.com/coder/quartz" ) const activeToolName = "read_file" +func awaitRunResult(ctx context.Context, t *testing.T, done <-chan error) error { + t.Helper() + + select { + case err := <-done: + return err + case <-ctx.Done(): + t.Fatal("timed out waiting for Run to complete") + return nil + } +} + func TestRun_ActiveToolsPrepareBehavior(t *testing.T) { t.Parallel() @@ -202,7 +216,7 @@ func TestStartupGuard_DisarmAndFireRace(t *testing.T) { for range 128 { var cancels atomic.Int32 - guard := newStartupGuard(time.Hour, func(err error) { + guard := newStartupGuard(quartz.NewReal(), time.Hour, func(err error) { if errors.Is(err, errStartupTimeout) { cancels.Add(1) } @@ -240,7 +254,7 @@ func TestStartupGuard_DisarmPreservesPermanentError(t *testing.T) { attemptCtx, cancelAttempt := context.WithCancelCause(context.Background()) defer cancelAttempt(nil) - guard := newStartupGuard(time.Hour, cancelAttempt) + guard := newStartupGuard(quartz.NewReal(), time.Hour, cancelAttempt) guard.Disarm() guard.onTimeout() @@ -259,6 +273,16 @@ func TestRun_RetriesStartupTimeoutWhileOpeningStream(t *testing.T) { const startupTimeout = 5 * time.Millisecond + ctx, cancel := context.WithTimeout( + context.Background(), + testutil.WaitShort, + ) + defer cancel() + + mClock := quartz.NewMock(t) + trap := mClock.Trap().AfterFunc("startupGuard") + defer trap.Close() + attempts := 0 attemptCause := make(chan error, 1) var retries []chatretry.ClassifiedError @@ -278,23 +302,32 @@ func TestRun_RetriesStartupTimeoutWhileOpeningStream(t *testing.T) { }, } - err := Run(context.Background(), RunOptions{ - Model: model, - MaxSteps: 1, - StartupTimeout: startupTimeout, - PersistStep: func(_ context.Context, _ PersistedStep) error { - return nil - }, - OnRetry: func( - _ int, - _ error, - classified chatretry.ClassifiedError, - _ time.Duration, - ) { - retries = append(retries, classified) - }, - }) - require.NoError(t, err) + done := make(chan error, 1) + go func() { + done <- Run(context.Background(), RunOptions{ + Model: model, + MaxSteps: 1, + StartupTimeout: startupTimeout, + Clock: mClock, + PersistStep: func(_ context.Context, _ PersistedStep) error { + return nil + }, + OnRetry: func( + _ int, + _ error, + classified chatretry.ClassifiedError, + _ time.Duration, + ) { + retries = append(retries, classified) + }, + }) + }() + + trap.MustWait(ctx).MustRelease(ctx) + mClock.Advance(startupTimeout).MustWait(ctx) + trap.MustWait(ctx).MustRelease(ctx) + + require.NoError(t, awaitRunResult(ctx, t, done)) require.Equal(t, 2, attempts) require.Len(t, retries, 1) require.Equal(t, chaterror.KindStartupTimeout, retries[0].Kind) @@ -305,7 +338,12 @@ func TestRun_RetriesStartupTimeoutWhileOpeningStream(t *testing.T) { "OpenAI did not start responding in time.", retries[0].Message, ) - require.ErrorIs(t, <-attemptCause, errStartupTimeout) + select { + case cause := <-attemptCause: + require.ErrorIs(t, cause, errStartupTimeout) + case <-ctx.Done(): + t.Fatal("timed out waiting for startup timeout cause") + } } func TestRun_RetriesStartupTimeoutBeforeFirstPart(t *testing.T) { @@ -313,6 +351,16 @@ func TestRun_RetriesStartupTimeoutBeforeFirstPart(t *testing.T) { const startupTimeout = 5 * time.Millisecond + ctx, cancel := context.WithTimeout( + context.Background(), + testutil.WaitShort, + ) + defer cancel() + + mClock := quartz.NewMock(t) + trap := mClock.Trap().AfterFunc("startupGuard") + defer trap.Close() + attempts := 0 attemptCause := make(chan error, 1) var retries []chatretry.ClassifiedError @@ -337,23 +385,32 @@ func TestRun_RetriesStartupTimeoutBeforeFirstPart(t *testing.T) { }, } - err := Run(context.Background(), RunOptions{ - Model: model, - MaxSteps: 1, - StartupTimeout: startupTimeout, - PersistStep: func(_ context.Context, _ PersistedStep) error { - return nil - }, - OnRetry: func( - _ int, - _ error, - classified chatretry.ClassifiedError, - _ time.Duration, - ) { - retries = append(retries, classified) - }, - }) - require.NoError(t, err) + done := make(chan error, 1) + go func() { + done <- Run(context.Background(), RunOptions{ + Model: model, + MaxSteps: 1, + StartupTimeout: startupTimeout, + Clock: mClock, + PersistStep: func(_ context.Context, _ PersistedStep) error { + return nil + }, + OnRetry: func( + _ int, + _ error, + classified chatretry.ClassifiedError, + _ time.Duration, + ) { + retries = append(retries, classified) + }, + }) + }() + + trap.MustWait(ctx).MustRelease(ctx) + mClock.Advance(startupTimeout).MustWait(ctx) + trap.MustWait(ctx).MustRelease(ctx) + + require.NoError(t, awaitRunResult(ctx, t, done)) require.Equal(t, 2, attempts) require.Len(t, retries, 1) require.Equal(t, chaterror.KindStartupTimeout, retries[0].Kind) @@ -364,7 +421,12 @@ func TestRun_RetriesStartupTimeoutBeforeFirstPart(t *testing.T) { "OpenAI did not start responding in time.", retries[0].Message, ) - require.ErrorIs(t, <-attemptCause, errStartupTimeout) + select { + case cause := <-attemptCause: + require.ErrorIs(t, cause, errStartupTimeout) + case <-ctx.Done(): + t.Fatal("timed out waiting for startup timeout cause") + } } func TestRun_FirstPartDisarmsStartupTimeout(t *testing.T) { @@ -372,8 +434,19 @@ func TestRun_FirstPartDisarmsStartupTimeout(t *testing.T) { const startupTimeout = 5 * time.Millisecond + ctx, cancel := context.WithTimeout( + context.Background(), + testutil.WaitShort, + ) + defer cancel() + + mClock := quartz.NewMock(t) + trap := mClock.Trap().AfterFunc("startupGuard") + attempts := 0 retried := false + firstPartYielded := make(chan struct{}, 1) + continueStream := make(chan struct{}) model := &loopTestModel{ provider: "openai", streamFn: func(ctx context.Context, _ fantasy.Call) (fantasy.StreamResponse, error) { @@ -382,18 +455,19 @@ func TestRun_FirstPartDisarmsStartupTimeout(t *testing.T) { if !yield(fantasy.StreamPart{Type: fantasy.StreamPartTypeTextStart, ID: "text-1"}) { return } - - timer := time.NewTimer(startupTimeout * 2) - defer timer.Stop() + select { + case firstPartYielded <- struct{}{}: + default: + } select { + case <-continueStream: case <-ctx.Done(): _ = yield(fantasy.StreamPart{ Type: fantasy.StreamPartTypeError, Error: ctx.Err(), }) return - case <-timer.C: } parts := []fantasy.StreamPart{ @@ -410,23 +484,40 @@ func TestRun_FirstPartDisarmsStartupTimeout(t *testing.T) { }, } - err := Run(context.Background(), RunOptions{ - Model: model, - MaxSteps: 1, - StartupTimeout: startupTimeout, - PersistStep: func(_ context.Context, _ PersistedStep) error { - return nil - }, - OnRetry: func( - _ int, - _ error, - _ chatretry.ClassifiedError, - _ time.Duration, - ) { - retried = true - }, - }) - require.NoError(t, err) + done := make(chan error, 1) + go func() { + done <- Run(context.Background(), RunOptions{ + Model: model, + MaxSteps: 1, + StartupTimeout: startupTimeout, + Clock: mClock, + PersistStep: func(_ context.Context, _ PersistedStep) error { + return nil + }, + OnRetry: func( + _ int, + _ error, + _ chatretry.ClassifiedError, + _ time.Duration, + ) { + retried = true + }, + }) + }() + + trap.MustWait(ctx).MustRelease(ctx) + trap.Close() + + select { + case <-firstPartYielded: + case <-ctx.Done(): + t.Fatal("timed out waiting for first stream part") + } + + mClock.Advance(startupTimeout).MustWait(ctx) + close(continueStream) + + require.NoError(t, awaitRunResult(ctx, t, done)) require.Equal(t, 1, attempts) require.False(t, retried) } @@ -479,6 +570,16 @@ func TestRun_RetriesStartupTimeoutWhenStreamClosesSilently(t *testing.T) { const startupTimeout = 5 * time.Millisecond + ctx, cancel := context.WithTimeout( + context.Background(), + testutil.WaitShort, + ) + defer cancel() + + mClock := quartz.NewMock(t) + trap := mClock.Trap().AfterFunc("startupGuard") + defer trap.Close() + attempts := 0 attemptCause := make(chan error, 1) var retries []chatretry.ClassifiedError @@ -499,23 +600,32 @@ func TestRun_RetriesStartupTimeoutWhenStreamClosesSilently(t *testing.T) { }, } - err := Run(context.Background(), RunOptions{ - Model: model, - MaxSteps: 1, - StartupTimeout: startupTimeout, - PersistStep: func(_ context.Context, _ PersistedStep) error { - return nil - }, - OnRetry: func( - _ int, - _ error, - classified chatretry.ClassifiedError, - _ time.Duration, - ) { - retries = append(retries, classified) - }, - }) - require.NoError(t, err) + done := make(chan error, 1) + go func() { + done <- Run(context.Background(), RunOptions{ + Model: model, + MaxSteps: 1, + StartupTimeout: startupTimeout, + Clock: mClock, + PersistStep: func(_ context.Context, _ PersistedStep) error { + return nil + }, + OnRetry: func( + _ int, + _ error, + classified chatretry.ClassifiedError, + _ time.Duration, + ) { + retries = append(retries, classified) + }, + }) + }() + + trap.MustWait(ctx).MustRelease(ctx) + mClock.Advance(startupTimeout).MustWait(ctx) + trap.MustWait(ctx).MustRelease(ctx) + + require.NoError(t, awaitRunResult(ctx, t, done)) require.Equal(t, 2, attempts) require.Len(t, retries, 1) require.Equal(t, chaterror.KindStartupTimeout, retries[0].Kind) @@ -526,7 +636,12 @@ func TestRun_RetriesStartupTimeoutWhenStreamClosesSilently(t *testing.T) { "OpenAI did not start responding in time.", retries[0].Message, ) - require.ErrorIs(t, <-attemptCause, errStartupTimeout) + select { + case cause := <-attemptCause: + require.ErrorIs(t, cause, errStartupTimeout) + case <-ctx.Done(): + t.Fatal("timed out waiting for startup timeout cause") + } } func TestRun_InterruptedStepPersistsSyntheticToolResult(t *testing.T) {