From 4f74a7adee879068901adcc516f856b0bfd38b82 Mon Sep 17 00:00:00 2001 From: Hugo Dutka Date: Tue, 16 Jun 2026 14:35:40 +0200 Subject: [PATCH] fix: enable goleak in chatd tests (#26335) Enable goleak in chatd tests and fix some leaks. Addresses https://github.com/coder/coder/pull/26109#discussion_r3380039964 --- coderd/x/chatd/chatadvisor/main_test.go | 13 +++++ coderd/x/chatd/chatdebug/main_test.go | 13 +++++ coderd/x/chatd/chatloop/main_test.go | 13 +++++ coderd/x/chatd/chatopenai/main_test.go | 13 +++++ coderd/x/chatd/chatprompt/main_test.go | 13 +++++ coderd/x/chatd/chatprovider/main_test.go | 13 +++++ coderd/x/chatd/chatretry/main_test.go | 13 +++++ coderd/x/chatd/chatsanitize/main_test.go | 13 +++++ coderd/x/chatd/chatstate/main_test.go | 13 +++++ coderd/x/chatd/chattest/main_test.go | 13 +++++ coderd/x/chatd/chattool/main_test.go | 13 +++++ coderd/x/chatd/main_test.go | 13 +++++ coderd/x/chatd/messagepartbuffer/main_test.go | 13 +++++ .../message_part_buffer_test.go | 19 +++++++ coderd/x/chatd/stream_parts_internal_test.go | 11 +++- coderd/x/chatd/tasks_test.go | 50 ++++++++++--------- enterprise/coderd/x/chatd/main_test.go | 13 +++++ 17 files changed, 237 insertions(+), 25 deletions(-) create mode 100644 coderd/x/chatd/chatadvisor/main_test.go create mode 100644 coderd/x/chatd/chatdebug/main_test.go create mode 100644 coderd/x/chatd/chatloop/main_test.go create mode 100644 coderd/x/chatd/chatopenai/main_test.go create mode 100644 coderd/x/chatd/chatprompt/main_test.go create mode 100644 coderd/x/chatd/chatprovider/main_test.go create mode 100644 coderd/x/chatd/chatretry/main_test.go create mode 100644 coderd/x/chatd/chatsanitize/main_test.go create mode 100644 coderd/x/chatd/chatstate/main_test.go create mode 100644 coderd/x/chatd/chattest/main_test.go create mode 100644 coderd/x/chatd/chattool/main_test.go create mode 100644 coderd/x/chatd/main_test.go create mode 100644 coderd/x/chatd/messagepartbuffer/main_test.go create mode 100644 enterprise/coderd/x/chatd/main_test.go diff --git a/coderd/x/chatd/chatadvisor/main_test.go b/coderd/x/chatd/chatadvisor/main_test.go new file mode 100644 index 0000000000..3af91330f5 --- /dev/null +++ b/coderd/x/chatd/chatadvisor/main_test.go @@ -0,0 +1,13 @@ +package chatadvisor_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +} diff --git a/coderd/x/chatd/chatdebug/main_test.go b/coderd/x/chatd/chatdebug/main_test.go new file mode 100644 index 0000000000..4bc96c1d4c --- /dev/null +++ b/coderd/x/chatd/chatdebug/main_test.go @@ -0,0 +1,13 @@ +package chatdebug_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +} diff --git a/coderd/x/chatd/chatloop/main_test.go b/coderd/x/chatd/chatloop/main_test.go new file mode 100644 index 0000000000..dac51e9081 --- /dev/null +++ b/coderd/x/chatd/chatloop/main_test.go @@ -0,0 +1,13 @@ +package chatloop_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +} diff --git a/coderd/x/chatd/chatopenai/main_test.go b/coderd/x/chatd/chatopenai/main_test.go new file mode 100644 index 0000000000..6e948907e6 --- /dev/null +++ b/coderd/x/chatd/chatopenai/main_test.go @@ -0,0 +1,13 @@ +package chatopenai_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +} diff --git a/coderd/x/chatd/chatprompt/main_test.go b/coderd/x/chatd/chatprompt/main_test.go new file mode 100644 index 0000000000..111aa2ade1 --- /dev/null +++ b/coderd/x/chatd/chatprompt/main_test.go @@ -0,0 +1,13 @@ +package chatprompt_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +} diff --git a/coderd/x/chatd/chatprovider/main_test.go b/coderd/x/chatd/chatprovider/main_test.go new file mode 100644 index 0000000000..2b9ff2031f --- /dev/null +++ b/coderd/x/chatd/chatprovider/main_test.go @@ -0,0 +1,13 @@ +package chatprovider_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +} diff --git a/coderd/x/chatd/chatretry/main_test.go b/coderd/x/chatd/chatretry/main_test.go new file mode 100644 index 0000000000..31eab2426b --- /dev/null +++ b/coderd/x/chatd/chatretry/main_test.go @@ -0,0 +1,13 @@ +package chatretry_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +} diff --git a/coderd/x/chatd/chatsanitize/main_test.go b/coderd/x/chatd/chatsanitize/main_test.go new file mode 100644 index 0000000000..b4a11886eb --- /dev/null +++ b/coderd/x/chatd/chatsanitize/main_test.go @@ -0,0 +1,13 @@ +package chatsanitize_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +} diff --git a/coderd/x/chatd/chatstate/main_test.go b/coderd/x/chatd/chatstate/main_test.go new file mode 100644 index 0000000000..9cb06ca3ce --- /dev/null +++ b/coderd/x/chatd/chatstate/main_test.go @@ -0,0 +1,13 @@ +package chatstate_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +} diff --git a/coderd/x/chatd/chattest/main_test.go b/coderd/x/chatd/chattest/main_test.go new file mode 100644 index 0000000000..8a52fa1da3 --- /dev/null +++ b/coderd/x/chatd/chattest/main_test.go @@ -0,0 +1,13 @@ +package chattest_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +} diff --git a/coderd/x/chatd/chattool/main_test.go b/coderd/x/chatd/chattool/main_test.go new file mode 100644 index 0000000000..3b3179a847 --- /dev/null +++ b/coderd/x/chatd/chattool/main_test.go @@ -0,0 +1,13 @@ +package chattool_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +} diff --git a/coderd/x/chatd/main_test.go b/coderd/x/chatd/main_test.go new file mode 100644 index 0000000000..701f2b0234 --- /dev/null +++ b/coderd/x/chatd/main_test.go @@ -0,0 +1,13 @@ +package chatd_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +} diff --git a/coderd/x/chatd/messagepartbuffer/main_test.go b/coderd/x/chatd/messagepartbuffer/main_test.go new file mode 100644 index 0000000000..036a9a18d7 --- /dev/null +++ b/coderd/x/chatd/messagepartbuffer/main_test.go @@ -0,0 +1,13 @@ +package messagepartbuffer_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +} diff --git a/coderd/x/chatd/messagepartbuffer/message_part_buffer_test.go b/coderd/x/chatd/messagepartbuffer/message_part_buffer_test.go index f1fcf300b2..475f2dfa1b 100644 --- a/coderd/x/chatd/messagepartbuffer/message_part_buffer_test.go +++ b/coderd/x/chatd/messagepartbuffer/message_part_buffer_test.go @@ -19,6 +19,7 @@ func TestBuffer_CreateEpisodeRejectsDuplicate(t *testing.T) { t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() key := testEpisodeKey() require.NoError(t, buffer.CreateEpisode(key)) require.ErrorIs(t, buffer.CreateEpisode(key), messagepartbuffer.ErrEpisodeExists) @@ -28,6 +29,7 @@ func TestBuffer_AddPartAndGetParts(t *testing.T) { t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() key := testEpisodeKey() require.NoError(t, buffer.CreateEpisode(key)) require.NoError(t, buffer.AddPart(key, codersdk.ChatMessageRoleAssistant, codersdk.ChatMessageText("hello"))) @@ -44,6 +46,7 @@ func TestBuffer_AddPartMissingEpisodeReturnsNotFound(t *testing.T) { t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() err := buffer.AddPart(testEpisodeKey(), codersdk.ChatMessageRoleAssistant, codersdk.ChatMessageText("hello")) require.ErrorIs(t, err, messagepartbuffer.ErrEpisodeNotFound) } @@ -52,6 +55,7 @@ func TestBuffer_GetPartsMissingEpisodeReturnsNotFound(t *testing.T) { t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() _, err := buffer.GetParts(testEpisodeKey()) require.ErrorIs(t, err, messagepartbuffer.ErrEpisodeNotFound) } @@ -60,6 +64,7 @@ func TestBuffer_AddPartFullEpisodeReturnsFull(t *testing.T) { t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{MaxEpisodeBytes: 1}) + defer buffer.Close() key := testEpisodeKey() require.NoError(t, buffer.CreateEpisode(key)) err := buffer.AddPart(key, codersdk.ChatMessageRoleAssistant, codersdk.ChatMessageText("hello")) @@ -73,6 +78,7 @@ func TestBuffer_CloseEpisodeMissingCreatesClosedEpisode(t *testing.T) { t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() key := testEpisodeKey() require.NoError(t, buffer.CloseEpisode(key)) parts, err := buffer.GetParts(key) @@ -86,6 +92,7 @@ func TestBuffer_CloseEpisodeIdempotent(t *testing.T) { t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() key := testEpisodeKey() require.NoError(t, buffer.CreateEpisode(key)) require.NoError(t, buffer.CloseEpisode(key)) @@ -96,6 +103,7 @@ func TestBuffer_SubscribeExistingReplaysThenStreamsLiveParts(t *testing.T) { t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() key := testEpisodeKey() require.NoError(t, buffer.CreateEpisode(key)) require.NoError(t, buffer.AddPart(key, codersdk.ChatMessageRoleAssistant, codersdk.ChatMessageText("before"))) @@ -114,6 +122,7 @@ func TestBuffer_SubscribeClosedEpisodeReplaysThenCloses(t *testing.T) { t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() key := testEpisodeKey() require.NoError(t, buffer.CreateEpisode(key)) require.NoError(t, buffer.AddPart(key, codersdk.ChatMessageRoleAssistant, codersdk.ChatMessageText("before"))) @@ -131,6 +140,7 @@ func TestBuffer_SubscribeBeforeCreateReturnsAndWaitsWithoutNotFound(t *testing.T t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() key := testEpisodeKey() ctx := testutil.Context(t, testutil.WaitLong) ch, cancel, err := buffer.SubscribeToEpisode(ctx, key) @@ -152,6 +162,7 @@ func TestBuffer_AddPartAssignsContiguousSeq(t *testing.T) { t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() key := testEpisodeKey() require.NoError(t, buffer.CreateEpisode(key)) for i := range 3 { @@ -168,6 +179,7 @@ func TestBuffer_EpisodeByteLimitUsesJSONAccounting(t *testing.T) { part := codersdk.ChatMessageText("hello") limit := serializedPartBytes(t, messagepartbuffer.Part{Seq: 1, Role: codersdk.ChatMessageRoleAssistant, MessagePart: part}) buffer := messagepartbuffer.New(messagepartbuffer.Options{MaxEpisodeBytes: limit}) + defer buffer.Close() key := testEpisodeKey() require.NoError(t, buffer.CreateEpisode(key)) require.NoError(t, buffer.AddPart(key, codersdk.ChatMessageRoleAssistant, part)) @@ -186,6 +198,7 @@ func TestBuffer_GCClosedEpisodeAfterGraceAndNoSubscribers(t *testing.T) { ClosedEpisodeRetention: time.Minute, SubscriberSendTimeout: 10 * time.Minute, }) + defer buffer.Close() key := testEpisodeKey() require.NoError(t, buffer.CreateEpisode(key)) require.NoError(t, buffer.AddPart(key, codersdk.ChatMessageRoleAssistant, codersdk.ChatMessageText("held"))) @@ -217,6 +230,7 @@ func TestBuffer_GCRetainedSubscribedEpisodeDoesNotBlockOtherExpiredEpisodes(t *t ClosedEpisodeRetention: time.Minute, SubscriberSendTimeout: 10 * time.Minute, }) + defer buffer.Close() retainedKey := testEpisodeKey() collectedKey := testEpisodeKey() require.NoError(t, buffer.CreateEpisode(retainedKey)) @@ -257,6 +271,7 @@ func TestBuffer_SlowSubscriberClosed(t *testing.T) { Clock: clock, SubscriberSendTimeout: time.Second, }) + defer buffer.Close() key := testEpisodeKey() require.NoError(t, buffer.CreateEpisode(key)) ctx := testutil.Context(t, testutil.WaitLong) @@ -277,6 +292,7 @@ func TestBuffer_BurstyOutputDoesNotCloseSubscriberBeforeSendTimeout(t *testing.T t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{SubscriberChannelSize: 1}) + defer buffer.Close() key := testEpisodeKey() require.NoError(t, buffer.CreateEpisode(key)) ctx := testutil.Context(t, testutil.WaitLong) @@ -297,6 +313,7 @@ func TestBuffer_SubscribeCanceledBeforeCreateCanCreateEpisode(t *testing.T) { t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() key := testEpisodeKey() ctx, cancel := context.WithCancel(context.Background()) ch, cancelSub, err := buffer.SubscribeToEpisode(ctx, key) @@ -311,6 +328,7 @@ func TestBuffer_SubscribeCanceledWithoutCreateReclaimsEpisode(t *testing.T) { t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() key := testEpisodeKey() ctx := testutil.Context(t, testutil.WaitLong) ch, cancelSub, err := buffer.SubscribeToEpisode(ctx, key) @@ -329,6 +347,7 @@ func TestBuffer_CloseClosesPendingSubscriptionAndRejectsOperations(t *testing.T) t.Parallel() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() key := testEpisodeKey() ctx := testutil.Context(t, testutil.WaitLong) ch, cancel, err := buffer.SubscribeToEpisode(ctx, key) diff --git a/coderd/x/chatd/stream_parts_internal_test.go b/coderd/x/chatd/stream_parts_internal_test.go index cc855daddd..a598ffcfd8 100644 --- a/coderd/x/chatd/stream_parts_internal_test.go +++ b/coderd/x/chatd/stream_parts_internal_test.go @@ -24,6 +24,7 @@ func TestStreamPartsEndpointReplayLiveAndReselect(t *testing.T) { ctx := testutil.Context(t, testutil.WaitShort) chatID := uuid.New() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() endpoint := streamPartsEndpoint{ chatID: chatID, buffer: buffer, @@ -80,6 +81,7 @@ func TestStreamPartsEndpointWaitsForMissingEpisode(t *testing.T) { ctx := testutil.Context(t, testutil.WaitShort) chatID := uuid.New() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() endpoint := streamPartsEndpoint{ chatID: chatID, buffer: buffer, @@ -110,6 +112,7 @@ func TestStreamPartsEndpointReselectsWhileEpisodeMissing(t *testing.T) { ctx := testutil.Context(t, testutil.WaitShort) chatID := uuid.New() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() endpoint := streamPartsEndpoint{ chatID: chatID, buffer: buffer, @@ -141,6 +144,7 @@ func TestStreamPartsEndpointClientDisconnectCancels(t *testing.T) { ctx := testutil.Context(t, testutil.WaitShort) chatID := uuid.New() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() endpoint := streamPartsEndpoint{ chatID: chatID, buffer: buffer, @@ -163,6 +167,7 @@ func TestStreamPartsEndpointWebSocket(t *testing.T) { ctx := testutil.Context(t, testutil.WaitShort) chatID := uuid.New() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() endpoint := streamPartsEndpoint{ chatID: chatID, buffer: buffer, @@ -171,7 +176,7 @@ func TestStreamPartsEndpointWebSocket(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) { _ = endpoint.serveWebSocket(rw, r) })) - t.Cleanup(server.Close) + defer server.Close() key := messagepartbuffer.Key{ChatID: chatID, HistoryVersion: 1, GenerationAttempt: 1} require.NoError(t, buffer.CreateEpisode(key)) @@ -196,6 +201,7 @@ func TestStreamPartsWebSocketSession(t *testing.T) { ctx := testutil.Context(t, testutil.WaitShort) chatID := uuid.New() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() endpoint := streamPartsEndpoint{ chatID: chatID, buffer: buffer, @@ -204,7 +210,7 @@ func TestStreamPartsWebSocketSession(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) { _ = endpoint.serveWebSocket(rw, r) })) - t.Cleanup(server.Close) + defer server.Close() key := messagepartbuffer.Key{ChatID: chatID, HistoryVersion: 4, GenerationAttempt: 2} require.NoError(t, buffer.CreateEpisode(key)) @@ -232,6 +238,7 @@ func TestLocalStreamPartsDialerReplayLiveAndClose(t *testing.T) { ctx := testutil.Context(t, testutil.WaitShort) chatID := uuid.New() buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + defer buffer.Close() dialer := NewLocalStreamPartsDialer(LocalStreamPartsDialerConfig{ Buffer: buffer, Logger: slogtest.Make(t, nil), diff --git a/coderd/x/chatd/tasks_test.go b/coderd/x/chatd/tasks_test.go index f844d1bb0b..a4540a1537 100644 --- a/coderd/x/chatd/tasks_test.go +++ b/coderd/x/chatd/tasks_test.go @@ -113,7 +113,9 @@ func TestInterruptTask_FinishInterruptionOnly(t *testing.T) { workerID := uuid.New() runnerID := uuid.New() acquired := f.acquireChat(t, chat.ID, workerID, runnerID) - buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + recorder := newTaskSideEffectRecorder() + starter := newTestTaskStarter(t, f, recorder) + buffer := starter.opts.MessagePartBuffer key := messagepartbuffer.Key{ ChatID: chat.ID, HistoryVersion: acquired.HistoryVersion, @@ -123,8 +125,6 @@ func TestInterruptTask_FinishInterruptionOnly(t *testing.T) { require.NoError(t, buffer.AddPart(key, codersdk.ChatMessageRoleAssistant, codersdk.ChatMessageText("partial answer"))) interrupting := f.interruptChat(t, chat.ID) require.Equal(t, database.ChatStatusInterrupting, interrupting.Status) - recorder := newTaskSideEffectRecorder() - starter := newTestTaskStarter(t, f, buffer, recorder) err := starter.StartInterrupt(testutil.Context(t, testutil.WaitLong), chatWorkerTaskStartInput{ ChatID: chat.ID, @@ -166,7 +166,7 @@ func TestInterruptTask_StaleFenceExits(t *testing.T) { otherRunnerID := uuid.New() f.acquireChat(t, chat.ID, otherWorkerID, otherRunnerID) recorder := newTaskSideEffectRecorder() - starter := newTestTaskStarter(t, f, messagepartbuffer.New(messagepartbuffer.Options{}), recorder) + starter := newTestTaskStarter(t, f, recorder) err := starter.StartInterrupt(testutil.Context(t, testutil.WaitLong), chatWorkerTaskStartInput{ ChatID: chat.ID, @@ -197,7 +197,7 @@ func TestInterruptTask_MissingEpisodePersistsNilPartials(t *testing.T) { f.acquireChat(t, chat.ID, workerID, runnerID) interrupting := f.forceExecutionState(t, chat.ID, database.ChatStatusInterrupting, false, sql.NullTime{}) recorder := newTaskSideEffectRecorder() - starter := newTestTaskStarter(t, f, messagepartbuffer.New(messagepartbuffer.Options{}), recorder) + starter := newTestTaskStarter(t, f, recorder) err := starter.StartInterrupt(testutil.Context(t, testutil.WaitLong), chatWorkerTaskStartInput{ ChatID: chat.ID, @@ -227,7 +227,9 @@ func TestInterruptTask_BufferedPartsBecomePartialMessages(t *testing.T) { workerID := uuid.New() runnerID := uuid.New() acquired := f.acquireChat(t, chat.ID, workerID, runnerID) - buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + recorder := newTaskSideEffectRecorder() + starter := newTestTaskStarter(t, f, recorder) + buffer := starter.opts.MessagePartBuffer key := messagepartbuffer.Key{ChatID: chat.ID, HistoryVersion: acquired.HistoryVersion, GenerationAttempt: acquired.GenerationAttempt} require.NoError(t, buffer.CreateEpisode(key)) callID := "call_" + uuid.NewString() @@ -238,8 +240,6 @@ func TestInterruptTask_BufferedPartsBecomePartialMessages(t *testing.T) { Args: json.RawMessage(`{"value":1}`), })) interrupting := f.interruptChat(t, chat.ID) - recorder := newTaskSideEffectRecorder() - starter := newTestTaskStarter(t, f, buffer, recorder) err := starter.StartInterrupt(testutil.Context(t, testutil.WaitLong), chatWorkerTaskStartInput{ ChatID: chat.ID, @@ -276,7 +276,7 @@ func TestRequiresActionTimeout_ExpiredCancelsOnly(t *testing.T) { acquired := f.acquireChat(t, chat.ID, workerID, runnerID) expired := f.setRequiresActionDeadline(t, chat.ID, sql.NullTime{Time: time.Now().Add(-time.Minute), Valid: true}) recorder := newTaskSideEffectRecorder() - starter := newTestTaskStarter(t, f, messagepartbuffer.New(messagepartbuffer.Options{}), recorder) + starter := newTestTaskStarter(t, f, recorder) err := starter.StartRequiresActionTimeout(testutil.Context(t, testutil.WaitLong), chatWorkerTaskStartInput{ ChatID: chat.ID, @@ -306,7 +306,7 @@ func TestRequiresActionTimeout_NullDeadlineCancelsImmediately(t *testing.T) { acquired := f.acquireChat(t, chat.ID, workerID, runnerID) nullDeadline := f.setRequiresActionDeadline(t, chat.ID, sql.NullTime{}) recorder := newTaskSideEffectRecorder() - starter := newTestTaskStarter(t, f, messagepartbuffer.New(messagepartbuffer.Options{}), recorder) + starter := newTestTaskStarter(t, f, recorder) err := starter.StartRequiresActionTimeout(testutil.Context(t, testutil.WaitLong), chatWorkerTaskStartInput{ ChatID: chat.ID, @@ -335,7 +335,7 @@ func TestRequiresActionTimeout_StaleFenceExitsAfterToolResult(t *testing.T) { expired := f.setRequiresActionDeadline(t, chat.ID, sql.NullTime{Time: time.Now().Add(-time.Minute), Valid: true}) f.forceExecutionState(t, chat.ID, database.ChatStatusRunning, false, sql.NullTime{}) recorder := newTaskSideEffectRecorder() - starter := newTestTaskStarter(t, f, messagepartbuffer.New(messagepartbuffer.Options{}), recorder) + starter := newTestTaskStarter(t, f, recorder) err := starter.StartRequiresActionTimeout(testutil.Context(t, testutil.WaitLong), chatWorkerTaskStartInput{ ChatID: chat.ID, @@ -363,7 +363,7 @@ func TestAbandonTask_AbandonOnly(t *testing.T) { runnerID := uuid.New() acquired := f.acquireChat(t, chat.ID, workerID, runnerID) recorder := newTaskSideEffectRecorder() - starter := newTestTaskStarter(t, f, messagepartbuffer.New(messagepartbuffer.Options{}), recorder) + starter := newTestTaskStarter(t, f, recorder) err := starter.StartAbandon(testutil.Context(t, testutil.WaitLong), chatWorkerTaskStartInput{ ChatID: chat.ID, @@ -395,7 +395,7 @@ func TestAbandonTask_OwnershipMismatchRequestsCleanup(t *testing.T) { otherRunnerID := uuid.New() latestOwner := f.acquireChat(t, chat.ID, otherWorkerID, otherRunnerID) recorder := newTaskSideEffectRecorder() - starter := newTestTaskStarter(t, f, messagepartbuffer.New(messagepartbuffer.Options{}), recorder) + starter := newTestTaskStarter(t, f, recorder) err := starter.StartAbandon(testutil.Context(t, testutil.WaitLong), chatWorkerTaskStartInput{ ChatID: chat.ID, @@ -423,7 +423,7 @@ func TestAbandonTask_StaleStatusFenceExits(t *testing.T) { acquired := f.acquireChat(t, chat.ID, workerID, runnerID) f.forceExecutionState(t, chat.ID, database.ChatStatusInterrupting, false, sql.NullTime{}) recorder := newTaskSideEffectRecorder() - starter := newTestTaskStarter(t, f, messagepartbuffer.New(messagepartbuffer.Options{}), recorder) + starter := newTestTaskStarter(t, f, recorder) err := starter.StartAbandon(testutil.Context(t, testutil.WaitLong), chatWorkerTaskStartInput{ ChatID: chat.ID, @@ -451,7 +451,7 @@ func TestGenerationTask_RecordRetryState(t *testing.T) { runnerID := uuid.New() acquired := f.acquireChat(t, chat.ID, workerID, runnerID) recorder := newTaskSideEffectRecorder() - starter := newTestTaskStarter(t, f, messagepartbuffer.New(messagepartbuffer.Options{}), recorder) + starter := newTestTaskStarter(t, f, recorder) attempt, _, _, closeEpisode, err := starter.beginGenerationAttempt( testutil.Context(t, testutil.WaitLong), @@ -521,7 +521,7 @@ func TestGenerationTask_RecordRetryStateUsesDurableGenerationAttempt(t *testing. workerID := uuid.New() runnerID := uuid.New() acquired := f.acquireChat(t, chat.ID, workerID, runnerID) - starter := newTestTaskStarter(t, f, messagepartbuffer.New(messagepartbuffer.Options{}), newTaskSideEffectRecorder()) + starter := newTestTaskStarter(t, f, newTaskSideEffectRecorder()) machine := chatstate.NewChatMachine(f.db, f.pubsub, chat.ID) for range 3 { @@ -579,7 +579,7 @@ func TestGenerationTask_RecordRetryStateClearedByNextAttempt(t *testing.T) { workerID := uuid.New() runnerID := uuid.New() acquired := f.acquireChat(t, chat.ID, workerID, runnerID) - starter := newTestTaskStarter(t, f, messagepartbuffer.New(messagepartbuffer.Options{}), newTaskSideEffectRecorder()) + starter := newTestTaskStarter(t, f, newTaskSideEffectRecorder()) machine := chatstate.NewChatMachine(f.db, f.pubsub, chat.ID) input := chatWorkerTaskStartInput{ ChatID: chat.ID, @@ -628,7 +628,7 @@ func TestGenerationTask_RecordRetryStateStaleFenceExits(t *testing.T) { workerID := uuid.New() runnerID := uuid.New() acquired := f.acquireChat(t, chat.ID, workerID, runnerID) - starter := newTestTaskStarter(t, f, messagepartbuffer.New(messagepartbuffer.Options{}), newTaskSideEffectRecorder()) + starter := newTestTaskStarter(t, f, newTaskSideEffectRecorder()) machine := chatstate.NewChatMachine(f.db, f.pubsub, chat.ID) attempt, _, _, closeEpisode, err := starter.beginGenerationAttempt( testutil.Context(t, testutil.WaitLong), @@ -678,7 +678,7 @@ func TestRunner_StartsRealInterruptTask(t *testing.T) { f := newTaskTestFixture(t) chat := f.createRunningChat(t) - worker := startRealTaskWorker(t, f, messagepartbuffer.New(messagepartbuffer.Options{})) + worker := startRealTaskWorker(t, f) waitOwnedChat(t, f, chat.ID, worker.chatWorkerID()) interrupting := f.interruptChat(t, chat.ID) @@ -699,7 +699,7 @@ func TestRunner_StartsRealRequiresActionTimeoutTask(t *testing.T) { f := newTaskTestFixture(t) chat := f.createRequiresActionChat(t) f.setRequiresActionDeadline(t, chat.ID, sql.NullTime{Time: time.Now().Add(-time.Minute), Valid: true}) - worker := startRealTaskWorker(t, f, messagepartbuffer.New(messagepartbuffer.Options{})) + worker := startRealTaskWorker(t, f) testutil.Eventually(testutil.Context(t, testutil.WaitLong), t, func(ctx context.Context) bool { latest, err := f.db.GetChatByID(ctx, chat.ID) @@ -716,7 +716,7 @@ func TestRunner_StartsRealAbandonTask(t *testing.T) { f := newTaskTestFixture(t) chat := f.createRunningChat(t) - worker := startRealTaskWorker(t, f, messagepartbuffer.New(messagepartbuffer.Options{})) + worker := startRealTaskWorker(t, f) waitOwnedChat(t, f, chat.ID, worker.chatWorkerID()) updated := f.forceExecutionState(t, chat.ID, database.ChatStatusError, false, sql.NullTime{}) @@ -995,8 +995,10 @@ func (p *taskRecordingPubsub) watchEvents(t *testing.T) []codersdk.ChatWatchEven return out } -func startRealTaskWorker(t *testing.T, f *taskTestFixture, buffer *messagepartbuffer.Buffer) *chatWorker { +func startRealTaskWorker(t *testing.T, f *taskTestFixture) *chatWorker { t.Helper() + buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + t.Cleanup(buffer.Close) worker, err := newChatWorker(nil, chatWorkerOptions{ WorkerID: uuid.New(), Store: f.db, @@ -1115,8 +1117,10 @@ func (r *taskSideEffectRecorder) requireInterruptionOutcome(t *testing.T, chatID t.Fatalf("missing interruption outcome chat_id=%s status=%s outcomes=%v", chatID, status, r.interrupts) } -func newTestTaskStarter(t *testing.T, f *taskTestFixture, buffer *messagepartbuffer.Buffer, recorder *taskSideEffectRecorder) *taskStarter { +func newTestTaskStarter(t *testing.T, f *taskTestFixture, recorder *taskSideEffectRecorder) *taskStarter { t.Helper() + buffer := messagepartbuffer.New(messagepartbuffer.Options{}) + t.Cleanup(buffer.Close) starter, err := newTaskStarter(nil, chatWorkerOptions{ Store: f.db, Pubsub: f.pubsub, diff --git a/enterprise/coderd/x/chatd/main_test.go b/enterprise/coderd/x/chatd/main_test.go new file mode 100644 index 0000000000..701f2b0234 --- /dev/null +++ b/enterprise/coderd/x/chatd/main_test.go @@ -0,0 +1,13 @@ +package chatd_test + +import ( + "testing" + + "go.uber.org/goleak" + + "github.com/coder/coder/v2/testutil" +) + +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m, testutil.GoleakOptions...) +}