mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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
This commit is contained in:
@@ -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...)
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user