From 2abe55549c82d7f1bd8e0b7eb6550081bc5e67ec Mon Sep 17 00:00:00 2001 From: Kyle Carberry Date: Sat, 28 Feb 2026 17:14:11 -0500 Subject: [PATCH] fix: return in-flight chats to pending on server shutdown (#22443) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit When a chatd server shuts down (`Close()`), the server context is canceled. Previously, in-flight chats would be marked as `error` because the `context.Canceled` error was not distinguished from actual processing failures. This adds `isShutdownCancellation()` to detect when the error is caused by the server context being canceled (as opposed to a chat-specific cancellation like `ErrInterrupted`). When detected, the chat status is set to `pending` with no `last_error`, allowing another replica to pick it up and retry. Extracted from #22440 — only the context cancellation bug fix, no chattest changes. --- coderd/chatd/chatd.go | 25 +++++++ coderd/chatd/chatd_test.go | 131 +++++++++++++++++++++++++++++++++++++ 2 files changed, 156 insertions(+) diff --git a/coderd/chatd/chatd.go b/coderd/chatd/chatd.go index aa698fa068..5b5cec2e53 100644 --- a/coderd/chatd/chatd.go +++ b/coderd/chatd/chatd.go @@ -1757,6 +1757,12 @@ func (p *Server) processChat(ctx context.Context, chat database.Chat) { status = database.ChatStatusWaiting return } + if isShutdownCancellation(ctx, chatCtx, err) { + logger.Info(ctx, "chat canceled during shutdown; returning to pending") + status = database.ChatStatusPending + lastError = "" + return + } logger.Error(ctx, "failed to process chat", slog.Error(err)) if reason, ok := processingFailureReason(err); ok { lastError = reason @@ -1767,6 +1773,25 @@ func (p *Server) processChat(ctx context.Context, chat database.Chat) { } } +func isShutdownCancellation( + serverCtx context.Context, + chatCtx context.Context, + err error, +) bool { + if err == nil { + return false + } + // During Close(), the server context is canceled. In-flight chats should + // be returned to pending so another replica can retry them. + if serverCtx.Err() == nil { + return false + } + if errors.Is(err, context.Canceled) { + return true + } + return errors.Is(context.Cause(chatCtx), context.Canceled) +} + func (p *Server) runChat( ctx context.Context, chat database.Chat, diff --git a/coderd/chatd/chatd_test.go b/coderd/chatd/chatd_test.go index 294af89a7c..d5bb3423d0 100644 --- a/coderd/chatd/chatd_test.go +++ b/coderd/chatd/chatd_test.go @@ -5,6 +5,7 @@ import ( "database/sql" "encoding/json" "errors" + "sync/atomic" "testing" "time" @@ -15,6 +16,7 @@ import ( "cdr.dev/slog/v3/sloggers/slogtest" "github.com/coder/coder/v2/coderd/chatd" + "github.com/coder/coder/v2/coderd/chatd/chattest" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/db2sdk" "github.com/coder/coder/v2/coderd/database/dbgen" @@ -755,3 +757,132 @@ func seedChatDependencies( require.NoError(t, err) return user, model } + +func setOpenAIProviderBaseURL( + ctx context.Context, + t *testing.T, + db database.Store, + baseURL string, +) { + t.Helper() + + provider, err := db.GetChatProviderByProvider(ctx, "openai") + require.NoError(t, err) + + _, err = db.UpdateChatProvider(ctx, database.UpdateChatProviderParams{ + ID: provider.ID, + DisplayName: provider.DisplayName, + APIKey: provider.APIKey, + BaseUrl: baseURL, + ApiKeyKeyID: provider.ApiKeyKeyID, + Enabled: provider.Enabled, + }) + require.NoError(t, err) +} + +func TestCloseDuringShutdownContextCanceledShouldRetryOnNewReplica(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + ctx := testutil.Context(t, testutil.WaitLong) + + var requestCount atomic.Int32 + streamStarted := make(chan struct{}) + openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { + if requestCount.Add(1) == 1 { + chunks := make(chan chattest.OpenAIChunk, 1) + go func() { + defer close(chunks) + chunks <- chattest.OpenAITextChunks("partial")[0] + select { + case <-streamStarted: + default: + close(streamStarted) + } + <-req.Context().Done() + }() + return chattest.OpenAIResponse{StreamingChunks: chunks} + } + return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("retry", " complete")...) + }) + + loggerA := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + serverA := chatd.New(chatd.Config{ + Logger: loggerA, + Database: db, + ReplicaID: uuid.New(), + Pubsub: ps, + PendingChatAcquireInterval: 10 * time.Millisecond, + InFlightChatStaleAfter: testutil.WaitSuperLong, + }) + t.Cleanup(func() { + require.NoError(t, serverA.Close()) + }) + + user, model := seedChatDependencies(ctx, t, db) + setOpenAIProviderBaseURL(ctx, t, db, openAIURL) + + chat, err := serverA.CreateChat(ctx, chatd.CreateOptions{ + OwnerID: user.ID, + Title: "shutdown-retry", + ModelConfigID: model.ID, + InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}}, + }) + require.NoError(t, err) + + require.Eventually(t, func() bool { + fromDB, dbErr := db.GetChatByID(ctx, chat.ID) + if dbErr != nil { + return false + } + return fromDB.Status == database.ChatStatusRunning && fromDB.WorkerID.Valid + }, testutil.WaitMedium, testutil.IntervalFast) + + require.Eventually(t, func() bool { + select { + case <-streamStarted: + return true + default: + return false + } + }, testutil.WaitMedium, testutil.IntervalFast) + + require.NoError(t, serverA.Close()) + + require.Eventually(t, func() bool { + fromDB, dbErr := db.GetChatByID(ctx, chat.ID) + if dbErr != nil { + return false + } + return fromDB.Status == database.ChatStatusPending && + !fromDB.WorkerID.Valid && + !fromDB.LastError.Valid + }, testutil.WaitMedium, testutil.IntervalFast) + + loggerB := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + serverB := chatd.New(chatd.Config{ + Logger: loggerB, + Database: db, + ReplicaID: uuid.New(), + Pubsub: ps, + PendingChatAcquireInterval: 10 * time.Millisecond, + InFlightChatStaleAfter: testutil.WaitSuperLong, + }) + t.Cleanup(func() { + require.NoError(t, serverB.Close()) + }) + + require.Eventually(t, func() bool { + return requestCount.Load() >= 2 + }, testutil.WaitMedium, testutil.IntervalFast) + + require.Eventually(t, func() bool { + fromDB, dbErr := db.GetChatByID(ctx, chat.ID) + if dbErr != nil { + return false + } + return fromDB.Status == database.ChatStatusWaiting && + !fromDB.WorkerID.Valid && + !fromDB.LastError.Valid + }, testutil.WaitMedium, testutil.IntervalFast) +}