Thanks to visit codestin.com
Credit goes to github.com

Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
84 changes: 48 additions & 36 deletions coderd/x/chatd/chatd.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Comment thread
ThomasK33 marked this conversation as resolved.
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(
Expand All @@ -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,
Expand All @@ -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(
Expand All @@ -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
}

Expand All @@ -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(
Expand All @@ -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
}

Expand All @@ -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)
Expand Down
17 changes: 5 additions & 12 deletions coderd/x/chatd/chatd_internal_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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 {
Expand Down
145 changes: 142 additions & 3 deletions coderd/x/chatd/chatd_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5859,6 +5859,138 @@ func newActiveTestServer(
return server
}

func TestProposeChatTitle_DebugRun(t *testing.T) {
Comment thread
ThomasK33 marked this conversation as resolved.
Comment thread
ThomasK33 marked this conversation as resolved.
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,
Expand Down Expand Up @@ -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},
Expand All @@ -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
Expand Down
Loading
Loading