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:
Kyle Carberry
2026-03-06 21:02:25 +00:00
committed by GitHub
parent 2cd871e88f
commit eecb7d0b66
7 changed files with 529 additions and 635 deletions
+428 -241
View File
@@ -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))
}
}