diff --git a/coderd/chatd/chatd.go b/coderd/chatd/chatd.go index 3e33eacd33..4f4f79be88 100644 --- a/coderd/chatd/chatd.go +++ b/coderd/chatd/chatd.go @@ -1968,6 +1968,13 @@ func (p *Server) runChat( streamCall.MaxOutputTokens = &maxOutputTokens } + // Generate the tool call ID up front so that the OnStart + // streaming part and the Persist durable messages share + // the same identifier. Without this the client cannot + // correlate the "Summarizing..." tool call with the + // "Summarized" tool result. + compactionToolCallID := "chat_summarized_" + uuid.NewString() + compactionOptions := &chatloop.CompactionOptions{ ThresholdPercent: modelConfig.CompressionThreshold, ContextLimit: modelConfig.ContextLimit, @@ -1979,6 +1986,7 @@ func (p *Server) runChat( persistCtx, chat.ID, modelConfig.ID, + compactionToolCallID, result, ); err != nil { return xerrors.Errorf("persist context summary: %w", err) @@ -1992,6 +2000,16 @@ func (p *Server) runChat( ) return nil }, + OnStart: func() { + // Publish a streaming tool-call part immediately so + // connected clients see "Summarizing..." while the + // LLM generates the summary. + p.publishMessagePart(chat.ID, string(fantasy.MessageRoleAssistant), codersdk.ChatMessagePart{ + Type: codersdk.ChatMessagePartTypeToolCall, + ToolCallID: compactionToolCallID, + ToolName: "chat_summarized", + }) + }, OnError: func(err error) { logger.Warn(ctx, "failed to compact chat context", slog.Error(err)) }, @@ -2063,6 +2081,7 @@ func (p *Server) persistChatContextSummary( ctx context.Context, chatID uuid.UUID, modelConfigID uuid.UUID, + toolCallID string, result chatloop.CompactionResult, ) error { if strings.TrimSpace(result.SystemSummary) == "" || @@ -2097,7 +2116,6 @@ func (p *Server) persistChatContextSummary( return xerrors.Errorf("insert hidden summary message: %w", err) } - toolCallID := "chat_summarized_" + uuid.NewString() args, err := json.Marshal(map[string]any{ "source": "automatic", "threshold_percent": result.ThresholdPercent, @@ -2182,6 +2200,16 @@ func (p *Server) persistChatContextSummary( return xerrors.Errorf("insert summary tool result message: %w", err) } + // Publish a streaming tool-result part so connected clients + // transition from "Summarizing..." to "Summarized" before the + // durable messages and status change arrive. + p.publishMessagePart(chatID, string(fantasy.MessageRoleTool), codersdk.ChatMessagePart{ + Type: codersdk.ChatMessagePartTypeToolResult, + ToolCallID: toolCallID, + ToolName: "chat_summarized", + Result: summaryResult, + }) + p.publishMessage(chatID, assistantMessage) p.publishMessage(chatID, toolMessage) return nil diff --git a/coderd/chatd/chatloop/compaction.go b/coderd/chatd/chatloop/compaction.go index 28ccbc2abc..610682169f 100644 --- a/coderd/chatd/chatloop/compaction.go +++ b/coderd/chatd/chatloop/compaction.go @@ -30,6 +30,7 @@ type CompactionOptions struct { SystemSummaryPrefix string Timeout time.Duration Persist func(context.Context, CompactionResult) error + OnStart func() OnError func(error) } @@ -134,6 +135,10 @@ func maybeCompact( return nil } + if config.OnStart != nil { + config.OnStart() + } + summary, err := generateCompactionSummary( ctx, runOpts.Model, diff --git a/coderd/chatd/chatloop/compaction_test.go b/coderd/chatd/chatloop/compaction_test.go index 152c543d72..f2f4df18dc 100644 --- a/coderd/chatd/chatloop/compaction_test.go +++ b/coderd/chatd/chatloop/compaction_test.go @@ -83,6 +83,113 @@ func TestRun_Compaction(t *testing.T) { require.InDelta(t, 80.0, persistedCompaction.UsagePercent, 0.0001) }) + t.Run("OnStartFiresBeforePersist", func(t *testing.T) { + t.Parallel() + + const summaryText = "compaction summary for ordering test" + + // Track the order of callbacks to verify OnStart fires + // before the Generate call (summary generation) and + // before Persist. + var callOrder []string + + model := &loopTestModel{ + provider: "fake", + streamFn: func(_ context.Context, _ fantasy.Call) (fantasy.StreamResponse, error) { + return streamFromParts([]fantasy.StreamPart{ + {Type: fantasy.StreamPartTypeTextStart, ID: "text-1"}, + {Type: fantasy.StreamPartTypeTextDelta, ID: "text-1", Delta: "done"}, + {Type: fantasy.StreamPartTypeTextEnd, ID: "text-1"}, + { + Type: fantasy.StreamPartTypeFinish, + FinishReason: fantasy.FinishReasonStop, + Usage: fantasy.Usage{ + InputTokens: 80, + TotalTokens: 85, + }, + }, + }), nil + }, + generateFn: func(_ context.Context, _ fantasy.Call) (*fantasy.Response, error) { + callOrder = append(callOrder, "generate") + return &fantasy.Response{ + Content: []fantasy.Content{ + fantasy.TextContent{Text: summaryText}, + }, + }, nil + }, + } + + _, err := Run(context.Background(), RunOptions{ + Model: model, + Messages: []fantasy.Message{ + textMessage(fantasy.MessageRoleUser, "hello"), + }, + MaxSteps: 1, + PersistStep: func(_ context.Context, _ PersistedStep) error { + return nil + }, + ContextLimitFallback: 100, + Compaction: &CompactionOptions{ + ThresholdPercent: 70, + SummaryPrompt: "summarize now", + OnStart: func() { + callOrder = append(callOrder, "on_start") + }, + Persist: func(_ context.Context, _ CompactionResult) error { + callOrder = append(callOrder, "persist") + return nil + }, + }, + }) + require.NoError(t, err) + require.Equal(t, []string{"on_start", "generate", "persist"}, callOrder) + }) + + t.Run("OnStartNotCalledBelowThreshold", func(t *testing.T) { + t.Parallel() + + onStartCalled := false + + model := &loopTestModel{ + provider: "fake", + streamFn: func(_ context.Context, _ fantasy.Call) (fantasy.StreamResponse, error) { + return streamFromParts([]fantasy.StreamPart{ + { + Type: fantasy.StreamPartTypeFinish, + FinishReason: fantasy.FinishReasonStop, + Usage: fantasy.Usage{ + InputTokens: 10, + }, + }, + }), nil + }, + } + + _, err := Run(context.Background(), RunOptions{ + Model: model, + Messages: []fantasy.Message{ + textMessage(fantasy.MessageRoleUser, "hello"), + }, + MaxSteps: 1, + PersistStep: func(_ context.Context, _ PersistedStep) error { + return nil + }, + ContextLimitFallback: 100, + Compaction: &CompactionOptions{ + ThresholdPercent: 70, + OnStart: func() { + onStartCalled = true + }, + Persist: func(_ context.Context, _ CompactionResult) error { + return nil + }, + }, + }) + require.NoError(t, err) + require.False(t, onStartCalled, "OnStart should not fire when usage is below threshold") + }) + t.Run("ErrorsAreReported", func(t *testing.T) { t.Parallel()