mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix(chatd): deliver retry control events via pubsub (#23349)
> **PR Stack** > 1. #23351 ← `#23282` > 2. #23282 ← `#23275` > 3. #23275 ← `#23349` > 4. **#23349** ← `main` *(you are here)* --- Retry events were published only to the local in-process stream via `publishEvent()`. When pubsub is active, `Subscribe()`'s merge loop only forwarded durable events (messages, status, errors) from pubsub notifications, so retry events were silently dropped for cross-replica subscribers. This adds a `publishRetry()` helper that publishes both locally and via pubsub, and extends the `Subscribe()` notification handler to forward retry events. **Changes:** - `coderd/pubsub/chatstreamnotify.go`: add `Retry` field to notify message - `coderd/chatd/chatd.go`: add `publishRetry()`, update `OnRetry` callback, extend `Subscribe()` to forward `notify.Retry` - `coderd/chatd/chatd_internal_test.go`: focused pubsub delivery test - `enterprise/coderd/chatd/chatd_test.go`: cross-replica end-to-end test
This commit is contained in:
+34
-12
@@ -1695,8 +1695,9 @@ func (p *Server) Subscribe(
|
||||
var allCancels []func()
|
||||
allCancels = append(allCancels, localCancel)
|
||||
|
||||
// Subscribe to pubsub for durable events (status, messages,
|
||||
// queue updates, errors). When pubsub is nil (e.g. in-memory
|
||||
// Subscribe to pubsub for durable and structured control
|
||||
// events (status, messages, queue updates, retry, errors).
|
||||
// When pubsub is nil (e.g. in-memory
|
||||
// single-instance) we skip this and deliver all local events.
|
||||
//
|
||||
// This MUST happen before the DB queries below so that any
|
||||
@@ -1964,6 +1965,17 @@ func (p *Server) Subscribe(
|
||||
}
|
||||
}
|
||||
}
|
||||
if notify.Retry != nil {
|
||||
select {
|
||||
case <-mergedCtx.Done():
|
||||
return
|
||||
case mergedEvents <- codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeRetry,
|
||||
ChatID: chatID,
|
||||
Retry: notify.Retry,
|
||||
}:
|
||||
}
|
||||
}
|
||||
if notify.Error != "" {
|
||||
select {
|
||||
case <-mergedCtx.Done():
|
||||
@@ -2072,7 +2084,8 @@ func (p *Server) publishStatus(chatID uuid.UUID, status database.ChatStatus, wor
|
||||
}
|
||||
|
||||
// publishChatStreamNotify broadcasts a per-chat stream notification via
|
||||
// PostgreSQL pubsub so that all replicas can read updates from the database.
|
||||
// PostgreSQL pubsub so that all replicas can merge durable database updates
|
||||
// with transient control events.
|
||||
func (p *Server) publishChatStreamNotify(chatID uuid.UUID, notify coderdpubsub.ChatStreamNotifyMessage) {
|
||||
if p.pubsub == nil {
|
||||
return
|
||||
@@ -2169,6 +2182,19 @@ func (p *Server) PublishDiffStatusChange(ctx context.Context, chatID uuid.UUID)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *Server) publishRetry(chatID uuid.UUID, payload *codersdk.ChatStreamRetry) {
|
||||
if payload == nil {
|
||||
return
|
||||
}
|
||||
p.publishEvent(chatID, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeRetry,
|
||||
Retry: payload,
|
||||
})
|
||||
p.publishChatStreamNotify(chatID, coderdpubsub.ChatStreamNotifyMessage{
|
||||
Retry: payload,
|
||||
})
|
||||
}
|
||||
|
||||
func (p *Server) publishError(chatID uuid.UUID, message string) {
|
||||
message = strings.TrimSpace(message)
|
||||
if message == "" {
|
||||
@@ -3263,15 +3289,11 @@ func (p *Server) runChat(
|
||||
slog.F("delay", delay.String()),
|
||||
slog.Error(retryErr),
|
||||
)
|
||||
p.publishEvent(chat.ID, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeRetry,
|
||||
ChatID: chat.ID,
|
||||
Retry: &codersdk.ChatStreamRetry{
|
||||
Attempt: attempt,
|
||||
DelayMs: delay.Milliseconds(),
|
||||
Error: retryErr.Error(),
|
||||
RetryingAt: time.Now().Add(delay),
|
||||
},
|
||||
p.publishRetry(chat.ID, &codersdk.ChatStreamRetry{
|
||||
Attempt: attempt,
|
||||
DelayMs: delay.Milliseconds(),
|
||||
Error: retryErr.Error(),
|
||||
RetryingAt: time.Now().Add(delay),
|
||||
})
|
||||
},
|
||||
|
||||
|
||||
@@ -420,6 +420,47 @@ func TestSubscribeFullRefreshStillUsesDatabaseCatchup(t *testing.T) {
|
||||
requireNoStreamEvent(t, events, 200*time.Millisecond)
|
||||
}
|
||||
|
||||
func TestSubscribeDeliversRetryEventViaPubsubOnce(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancelCtx := context.WithCancel(context.Background())
|
||||
defer cancelCtx()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
|
||||
chatID := uuid.New()
|
||||
chat := database.Chat{ID: chatID, Status: database.ChatStatusPending}
|
||||
|
||||
gomock.InOrder(
|
||||
db.EXPECT().GetChatMessagesByChatID(gomock.Any(), database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chatID,
|
||||
AfterID: 0,
|
||||
}).Return(nil, nil),
|
||||
db.EXPECT().GetChatQueuedMessages(gomock.Any(), chatID).Return(nil, nil),
|
||||
db.EXPECT().GetChatByID(gomock.Any(), chatID).Return(chat, nil),
|
||||
)
|
||||
|
||||
server := newSubscribeTestServer(t, db)
|
||||
_, events, cancel, ok := server.Subscribe(ctx, chatID, nil, 0)
|
||||
require.True(t, ok)
|
||||
defer cancel()
|
||||
|
||||
retryingAt := time.Unix(1_700_000_000, 0).UTC()
|
||||
expected := &codersdk.ChatStreamRetry{
|
||||
Attempt: 1,
|
||||
DelayMs: (1500 * time.Millisecond).Milliseconds(),
|
||||
Error: "rate limit exceeded",
|
||||
RetryingAt: retryingAt,
|
||||
}
|
||||
|
||||
server.publishRetry(chatID, expected)
|
||||
|
||||
event := requireStreamRetryEvent(t, events)
|
||||
require.Equal(t, expected, event.Retry)
|
||||
requireNoStreamEvent(t, events, 200*time.Millisecond)
|
||||
}
|
||||
|
||||
func newSubscribeTestServer(t *testing.T, db database.Store) *Server {
|
||||
t.Helper()
|
||||
|
||||
@@ -445,6 +486,21 @@ func requireStreamMessageEvent(t *testing.T, events <-chan codersdk.ChatStreamEv
|
||||
}
|
||||
}
|
||||
|
||||
func requireStreamRetryEvent(t *testing.T, events <-chan codersdk.ChatStreamEvent) codersdk.ChatStreamEvent {
|
||||
t.Helper()
|
||||
|
||||
select {
|
||||
case event, ok := <-events:
|
||||
require.True(t, ok, "chat stream closed before delivering an event")
|
||||
require.Equal(t, codersdk.ChatStreamEventTypeRetry, event.Type)
|
||||
require.NotNil(t, event.Retry)
|
||||
return event
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("timed out waiting for chat stream retry event")
|
||||
return codersdk.ChatStreamEvent{}
|
||||
}
|
||||
}
|
||||
|
||||
func requireNoStreamEvent(t *testing.T, events <-chan codersdk.ChatStreamEvent, wait time.Duration) {
|
||||
t.Helper()
|
||||
|
||||
|
||||
@@ -4,6 +4,8 @@ import (
|
||||
"fmt"
|
||||
|
||||
"github.com/google/uuid"
|
||||
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
// ChatStreamNotifyChannel returns the pubsub channel for per-chat
|
||||
@@ -14,8 +16,9 @@ func ChatStreamNotifyChannel(chatID uuid.UUID) string {
|
||||
}
|
||||
|
||||
// ChatStreamNotifyMessage is the payload published on the per-chat
|
||||
// stream notification channel. The actual message content is read
|
||||
// from the database by subscribers.
|
||||
// stream notification channel. Durable message content is still read
|
||||
// from the database, while transient control events can be carried
|
||||
// inline for cross-replica delivery.
|
||||
type ChatStreamNotifyMessage struct {
|
||||
// AfterMessageID tells subscribers to query messages after this
|
||||
// ID. Set when a new message is persisted.
|
||||
@@ -29,6 +32,11 @@ type ChatStreamNotifyMessage struct {
|
||||
// by enterprise relay to know where to connect.
|
||||
WorkerID string `json:"worker_id,omitempty"`
|
||||
|
||||
// Retry carries a structured retry event for cross-replica live
|
||||
// delivery. This is transient stream state and is not read back
|
||||
// from the database.
|
||||
Retry *codersdk.ChatStreamRetry `json:"retry,omitempty"`
|
||||
|
||||
// Error is set when a processing error occurs.
|
||||
Error string `json:"error,omitempty"`
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"math"
|
||||
"net/http"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
@@ -16,6 +17,7 @@ import (
|
||||
|
||||
"cdr.dev/slog/v3/sloggers/slogtest"
|
||||
osschatd "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/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
@@ -97,6 +99,50 @@ func seedChatDependencies(
|
||||
return user, model
|
||||
}
|
||||
|
||||
func newActiveWorkerServer(
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
ps dbpubsub.Pubsub,
|
||||
replicaID uuid.UUID,
|
||||
) *osschatd.Server {
|
||||
t.Helper()
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
server := osschatd.New(osschatd.Config{
|
||||
Logger: logger,
|
||||
Database: db,
|
||||
ReplicaID: replicaID,
|
||||
Pubsub: ps,
|
||||
PendingChatAcquireInterval: 10 * time.Millisecond,
|
||||
InFlightChatStaleAfter: testutil.WaitSuperLong,
|
||||
})
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, server.Close())
|
||||
})
|
||||
return server
|
||||
}
|
||||
|
||||
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 TestSubscribeRelayReconnectsOnDrop(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -400,6 +446,129 @@ func TestSubscribeRelaySnapshotDelivered(t *testing.T) {
|
||||
require.True(t, hasStatus, "initial snapshot should contain status event")
|
||||
}
|
||||
|
||||
func TestSubscribeRetryEventAcrossInstances(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
workerID := uuid.New()
|
||||
subscriberID := uuid.New()
|
||||
|
||||
var streamCalls atomic.Int32
|
||||
firstStreamStarted := make(chan struct{})
|
||||
allowFirstFailure := make(chan struct{})
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("retry-across-instances")
|
||||
}
|
||||
if streamCalls.Add(1) == 1 {
|
||||
select {
|
||||
case <-firstStreamStarted:
|
||||
default:
|
||||
close(firstStreamStarted)
|
||||
}
|
||||
<-allowFirstFailure
|
||||
return chattest.OpenAIRateLimitResponse()
|
||||
}
|
||||
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("retry", " complete")...)
|
||||
})
|
||||
|
||||
worker := newActiveWorkerServer(t, db, ps, workerID)
|
||||
subscriber := newTestServer(t, db, ps, subscriberID, func(
|
||||
ctx context.Context,
|
||||
chatID uuid.UUID,
|
||||
targetWorkerID uuid.UUID,
|
||||
requestHeader http.Header,
|
||||
) (
|
||||
[]codersdk.ChatStreamEvent,
|
||||
<-chan codersdk.ChatStreamEvent,
|
||||
func(),
|
||||
error,
|
||||
) {
|
||||
if targetWorkerID != workerID {
|
||||
return nil, nil, nil, xerrors.Errorf("unexpected relay target %s", targetWorkerID)
|
||||
}
|
||||
snapshot, events, cancel, ok := worker.Subscribe(ctx, chatID, requestHeader, math.MaxInt64)
|
||||
if !ok {
|
||||
return nil, nil, nil, xerrors.New("worker subscribe failed")
|
||||
}
|
||||
return snapshot, events, cancel, nil
|
||||
}, nil)
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, model := seedChatDependencies(ctx, t, db)
|
||||
setOpenAIProviderBaseURL(ctx, t, db, openAIURL)
|
||||
|
||||
chat, err := worker.CreateChat(ctx, osschatd.CreateOptions{
|
||||
OwnerID: user.ID,
|
||||
Title: "retry-across-instances",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("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 && fromDB.WorkerID.UUID == workerID
|
||||
}, testutil.WaitMedium, testutil.IntervalFast)
|
||||
|
||||
select {
|
||||
case <-firstStreamStarted:
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out waiting for first streaming attempt")
|
||||
}
|
||||
|
||||
_, events, cancel, ok := subscriber.Subscribe(ctx, chat.ID, nil, 0)
|
||||
require.True(t, ok)
|
||||
defer cancel()
|
||||
|
||||
close(allowFirstFailure)
|
||||
|
||||
var retryEvent *codersdk.ChatStreamRetry
|
||||
var waitingSeen bool
|
||||
var waitingBeforeRetry bool
|
||||
var assistantMessageBeforeRetry bool
|
||||
require.Eventually(t, func() bool {
|
||||
select {
|
||||
case event, ok := <-events:
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
switch event.Type {
|
||||
case codersdk.ChatStreamEventTypeRetry:
|
||||
retryEvent = event.Retry
|
||||
case codersdk.ChatStreamEventTypeMessage:
|
||||
if event.Message != nil && event.Message.Role == codersdk.ChatMessageRoleAssistant {
|
||||
if retryEvent == nil {
|
||||
assistantMessageBeforeRetry = true
|
||||
}
|
||||
}
|
||||
case codersdk.ChatStreamEventTypeStatus:
|
||||
if event.Status != nil && event.Status.Status == codersdk.ChatStatusWaiting {
|
||||
if retryEvent == nil {
|
||||
waitingBeforeRetry = true
|
||||
}
|
||||
waitingSeen = true
|
||||
}
|
||||
}
|
||||
return retryEvent != nil && waitingSeen
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}, testutil.WaitLong, testutil.IntervalFast)
|
||||
|
||||
require.NotNil(t, retryEvent)
|
||||
require.Equal(t, 1, retryEvent.Attempt)
|
||||
require.Greater(t, retryEvent.DelayMs, int64(0))
|
||||
require.Contains(t, retryEvent.Error, "Rate limit exceeded")
|
||||
require.False(t, assistantMessageBeforeRetry)
|
||||
require.False(t, waitingBeforeRetry)
|
||||
require.GreaterOrEqual(t, streamCalls.Load(), int32(2))
|
||||
}
|
||||
|
||||
// TestSubscribeRelayStaleDialDiscardedAfterInterrupt verifies that when a
|
||||
// user interrupts a streaming chat and sends a new message (which gets
|
||||
// picked up by a different replica), an in-flight relay dial to the
|
||||
|
||||
Reference in New Issue
Block a user