mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix: preserve stream state on interrupt, fix auto-promote error handling (#24314)
When tryAutoPromoteQueuedMessage's insert fails, return the error instead of swallowing it so the transaction rolls back and the queued message survives. Previously the POP DELETE committed while the INSERT silently failed, permanently losing the message. Remove clearStreamState() from the pending/waiting status handler in the frontend. The durable message event clears stream state via the existing needsStreamReset path, eliminating the visual gap where content vanishes before the persisted message arrives. Fixes CODAGT-61
This commit is contained in:
@@ -5040,9 +5040,7 @@ func (p *Server) tryAutoPromoteQueuedMessage(
|
||||
).withCreatedBy(chat.OwnerID))
|
||||
msgs, err := insertChatMessageWithStore(ctx, tx, msgParams)
|
||||
if err != nil {
|
||||
logger.Error(ctx, "failed to promote queued message",
|
||||
slog.F("queued_message_id", nextQueued.ID), slog.Error(err))
|
||||
return nil, nil, false, nil
|
||||
return nil, nil, false, xerrors.Errorf("insert promoted message: %w", err)
|
||||
}
|
||||
msg := msgs[0]
|
||||
|
||||
@@ -5148,7 +5146,8 @@ func (p *Server) finishActiveChat(
|
||||
var promoteErr error
|
||||
result.promotedMessage, result.remainingQueuedMessages, result.shouldPublishQueueUpdate, promoteErr = p.tryAutoPromoteQueuedMessage(ctx, tx, latestChat)
|
||||
if promoteErr != nil {
|
||||
logger.Error(ctx, "failed to auto-promote queued message", slog.Error(promoteErr))
|
||||
logger.Error(ctx, "auto-promote queued message failed, rolling back", slog.Error(promoteErr))
|
||||
return xerrors.Errorf("auto-promote queued message: %w", promoteErr)
|
||||
} else if result.promotedMessage != nil {
|
||||
status = database.ChatStatusPending
|
||||
}
|
||||
|
||||
@@ -3524,9 +3524,9 @@ func TestProcessChat_IgnoresStaleControlNotification(t *testing.T) {
|
||||
clock: clock,
|
||||
workerID: workerID,
|
||||
chatHeartbeatInterval: time.Minute,
|
||||
metrics: chatloop.NopMetrics(),
|
||||
configCache: newChatConfigCache(ctx, db, clock),
|
||||
heartbeatRegistry: make(map[uuid.UUID]*heartbeatEntry),
|
||||
metrics: chatloop.NopMetrics(),
|
||||
}
|
||||
|
||||
// Publish a stale "pending" notification on the control channel
|
||||
@@ -3680,6 +3680,7 @@ func TestHeartbeatTick_StolenChatIsInterrupted(t *testing.T) {
|
||||
clock: clock,
|
||||
workerID: workerID,
|
||||
chatHeartbeatInterval: time.Minute,
|
||||
metrics: chatloop.NopMetrics(),
|
||||
heartbeatRegistry: make(map[uuid.UUID]*heartbeatEntry),
|
||||
}
|
||||
|
||||
@@ -3760,6 +3761,7 @@ func TestHeartbeatTick_DBErrorDoesNotInterruptChats(t *testing.T) {
|
||||
clock: clock,
|
||||
workerID: uuid.New(),
|
||||
chatHeartbeatInterval: time.Minute,
|
||||
metrics: chatloop.NopMetrics(),
|
||||
heartbeatRegistry: make(map[uuid.UUID]*heartbeatEntry),
|
||||
}
|
||||
|
||||
@@ -4717,3 +4719,224 @@ func TestGetWorkspaceConn_DialErrorNotMisclassifiedAsTimeout(t *testing.T) {
|
||||
// The original dial error should propagate.
|
||||
require.ErrorContains(t, err, "authentication failed")
|
||||
}
|
||||
|
||||
// TestAutoPromote_InsertFailureRollsBackTransaction verifies that when
|
||||
// tryAutoPromoteQueuedMessage pops a queued message but the subsequent
|
||||
// insert fails, the error propagates to the InTx callback, causing the
|
||||
// transaction to roll back and preserving the queued message.
|
||||
func TestAutoPromote_InsertFailureRollsBackTransaction(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
tx := dbmock.NewMockStore(ctrl)
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
ps := dbpubsub.NewInMemory()
|
||||
clock := quartz.NewReal()
|
||||
|
||||
chatID := uuid.New()
|
||||
workerID := uuid.New()
|
||||
ownerID := uuid.New()
|
||||
modelConfigID := uuid.New()
|
||||
|
||||
waitingChat := database.Chat{
|
||||
ID: chatID,
|
||||
OwnerID: ownerID,
|
||||
LastModelConfigID: modelConfigID,
|
||||
Status: database.ChatStatusWaiting,
|
||||
WorkerID: uuid.NullUUID{UUID: workerID, Valid: true},
|
||||
}
|
||||
queuedMsg := database.ChatQueuedMessage{
|
||||
ID: 1,
|
||||
ChatID: chatID,
|
||||
Content: []byte(`[{"type":"text","text":"queued"}]`),
|
||||
}
|
||||
insertErr := xerrors.New("insert failed")
|
||||
|
||||
server := &Server{
|
||||
db: db,
|
||||
logger: logger,
|
||||
pubsub: ps,
|
||||
configCache: newChatConfigCache(ctx, db, clock),
|
||||
}
|
||||
|
||||
// The caller runs tryAutoPromoteQueuedMessage inside InTx.
|
||||
// Wire the mock to execute the callback against the TX mock.
|
||||
var txErr error
|
||||
db.EXPECT().InTx(gomock.Any(), gomock.Any()).DoAndReturn(
|
||||
func(fn func(database.Store) error, _ *database.TxOptions) error {
|
||||
txErr = fn(tx)
|
||||
return txErr
|
||||
},
|
||||
)
|
||||
|
||||
// Inside the TX: lock chat, get queued messages, resolve model
|
||||
// config, pop queued message, insert fails.
|
||||
tx.EXPECT().GetChatByIDForUpdate(gomock.Any(), chatID).Return(waitingChat, nil)
|
||||
tx.EXPECT().GetChatQueuedMessages(gomock.Any(), chatID).Return([]database.ChatQueuedMessage{queuedMsg}, nil)
|
||||
tx.EXPECT().GetChatModelConfigByID(gomock.Any(), modelConfigID).Return(database.ChatModelConfig{ID: modelConfigID}, nil)
|
||||
tx.EXPECT().PopNextQueuedMessage(gomock.Any(), chatID).Return(queuedMsg, nil)
|
||||
tx.EXPECT().InsertChatMessages(gomock.Any(), gomock.Any()).Return(nil, insertErr)
|
||||
|
||||
// Invoke tryAutoPromoteQueuedMessage through the same InTx
|
||||
// pattern the processChat defer uses. The test directly calls
|
||||
// the production path to verify error propagation.
|
||||
_ = db.InTx(func(txStore database.Store) error {
|
||||
latestChat, err := txStore.GetChatByIDForUpdate(ctx, chatID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
_, _, _, promoteErr := server.tryAutoPromoteQueuedMessage(ctx, txStore, latestChat)
|
||||
if promoteErr != nil {
|
||||
return promoteErr
|
||||
}
|
||||
|
||||
// This code path should not be reached when the insert
|
||||
// fails, because promoteErr should be non-nil.
|
||||
return nil
|
||||
}, nil)
|
||||
|
||||
// The InTx callback must return a non-nil error so the
|
||||
// transaction rolls back, preserving the queued message.
|
||||
require.Error(t, txErr, "InTx callback should return error when insert fails")
|
||||
}
|
||||
|
||||
// TestAutoPromote_WakesRunLoopAfterPromotion verifies that after the
|
||||
func TestAutoPromote_InsertFailureSkipsStatusUpdate(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
tx := dbmock.NewMockStore(ctrl)
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
ps := dbpubsub.NewInMemory()
|
||||
clock := quartz.NewReal()
|
||||
|
||||
chatID := uuid.New()
|
||||
workerID := uuid.New()
|
||||
ownerID := uuid.New()
|
||||
modelConfigID := uuid.New()
|
||||
|
||||
waitingChat := database.Chat{
|
||||
ID: chatID,
|
||||
OwnerID: ownerID,
|
||||
LastModelConfigID: modelConfigID,
|
||||
Status: database.ChatStatusWaiting,
|
||||
WorkerID: uuid.NullUUID{UUID: workerID, Valid: true},
|
||||
}
|
||||
queuedMsg := database.ChatQueuedMessage{
|
||||
ID: 1,
|
||||
ChatID: chatID,
|
||||
Content: []byte(`[{"type":"text","text":"queued"}]`),
|
||||
}
|
||||
|
||||
wakeCh := make(chan struct{}, 1)
|
||||
server := &Server{
|
||||
db: db,
|
||||
logger: logger,
|
||||
pubsub: ps,
|
||||
clock: clock,
|
||||
workerID: workerID,
|
||||
wakeCh: wakeCh,
|
||||
chatHeartbeatInterval: time.Minute,
|
||||
metrics: chatloop.NopMetrics(),
|
||||
configCache: newChatConfigCache(ctx, db, clock),
|
||||
heartbeatRegistry: make(map[uuid.UUID]*heartbeatEntry),
|
||||
}
|
||||
|
||||
// Block model resolution until the control subscriber fires.
|
||||
modelBlocked := make(chan struct{})
|
||||
db.EXPECT().GetChatModelConfigByID(gomock.Any(), gomock.Any()).DoAndReturn(
|
||||
func(ctx context.Context, _ uuid.UUID) (database.ChatModelConfig, error) {
|
||||
<-modelBlocked
|
||||
return database.ChatModelConfig{}, xerrors.New("no model")
|
||||
},
|
||||
).AnyTimes()
|
||||
db.EXPECT().GetEnabledChatProviders(gomock.Any()).Return(nil, nil).AnyTimes()
|
||||
db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return(nil, nil).AnyTimes()
|
||||
db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(
|
||||
database.ChatUsageLimitConfig{}, sql.ErrNoRows,
|
||||
).AnyTimes()
|
||||
db.EXPECT().GetChatMessagesForPromptByChatID(gomock.Any(), chatID).Return(nil, nil).AnyTimes()
|
||||
|
||||
// The deferred cleanup transaction: InsertChatMessages fails,
|
||||
// so UpdateChatStatus must NOT be called.
|
||||
db.EXPECT().InTx(gomock.Any(), gomock.Any()).DoAndReturn(
|
||||
func(fn func(database.Store) error, _ *database.TxOptions) error {
|
||||
return fn(tx)
|
||||
},
|
||||
)
|
||||
tx.EXPECT().GetChatByIDForUpdate(gomock.Any(), chatID).Return(waitingChat, nil)
|
||||
tx.EXPECT().GetChatQueuedMessages(gomock.Any(), chatID).Return([]database.ChatQueuedMessage{queuedMsg}, nil)
|
||||
tx.EXPECT().GetChatModelConfigByID(gomock.Any(), modelConfigID).Return(database.ChatModelConfig{ID: modelConfigID}, nil)
|
||||
tx.EXPECT().PopNextQueuedMessage(gomock.Any(), chatID).Return(queuedMsg, nil)
|
||||
tx.EXPECT().InsertChatMessages(gomock.Any(), gomock.Any()).Return(
|
||||
nil, xerrors.New("insert failed"),
|
||||
)
|
||||
tx.EXPECT().UpdateChatStatus(gomock.Any(), gomock.Any()).Times(0)
|
||||
|
||||
// Subscribe BEFORE launching the goroutine.
|
||||
runningCh := make(chan struct{}, 1)
|
||||
unsubRunning, err := ps.SubscribeWithErr(
|
||||
coderdpubsub.ChatStreamNotifyChannel(chatID),
|
||||
func(_ context.Context, msg []byte, err error) {
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
var notify coderdpubsub.ChatStreamNotifyMessage
|
||||
if json.Unmarshal(msg, ¬ify) != nil {
|
||||
return
|
||||
}
|
||||
if notify.Status == string(database.ChatStatusRunning) {
|
||||
select {
|
||||
case runningCh <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
defer unsubRunning()
|
||||
|
||||
chat := database.Chat{ID: chatID, OwnerID: ownerID, LastModelConfigID: modelConfigID}
|
||||
processDone := make(chan struct{})
|
||||
go func() {
|
||||
defer close(processDone)
|
||||
server.processChat(ctx, chat)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-runningCh:
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out waiting for running status")
|
||||
}
|
||||
|
||||
// Publish an interrupt so processChat exits runChat.
|
||||
interruptMsg, err := json.Marshal(coderdpubsub.ChatStreamNotifyMessage{
|
||||
Status: string(database.ChatStatusWaiting),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
err = ps.Publish(coderdpubsub.ChatStreamNotifyChannel(chatID), interruptMsg)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Unblock model resolution so runChat can exit.
|
||||
close(modelBlocked)
|
||||
|
||||
select {
|
||||
case <-processDone:
|
||||
case <-ctx.Done():
|
||||
t.Fatal("processChat did not complete")
|
||||
}
|
||||
|
||||
// The wake channel should NOT have a signal because the
|
||||
// transaction failed before reaching UpdateChatStatus.
|
||||
select {
|
||||
case <-wakeCh:
|
||||
t.Fatal("wake channel should not have a signal after insert failure")
|
||||
default:
|
||||
// No signal, as expected.
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user