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))
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user