diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index 7393722e6d6..4f76d449e5a 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -2236,6 +2236,13 @@ var ErrManualTitleRegenerationInProgress = xerrors.New( "manual title regeneration already in progress", ) +type manualTitleCandidateResult struct { + title string + modelConfig database.ChatModelConfig + usage fantasy.Usage + hasMessages bool +} + type manualTitleGenerationError struct { cause error modelConfig database.ChatModelConfig @@ -2483,16 +2490,17 @@ func (p *Server) recordManualTitleGenerationFailure( return generationErr } -//nolint:revive // flag-parameter: enableDebug toggles optional debug capture on a shared code path; splitting would duplicate message fetch and model resolution. -func (p *Server) fetchAndGenerateManualTitle( +// generateManualTitleCandidate performs only model generation and returns the +// candidate plus accounting metadata. Endpoint-specific commit paths are +// responsible for recording usage and deciding whether to persist the title. +func (p *Server) generateManualTitleCandidate( ctx context.Context, store database.Store, chat database.Chat, keys chatprovider.ProviderAPIKeys, - enableDebug bool, -) (title string, modelConfig database.ChatModelConfig, usage fantasy.Usage, hasMessages bool, err error) { +) (manualTitleCandidateResult, error) { if limitErr := p.checkUsageLimit(ctx, store, chat.OwnerID, uuid.NullUUID{UUID: chat.OrganizationID, Valid: true}); limitErr != nil { - return "", database.ChatModelConfig{}, fantasy.Usage{}, false, limitErr + return manualTitleCandidateResult{}, limitErr } headMessages, err := store.GetChatMessagesByChatIDAscPaginated( @@ -2504,7 +2512,7 @@ func (p *Server) fetchAndGenerateManualTitle( }, ) if err != nil { - return "", database.ChatModelConfig{}, fantasy.Usage{}, false, xerrors.Errorf("get head chat messages: %w", err) + return manualTitleCandidateResult{}, xerrors.Errorf("get head chat messages: %w", err) } tailMessages, err := store.GetChatMessagesByChatIDDescPaginated( ctx, @@ -2515,50 +2523,54 @@ func (p *Server) fetchAndGenerateManualTitle( }, ) if err != nil { - return "", database.ChatModelConfig{}, fantasy.Usage{}, false, xerrors.Errorf("get tail chat messages: %w", err) + return manualTitleCandidateResult{}, xerrors.Errorf("get tail chat messages: %w", err) } messages := mergeManualTitleMessages(headMessages, tailMessages) if len(messages) == 0 { - return "", database.ChatModelConfig{}, fantasy.Usage{}, false, nil + return manualTitleCandidateResult{}, nil } model, modelConfig, err := p.resolveManualTitleModel(ctx, store, chat, keys) + result := manualTitleCandidateResult{ + modelConfig: modelConfig, + hasMessages: true, + } if err != nil { - return "", database.ChatModelConfig{}, fantasy.Usage{}, true, err + return result, err } titleCtx := ctx titleModel := model finishDebugRun := func(error) {} - if enableDebug { - if debugSvc := p.debugService(); debugSvc != nil && debugSvc.IsEnabled(ctx, chat.ID, chat.OwnerID) { - titleCtx, titleModel, finishDebugRun = p.prepareManualTitleDebugRun( - ctx, - debugSvc, - chat, - modelConfig, - keys, - messages, - model, - ) - } + if debugSvc := p.debugService(); debugSvc != nil && debugSvc.IsEnabled(ctx, chat.ID, chat.OwnerID) { + titleCtx, titleModel, finishDebugRun = p.prepareManualTitleDebugRun( + ctx, + debugSvc, + chat, + modelConfig, + keys, + messages, + model, + ) } - title, usage, err = generateManualTitle(titleCtx, messages, titleModel) + title, usage, err := generateManualTitle(titleCtx, messages, titleModel) finishDebugRun(err) + result.title = title + result.usage = usage if err != nil { wrappedErr := xerrors.Errorf("generate manual title: %w", err) if usage == (fantasy.Usage{}) { - return "", modelConfig, fantasy.Usage{}, true, wrappedErr + return result, wrappedErr } - return "", modelConfig, usage, true, &manualTitleGenerationError{ + return result, &manualTitleGenerationError{ cause: wrappedErr, modelConfig: modelConfig, usage: usage, } } - return title, modelConfig, usage, true, nil + return result, nil } func (p *Server) proposeChatTitleWithStore( @@ -2567,11 +2579,11 @@ func (p *Server) proposeChatTitleWithStore( chat database.Chat, keys chatprovider.ProviderAPIKeys, ) (string, error) { - title, modelConfig, usage, hasMessages, err := p.fetchAndGenerateManualTitle(ctx, store, chat, keys, false) + result, err := p.generateManualTitleCandidate(ctx, store, chat, keys) if err != nil { return "", err } - if !hasMessages { + if !result.hasMessages { return "", nil } @@ -2581,13 +2593,13 @@ func (p *Server) proposeChatTitleWithStore( recordCtx, store, chat, - modelConfig, - usage, + result.modelConfig, + result.usage, "", ); recordErr != nil { return "", xerrors.Errorf("record manual title usage: %w", recordErr) } - return title, nil + return result.title, nil } func (p *Server) regenerateChatTitleWithStore( @@ -2596,11 +2608,11 @@ func (p *Server) regenerateChatTitleWithStore( chat database.Chat, keys chatprovider.ProviderAPIKeys, ) (database.Chat, error) { - title, modelConfig, usage, hasMessages, err := p.fetchAndGenerateManualTitle(ctx, store, chat, keys, true) + result, err := p.generateManualTitleCandidate(ctx, store, chat, keys) if err != nil { return database.Chat{}, err } - if !hasMessages { + if !result.hasMessages { return chat, nil } @@ -2611,12 +2623,12 @@ func (p *Server) regenerateChatTitleWithStore( recordCtx, store, chat, - modelConfig, - usage, - title, + result.modelConfig, + result.usage, + result.title, ) if recordErr != nil { - if title != "" { + if result.title != "" { return database.Chat{}, xerrors.Errorf("record manual title usage and update chat title: %w", recordErr) } return database.Chat{}, xerrors.Errorf("record manual title usage: %w", recordErr) diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index 6cc6956ab52..1f77d8f0396 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -4847,12 +4847,14 @@ func TestAutoPromote_InsertFailureSkipsStatusUpdate(t *testing.T) { heartbeatRegistry: make(map[uuid.UUID]*heartbeatEntry), } - // Block model resolution until the control subscriber fires. + // Block model resolution until the running status has been + // published. Returning ErrInterrupted makes processChat enter the + // waiting-state auto-promotion path deterministically. modelBlocked := make(chan struct{}) db.EXPECT().GetChatModelConfigByID(gomock.Any(), gomock.Any()).DoAndReturn( - func(ctx context.Context, _ uuid.UUID) (database.ChatModelConfig, error) { + func(context.Context, uuid.UUID) (database.ChatModelConfig, error) { <-modelBlocked - return database.ChatModelConfig{}, xerrors.New("no model") + return database.ChatModelConfig{}, chatloop.ErrInterrupted }, ).AnyTimes() db.EXPECT().GetEnabledChatProviders(gomock.Any()).Return(nil, nil).AnyTimes() @@ -4914,15 +4916,6 @@ func TestAutoPromote_InsertFailureSkipsStatusUpdate(t *testing.T) { t.Fatal("timed out waiting for running status") } - // Publish an interrupt so processChat exits runChat. - interruptMsg, err := json.Marshal(coderdpubsub.ChatStreamNotifyMessage{ - Status: string(database.ChatStatusWaiting), - }) - require.NoError(t, err) - err = ps.Publish(coderdpubsub.ChatStreamNotifyChannel(chatID), interruptMsg) - require.NoError(t, err) - - // Unblock model resolution so runChat can exit. close(modelBlocked) select { diff --git a/coderd/x/chatd/chatd_test.go b/coderd/x/chatd/chatd_test.go index 35099214568..5de377ac2e9 100644 --- a/coderd/x/chatd/chatd_test.go +++ b/coderd/x/chatd/chatd_test.go @@ -5859,6 +5859,138 @@ func newActiveTestServer( return server } +func TestProposeChatTitle_DebugRun(t *testing.T) { + t.Parallel() + + wantTitle := "Debug proposal title" + tests := []struct { + name string + alwaysEnableDebugLogs bool + response func() chattest.OpenAIResponse + wantErr bool + wantTitle string + wantTitleGenerationRuns int + wantDebugStatus codersdk.ChatDebugStatus + }{ + { + name: "Enabled", + alwaysEnableDebugLogs: true, + response: func() chattest.OpenAIResponse { + return chattest.OpenAINonStreamingResponse( + "{\"title\":\"" + wantTitle + "\"}", + ) + }, + wantTitle: wantTitle, + wantTitleGenerationRuns: 1, + wantDebugStatus: codersdk.ChatDebugStatusCompleted, + }, + { + name: "Disabled", + alwaysEnableDebugLogs: false, + response: func() chattest.OpenAIResponse { + return chattest.OpenAINonStreamingResponse( + "{\"title\":\"" + wantTitle + "\"}", + ) + }, + wantTitle: wantTitle, + }, + { + name: "GenerationErrorFinalizesDebugRun", + alwaysEnableDebugLogs: true, + response: func() chattest.OpenAIResponse { + return chattest.OpenAINonStreamingResponse("not json") + }, + wantErr: true, + wantTitleGenerationRuns: 1, + wantDebugStatus: codersdk.ChatDebugStatusError, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + db, ps, rawDB := dbtestutil.NewDBWithSQLDB(t) + openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { + require.False(t, req.Stream) + return tt.response() + }) + user, org, model := seedChatDependenciesWithProvider( + ctx, + t, + db, + "openai", + openAIURL, + ) + server := chatd.New(chatd.Config{ + Logger: slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}), + Database: db, + ReplicaID: uuid.New(), + Pubsub: ps, + PendingChatAcquireInterval: testutil.WaitLong, + AlwaysEnableDebugLogs: tt.alwaysEnableDebugLogs, + }) + t.Cleanup(func() { + require.NoError(t, server.Close()) + }) + + chat, err := db.InsertChat(ctx, database.InsertChatParams{ + OrganizationID: org.ID, + Status: database.ChatStatusCompleted, + ClientType: database.ChatClientTypeUi, + OwnerID: user.ID, + Title: "original title", + LastModelConfigID: model.ID, + }) + require.NoError(t, err) + messages := insertUserTextMessage( + ctx, + t, + db, + chat.ID, + user.ID, + model.ID, + "summarize debug title generation", + model.ContextLimit, + ) + require.Len(t, messages, 1) + + gotTitle, err := server.ProposeChatTitle(ctx, chat) + if tt.wantErr { + require.Error(t, err) + } else { + require.NoError(t, err) + require.Equal(t, tt.wantTitle, gotTitle) + } + + runs, err := db.GetChatDebugRunsByChatID(ctx, database.GetChatDebugRunsByChatIDParams{ + ChatID: chat.ID, + LimitVal: 100, + }) + require.NoError(t, err) + require.Len(t, runs, tt.wantTitleGenerationRuns) + if tt.wantTitleGenerationRuns > 0 { + require.Equal(t, string(codersdk.ChatDebugRunKindTitleGeneration), runs[0].Kind) + require.Equal(t, string(tt.wantDebugStatus), runs[0].Status) + require.True(t, runs[0].FinishedAt.Valid) + require.True(t, runs[0].HistoryTipMessageID.Valid) + require.Equal(t, messages[0].ID, runs[0].HistoryTipMessageID.Int64) + } + if !tt.wantErr { + var usageMessages int + err = rawDB.QueryRowContext( + ctx, + `SELECT count(*) FROM chat_messages WHERE chat_id = $1 AND visibility = 'model' AND deleted = true`, + chat.ID, + ).Scan(&usageMessages) + require.NoError(t, err) + require.Equal(t, 1, usageMessages) + } + }) + } +} + func seedChatDependencies( ctx context.Context, t *testing.T, @@ -6052,13 +6184,19 @@ func insertUserTextMessage( userID uuid.UUID, modelConfigID uuid.UUID, text string, -) { + contextLimit ...int64, +) []database.ChatMessage { t.Helper() + require.LessOrEqual(t, len(contextLimit), 1) + contextLimitValue := int64(0) + if len(contextLimit) == 1 { + contextLimitValue = contextLimit[0] + } content, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{codersdk.ChatMessageText(text)}) require.NoError(t, err) - _, err = db.InsertChatMessages(ctx, database.InsertChatMessagesParams{ + messages, err := db.InsertChatMessages(ctx, database.InsertChatMessagesParams{ ChatID: chatID, CreatedBy: []uuid.UUID{userID}, ModelConfigID: []uuid.UUID{modelConfigID}, @@ -6072,13 +6210,14 @@ func insertUserTextMessage( ReasoningTokens: []int64{0}, CacheCreationTokens: []int64{0}, CacheReadTokens: []int64{0}, - ContextLimit: []int64{0}, + ContextLimit: []int64{contextLimitValue}, Compressed: []bool{false}, TotalCostMicros: []int64{0}, RuntimeMs: []int64{0}, ProviderResponseID: []string{""}, }) require.NoError(t, err) + return messages } // seedWorkspaceWithAgent creates a full workspace chain with a connected diff --git a/site/src/api/queries/chats.test.ts b/site/src/api/queries/chats.test.ts index 54eb5765695..15bf01bd763 100644 --- a/site/src/api/queries/chats.test.ts +++ b/site/src/api/queries/chats.test.ts @@ -30,6 +30,7 @@ import { paginatedChatCostUsers, pinChat, promoteChatQueuedMessage, + proposeChatTitle, regenerateChatTitle, removeChildFromParentInCache, reorderPinnedChat, @@ -54,6 +55,7 @@ vi.mock("#/api/api", () => ({ editChatMessage: vi.fn(), interruptChat: vi.fn(), promoteChatQueuedMessage: vi.fn(), + proposeChatTitle: vi.fn(), regenerateChatTitle: vi.fn(), }, }, @@ -1293,6 +1295,39 @@ describe("mutation invalidation scope", () => { } }); + for (const { label, error } of [ + { label: "success", error: undefined }, + { label: "failure", error: new Error("proposal failed") }, + ]) { + it(`proposeChatTitle invalidates debug runs on ${label} without touching unrelated queries`, async () => { + const queryClient = createTestQueryClient(); + const chatId = "chat-1"; + seedAllActiveQueries(queryClient, chatId); + + const mutation = proposeChatTitle(queryClient); + await mutation.onSettled(undefined, error, chatId); + + expect( + queryClient.getQueryState(chatDebugRunsKey(chatId))?.isInvalidated, + "chatDebugRunsKey should be invalidated", + ).toBe(true); + + for (const { label, key } of [ + { label: "flat chats", key: chatsKey }, + { label: "infinite chats", key: [...chatsKey, { archived: false }] }, + { label: "chat detail", key: chatKey(chatId) }, + { label: "messages", key: chatMessagesKey(chatId) }, + ...unrelatedKeys(chatId), + ]) { + const state = queryClient.getQueryState(key); + expect( + state?.isInvalidated, + `${label} should NOT be invalidated by proposeChatTitle`, + ).not.toBe(true); + } + }); + } + it("createChat invalidates only sidebar queries on success", async () => { const queryClient = createTestQueryClient(); const chatId = "chat-1"; diff --git a/site/src/api/queries/chats.ts b/site/src/api/queries/chats.ts index 711d28a89ac..c2124175d3e 100644 --- a/site/src/api/queries/chats.ts +++ b/site/src/api/queries/chats.ts @@ -1002,6 +1002,18 @@ export const regenerateChatTitle = (queryClient: QueryClient) => ({ }, }); +export const proposeChatTitle = (queryClient: QueryClient) => ({ + mutationFn: (chatId: string) => API.experimental.proposeChatTitle(chatId), + + onSettled: ( + _data: { title: string } | undefined, + _error: unknown, + chatId: string, + ) => { + void invalidateChatDebugRuns(queryClient, chatId); + }, +}); + type UpdateChatTitleVariables = { chatId: string; title: string; diff --git a/site/src/pages/AgentsPage/AgentsPage.tsx b/site/src/pages/AgentsPage/AgentsPage.tsx index 5521642a389..ad2097c2beb 100644 --- a/site/src/pages/AgentsPage/AgentsPage.tsx +++ b/site/src/pages/AgentsPage/AgentsPage.tsx @@ -23,6 +23,7 @@ import { mergeWatchedChatIntoCaches, pinChat, prependToInfiniteChatsCache, + proposeChatTitle, readInfiniteChatsCache, regenerateChatTitle, removeChildFromParentInCache, @@ -246,6 +247,7 @@ const AgentsPage: FC = () => { toast.error(getErrorMessage(error, "Failed to generate new title.")); }, }); + const proposeTitleMutation = useMutation(proposeChatTitle(queryClient)); const renameTitleMutation = useMutation({ ...updateChatTitle(queryClient), onError: (error: unknown) => { @@ -438,7 +440,7 @@ const AgentsPage: FC = () => { return promise; }; const requestProposeTitle = async (chatId: string): Promise => { - const result = await API.experimental.proposeChatTitle(chatId); + const result = await proposeTitleMutation.mutateAsync(chatId); return result.title; }; const requestRenameTitle = async (chatId: string, title: string) => {