mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix: resolve bugs in chatd streaming system (#22720)
Split from #22693 per review feedback. Fixes multiple bugs in coderd/chatd and sub-packages including race conditions, transaction safety, stream buffer bounds, retry limits, and enterprise relay improvements. See commit message for full list.
This commit is contained in:
+428
-241
@@ -43,6 +43,10 @@ const (
|
||||
instructionCacheTTL = 5 * time.Minute
|
||||
chatHeartbeatInterval = 30 * time.Second
|
||||
maxChatSteps = 1200
|
||||
// maxStreamBufferSize caps the number of events buffered
|
||||
// per chat during a single LLM step. When exceeded the
|
||||
// oldest event is evicted so memory stays bounded.
|
||||
maxStreamBufferSize = 10000
|
||||
|
||||
// staleRecoveryIntervalDivisor determines how often the stale
|
||||
// recovery loop runs relative to the stale threshold. A value
|
||||
@@ -104,12 +108,14 @@ type AgentConnFunc func(ctx context.Context, agentID uuid.UUID) (workspacesdk.Ag
|
||||
// - ctx: subscription lifetime context (canceled on unsubscribe).
|
||||
// - params: all state needed to build the merged stream.
|
||||
//
|
||||
// Returns the merged event channel and a cleanup function.
|
||||
// Returns the merged event channel. Cleanup is driven by ctx
|
||||
// cancellation — the merge goroutine tears down all relay state
|
||||
// in its defer when ctx is done.
|
||||
// Set by enterprise for HA deployments. Nil in AGPL single-replica.
|
||||
type SubscribeFn func(
|
||||
ctx context.Context,
|
||||
params SubscribeFnParams,
|
||||
) (<-chan codersdk.ChatStreamEvent, func())
|
||||
) <-chan codersdk.ChatStreamEvent
|
||||
|
||||
// StatusNotification informs the enterprise relay manager of chat
|
||||
// status changes so it can open or close relay connections.
|
||||
@@ -521,7 +527,7 @@ func (p *Server) EditMessage(
|
||||
return EditMessageResult{}, txErr
|
||||
}
|
||||
|
||||
p.publishMessage(opts.ChatID, result.Message)
|
||||
p.publishEditedMessage(opts.ChatID, result.Message)
|
||||
p.publishEvent(opts.ChatID, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeQueueUpdate,
|
||||
QueuedMessages: []codersdk.ChatQueuedMessage{},
|
||||
@@ -554,6 +560,26 @@ func (p *Server) ArchiveChat(ctx context.Context, chatID uuid.UUID) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// UnarchiveChat unarchives a chat and publishes a created event so sidebar
|
||||
// clients are notified that the chat has reappeared.
|
||||
func (p *Server) UnarchiveChat(ctx context.Context, chatID uuid.UUID) error {
|
||||
if chatID == uuid.Nil {
|
||||
return xerrors.New("chat_id is required")
|
||||
}
|
||||
|
||||
chat, err := p.db.GetChatByID(ctx, chatID)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get chat: %w", err)
|
||||
}
|
||||
|
||||
if err := p.db.UnarchiveChatByID(ctx, chatID); err != nil {
|
||||
return xerrors.Errorf("unarchive chat: %w", err)
|
||||
}
|
||||
|
||||
p.publishChatPubsubEvent(chat, coderdpubsub.ChatEventKindCreated)
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteQueued removes a queued user message and publishes the queue update.
|
||||
func (p *Server) DeleteQueued(
|
||||
ctx context.Context,
|
||||
@@ -564,28 +590,51 @@ func (p *Server) DeleteQueued(
|
||||
return xerrors.New("chat_id is required")
|
||||
}
|
||||
|
||||
err := p.db.DeleteChatQueuedMessage(ctx, database.DeleteChatQueuedMessageParams{
|
||||
ID: queuedMessageID,
|
||||
ChatID: chatID,
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("delete queued message: %w", err)
|
||||
}
|
||||
var queuedMessages []database.ChatQueuedMessage
|
||||
var queueLoadedOK bool
|
||||
|
||||
txErr := p.db.InTx(func(tx database.Store) error {
|
||||
// Lock the chat row to prevent processChat from
|
||||
// auto-promoting a message the user intended to delete.
|
||||
if _, err := tx.GetChatByIDForUpdate(ctx, chatID); err != nil {
|
||||
return xerrors.Errorf("lock chat: %w", err)
|
||||
}
|
||||
|
||||
err := tx.DeleteChatQueuedMessage(ctx, database.DeleteChatQueuedMessageParams{
|
||||
ID: queuedMessageID,
|
||||
ChatID: chatID,
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("delete queued message: %w", err)
|
||||
}
|
||||
|
||||
var err2 error
|
||||
queuedMessages, err2 = tx.GetChatQueuedMessages(ctx, chatID)
|
||||
if err2 != nil {
|
||||
p.logger.Warn(ctx, "failed to load queued messages after delete",
|
||||
slog.F("chat_id", chatID),
|
||||
slog.F("queued_message_id", queuedMessageID),
|
||||
slog.Error(err2),
|
||||
)
|
||||
// Non-fatal: the delete succeeded, so we still commit.
|
||||
return nil
|
||||
}
|
||||
queueLoadedOK = true
|
||||
|
||||
queuedMessages, err := p.db.GetChatQueuedMessages(ctx, chatID)
|
||||
if err != nil {
|
||||
p.logger.Warn(ctx, "failed to load queued messages after delete",
|
||||
slog.F("chat_id", chatID),
|
||||
slog.F("queued_message_id", queuedMessageID),
|
||||
slog.Error(err),
|
||||
)
|
||||
return nil
|
||||
}, nil)
|
||||
if txErr != nil {
|
||||
return txErr
|
||||
}
|
||||
|
||||
p.publishEvent(chatID, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeQueueUpdate,
|
||||
QueuedMessages: db2sdk.ChatQueuedMessages(queuedMessages),
|
||||
})
|
||||
if queueLoadedOK {
|
||||
p.publishEvent(chatID, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeQueueUpdate,
|
||||
QueuedMessages: db2sdk.ChatQueuedMessages(queuedMessages),
|
||||
})
|
||||
}
|
||||
// Always notify subscribers so they can re-fetch, even if we
|
||||
// failed to load the updated queue payload above.
|
||||
p.publishChatStreamNotify(chatID, coderdpubsub.ChatStreamNotifyMessage{
|
||||
QueueUpdate: true,
|
||||
})
|
||||
@@ -681,6 +730,7 @@ func (p *Server) PromoteQueued(
|
||||
})
|
||||
p.publishMessage(opts.ChatID, promoted)
|
||||
p.publishStatus(opts.ChatID, updatedChat.Status, updatedChat.WorkerID)
|
||||
p.publishChatPubsubEvent(updatedChat, coderdpubsub.ChatEventKindStatusChange)
|
||||
|
||||
return result, nil
|
||||
}
|
||||
@@ -964,9 +1014,15 @@ func (p *Server) publishToStream(chatID uuid.UUID, event codersdk.ChatStreamEven
|
||||
state := p.streamStateLocked(chatID)
|
||||
if event.Type == codersdk.ChatStreamEventTypeMessagePart {
|
||||
if !state.buffering {
|
||||
p.cleanupStreamIfIdleLocked(chatID, state)
|
||||
p.streamMu.Unlock()
|
||||
return
|
||||
}
|
||||
if len(state.buffer) >= maxStreamBufferSize {
|
||||
p.logger.Warn(context.Background(), "chat stream buffer full, dropping oldest event",
|
||||
slog.F("chat_id", chatID), slog.F("buffer_size", len(state.buffer)))
|
||||
state.buffer = state.buffer[1:]
|
||||
}
|
||||
state.buffer = append(state.buffer, event)
|
||||
}
|
||||
subscribers := make([]chan codersdk.ChatStreamEvent, 0, len(state.subscribers))
|
||||
@@ -983,6 +1039,15 @@ func (p *Server) publishToStream(chatID uuid.UUID, event codersdk.ChatStreamEven
|
||||
slog.F("chat_id", chatID), slog.F("type", event.Type))
|
||||
}
|
||||
}
|
||||
|
||||
// Clean up the stream entry if it was created by
|
||||
// streamStateLocked but has no subscribers and is not
|
||||
// actively buffering (e.g. publish with no watchers).
|
||||
p.streamMu.Lock()
|
||||
if cur, ok := p.chatStreams[chatID]; ok {
|
||||
p.cleanupStreamIfIdleLocked(chatID, cur)
|
||||
}
|
||||
p.streamMu.Unlock()
|
||||
}
|
||||
|
||||
func (p *Server) subscribeToStream(chatID uuid.UUID) (
|
||||
@@ -1055,73 +1120,6 @@ func (p *Server) Subscribe(
|
||||
// Subscribe to local stream for message_parts (ephemeral).
|
||||
localSnapshot, localParts, localCancel := p.subscribeToStream(chatID)
|
||||
|
||||
// Build initial snapshot synchronously.
|
||||
initialSnapshot := make([]codersdk.ChatStreamEvent, 0)
|
||||
// Add local message_parts to snapshot
|
||||
for _, event := range localSnapshot {
|
||||
if event.Type == codersdk.ChatStreamEventTypeMessagePart {
|
||||
initialSnapshot = append(initialSnapshot, event)
|
||||
}
|
||||
}
|
||||
|
||||
// Load initial messages from DB. When afterMessageID > 0 the
|
||||
// caller already has messages up to that ID (e.g. from the REST
|
||||
// endpoint), so we only fetch newer ones to avoid sending
|
||||
// duplicate data.
|
||||
messages, err := p.db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chatID,
|
||||
AfterID: afterMessageID,
|
||||
})
|
||||
if err == nil {
|
||||
for _, msg := range messages {
|
||||
sdkMsg := db2sdk.ChatMessage(msg)
|
||||
initialSnapshot = append(initialSnapshot, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeMessage,
|
||||
ChatID: chatID,
|
||||
Message: &sdkMsg,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Load initial queue.
|
||||
queued, err := p.db.GetChatQueuedMessages(ctx, chatID)
|
||||
if err == nil && len(queued) > 0 {
|
||||
initialSnapshot = append(initialSnapshot, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeQueueUpdate,
|
||||
ChatID: chatID,
|
||||
QueuedMessages: db2sdk.ChatQueuedMessages(queued),
|
||||
})
|
||||
}
|
||||
|
||||
// Get initial chat state to determine if we need a relay.
|
||||
chat, err := p.db.GetChatByID(ctx, chatID)
|
||||
|
||||
// Include the current chat status in the snapshot so the
|
||||
// frontend can gate message_part processing correctly from
|
||||
// the very first batch, without waiting for a separate REST
|
||||
// query.
|
||||
if err == nil {
|
||||
statusEvent := codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeStatus,
|
||||
ChatID: chatID,
|
||||
Status: &codersdk.ChatStreamStatus{
|
||||
Status: codersdk.ChatStatus(chat.Status),
|
||||
},
|
||||
}
|
||||
// Prepend so the frontend sees the status before any
|
||||
// message_part events.
|
||||
initialSnapshot = append([]codersdk.ChatStreamEvent{statusEvent}, initialSnapshot...)
|
||||
}
|
||||
|
||||
// Track the last message ID we've seen for DB queries.
|
||||
// Initialize from afterMessageID so that when the caller passes
|
||||
// afterMessageID > 0 but no new messages exist yet, the first
|
||||
// pubsub catch-up doesn't re-fetch already-seen messages.
|
||||
lastMessageID := afterMessageID
|
||||
if len(messages) > 0 {
|
||||
lastMessageID = messages[len(messages)-1].ID
|
||||
}
|
||||
|
||||
// Merge all event sources.
|
||||
mergedCtx, mergedCancel := context.WithCancel(ctx)
|
||||
mergedEvents := make(chan codersdk.ChatStreamEvent, 128)
|
||||
@@ -1132,6 +1130,10 @@ func (p *Server) Subscribe(
|
||||
// Subscribe to pubsub for durable events (status, messages,
|
||||
// queue updates, 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
|
||||
// notification published between the query and the subscription
|
||||
// is not lost (subscribe-first-then-query pattern).
|
||||
var notifications <-chan coderdpubsub.ChatStreamNotifyMessage
|
||||
var errCh <-chan error
|
||||
if p.pubsub != nil {
|
||||
@@ -1175,18 +1177,115 @@ func (p *Server) Subscribe(
|
||||
}
|
||||
}
|
||||
|
||||
// Build initial snapshot synchronously. The pubsub subscription
|
||||
// is already active so no notifications can be lost during this
|
||||
// window.
|
||||
initialSnapshot := make([]codersdk.ChatStreamEvent, 0)
|
||||
// Add local message_parts to snapshot
|
||||
for _, event := range localSnapshot {
|
||||
if event.Type == codersdk.ChatStreamEventTypeMessagePart {
|
||||
initialSnapshot = append(initialSnapshot, event)
|
||||
}
|
||||
}
|
||||
|
||||
// Load initial messages from DB. When afterMessageID > 0 the
|
||||
// caller already has messages up to that ID (e.g. from the REST
|
||||
// endpoint), so we only fetch newer ones to avoid sending
|
||||
// duplicate data.
|
||||
messages, err := p.db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chatID,
|
||||
AfterID: afterMessageID,
|
||||
})
|
||||
if err != nil {
|
||||
p.logger.Error(ctx, "failed to load initial chat messages",
|
||||
slog.Error(err),
|
||||
slog.F("chat_id", chatID),
|
||||
)
|
||||
initialSnapshot = append(initialSnapshot, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeError,
|
||||
ChatID: chatID,
|
||||
Error: &codersdk.ChatStreamError{Message: "failed to load initial snapshot"},
|
||||
})
|
||||
} else {
|
||||
for _, msg := range messages {
|
||||
sdkMsg := db2sdk.ChatMessage(msg)
|
||||
initialSnapshot = append(initialSnapshot, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeMessage,
|
||||
ChatID: chatID,
|
||||
Message: &sdkMsg,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Load initial queue.
|
||||
queued, err := p.db.GetChatQueuedMessages(ctx, chatID)
|
||||
if err != nil {
|
||||
p.logger.Error(ctx, "failed to load initial queued messages",
|
||||
slog.Error(err),
|
||||
slog.F("chat_id", chatID),
|
||||
)
|
||||
initialSnapshot = append(initialSnapshot, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeError,
|
||||
ChatID: chatID,
|
||||
Error: &codersdk.ChatStreamError{Message: "failed to load initial snapshot"},
|
||||
})
|
||||
} else if len(queued) > 0 {
|
||||
initialSnapshot = append(initialSnapshot, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeQueueUpdate,
|
||||
ChatID: chatID,
|
||||
QueuedMessages: db2sdk.ChatQueuedMessages(queued),
|
||||
})
|
||||
}
|
||||
|
||||
// Get initial chat state to determine if we need a relay.
|
||||
chat, chatErr := p.db.GetChatByID(ctx, chatID)
|
||||
|
||||
// Include the current chat status in the snapshot so the
|
||||
// frontend can gate message_part processing correctly from
|
||||
// the very first batch, without waiting for a separate REST
|
||||
// query.
|
||||
if chatErr != nil {
|
||||
p.logger.Error(ctx, "failed to load initial chat state",
|
||||
slog.Error(chatErr),
|
||||
slog.F("chat_id", chatID),
|
||||
)
|
||||
initialSnapshot = append(initialSnapshot, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeError,
|
||||
ChatID: chatID,
|
||||
Error: &codersdk.ChatStreamError{Message: "failed to load initial snapshot"},
|
||||
})
|
||||
} else {
|
||||
statusEvent := codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeStatus,
|
||||
ChatID: chatID,
|
||||
Status: &codersdk.ChatStreamStatus{
|
||||
Status: codersdk.ChatStatus(chat.Status),
|
||||
},
|
||||
}
|
||||
// Prepend so the frontend sees the status before any
|
||||
// message_part events.
|
||||
initialSnapshot = append([]codersdk.ChatStreamEvent{statusEvent}, initialSnapshot...)
|
||||
}
|
||||
|
||||
// Track the last message ID we've seen for DB queries.
|
||||
// Initialize from afterMessageID so that when the caller passes
|
||||
// afterMessageID > 0 but no new messages exist yet, the first
|
||||
// pubsub catch-up doesn't re-fetch already-seen messages.
|
||||
lastMessageID := afterMessageID
|
||||
if len(messages) > 0 {
|
||||
lastMessageID = messages[len(messages)-1].ID
|
||||
}
|
||||
|
||||
// When an enterprise SubscribeFn is provided and the chat
|
||||
// lookup succeeded, call it to get relay events (message_parts
|
||||
// from remote replicas). OSS now owns pubsub subscription,
|
||||
// message catch-up, queue updates, and status forwarding;
|
||||
// enterprise only manages relay dialing.
|
||||
var relayEvents <-chan codersdk.ChatStreamEvent
|
||||
var relayCleanup func()
|
||||
var statusNotifications chan StatusNotification
|
||||
if p.subscribeFn != nil && err == nil {
|
||||
if p.subscribeFn != nil && chatErr == nil {
|
||||
statusNotifications = make(chan StatusNotification, 10)
|
||||
var relayEvCh <-chan codersdk.ChatStreamEvent
|
||||
relayEvCh, relayCleanup = p.subscribeFn(mergedCtx, SubscribeFnParams{
|
||||
relayEvents = p.subscribeFn(mergedCtx, SubscribeFnParams{
|
||||
ChatID: chatID,
|
||||
Chat: chat,
|
||||
WorkerID: p.workerID,
|
||||
@@ -1195,9 +1294,7 @@ func (p *Server) Subscribe(
|
||||
DB: p.db,
|
||||
Logger: p.logger,
|
||||
})
|
||||
relayEvents = relayEvCh
|
||||
}
|
||||
|
||||
hasPubsub := false
|
||||
if p.pubsub != nil {
|
||||
// hasPubsub is only true when we actually subscribed
|
||||
@@ -1367,9 +1464,6 @@ func (p *Server) Subscribe(
|
||||
cancelFn()
|
||||
}
|
||||
}
|
||||
if relayCleanup != nil {
|
||||
relayCleanup()
|
||||
}
|
||||
}
|
||||
return initialSnapshot, mergedEvents, cancel, true
|
||||
}
|
||||
@@ -1538,6 +1632,20 @@ func (p *Server) publishMessage(chatID uuid.UUID, message database.ChatMessage)
|
||||
})
|
||||
}
|
||||
|
||||
// publishEditedMessage is like publishMessage but uses
|
||||
// AfterMessageID=0 so remote subscribers re-fetch from the
|
||||
// beginning, ensuring the edit is never silently dropped.
|
||||
func (p *Server) publishEditedMessage(chatID uuid.UUID, message database.ChatMessage) {
|
||||
sdkMessage := db2sdk.ChatMessage(message)
|
||||
p.publishEvent(chatID, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeMessage,
|
||||
Message: &sdkMessage,
|
||||
})
|
||||
p.publishChatStreamNotify(chatID, coderdpubsub.ChatStreamNotifyMessage{
|
||||
AfterMessageID: 0,
|
||||
})
|
||||
}
|
||||
|
||||
func (p *Server) publishMessagePart(chatID uuid.UUID, role string, part codersdk.ChatMessagePart) {
|
||||
if part.Type == "" {
|
||||
return
|
||||
@@ -1704,6 +1812,7 @@ func (p *Server) processChat(ctx context.Context, chat database.Chat) {
|
||||
lastError := ""
|
||||
remainingQueuedMessages := []database.ChatQueuedMessage{}
|
||||
shouldPublishQueueUpdate := false
|
||||
var promotedMessage *database.ChatMessage
|
||||
|
||||
defer func() {
|
||||
// Use a context that is not canceled by Close() so we can
|
||||
@@ -1744,7 +1853,7 @@ func (p *Server) processChat(ctx context.Context, chat database.Chat) {
|
||||
// If someone else already set the chat to pending (e.g.
|
||||
// the promote endpoint), don't overwrite it — just clear
|
||||
// the worker and let the processor pick it back up.
|
||||
if latestChat.Status == database.ChatStatusPending && status == database.ChatStatusWaiting {
|
||||
if latestChat.Status == database.ChatStatusPending {
|
||||
status = database.ChatStatusPending
|
||||
} else if status == database.ChatStatusWaiting {
|
||||
// Try to auto-promote the next queued message.
|
||||
@@ -1773,8 +1882,7 @@ func (p *Server) processChat(ctx context.Context, chat database.Chat) {
|
||||
slog.F("queued_message_id", nextQueued.ID), slog.Error(insertErr))
|
||||
} else {
|
||||
status = database.ChatStatusPending
|
||||
|
||||
p.publishMessage(chat.ID, msg)
|
||||
promotedMessage = &msg
|
||||
|
||||
remaining, qErr := tx.GetChatQueuedMessages(cleanupCtx, chat.ID)
|
||||
if qErr == nil {
|
||||
@@ -1803,8 +1911,13 @@ func (p *Server) processChat(ctx context.Context, chat database.Chat) {
|
||||
}
|
||||
if err != nil {
|
||||
logger.Error(cleanupCtx, "failed to release chat", slog.Error(err))
|
||||
return
|
||||
}
|
||||
if err == nil && shouldPublishQueueUpdate {
|
||||
|
||||
if promotedMessage != nil {
|
||||
p.publishMessage(chat.ID, *promotedMessage)
|
||||
}
|
||||
if shouldPublishQueueUpdate {
|
||||
p.publishEvent(chat.ID, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeQueueUpdate,
|
||||
QueuedMessages: db2sdk.ChatQueuedMessages(remainingQueuedMessages),
|
||||
@@ -1840,7 +1953,6 @@ func (p *Server) processChat(ctx context.Context, chat database.Chat) {
|
||||
Body: "Agent has finished running.",
|
||||
Icon: "/favicon.ico",
|
||||
Data: map[string]string{"url": fmt.Sprintf("/agents/%s", chat.ID)},
|
||||
Tag: chat.ID.String(),
|
||||
}
|
||||
if status == database.ChatStatusError {
|
||||
pushMsg.Body = "Agent encountered an error."
|
||||
@@ -2033,10 +2145,7 @@ func (p *Server) runChat(
|
||||
if conn == nil {
|
||||
conn = agentConn
|
||||
releaseConn = agentRelease
|
||||
chatStateMu.Unlock()
|
||||
|
||||
// Inject chat identity headers so agent-side
|
||||
// handlers can track which paths this chat edits.
|
||||
var ancestorIDs []string
|
||||
if chatSnapshot.ParentChatID.Valid {
|
||||
ancestorIDs = append(ancestorIDs, chatSnapshot.ParentChatID.UUID.String())
|
||||
@@ -2051,6 +2160,7 @@ func (p *Server) runChat(
|
||||
workspacesdk.CoderAncestorChatIDsHeader: {string(ancestorJSON)},
|
||||
})
|
||||
|
||||
chatStateMu.Unlock()
|
||||
return agentConn, nil
|
||||
}
|
||||
currentConn := conn
|
||||
@@ -2070,6 +2180,16 @@ func (p *Server) runChat(
|
||||
modelConfigContextLimit := modelConfig.ContextLimit
|
||||
|
||||
persistStep := func(persistCtx context.Context, step chatloop.PersistedStep) error {
|
||||
// If the chat context has been canceled (e.g. by an
|
||||
// EditMessage call), bail out before inserting any
|
||||
// messages. This closes the race window between
|
||||
// EditMessage committing its transaction (which deletes
|
||||
// messages after the edit point) and the cancellation
|
||||
// propagating to the processing loop.
|
||||
if persistCtx.Err() != nil {
|
||||
return chatloop.ErrInterrupted
|
||||
}
|
||||
|
||||
// Split the step content into assistant blocks and tool
|
||||
// result blocks so they can be stored as separate messages
|
||||
// with the appropriate roles.
|
||||
@@ -2087,66 +2207,77 @@ func (p *Server) runChat(
|
||||
assistantBlocks = append(assistantBlocks, block)
|
||||
}
|
||||
|
||||
if len(assistantBlocks) > 0 {
|
||||
assistantContent, err := chatprompt.MarshalContent(assistantBlocks, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
var insertedMessages []database.ChatMessage
|
||||
err := p.db.InTx(func(tx database.Store) error {
|
||||
if len(assistantBlocks) > 0 {
|
||||
assistantContent, marshalErr := chatprompt.MarshalContent(assistantBlocks, nil)
|
||||
if marshalErr != nil {
|
||||
return marshalErr
|
||||
}
|
||||
|
||||
hasUsage := step.Usage != (fantasy.Usage{})
|
||||
assistantMessage, insertErr := tx.InsertChatMessage(persistCtx, database.InsertChatMessageParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleAssistant),
|
||||
Content: assistantContent,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
InputTokens: usageNullInt64(step.Usage.InputTokens, hasUsage),
|
||||
OutputTokens: usageNullInt64(step.Usage.OutputTokens, hasUsage),
|
||||
TotalTokens: usageNullInt64(step.Usage.TotalTokens, hasUsage),
|
||||
ReasoningTokens: usageNullInt64(
|
||||
step.Usage.ReasoningTokens,
|
||||
hasUsage,
|
||||
),
|
||||
CacheCreationTokens: usageNullInt64(
|
||||
step.Usage.CacheCreationTokens,
|
||||
hasUsage,
|
||||
),
|
||||
CacheReadTokens: usageNullInt64(step.Usage.CacheReadTokens, hasUsage),
|
||||
ContextLimit: step.ContextLimit,
|
||||
Compressed: sql.NullBool{},
|
||||
})
|
||||
if insertErr != nil {
|
||||
return xerrors.Errorf("insert assistant message: %w", insertErr)
|
||||
}
|
||||
insertedMessages = append(insertedMessages, assistantMessage)
|
||||
}
|
||||
|
||||
hasUsage := step.Usage != (fantasy.Usage{})
|
||||
assistantMessage, err := p.db.InsertChatMessage(persistCtx, database.InsertChatMessageParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleAssistant),
|
||||
Content: assistantContent,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
InputTokens: usageNullInt64(step.Usage.InputTokens, hasUsage),
|
||||
OutputTokens: usageNullInt64(step.Usage.OutputTokens, hasUsage),
|
||||
TotalTokens: usageNullInt64(step.Usage.TotalTokens, hasUsage),
|
||||
ReasoningTokens: usageNullInt64(
|
||||
step.Usage.ReasoningTokens,
|
||||
hasUsage,
|
||||
),
|
||||
CacheCreationTokens: usageNullInt64(
|
||||
step.Usage.CacheCreationTokens,
|
||||
hasUsage,
|
||||
),
|
||||
CacheReadTokens: usageNullInt64(step.Usage.CacheReadTokens, hasUsage),
|
||||
ContextLimit: step.ContextLimit,
|
||||
Compressed: sql.NullBool{},
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("insert assistant message: %w", err)
|
||||
for _, tr := range toolResults {
|
||||
resultContent, marshalErr := chatprompt.MarshalToolResultContent(tr)
|
||||
if marshalErr != nil {
|
||||
return marshalErr
|
||||
}
|
||||
|
||||
toolMessage, insertErr := tx.InsertChatMessage(persistCtx, database.InsertChatMessageParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleTool),
|
||||
Content: resultContent,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
InputTokens: sql.NullInt64{},
|
||||
OutputTokens: sql.NullInt64{},
|
||||
TotalTokens: sql.NullInt64{},
|
||||
ReasoningTokens: sql.NullInt64{},
|
||||
CacheCreationTokens: sql.NullInt64{},
|
||||
CacheReadTokens: sql.NullInt64{},
|
||||
ContextLimit: sql.NullInt64{},
|
||||
Compressed: sql.NullBool{},
|
||||
})
|
||||
if insertErr != nil {
|
||||
return xerrors.Errorf("insert tool result: %w", insertErr)
|
||||
}
|
||||
insertedMessages = append(insertedMessages, toolMessage)
|
||||
}
|
||||
p.publishMessage(chat.ID, assistantMessage)
|
||||
|
||||
return nil
|
||||
}, nil)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("persist step transaction: %w", err)
|
||||
}
|
||||
|
||||
for _, tr := range toolResults {
|
||||
resultContent, err := chatprompt.MarshalToolResultContent(tr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
toolMessage, err := p.db.InsertChatMessage(persistCtx, database.InsertChatMessageParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleTool),
|
||||
Content: resultContent,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
InputTokens: sql.NullInt64{},
|
||||
OutputTokens: sql.NullInt64{},
|
||||
TotalTokens: sql.NullInt64{},
|
||||
ReasoningTokens: sql.NullInt64{},
|
||||
CacheCreationTokens: sql.NullInt64{},
|
||||
CacheReadTokens: sql.NullInt64{},
|
||||
ContextLimit: sql.NullInt64{},
|
||||
Compressed: sql.NullBool{},
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("insert tool result: %w", err)
|
||||
}
|
||||
|
||||
p.publishMessage(chat.ID, toolMessage)
|
||||
for _, msg := range insertedMessages {
|
||||
p.publishMessage(chat.ID, msg)
|
||||
}
|
||||
|
||||
// Clear the stream buffer now that the step is
|
||||
@@ -2160,7 +2291,6 @@ func (p *Server) runChat(
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Apply the default MaxOutputTokens if the model config
|
||||
// does not specify one.
|
||||
if callConfig.MaxOutputTokens == nil {
|
||||
@@ -2308,6 +2438,12 @@ func (p *Server) runChat(
|
||||
},
|
||||
|
||||
OnRetry: func(attempt int, retryErr error, delay time.Duration) {
|
||||
p.streamMu.Lock()
|
||||
if state, ok := p.chatStreams[chat.ID]; ok {
|
||||
state.buffer = nil
|
||||
}
|
||||
p.streamMu.Unlock()
|
||||
|
||||
logger.Warn(ctx, "retrying LLM stream",
|
||||
slog.F("attempt", attempt),
|
||||
slog.F("delay", delay.String()),
|
||||
@@ -2351,28 +2487,6 @@ func (p *Server) persistChatContextSummary(
|
||||
return xerrors.Errorf("encode system summary: %w", err)
|
||||
}
|
||||
|
||||
_, err = p.db.InsertChatMessage(ctx, database.InsertChatMessageParams{
|
||||
ChatID: chatID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleSystem),
|
||||
Content: pqtype.NullRawMessage{
|
||||
RawMessage: systemContent,
|
||||
Valid: len(systemContent) > 0,
|
||||
},
|
||||
Visibility: database.ChatMessageVisibilityModel,
|
||||
Compressed: sql.NullBool{Bool: true, Valid: true},
|
||||
InputTokens: sql.NullInt64{},
|
||||
OutputTokens: sql.NullInt64{},
|
||||
TotalTokens: sql.NullInt64{},
|
||||
ReasoningTokens: sql.NullInt64{},
|
||||
CacheCreationTokens: sql.NullInt64{},
|
||||
CacheReadTokens: sql.NullInt64{},
|
||||
ContextLimit: sql.NullInt64{},
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("insert hidden summary message: %w", err)
|
||||
}
|
||||
|
||||
args, err := json.Marshal(map[string]any{
|
||||
"source": "automatic",
|
||||
"threshold_percent": result.ThresholdPercent,
|
||||
@@ -2392,29 +2506,7 @@ func (p *Server) persistChatContextSummary(
|
||||
return xerrors.Errorf("encode summary tool call: %w", err)
|
||||
}
|
||||
|
||||
assistantMessage, err := p.db.InsertChatMessage(ctx, database.InsertChatMessageParams{
|
||||
ChatID: chatID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleAssistant),
|
||||
Content: assistantContent,
|
||||
Visibility: database.ChatMessageVisibilityUser,
|
||||
Compressed: sql.NullBool{
|
||||
Bool: true,
|
||||
Valid: true,
|
||||
},
|
||||
InputTokens: sql.NullInt64{},
|
||||
OutputTokens: sql.NullInt64{},
|
||||
TotalTokens: sql.NullInt64{},
|
||||
ReasoningTokens: sql.NullInt64{},
|
||||
CacheCreationTokens: sql.NullInt64{},
|
||||
CacheReadTokens: sql.NullInt64{},
|
||||
ContextLimit: sql.NullInt64{},
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("insert summary tool call message: %w", err)
|
||||
}
|
||||
|
||||
summaryResult, marshalErr := json.Marshal(map[string]any{
|
||||
summaryResult, err := json.Marshal(map[string]any{
|
||||
"summary": result.SummaryReport,
|
||||
"source": "automatic",
|
||||
"threshold_percent": result.ThresholdPercent,
|
||||
@@ -2422,8 +2514,8 @@ func (p *Server) persistChatContextSummary(
|
||||
"context_tokens": result.ContextTokens,
|
||||
"context_limit_tokens": result.ContextLimit,
|
||||
})
|
||||
if marshalErr != nil {
|
||||
return xerrors.Errorf("encode summary result payload: %w", marshalErr)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("encode summary result payload: %w", err)
|
||||
}
|
||||
toolResult, err := chatprompt.MarshalToolResult(
|
||||
toolCallID,
|
||||
@@ -2435,30 +2527,88 @@ func (p *Server) persistChatContextSummary(
|
||||
return xerrors.Errorf("encode summary tool result: %w", err)
|
||||
}
|
||||
|
||||
toolMessage, err := p.db.InsertChatMessage(ctx, database.InsertChatMessageParams{
|
||||
ChatID: chatID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleTool),
|
||||
Content: toolResult,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
Compressed: sql.NullBool{
|
||||
Bool: true,
|
||||
Valid: true,
|
||||
},
|
||||
InputTokens: sql.NullInt64{},
|
||||
OutputTokens: sql.NullInt64{},
|
||||
TotalTokens: sql.NullInt64{},
|
||||
ReasoningTokens: sql.NullInt64{},
|
||||
CacheCreationTokens: sql.NullInt64{},
|
||||
CacheReadTokens: sql.NullInt64{},
|
||||
ContextLimit: sql.NullInt64{},
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("insert summary tool result message: %w", err)
|
||||
var insertedMessages []database.ChatMessage
|
||||
|
||||
txErr := p.db.InTx(func(tx database.Store) error {
|
||||
_, txErr := tx.InsertChatMessage(ctx, database.InsertChatMessageParams{
|
||||
ChatID: chatID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleSystem),
|
||||
Content: pqtype.NullRawMessage{
|
||||
RawMessage: systemContent,
|
||||
Valid: len(systemContent) > 0,
|
||||
},
|
||||
Visibility: database.ChatMessageVisibilityModel,
|
||||
Compressed: sql.NullBool{Bool: true, Valid: true},
|
||||
InputTokens: sql.NullInt64{},
|
||||
OutputTokens: sql.NullInt64{},
|
||||
TotalTokens: sql.NullInt64{},
|
||||
ReasoningTokens: sql.NullInt64{},
|
||||
CacheCreationTokens: sql.NullInt64{},
|
||||
CacheReadTokens: sql.NullInt64{},
|
||||
ContextLimit: sql.NullInt64{},
|
||||
})
|
||||
if txErr != nil {
|
||||
return xerrors.Errorf("insert hidden summary message: %w", txErr)
|
||||
}
|
||||
|
||||
assistantMessage, txErr := tx.InsertChatMessage(ctx, database.InsertChatMessageParams{
|
||||
ChatID: chatID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleAssistant),
|
||||
Content: assistantContent,
|
||||
Visibility: database.ChatMessageVisibilityUser,
|
||||
Compressed: sql.NullBool{
|
||||
Bool: true,
|
||||
Valid: true,
|
||||
},
|
||||
InputTokens: sql.NullInt64{},
|
||||
OutputTokens: sql.NullInt64{},
|
||||
TotalTokens: sql.NullInt64{},
|
||||
ReasoningTokens: sql.NullInt64{},
|
||||
CacheCreationTokens: sql.NullInt64{},
|
||||
CacheReadTokens: sql.NullInt64{},
|
||||
ContextLimit: sql.NullInt64{},
|
||||
})
|
||||
if txErr != nil {
|
||||
return xerrors.Errorf("insert summary tool call message: %w", txErr)
|
||||
}
|
||||
insertedMessages = append(insertedMessages, assistantMessage)
|
||||
|
||||
toolMessage, txErr := tx.InsertChatMessage(ctx, database.InsertChatMessageParams{
|
||||
ChatID: chatID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleTool),
|
||||
Content: toolResult,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
Compressed: sql.NullBool{
|
||||
Bool: true,
|
||||
Valid: true,
|
||||
},
|
||||
InputTokens: sql.NullInt64{},
|
||||
OutputTokens: sql.NullInt64{},
|
||||
TotalTokens: sql.NullInt64{},
|
||||
ReasoningTokens: sql.NullInt64{},
|
||||
CacheCreationTokens: sql.NullInt64{},
|
||||
CacheReadTokens: sql.NullInt64{},
|
||||
ContextLimit: sql.NullInt64{},
|
||||
})
|
||||
if txErr != nil {
|
||||
return xerrors.Errorf("insert summary tool result message: %w", txErr)
|
||||
}
|
||||
insertedMessages = append(insertedMessages, toolMessage)
|
||||
|
||||
return nil
|
||||
}, nil)
|
||||
if txErr != nil {
|
||||
return txErr
|
||||
}
|
||||
|
||||
p.publishMessage(chatID, assistantMessage)
|
||||
p.publishMessage(chatID, toolMessage)
|
||||
// Publish after transaction commits to avoid notifying
|
||||
// subscribers about messages that could be rolled back.
|
||||
for _, msg := range insertedMessages {
|
||||
p.publishMessage(chatID, msg)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -2666,26 +2816,63 @@ func (p *Server) recoverStaleChats(ctx context.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
recovered := 0
|
||||
for _, chat := range staleChats {
|
||||
p.logger.Info(ctx, "recovering stale chat", slog.F("chat_id", chat.ID))
|
||||
|
||||
// Reset to pending so any replica can pick it up.
|
||||
_, err := p.db.UpdateChatStatus(ctx, database.UpdateChatStatusParams{
|
||||
ID: chat.ID,
|
||||
Status: database.ChatStatusPending,
|
||||
WorkerID: uuid.NullUUID{},
|
||||
StartedAt: sql.NullTime{},
|
||||
HeartbeatAt: sql.NullTime{},
|
||||
LastError: sql.NullString{},
|
||||
})
|
||||
// Use a transaction with FOR UPDATE to avoid a TOCTOU race:
|
||||
// between GetStaleChats (a bare SELECT) and here, the chat's
|
||||
// heartbeat may have been refreshed. We re-check freshness
|
||||
// under the row lock before resetting.
|
||||
err := p.db.InTx(func(tx database.Store) error {
|
||||
locked, lockErr := tx.GetChatByIDForUpdate(ctx, chat.ID)
|
||||
if lockErr != nil {
|
||||
return xerrors.Errorf("lock chat for recovery: %w", lockErr)
|
||||
}
|
||||
|
||||
// Only recover chats that are still running.
|
||||
// Between GetStaleChats and this lock, the chat
|
||||
// may have completed normally.
|
||||
if locked.Status != database.ChatStatusRunning {
|
||||
p.logger.Debug(ctx, "chat status changed since snapshot, skipping recovery",
|
||||
slog.F("chat_id", chat.ID),
|
||||
slog.F("status", locked.Status))
|
||||
return nil
|
||||
}
|
||||
|
||||
// Re-check: only recover if the chat is still stale.
|
||||
// A valid heartbeat that is at or after the stale
|
||||
// threshold means the chat was refreshed after our
|
||||
// initial snapshot — skip it.
|
||||
if locked.HeartbeatAt.Valid && !locked.HeartbeatAt.Time.Before(staleAfter) {
|
||||
p.logger.Debug(ctx, "chat heartbeat refreshed since snapshot, skipping recovery",
|
||||
slog.F("chat_id", chat.ID))
|
||||
return nil
|
||||
}
|
||||
|
||||
// Reset to pending so any replica can pick it up.
|
||||
_, updateErr := tx.UpdateChatStatus(ctx, database.UpdateChatStatusParams{
|
||||
ID: chat.ID,
|
||||
Status: database.ChatStatusPending,
|
||||
WorkerID: uuid.NullUUID{},
|
||||
StartedAt: sql.NullTime{},
|
||||
HeartbeatAt: sql.NullTime{},
|
||||
LastError: sql.NullString{},
|
||||
})
|
||||
if updateErr != nil {
|
||||
return updateErr
|
||||
}
|
||||
recovered++
|
||||
return nil
|
||||
}, nil)
|
||||
if err != nil {
|
||||
p.logger.Error(ctx, "failed to recover stale chat",
|
||||
slog.F("chat_id", chat.ID), slog.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
if len(staleChats) > 0 {
|
||||
p.logger.Info(ctx, "recovered stale chats", slog.F("count", len(staleChats)))
|
||||
if recovered > 0 {
|
||||
p.logger.Info(ctx, "recovered stale chats", slog.F("count", recovered))
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+20
-347
@@ -6,7 +6,6 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
@@ -17,8 +16,6 @@ import (
|
||||
"github.com/google/uuid"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog/v3/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/agent/agenttest"
|
||||
@@ -32,8 +29,6 @@ import (
|
||||
dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/util/slice"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk/agentconnmock"
|
||||
"github.com/coder/coder/v2/provisioner/echo"
|
||||
proto "github.com/coder/coder/v2/provisionersdk/proto"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
@@ -78,10 +73,11 @@ func TestInterruptChatBroadcastsStatusAcrossInstances(t *testing.T) {
|
||||
require.Eventually(t, func() bool {
|
||||
select {
|
||||
case event := <-events:
|
||||
if event.Type != codersdk.ChatStreamEventTypeStatus || event.Status == nil {
|
||||
return false
|
||||
if event.Type == codersdk.ChatStreamEventTypeStatus && event.Status != nil {
|
||||
return event.Status.Status == codersdk.ChatStatusWaiting
|
||||
}
|
||||
return event.Status.Status == codersdk.ChatStatusWaiting
|
||||
t.Logf("skipping unexpected event: type=%s", event.Type)
|
||||
return false
|
||||
default:
|
||||
return false
|
||||
}
|
||||
@@ -870,15 +866,15 @@ func TestSubscribeNoPubsubNoDuplicateMessageParts(t *testing.T) {
|
||||
// events — the snapshot already contained everything. Before
|
||||
// the fix, localSnapshot was replayed into the channel,
|
||||
// causing duplicates.
|
||||
select {
|
||||
case event, ok := <-events:
|
||||
if ok {
|
||||
t.Fatalf("unexpected event from channel (would be a duplicate): type=%s", event.Type)
|
||||
require.Never(t, func() bool {
|
||||
select {
|
||||
case <-events:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
// Channel closed without events is fine.
|
||||
case <-time.After(200 * time.Millisecond):
|
||||
// No events — correct behavior.
|
||||
}
|
||||
}, 200*time.Millisecond, testutil.IntervalFast,
|
||||
"expected no duplicate events after snapshot")
|
||||
}
|
||||
|
||||
func TestSubscribeAfterMessageID(t *testing.T) {
|
||||
@@ -1533,13 +1529,16 @@ func TestSuccessfulChatSendsWebPushWithNavigationData(t *testing.T) {
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Wait for a web push notification to be dispatched. The dispatch
|
||||
// happens asynchronously after the DB status is updated, so we need
|
||||
// to poll rather than assert immediately.
|
||||
testutil.Eventually(ctx, t, func(_ context.Context) bool {
|
||||
return mockPush.dispatchCount.Load() >= 1
|
||||
// Wait for the chat to complete and return to waiting status.
|
||||
testutil.Eventually(ctx, t, func(ctx context.Context) bool {
|
||||
fromDB, dbErr := db.GetChatByID(ctx, chat.ID)
|
||||
if dbErr != nil {
|
||||
return false
|
||||
}
|
||||
return fromDB.Status == database.ChatStatusWaiting && !fromDB.WorkerID.Valid && mockPush.dispatchCount.Load() == 1
|
||||
}, testutil.IntervalFast)
|
||||
|
||||
// Verify a web push notification was dispatched exactly once.
|
||||
require.Equal(t, int32(1), mockPush.dispatchCount.Load(),
|
||||
"expected exactly one web push dispatch for a completed chat")
|
||||
|
||||
@@ -1558,75 +1557,6 @@ func TestSuccessfulChatSendsWebPushWithNavigationData(t *testing.T) {
|
||||
"web push Data should contain the chat navigation URL")
|
||||
}
|
||||
|
||||
func TestSuccessfulChatSendsWebPushWithTag(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
// Set up a mock OpenAI that returns a simple streaming response.
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("done")...)
|
||||
})
|
||||
|
||||
// Mock webpush dispatcher that captures calls.
|
||||
mockPush := &mockWebpushDispatcher{}
|
||||
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
server := chatd.New(chatd.Config{
|
||||
Logger: logger,
|
||||
Database: db,
|
||||
ReplicaID: uuid.New(),
|
||||
Pubsub: ps,
|
||||
PendingChatAcquireInterval: 10 * time.Millisecond,
|
||||
InFlightChatStaleAfter: testutil.WaitSuperLong,
|
||||
WebpushDispatcher: mockPush,
|
||||
})
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, server.Close())
|
||||
})
|
||||
|
||||
user, model := seedChatDependencies(ctx, t, db)
|
||||
setOpenAIProviderBaseURL(ctx, t, db, openAIURL)
|
||||
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OwnerID: user.ID,
|
||||
Title: "push-tag-test",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Wait for the web push notification to be dispatched.
|
||||
// We poll dispatchCount rather than DB status because the
|
||||
// push fires after the status update, creating a small race
|
||||
// window.
|
||||
testutil.Eventually(ctx, t, func(_ context.Context) bool {
|
||||
return mockPush.dispatchCount.Load() >= 1
|
||||
}, testutil.IntervalFast)
|
||||
|
||||
require.Equal(t, int32(1), mockPush.dispatchCount.Load(),
|
||||
"expected exactly one web push dispatch for a completed chat")
|
||||
|
||||
// Verify the push notification tag is set to the chat ID for dedup.
|
||||
mockPush.mu.Lock()
|
||||
capturedMsg := mockPush.lastMessage
|
||||
capturedUser := mockPush.lastUserID
|
||||
mockPush.mu.Unlock()
|
||||
|
||||
require.Equal(t, chat.ID.String(), capturedMsg.Tag,
|
||||
"push notification tag should equal the chat ID for deduplication")
|
||||
require.Equal(t, user.ID, capturedUser,
|
||||
"push notification should be dispatched to the chat owner")
|
||||
require.Equal(t, "push-tag-test", capturedMsg.Title,
|
||||
"push notification title should match the chat title")
|
||||
require.Equal(t, "Agent has finished running.", capturedMsg.Body,
|
||||
"push notification body should indicate the agent finished")
|
||||
}
|
||||
|
||||
func TestCloseDuringShutdownContextCanceledShouldRetryOnNewReplica(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -1733,260 +1663,3 @@ func TestCloseDuringShutdownContextCanceledShouldRetryOnNewReplica(t *testing.T)
|
||||
!fromDB.LastError.Valid
|
||||
}, testutil.WaitMedium, testutil.IntervalFast)
|
||||
}
|
||||
|
||||
func TestHeaderInjection(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// seedWorkspaceAgent creates the DB entities needed so that
|
||||
// GetWorkspaceAgentsInLatestBuildByWorkspaceID returns an
|
||||
// agent for the given workspace.
|
||||
seedWorkspaceAgent := func(
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
ps dbpubsub.Pubsub,
|
||||
ownerID uuid.UUID,
|
||||
orgID uuid.UUID,
|
||||
) (workspaceID uuid.UUID, agentID uuid.UUID) {
|
||||
t.Helper()
|
||||
|
||||
// TemplateVersion needs its own provisioner job.
|
||||
versionJob := dbgen.ProvisionerJob(t, db, ps, database.ProvisionerJob{
|
||||
OrganizationID: orgID,
|
||||
InitiatorID: ownerID,
|
||||
Type: database.ProvisionerJobTypeTemplateVersionImport,
|
||||
})
|
||||
tv := dbgen.TemplateVersion(t, db, database.TemplateVersion{
|
||||
OrganizationID: orgID,
|
||||
CreatedBy: ownerID,
|
||||
JobID: versionJob.ID,
|
||||
})
|
||||
templ := dbgen.Template(t, db, database.Template{
|
||||
OrganizationID: orgID,
|
||||
CreatedBy: ownerID,
|
||||
ActiveVersionID: tv.ID,
|
||||
})
|
||||
ws := dbgen.Workspace(t, db, database.WorkspaceTable{
|
||||
OwnerID: ownerID,
|
||||
OrganizationID: orgID,
|
||||
TemplateID: templ.ID,
|
||||
})
|
||||
buildJob := dbgen.ProvisionerJob(t, db, ps, database.ProvisionerJob{
|
||||
OrganizationID: orgID,
|
||||
InitiatorID: ownerID,
|
||||
Type: database.ProvisionerJobTypeWorkspaceBuild,
|
||||
})
|
||||
build := dbgen.WorkspaceBuild(t, db, database.WorkspaceBuild{
|
||||
WorkspaceID: ws.ID,
|
||||
JobID: buildJob.ID,
|
||||
BuildNumber: 1,
|
||||
InitiatorID: ownerID,
|
||||
TemplateVersionID: tv.ID,
|
||||
})
|
||||
resource := dbgen.WorkspaceResource(t, db, database.WorkspaceResource{
|
||||
JobID: build.JobID,
|
||||
})
|
||||
agent := dbgen.WorkspaceAgent(t, db, database.WorkspaceAgent{
|
||||
ResourceID: resource.ID,
|
||||
})
|
||||
return ws.ID, agent.ID
|
||||
}
|
||||
|
||||
t.Run("WithParentChat", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, model := seedChatDependencies(ctx, t, db)
|
||||
|
||||
org, err := db.GetDefaultOrganization(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
workspaceID, expectedAgentID := seedWorkspaceAgent(t, db, ps, user.ID, org.ID)
|
||||
|
||||
// Set up the mock OpenAI to return a simple text response
|
||||
// so the chat finishes cleanly.
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
return chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("done")...,
|
||||
)
|
||||
})
|
||||
setOpenAIProviderBaseURL(ctx, t, db, openAIURL)
|
||||
|
||||
// Wire up the mock agent connection so we can capture
|
||||
// the headers passed to SetExtraHeaders.
|
||||
ctrl := gomock.NewController(t)
|
||||
mockConn := agentconnmock.NewMockAgentConn(ctrl)
|
||||
|
||||
var capturedHeaders http.Header
|
||||
headersCaptured := make(chan struct{})
|
||||
|
||||
// SetExtraHeaders is called once when the connection
|
||||
// is first established.
|
||||
mockConn.EXPECT().SetExtraHeaders(gomock.Any()).Do(func(h http.Header) {
|
||||
capturedHeaders = h
|
||||
close(headersCaptured)
|
||||
})
|
||||
// resolveInstructions calls LS to look for instruction
|
||||
// files; return an error so it skips gracefully.
|
||||
mockConn.EXPECT().LS(gomock.Any(), gomock.Any(), gomock.Any()).Return(
|
||||
workspacesdk.LSResponse{}, xerrors.New("not found"),
|
||||
).AnyTimes()
|
||||
// The connection is closed when the chat finishes.
|
||||
mockConn.EXPECT().Close().Return(nil).AnyTimes()
|
||||
|
||||
agentConnFn := func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
|
||||
require.Equal(t, expectedAgentID, agentID)
|
||||
return mockConn, func() {}, nil
|
||||
}
|
||||
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
server := chatd.New(chatd.Config{
|
||||
Logger: logger,
|
||||
Database: db,
|
||||
ReplicaID: uuid.New(),
|
||||
Pubsub: ps,
|
||||
AgentConn: agentConnFn,
|
||||
PendingChatAcquireInterval: 10 * time.Millisecond,
|
||||
InFlightChatStaleAfter: testutil.WaitSuperLong,
|
||||
})
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, server.Close())
|
||||
})
|
||||
|
||||
// Create a real parent chat so the FK constraint is
|
||||
// satisfied.
|
||||
parentChat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OwnerID: user.ID,
|
||||
Title: "parent-chat",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []fantasy.Content{
|
||||
fantasy.TextContent{Text: "parent"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OwnerID: user.ID,
|
||||
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
|
||||
ParentChatID: uuid.NullUUID{UUID: parentChat.ID, Valid: true},
|
||||
Title: "header-injection-parent",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []fantasy.Content{
|
||||
fantasy.TextContent{Text: "hello"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Wait for the chat to be processed and headers to be
|
||||
// captured.
|
||||
select {
|
||||
case <-headersCaptured:
|
||||
case <-ctx.Done():
|
||||
require.FailNow(t, "timed out waiting for SetExtraHeaders")
|
||||
}
|
||||
|
||||
require.Equal(t,
|
||||
chat.ID.String(),
|
||||
capturedHeaders.Get(workspacesdk.CoderChatIDHeader),
|
||||
)
|
||||
|
||||
ancestorJSON := capturedHeaders.Get(workspacesdk.CoderAncestorChatIDsHeader)
|
||||
var ancestorIDs []string
|
||||
err = json.Unmarshal([]byte(ancestorJSON), &ancestorIDs)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []string{parentChat.ID.String()}, ancestorIDs)
|
||||
})
|
||||
|
||||
t.Run("WithoutParentChat", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, model := seedChatDependencies(ctx, t, db)
|
||||
|
||||
org, err := db.GetDefaultOrganization(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
workspaceID, expectedAgentID := seedWorkspaceAgent(t, db, ps, user.ID, org.ID)
|
||||
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
return chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("done")...,
|
||||
)
|
||||
})
|
||||
setOpenAIProviderBaseURL(ctx, t, db, openAIURL)
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
mockConn := agentconnmock.NewMockAgentConn(ctrl)
|
||||
|
||||
var capturedHeaders http.Header
|
||||
headersCaptured := make(chan struct{})
|
||||
|
||||
mockConn.EXPECT().SetExtraHeaders(gomock.Any()).Do(func(h http.Header) {
|
||||
capturedHeaders = h
|
||||
close(headersCaptured)
|
||||
})
|
||||
mockConn.EXPECT().LS(gomock.Any(), gomock.Any(), gomock.Any()).Return(
|
||||
workspacesdk.LSResponse{}, xerrors.New("not found"),
|
||||
).AnyTimes()
|
||||
mockConn.EXPECT().Close().Return(nil).AnyTimes()
|
||||
|
||||
agentConnFn := func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
|
||||
require.Equal(t, expectedAgentID, agentID)
|
||||
return mockConn, func() {}, nil
|
||||
}
|
||||
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
server := chatd.New(chatd.Config{
|
||||
Logger: logger,
|
||||
Database: db,
|
||||
ReplicaID: uuid.New(),
|
||||
Pubsub: ps,
|
||||
AgentConn: agentConnFn,
|
||||
PendingChatAcquireInterval: 10 * time.Millisecond,
|
||||
InFlightChatStaleAfter: testutil.WaitSuperLong,
|
||||
})
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, server.Close())
|
||||
})
|
||||
|
||||
// Create a chat without a parent — the ancestor header
|
||||
// should contain an empty JSON array.
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OwnerID: user.ID,
|
||||
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
|
||||
Title: "header-injection-no-parent",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []fantasy.Content{
|
||||
fantasy.TextContent{Text: "hello"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
select {
|
||||
case <-headersCaptured:
|
||||
case <-ctx.Done():
|
||||
require.FailNow(t, "timed out waiting for SetExtraHeaders")
|
||||
}
|
||||
|
||||
require.Equal(t,
|
||||
chat.ID.String(),
|
||||
capturedHeaders.Get(workspacesdk.CoderChatIDHeader),
|
||||
)
|
||||
|
||||
// When there is no parent, the code declares
|
||||
// var ancestorIDs []string and never appends to it,
|
||||
// so json.Marshal produces "null".
|
||||
ancestorJSON := capturedHeaders.Get(workspacesdk.CoderAncestorChatIDsHeader)
|
||||
var ancestorIDs []string
|
||||
err = json.Unmarshal([]byte(ancestorJSON), &ancestorIDs)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, ancestorIDs)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -73,7 +73,9 @@ type RunOptions struct {
|
||||
// OnRetry is called before each retry attempt when the LLM
|
||||
// stream fails with a retryable error. It provides the attempt
|
||||
// number, error, and backoff delay so callers can publish status
|
||||
// events to connected clients.
|
||||
// events to connected clients. Callers should also clear any
|
||||
// buffered stream state from the failed attempt in this callback
|
||||
// to avoid sending duplicated content.
|
||||
OnRetry chatretry.OnRetryFn
|
||||
|
||||
OnInterruptedPersistError func(error)
|
||||
@@ -209,6 +211,10 @@ func Run(ctx context.Context, opts RunOptions) error {
|
||||
var lastUsage fantasy.Usage
|
||||
var lastProviderMetadata fantasy.ProviderMetadata
|
||||
|
||||
totalSteps := 0
|
||||
// When totalSteps reaches MaxSteps the inner loop exits immediately
|
||||
// (its condition is false), stoppedByModel stays false, and the
|
||||
// post-loop guard breaks the outer compaction loop.
|
||||
for compactionAttempt := 0; ; compactionAttempt++ {
|
||||
alreadyCompacted := false
|
||||
// stoppedByModel is true when the inner step loop
|
||||
@@ -222,7 +228,8 @@ func Run(ctx context.Context, opts RunOptions) error {
|
||||
// agent never had a chance to use the compacted context.
|
||||
compactedOnFinalStep := false
|
||||
|
||||
for step := 0; step < opts.MaxSteps; step++ {
|
||||
for step := 0; totalSteps < opts.MaxSteps; step++ {
|
||||
totalSteps++
|
||||
// Copy messages so that provider-specific caching
|
||||
// mutations don't leak back to the caller's slice.
|
||||
// copy copies Message structs by value, so field
|
||||
@@ -321,6 +328,12 @@ func Run(ctx context.Context, opts RunOptions) error {
|
||||
lastUsage = result.usage
|
||||
lastProviderMetadata = result.providerMetadata
|
||||
|
||||
// Append the step's response messages so that both
|
||||
// inline and post-loop compaction see the full
|
||||
// conversation including the latest assistant reply.
|
||||
stepMessages := result.toResponseMessages()
|
||||
messages = append(messages, stepMessages...)
|
||||
|
||||
// Inline compaction.
|
||||
if opts.Compaction != nil && opts.ReloadMessages != nil {
|
||||
did, compactErr := tryCompact(
|
||||
@@ -354,17 +367,11 @@ func Run(ctx context.Context, opts RunOptions) error {
|
||||
// The agent is continuing with tool calls, so any
|
||||
// prior compaction has already been consumed.
|
||||
compactedOnFinalStep = false
|
||||
|
||||
// Build messages from the step for the next iteration.
|
||||
// toResponseMessages produces assistant-role content
|
||||
// (text, reasoning, tool calls) and tool-result content.
|
||||
stepMessages := result.toResponseMessages()
|
||||
messages = append(messages, stepMessages...)
|
||||
}
|
||||
|
||||
// Post-run compaction safety net: if we never compacted
|
||||
// during the loop, try once at the end.
|
||||
if !alreadyCompacted && opts.Compaction != nil {
|
||||
if !alreadyCompacted && opts.Compaction != nil && opts.ReloadMessages != nil {
|
||||
did, err := tryCompact(
|
||||
ctx,
|
||||
opts.Model,
|
||||
@@ -383,7 +390,6 @@ func Run(ctx context.Context, opts RunOptions) error {
|
||||
compactedOnFinalStep = true
|
||||
}
|
||||
}
|
||||
|
||||
// Re-enter the step loop when compaction fired on the
|
||||
// model's final step. This lets the agent continue
|
||||
// working with fresh summarized context instead of
|
||||
@@ -514,7 +520,6 @@ func processStepStream(
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
case fantasy.StreamPartTypeToolInputStart:
|
||||
activeToolCalls[part.ID] = &fantasy.ToolCallContent{
|
||||
ToolCallID: part.ID,
|
||||
|
||||
@@ -123,7 +123,8 @@ func tryCompact(
|
||||
config.SystemSummaryPrefix + "\n\n" + summary,
|
||||
)
|
||||
|
||||
err = config.Persist(ctx, CompactionResult{
|
||||
persistCtx := context.WithoutCancel(ctx)
|
||||
err = config.Persist(persistCtx, CompactionResult{
|
||||
SystemSummary: systemSummary,
|
||||
SummaryReport: summary,
|
||||
ThresholdPercent: config.ThresholdPercent,
|
||||
|
||||
@@ -76,9 +76,20 @@ func TestRun_Compaction(t *testing.T) {
|
||||
return nil
|
||||
},
|
||||
},
|
||||
ReloadMessages: func(_ context.Context) ([]fantasy.Message, error) {
|
||||
return []fantasy.Message{
|
||||
textMessage(fantasy.MessageRoleUser, "hello"),
|
||||
}, nil
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 1, persistCompactionCalls)
|
||||
// Compaction fires twice: once inline when the threshold is
|
||||
// reached on step 0 (the only step, since MaxSteps=1), and
|
||||
// once from the post-run safety net during the re-entry
|
||||
// iteration (where totalSteps already equals MaxSteps so the
|
||||
// inner loop doesn't execute, but lastUsage still exceeds
|
||||
// the threshold).
|
||||
require.Equal(t, 2, persistCompactionCalls)
|
||||
require.Contains(t, persistedCompaction.SystemSummary, summaryText)
|
||||
require.Equal(t, summaryText, persistedCompaction.SummaryReport)
|
||||
require.Equal(t, int64(80), persistedCompaction.ContextTokens)
|
||||
@@ -151,13 +162,25 @@ func TestRun_Compaction(t *testing.T) {
|
||||
return nil
|
||||
},
|
||||
},
|
||||
ReloadMessages: func(_ context.Context) ([]fantasy.Message, error) {
|
||||
return []fantasy.Message{
|
||||
textMessage(fantasy.MessageRoleUser, "hello"),
|
||||
}, nil
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
// Compaction fires twice (see PersistsWhenThresholdReached
|
||||
// for the full explanation). Each cycle follows the order:
|
||||
// publish_tool_call → generate → persist → publish_tool_result.
|
||||
require.Equal(t, []string{
|
||||
"publish_tool_call",
|
||||
"generate",
|
||||
"persist",
|
||||
"publish_tool_result",
|
||||
"publish_tool_call",
|
||||
"generate",
|
||||
"persist",
|
||||
"publish_tool_result",
|
||||
}, callOrder)
|
||||
})
|
||||
|
||||
@@ -457,6 +480,11 @@ func TestRun_Compaction(t *testing.T) {
|
||||
compactionErr = err
|
||||
},
|
||||
},
|
||||
ReloadMessages: func(_ context.Context) ([]fantasy.Message, error) {
|
||||
return []fantasy.Message{
|
||||
textMessage(fantasy.MessageRoleUser, "hello"),
|
||||
}, nil
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Error(t, compactionErr)
|
||||
|
||||
@@ -8,6 +8,8 @@ import (
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"golang.org/x/xerrors"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -18,6 +20,12 @@ const (
|
||||
// MaxDelay is the upper bound for the exponential backoff
|
||||
// duration. Matches the cap used in coder/mux.
|
||||
MaxDelay = 60 * time.Second
|
||||
|
||||
// MaxAttempts is the upper bound on retry attempts before
|
||||
// giving up. With a 60s max backoff this allows roughly
|
||||
// 25 minutes of retries, which is reasonable for transient
|
||||
// LLM provider issues.
|
||||
MaxAttempts = 25
|
||||
)
|
||||
|
||||
// nonRetryablePatterns are substrings that indicate a permanent error
|
||||
@@ -131,9 +139,8 @@ type RetryFn func(ctx context.Context) error
|
||||
type OnRetryFn func(attempt int, err error, delay time.Duration)
|
||||
|
||||
// Retry calls fn repeatedly until it succeeds, returns a
|
||||
// non-retryable error, or ctx is canceled. There is no max attempt
|
||||
// limit — retries continue indefinitely with exponential backoff
|
||||
// (capped at 60s), matching the behavior of coder/mux.
|
||||
// non-retryable error, ctx is canceled, or MaxAttempts is reached.
|
||||
// Retries use exponential backoff capped at MaxDelay.
|
||||
//
|
||||
// The onRetry callback (if non-nil) is called before each retry
|
||||
// attempt, giving the caller a chance to reset state, log, or
|
||||
@@ -156,10 +163,15 @@ func Retry(ctx context.Context, fn RetryFn, onRetry OnRetryFn) error {
|
||||
return ctx.Err()
|
||||
}
|
||||
|
||||
delay := Delay(attempt)
|
||||
attempt++
|
||||
if attempt >= MaxAttempts {
|
||||
return xerrors.Errorf("max retry attempts (%d) exceeded: %w", MaxAttempts, err)
|
||||
}
|
||||
|
||||
delay := Delay(attempt - 1)
|
||||
|
||||
if onRetry != nil {
|
||||
onRetry(attempt+1, err, delay)
|
||||
onRetry(attempt, err, delay)
|
||||
}
|
||||
|
||||
timer := time.NewTimer(delay)
|
||||
@@ -169,7 +181,5 @@ func Retry(ctx context.Context, fn RetryFn, onRetry OnRetryFn) error {
|
||||
return ctx.Err()
|
||||
case <-timer.C:
|
||||
}
|
||||
|
||||
attempt++
|
||||
}
|
||||
}
|
||||
|
||||
@@ -111,7 +111,7 @@ func (c MultiReplicaSubscribeConfig) clock() quartz.Clock {
|
||||
func NewMultiReplicaSubscribeFn(
|
||||
cfg MultiReplicaSubscribeConfig,
|
||||
) osschatd.SubscribeFn {
|
||||
return func(ctx context.Context, params osschatd.SubscribeFnParams) (<-chan codersdk.ChatStreamEvent, func()) {
|
||||
return func(ctx context.Context, params osschatd.SubscribeFnParams) <-chan codersdk.ChatStreamEvent {
|
||||
chatID := params.ChatID
|
||||
requestHeader := params.RequestHeader
|
||||
logger := params.Logger
|
||||
@@ -149,18 +149,13 @@ func NewMultiReplicaSubscribeFn(
|
||||
|
||||
// Merge all event sources.
|
||||
mergedEvents := make(chan codersdk.ChatStreamEvent, 128)
|
||||
var allCancels []func()
|
||||
if relayCancel != nil {
|
||||
allCancels = append(allCancels, relayCancel)
|
||||
}
|
||||
|
||||
// Channel for async relay establishment.
|
||||
type relayResult struct {
|
||||
parts <-chan codersdk.ChatStreamEvent
|
||||
cancel func()
|
||||
workerID uuid.UUID // the worker this dial targeted
|
||||
}
|
||||
relayReadyCh := make(chan relayResult, 1)
|
||||
relayReadyCh := make(chan relayResult, 4)
|
||||
|
||||
// Per-dial context so in-flight dials can be canceled when
|
||||
// a new dial is initiated or the relay is closed.
|
||||
@@ -182,15 +177,18 @@ func NewMultiReplicaSubscribeFn(
|
||||
dialCancel()
|
||||
dialCancel = nil
|
||||
}
|
||||
// Drain any buffered relay result from a canceled
|
||||
// dial.
|
||||
select {
|
||||
case result := <-relayReadyCh:
|
||||
if result.cancel != nil {
|
||||
result.cancel()
|
||||
// Drain all buffered relay results from canceled dials.
|
||||
for {
|
||||
select {
|
||||
case result := <-relayReadyCh:
|
||||
if result.cancel != nil {
|
||||
result.cancel()
|
||||
}
|
||||
default:
|
||||
goto drained
|
||||
}
|
||||
default:
|
||||
}
|
||||
drained:
|
||||
expectedWorkerID = uuid.Nil
|
||||
if relayCancel != nil {
|
||||
relayCancel()
|
||||
@@ -403,19 +401,11 @@ func NewMultiReplicaSubscribeFn(
|
||||
}
|
||||
}()
|
||||
|
||||
// The cancel function tears down the relay state
|
||||
// indirectly: the merge goroutine owns all relay state
|
||||
// (reconnectTimer, relayCancel, dialCancel, etc.) and
|
||||
// cleans it up via its defer closeRelay() when ctx is
|
||||
// canceled.
|
||||
cancel := func() {
|
||||
for _, cancelFn := range allCancels {
|
||||
if cancelFn != nil {
|
||||
cancelFn()
|
||||
}
|
||||
}
|
||||
}
|
||||
return mergedEvents, cancel
|
||||
// Cleanup is driven by ctx cancellation: the merge
|
||||
// goroutine owns all relay state (reconnectTimer,
|
||||
// relayCancel, dialCancel, etc.) and tears it down
|
||||
// via defer closeRelay() when ctx is done.
|
||||
return mergedEvents
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user