mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: allow renaming of agent chat title (#24489)
Co-authored-by: Coder Agents <noreply@coder.com>
This commit is contained in:
co-authored by
Coder Agents
parent
18a30a7a10
commit
410f9a5e19
+162
-48
@@ -2196,44 +2196,112 @@ func (p *Server) RegenerateChatTitle(
|
||||
keys,
|
||||
)
|
||||
if err != nil {
|
||||
var generationErr *manualTitleGenerationError
|
||||
if errors.As(err, &generationErr) {
|
||||
// Reuse chatd's scoped auth context for failure accounting while
|
||||
// detaching from request cancellation so usage is still recorded.
|
||||
//nolint:gocritic // Failure accounting still needs chatd-scoped config reads.
|
||||
recordCtx, recordCancel := context.WithTimeout(
|
||||
dbauthz.AsChatd(context.WithoutCancel(ctx)),
|
||||
5*time.Second,
|
||||
)
|
||||
defer recordCancel()
|
||||
if _, recordErr := recordManualTitleUsage(
|
||||
recordCtx,
|
||||
p.db,
|
||||
chat,
|
||||
generationErr.modelConfig,
|
||||
generationErr.usage,
|
||||
"",
|
||||
); recordErr != nil {
|
||||
return database.Chat{}, errors.Join(
|
||||
generationErr,
|
||||
xerrors.Errorf("record manual title usage: %w", recordErr),
|
||||
)
|
||||
}
|
||||
return database.Chat{}, generationErr
|
||||
}
|
||||
return database.Chat{}, err
|
||||
return database.Chat{}, p.recordManualTitleGenerationFailure(ctx, chat, err)
|
||||
}
|
||||
return updatedChat, nil
|
||||
}
|
||||
|
||||
func (p *Server) regenerateChatTitleWithStore(
|
||||
// RenameChatTitle persists a user-supplied chat title.
|
||||
func (p *Server) RenameChatTitle(
|
||||
ctx context.Context,
|
||||
chat database.Chat,
|
||||
newTitle string,
|
||||
) (updated database.Chat, wrote bool, err error) {
|
||||
//nolint:gocritic // Lock release needs chatd-scoped writes.
|
||||
chatdCtx := dbauthz.AsChatd(ctx)
|
||||
if err := p.acquireManualTitleLock(ctx, chat.ID); err != nil {
|
||||
return database.Chat{}, false, err
|
||||
}
|
||||
defer p.releaseManualTitleLock(chatdCtx, chat.ID)
|
||||
|
||||
currentChat, err := p.db.GetChatByID(ctx, chat.ID)
|
||||
if err != nil {
|
||||
return database.Chat{}, false, xerrors.Errorf("get chat for rename: %w", err)
|
||||
}
|
||||
if newTitle == currentChat.Title {
|
||||
return currentChat, false, nil
|
||||
}
|
||||
|
||||
updatedChat, err := p.db.UpdateChatTitleByID(ctx, database.UpdateChatTitleByIDParams{
|
||||
ID: chat.ID,
|
||||
Title: newTitle,
|
||||
})
|
||||
if err != nil {
|
||||
return database.Chat{}, false, xerrors.Errorf("update chat title: %w", err)
|
||||
}
|
||||
return updatedChat, true, nil
|
||||
}
|
||||
|
||||
// PublishTitleChange broadcasts a title_change event for the given chat.
|
||||
func (p *Server) PublishTitleChange(chat database.Chat) {
|
||||
p.publishChatPubsubEvent(chat, codersdk.ChatWatchEventKindTitleChange, nil)
|
||||
}
|
||||
|
||||
// ProposeChatTitle generates a title suggestion from the chat's visible messages without persisting it.
|
||||
func (p *Server) ProposeChatTitle(
|
||||
ctx context.Context,
|
||||
chat database.Chat,
|
||||
) (string, error) {
|
||||
//nolint:gocritic // Non-admin users need chatd-scoped config reads here.
|
||||
chatdCtx := dbauthz.AsChatd(ctx)
|
||||
keys, err := p.resolveUserProviderAPIKeys(chatdCtx, chat.OwnerID)
|
||||
if err != nil {
|
||||
return "", xerrors.Errorf("resolve chat providers: %w", err)
|
||||
}
|
||||
if err := p.acquireManualTitleLock(ctx, chat.ID); err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer p.releaseManualTitleLock(chatdCtx, chat.ID)
|
||||
|
||||
title, err := p.proposeChatTitleWithStore(chatdCtx, p.db, chat, keys)
|
||||
if err != nil {
|
||||
return "", p.recordManualTitleGenerationFailure(ctx, chat, err)
|
||||
}
|
||||
return title, nil
|
||||
}
|
||||
|
||||
func (p *Server) recordManualTitleGenerationFailure(
|
||||
ctx context.Context,
|
||||
chat database.Chat,
|
||||
err error,
|
||||
) error {
|
||||
var generationErr *manualTitleGenerationError
|
||||
if !errors.As(err, &generationErr) {
|
||||
return err
|
||||
}
|
||||
|
||||
//nolint:gocritic // Failure accounting still needs chatd-scoped config reads.
|
||||
recordCtx, recordCancel := context.WithTimeout(
|
||||
dbauthz.AsChatd(context.WithoutCancel(ctx)),
|
||||
5*time.Second,
|
||||
)
|
||||
defer recordCancel()
|
||||
if _, recordErr := recordManualTitleUsage(
|
||||
recordCtx,
|
||||
p.db,
|
||||
chat,
|
||||
generationErr.modelConfig,
|
||||
generationErr.usage,
|
||||
"",
|
||||
); recordErr != nil {
|
||||
return errors.Join(
|
||||
generationErr,
|
||||
xerrors.Errorf("record manual title usage: %w", recordErr),
|
||||
)
|
||||
}
|
||||
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(
|
||||
ctx context.Context,
|
||||
store database.Store,
|
||||
chat database.Chat,
|
||||
keys chatprovider.ProviderAPIKeys,
|
||||
) (database.Chat, error) {
|
||||
enableDebug bool,
|
||||
) (title string, modelConfig database.ChatModelConfig, usage fantasy.Usage, hasMessages bool, err error) {
|
||||
if limitErr := p.checkUsageLimit(ctx, store, chat.OwnerID, uuid.NullUUID{UUID: chat.OrganizationID, Valid: true}); limitErr != nil {
|
||||
return database.Chat{}, limitErr
|
||||
return "", database.ChatModelConfig{}, fantasy.Usage{}, false, limitErr
|
||||
}
|
||||
|
||||
headMessages, err := store.GetChatMessagesByChatIDAscPaginated(
|
||||
@@ -2245,7 +2313,7 @@ func (p *Server) regenerateChatTitleWithStore(
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
return database.Chat{}, xerrors.Errorf("get head chat messages: %w", err)
|
||||
return "", database.ChatModelConfig{}, fantasy.Usage{}, false, xerrors.Errorf("get head chat messages: %w", err)
|
||||
}
|
||||
tailMessages, err := store.GetChatMessagesByChatIDDescPaginated(
|
||||
ctx,
|
||||
@@ -2256,49 +2324,95 @@ func (p *Server) regenerateChatTitleWithStore(
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
return database.Chat{}, xerrors.Errorf("get tail chat messages: %w", err)
|
||||
return "", database.ChatModelConfig{}, fantasy.Usage{}, false, xerrors.Errorf("get tail chat messages: %w", err)
|
||||
}
|
||||
messages := mergeManualTitleMessages(headMessages, tailMessages)
|
||||
if len(messages) == 0 {
|
||||
return chat, nil
|
||||
return "", database.ChatModelConfig{}, fantasy.Usage{}, false, nil
|
||||
}
|
||||
|
||||
model, modelConfig, err := p.resolveManualTitleModel(ctx, store, chat, keys)
|
||||
if err != nil {
|
||||
return database.Chat{}, err
|
||||
return "", database.ChatModelConfig{}, fantasy.Usage{}, true, err
|
||||
}
|
||||
|
||||
debugSvc := p.debugService()
|
||||
debugEnabled := debugSvc != nil && debugSvc.IsEnabled(ctx, chat.ID, chat.OwnerID)
|
||||
titleCtx := ctx
|
||||
titleModel := model
|
||||
finishDebugRun := func(error) {}
|
||||
if debugEnabled {
|
||||
titleCtx, titleModel, finishDebugRun = p.prepareManualTitleDebugRun(
|
||||
ctx,
|
||||
debugSvc,
|
||||
chat,
|
||||
modelConfig,
|
||||
keys,
|
||||
messages,
|
||||
model,
|
||||
)
|
||||
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,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
title, usage, err := generateManualTitle(titleCtx, messages, titleModel)
|
||||
title, usage, err = generateManualTitle(titleCtx, messages, titleModel)
|
||||
finishDebugRun(err)
|
||||
if err != nil {
|
||||
wrappedErr := xerrors.Errorf("generate manual title: %w", err)
|
||||
if usage == (fantasy.Usage{}) {
|
||||
return database.Chat{}, wrappedErr
|
||||
return "", modelConfig, fantasy.Usage{}, true, wrappedErr
|
||||
}
|
||||
return database.Chat{}, &manualTitleGenerationError{
|
||||
return "", modelConfig, usage, true, &manualTitleGenerationError{
|
||||
cause: wrappedErr,
|
||||
modelConfig: modelConfig,
|
||||
usage: usage,
|
||||
}
|
||||
}
|
||||
|
||||
return title, modelConfig, usage, true, nil
|
||||
}
|
||||
|
||||
func (p *Server) proposeChatTitleWithStore(
|
||||
ctx context.Context,
|
||||
store database.Store,
|
||||
chat database.Chat,
|
||||
keys chatprovider.ProviderAPIKeys,
|
||||
) (string, error) {
|
||||
title, modelConfig, usage, hasMessages, err := p.fetchAndGenerateManualTitle(ctx, store, chat, keys, false)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if !hasMessages {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
recordCtx, recordCancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second)
|
||||
defer recordCancel()
|
||||
if _, recordErr := recordManualTitleUsage(
|
||||
recordCtx,
|
||||
store,
|
||||
chat,
|
||||
modelConfig,
|
||||
usage,
|
||||
"",
|
||||
); recordErr != nil {
|
||||
return "", xerrors.Errorf("record manual title usage: %w", recordErr)
|
||||
}
|
||||
return title, nil
|
||||
}
|
||||
|
||||
func (p *Server) regenerateChatTitleWithStore(
|
||||
ctx context.Context,
|
||||
store database.Store,
|
||||
chat database.Chat,
|
||||
keys chatprovider.ProviderAPIKeys,
|
||||
) (database.Chat, error) {
|
||||
title, modelConfig, usage, hasMessages, err := p.fetchAndGenerateManualTitle(ctx, store, chat, keys, true)
|
||||
if err != nil {
|
||||
return database.Chat{}, err
|
||||
}
|
||||
if !hasMessages {
|
||||
return chat, nil
|
||||
}
|
||||
|
||||
recordCtx, recordCancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second)
|
||||
defer recordCancel()
|
||||
|
||||
|
||||
@@ -287,6 +287,99 @@ func TestStopAfterBehaviorTools(t *testing.T) {
|
||||
// TestArchiveChatWaitsForEveryInterruptedChat were removed along with
|
||||
// the process-local activeChats mechanism. Archive cleanup is now
|
||||
// best-effort; stale finalization handles any orphaned rows.
|
||||
|
||||
func TestRenameChatTitle(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
setupRealWorkerLock := func(
|
||||
db *dbmock.MockStore,
|
||||
chatID uuid.UUID,
|
||||
lockedChat database.Chat,
|
||||
) {
|
||||
lockTx := dbmock.NewMockStore(gomock.NewController(t))
|
||||
unlockTx := dbmock.NewMockStore(gomock.NewController(t))
|
||||
gomock.InOrder(
|
||||
db.EXPECT().InTx(gomock.Any(), database.DefaultTXOptions().WithID("chat_title_regenerate_lock")).DoAndReturn(
|
||||
func(fn func(database.Store) error, _ *database.TxOptions) error {
|
||||
return fn(lockTx)
|
||||
},
|
||||
),
|
||||
db.EXPECT().InTx(gomock.Any(), database.DefaultTXOptions().WithID("chat_title_regenerate_unlock")).DoAndReturn(
|
||||
func(fn func(database.Store) error, _ *database.TxOptions) error {
|
||||
return fn(unlockTx)
|
||||
},
|
||||
),
|
||||
)
|
||||
lockTx.EXPECT().GetChatByIDForUpdate(gomock.Any(), chatID).Return(lockedChat, nil)
|
||||
unlockTx.EXPECT().GetChatByIDForUpdate(gomock.Any(), chatID).Return(lockedChat, nil)
|
||||
}
|
||||
|
||||
t.Run("WritesAndReturnsWroteTrue", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
|
||||
chatID := uuid.New()
|
||||
workerID := uuid.New()
|
||||
stored := database.Chat{
|
||||
ID: chatID,
|
||||
Status: database.ChatStatusRunning,
|
||||
WorkerID: uuid.NullUUID{UUID: workerID, Valid: true},
|
||||
Title: "original",
|
||||
}
|
||||
updated := stored
|
||||
updated.Title = "renamed"
|
||||
|
||||
server := &Server{db: db, logger: logger}
|
||||
|
||||
setupRealWorkerLock(db, chatID, stored)
|
||||
db.EXPECT().GetChatByID(gomock.Any(), chatID).Return(stored, nil)
|
||||
db.EXPECT().UpdateChatTitleByID(gomock.Any(), database.UpdateChatTitleByIDParams{
|
||||
ID: chatID,
|
||||
Title: "renamed",
|
||||
}).Return(updated, nil)
|
||||
|
||||
got, wrote, err := server.RenameChatTitle(ctx, stored, "renamed")
|
||||
require.NoError(t, err)
|
||||
require.True(t, wrote, "fresh rename must report wrote=true")
|
||||
require.Equal(t, updated, got)
|
||||
})
|
||||
|
||||
t.Run("SkipsWriteWhenAlreadyAtNewTitle", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
|
||||
chatID := uuid.New()
|
||||
workerID := uuid.New()
|
||||
stale := database.Chat{
|
||||
ID: chatID,
|
||||
Status: database.ChatStatusRunning,
|
||||
WorkerID: uuid.NullUUID{UUID: workerID, Valid: true},
|
||||
Title: "pre-race",
|
||||
}
|
||||
landed := stale
|
||||
landed.Title = "landed-concurrently"
|
||||
|
||||
server := &Server{db: db, logger: logger}
|
||||
|
||||
setupRealWorkerLock(db, chatID, landed)
|
||||
db.EXPECT().GetChatByID(gomock.Any(), chatID).Return(landed, nil)
|
||||
|
||||
got, wrote, err := server.RenameChatTitle(ctx, stale, "landed-concurrently")
|
||||
require.NoError(t, err)
|
||||
require.False(t, wrote,
|
||||
"must report wrote=false when the stored row already matches newTitle so the handler suppresses a redundant title_change event")
|
||||
require.Equal(t, landed, got)
|
||||
})
|
||||
}
|
||||
|
||||
func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user