diff --git a/coderd/chatd/chatd.go b/coderd/chatd/chatd.go index 011fbc71e6..eadd7e3569 100644 --- a/coderd/chatd/chatd.go +++ b/coderd/chatd/chatd.go @@ -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)) } } diff --git a/coderd/chatd/chatd_test.go b/coderd/chatd/chatd_test.go index 069f2ff7fa..a04b2b04b2 100644 --- a/coderd/chatd/chatd_test.go +++ b/coderd/chatd/chatd_test.go @@ -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) - }) -} diff --git a/coderd/chatd/chatloop/chatloop.go b/coderd/chatd/chatloop/chatloop.go index 5697f6651d..a2eb67a960 100644 --- a/coderd/chatd/chatloop/chatloop.go +++ b/coderd/chatd/chatloop/chatloop.go @@ -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, diff --git a/coderd/chatd/chatloop/compaction.go b/coderd/chatd/chatloop/compaction.go index 2238c92544..0d137eb7ab 100644 --- a/coderd/chatd/chatloop/compaction.go +++ b/coderd/chatd/chatloop/compaction.go @@ -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, diff --git a/coderd/chatd/chatloop/compaction_test.go b/coderd/chatd/chatloop/compaction_test.go index 33985fa1b0..4e3c6df7bd 100644 --- a/coderd/chatd/chatloop/compaction_test.go +++ b/coderd/chatd/chatloop/compaction_test.go @@ -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) diff --git a/coderd/chatd/chatretry/chatretry.go b/coderd/chatd/chatretry/chatretry.go index affd309ed9..6b51916f93 100644 --- a/coderd/chatd/chatretry/chatretry.go +++ b/coderd/chatd/chatretry/chatretry.go @@ -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++ } } diff --git a/enterprise/coderd/chatd/chatd.go b/enterprise/coderd/chatd/chatd.go index 20bee8ac58..525878d2cb 100644 --- a/enterprise/coderd/chatd/chatd.go +++ b/enterprise/coderd/chatd/chatd.go @@ -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 } }