mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat(coderd): wire debug logging into chat lifecycle (#23917)
This commit is contained in:
@@ -1871,15 +1871,15 @@ func (q *querier) DeleteChatDebugDataAfterMessageID(ctx context.Context, arg dat
|
||||
return q.db.DeleteChatDebugDataAfterMessageID(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) DeleteChatDebugDataByChatID(ctx context.Context, chatID uuid.UUID) (int64, error) {
|
||||
chat, err := q.db.GetChatByID(ctx, chatID)
|
||||
func (q *querier) DeleteChatDebugDataByChatID(ctx context.Context, arg database.DeleteChatDebugDataByChatIDParams) (int64, error) {
|
||||
chat, err := q.db.GetChatByID(ctx, arg.ChatID)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if err := q.authorizeContext(ctx, policy.ActionUpdate, chat); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return q.db.DeleteChatDebugDataByChatID(ctx, chatID)
|
||||
return q.db.DeleteChatDebugDataByChatID(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) DeleteChatModelConfigByID(ctx context.Context, id uuid.UUID) error {
|
||||
|
||||
@@ -463,16 +463,17 @@ func (s *MethodTestSuite) TestChats() {
|
||||
}))
|
||||
s.Run("DeleteChatDebugDataAfterMessageID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
chat := testutil.Fake(s.T(), faker, database.Chat{})
|
||||
arg := database.DeleteChatDebugDataAfterMessageIDParams{ChatID: chat.ID, MessageID: 123}
|
||||
arg := database.DeleteChatDebugDataAfterMessageIDParams{ChatID: chat.ID, StartedBefore: dbtime.Now(), MessageID: 123}
|
||||
dbm.EXPECT().GetChatByID(gomock.Any(), chat.ID).Return(chat, nil).AnyTimes()
|
||||
dbm.EXPECT().DeleteChatDebugDataAfterMessageID(gomock.Any(), arg).Return(int64(1), nil).AnyTimes()
|
||||
check.Args(arg).Asserts(chat, policy.ActionUpdate).Returns(int64(1))
|
||||
}))
|
||||
s.Run("DeleteChatDebugDataByChatID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
chat := testutil.Fake(s.T(), faker, database.Chat{})
|
||||
arg := database.DeleteChatDebugDataByChatIDParams{ChatID: chat.ID, StartedBefore: dbtime.Now()}
|
||||
dbm.EXPECT().GetChatByID(gomock.Any(), chat.ID).Return(chat, nil).AnyTimes()
|
||||
dbm.EXPECT().DeleteChatDebugDataByChatID(gomock.Any(), chat.ID).Return(int64(1), nil).AnyTimes()
|
||||
check.Args(chat.ID).Asserts(chat, policy.ActionUpdate).Returns(int64(1))
|
||||
dbm.EXPECT().DeleteChatDebugDataByChatID(gomock.Any(), arg).Return(int64(1), nil).AnyTimes()
|
||||
check.Args(arg).Asserts(chat, policy.ActionUpdate).Returns(int64(1))
|
||||
}))
|
||||
s.Run("FinalizeStaleChatDebugRows", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
|
||||
now := dbtime.Now()
|
||||
|
||||
@@ -424,7 +424,7 @@ func (m queryMetricsStore) DeleteChatDebugDataAfterMessageID(ctx context.Context
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) DeleteChatDebugDataByChatID(ctx context.Context, chatID uuid.UUID) (int64, error) {
|
||||
func (m queryMetricsStore) DeleteChatDebugDataByChatID(ctx context.Context, chatID database.DeleteChatDebugDataByChatIDParams) (int64, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.DeleteChatDebugDataByChatID(ctx, chatID)
|
||||
m.queryLatencies.WithLabelValues("DeleteChatDebugDataByChatID").Observe(time.Since(start).Seconds())
|
||||
|
||||
@@ -687,18 +687,18 @@ func (mr *MockStoreMockRecorder) DeleteChatDebugDataAfterMessageID(ctx, arg any)
|
||||
}
|
||||
|
||||
// DeleteChatDebugDataByChatID mocks base method.
|
||||
func (m *MockStore) DeleteChatDebugDataByChatID(ctx context.Context, chatID uuid.UUID) (int64, error) {
|
||||
func (m *MockStore) DeleteChatDebugDataByChatID(ctx context.Context, arg database.DeleteChatDebugDataByChatIDParams) (int64, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "DeleteChatDebugDataByChatID", ctx, chatID)
|
||||
ret := m.ctrl.Call(m, "DeleteChatDebugDataByChatID", ctx, arg)
|
||||
ret0, _ := ret[0].(int64)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// DeleteChatDebugDataByChatID indicates an expected call of DeleteChatDebugDataByChatID.
|
||||
func (mr *MockStoreMockRecorder) DeleteChatDebugDataByChatID(ctx, chatID any) *gomock.Call {
|
||||
func (mr *MockStoreMockRecorder) DeleteChatDebugDataByChatID(ctx, arg any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteChatDebugDataByChatID", reflect.TypeOf((*MockStore)(nil).DeleteChatDebugDataByChatID), ctx, chatID)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteChatDebugDataByChatID", reflect.TypeOf((*MockStore)(nil).DeleteChatDebugDataByChatID), ctx, arg)
|
||||
}
|
||||
|
||||
// DeleteChatModelConfigByID mocks base method.
|
||||
|
||||
@@ -102,8 +102,16 @@ type sqlcQuerier interface {
|
||||
// be recreated.
|
||||
DeleteAllWebpushSubscriptions(ctx context.Context) error
|
||||
DeleteApplicationConnectAPIKeysByUserID(ctx context.Context, userID uuid.UUID) error
|
||||
// Deletes debug runs (and their cascaded steps) whose message IDs
|
||||
// exceed the cutoff. The started_before bound prevents retried
|
||||
// cleanup from deleting runs created by a replacement turn that
|
||||
// raced ahead of the retry window.
|
||||
DeleteChatDebugDataAfterMessageID(ctx context.Context, arg DeleteChatDebugDataAfterMessageIDParams) (int64, error)
|
||||
DeleteChatDebugDataByChatID(ctx context.Context, chatID uuid.UUID) (int64, error)
|
||||
// The started_before bound prevents retried cleanup from deleting
|
||||
// runs created by a replacement turn that races ahead of the retry
|
||||
// window (for example, after an unarchive races with a pending
|
||||
// archive-cleanup retry).
|
||||
DeleteChatDebugDataByChatID(ctx context.Context, arg DeleteChatDebugDataByChatIDParams) (int64, error)
|
||||
DeleteChatModelConfigByID(ctx context.Context, id uuid.UUID) error
|
||||
DeleteChatProviderByID(ctx context.Context, id uuid.UUID) error
|
||||
DeleteChatQueuedMessage(ctx context.Context, arg DeleteChatQueuedMessageParams) error
|
||||
|
||||
@@ -11524,8 +11524,9 @@ func TestDeleteChatDebugDataAfterMessageIDIncludesTriggeredRuns(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
deletedRows, err := store.DeleteChatDebugDataAfterMessageID(ctx, database.DeleteChatDebugDataAfterMessageIDParams{
|
||||
ChatID: chat.ID,
|
||||
MessageID: cutoff,
|
||||
ChatID: chat.ID,
|
||||
MessageID: cutoff,
|
||||
StartedBefore: time.Now().Add(time.Minute),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 3, deletedRows)
|
||||
@@ -12406,8 +12407,9 @@ func TestDeleteChatDebugDataAfterMessageIDNullMessagesSurvive(t *testing.T) {
|
||||
// Delete with an arbitrary cutoff. The run and its step should
|
||||
// survive because NULL > cutoff evaluates to NULL, not TRUE.
|
||||
deletedRows, err := store.DeleteChatDebugDataAfterMessageID(ctx, database.DeleteChatDebugDataAfterMessageIDParams{
|
||||
ChatID: chat.ID,
|
||||
MessageID: 1,
|
||||
ChatID: chat.ID,
|
||||
MessageID: 1,
|
||||
StartedBefore: time.Now().Add(time.Minute),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 0, deletedRows, "rows with NULL message IDs must not be deleted")
|
||||
@@ -12424,6 +12426,215 @@ func TestDeleteChatDebugDataAfterMessageIDNullMessagesSurvive(t *testing.T) {
|
||||
require.Equal(t, nullMsgStep.ID, remainingSteps[0].ID)
|
||||
}
|
||||
|
||||
// TestDeleteChatDebugDataAfterMessageIDStartedBeforeFiltersNewerRuns
|
||||
// verifies the started_before bound on DeleteChatDebugDataAfterMessageID.
|
||||
// The bound exists so that retried cleanup (e.g. after edit or archive)
|
||||
// cannot delete runs started by a replacement turn that races ahead of
|
||||
// the retry window. Without this filter, a stale cleanup would wipe
|
||||
// fresh debug rows.
|
||||
func TestDeleteChatDebugDataAfterMessageIDStartedBeforeFiltersNewerRuns(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
store, _ := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitMedium)
|
||||
|
||||
org := dbgen.Organization(t, store, database.Organization{})
|
||||
user := dbgen.User(t, store, database.User{})
|
||||
|
||||
providerName := "openai"
|
||||
modelName := "debug-model-started-before-" + uuid.NewString()
|
||||
|
||||
_, err := store.InsertChatProvider(ctx, database.InsertChatProviderParams{
|
||||
Provider: providerName,
|
||||
DisplayName: "Debug Provider",
|
||||
APIKey: "test-key",
|
||||
Enabled: true,
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
modelCfg, err := store.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{
|
||||
Provider: providerName,
|
||||
Model: modelName,
|
||||
DisplayName: "Debug Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
IsDefault: true,
|
||||
ContextLimit: 128000,
|
||||
CompressionThreshold: 80,
|
||||
Options: json.RawMessage(`{}`),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
chat, err := store.InsertChat(ctx, database.InsertChatParams{
|
||||
OrganizationID: org.ID,
|
||||
Status: database.ChatStatusWaiting,
|
||||
ClientType: database.ChatClientTypeUi,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: modelCfg.ID,
|
||||
Title: "chat-debug-started-before-" + uuid.NewString(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
const cutoff int64 = 50
|
||||
|
||||
// oldRun started an hour ago: must be deleted because it started
|
||||
// before the bound.
|
||||
oldStartedAt := time.Now().Add(-1 * time.Hour).UTC().
|
||||
Truncate(time.Microsecond)
|
||||
oldRun, err := store.InsertChatDebugRun(ctx, database.InsertChatDebugRunParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelCfg.ID, Valid: true},
|
||||
TriggerMessageID: sql.NullInt64{Int64: cutoff + 1, Valid: true},
|
||||
HistoryTipMessageID: sql.NullInt64{Int64: cutoff + 1, Valid: true},
|
||||
Kind: "chat_turn",
|
||||
Status: "in_progress",
|
||||
Provider: sql.NullString{String: providerName, Valid: true},
|
||||
Model: sql.NullString{String: modelName, Valid: true},
|
||||
StartedAt: sql.NullTime{Time: oldStartedAt, Valid: true},
|
||||
UpdatedAt: sql.NullTime{Time: oldStartedAt, Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Bound sits between the two runs. Any run whose started_at is at
|
||||
// or after this instant must survive.
|
||||
cutoffTime := time.Now().Add(-30 * time.Minute).UTC().
|
||||
Truncate(time.Microsecond)
|
||||
|
||||
// newRun started after cutoffTime with identical message_id values
|
||||
// that would otherwise match the delete predicate. It must survive
|
||||
// because started_before excludes it.
|
||||
newStartedAt := time.Now().UTC().Truncate(time.Microsecond)
|
||||
newRun, err := store.InsertChatDebugRun(ctx, database.InsertChatDebugRunParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelCfg.ID, Valid: true},
|
||||
TriggerMessageID: sql.NullInt64{Int64: cutoff + 1, Valid: true},
|
||||
HistoryTipMessageID: sql.NullInt64{Int64: cutoff + 1, Valid: true},
|
||||
Kind: "chat_turn",
|
||||
Status: "in_progress",
|
||||
Provider: sql.NullString{String: providerName, Valid: true},
|
||||
Model: sql.NullString{String: modelName, Valid: true},
|
||||
StartedAt: sql.NullTime{Time: newStartedAt, Valid: true},
|
||||
UpdatedAt: sql.NullTime{Time: newStartedAt, Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
deletedRows, err := store.DeleteChatDebugDataAfterMessageID(ctx, database.DeleteChatDebugDataAfterMessageIDParams{
|
||||
ChatID: chat.ID,
|
||||
MessageID: cutoff,
|
||||
StartedBefore: cutoffTime,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, deletedRows,
|
||||
"only the pre-cutoff run should be deleted")
|
||||
|
||||
// oldRun must be gone.
|
||||
_, err = store.GetChatDebugRunByID(ctx, oldRun.ID)
|
||||
require.ErrorIs(t, err, sql.ErrNoRows)
|
||||
|
||||
// newRun must survive the retry window.
|
||||
remaining, err := store.GetChatDebugRunByID(ctx, newRun.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, newRun.ID, remaining.ID)
|
||||
}
|
||||
|
||||
// TestDeleteChatDebugDataByChatIDStartedBeforeFiltersNewerRuns verifies
|
||||
// the started_before bound on DeleteChatDebugDataByChatID. Archive
|
||||
// cleanup retries rely on this bound to avoid deleting runs created
|
||||
// by a replacement turn that starts after an unarchive races ahead of
|
||||
// the retry window.
|
||||
func TestDeleteChatDebugDataByChatIDStartedBeforeFiltersNewerRuns(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
store, _ := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitMedium)
|
||||
|
||||
org := dbgen.Organization(t, store, database.Organization{})
|
||||
user := dbgen.User(t, store, database.User{})
|
||||
|
||||
providerName := "openai"
|
||||
modelName := "debug-model-by-chat-started-before-" + uuid.NewString()
|
||||
|
||||
_, err := store.InsertChatProvider(ctx, database.InsertChatProviderParams{
|
||||
Provider: providerName,
|
||||
DisplayName: "Debug Provider",
|
||||
APIKey: "test-key",
|
||||
Enabled: true,
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
modelCfg, err := store.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{
|
||||
Provider: providerName,
|
||||
Model: modelName,
|
||||
DisplayName: "Debug Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
IsDefault: true,
|
||||
ContextLimit: 128000,
|
||||
CompressionThreshold: 80,
|
||||
Options: json.RawMessage(`{}`),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
chat, err := store.InsertChat(ctx, database.InsertChatParams{
|
||||
OrganizationID: org.ID,
|
||||
Status: database.ChatStatusWaiting,
|
||||
ClientType: database.ChatClientTypeUi,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: modelCfg.ID,
|
||||
Title: "chat-debug-by-chat-" + uuid.NewString(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
oldStartedAt := time.Now().Add(-1 * time.Hour).UTC().
|
||||
Truncate(time.Microsecond)
|
||||
oldRun, err := store.InsertChatDebugRun(ctx, database.InsertChatDebugRunParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelCfg.ID, Valid: true},
|
||||
Kind: "chat_turn",
|
||||
Status: "in_progress",
|
||||
Provider: sql.NullString{String: providerName, Valid: true},
|
||||
Model: sql.NullString{String: modelName, Valid: true},
|
||||
StartedAt: sql.NullTime{Time: oldStartedAt, Valid: true},
|
||||
UpdatedAt: sql.NullTime{Time: oldStartedAt, Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
cutoffTime := time.Now().Add(-30 * time.Minute).UTC().
|
||||
Truncate(time.Microsecond)
|
||||
|
||||
newStartedAt := time.Now().UTC().Truncate(time.Microsecond)
|
||||
newRun, err := store.InsertChatDebugRun(ctx, database.InsertChatDebugRunParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelCfg.ID, Valid: true},
|
||||
Kind: "chat_turn",
|
||||
Status: "in_progress",
|
||||
Provider: sql.NullString{String: providerName, Valid: true},
|
||||
Model: sql.NullString{String: modelName, Valid: true},
|
||||
StartedAt: sql.NullTime{Time: newStartedAt, Valid: true},
|
||||
UpdatedAt: sql.NullTime{Time: newStartedAt, Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
deletedRows, err := store.DeleteChatDebugDataByChatID(ctx, database.DeleteChatDebugDataByChatIDParams{
|
||||
ChatID: chat.ID,
|
||||
StartedBefore: cutoffTime,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, deletedRows,
|
||||
"only the pre-cutoff run should be deleted")
|
||||
|
||||
_, err = store.GetChatDebugRunByID(ctx, oldRun.ID)
|
||||
require.ErrorIs(t, err, sql.ErrNoRows)
|
||||
|
||||
remaining, err := store.GetChatDebugRunByID(ctx, newRun.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, newRun.ID, remaining.ID)
|
||||
}
|
||||
|
||||
func TestChatHasUnread(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -2905,19 +2905,23 @@ WITH affected_runs AS (
|
||||
SELECT DISTINCT run.id
|
||||
FROM chat_debug_runs run
|
||||
WHERE run.chat_id = $1::uuid
|
||||
AND run.started_at < $2::timestamptz
|
||||
AND (
|
||||
run.history_tip_message_id > $2::bigint
|
||||
OR run.trigger_message_id > $2::bigint
|
||||
run.history_tip_message_id > $3::bigint
|
||||
OR run.trigger_message_id > $3::bigint
|
||||
)
|
||||
|
||||
UNION
|
||||
|
||||
SELECT DISTINCT step.run_id AS id
|
||||
FROM chat_debug_steps step
|
||||
JOIN chat_debug_runs run ON run.id = step.run_id
|
||||
AND run.chat_id = step.chat_id
|
||||
WHERE step.chat_id = $1::uuid
|
||||
AND run.started_at < $2::timestamptz
|
||||
AND (
|
||||
step.assistant_message_id > $2::bigint
|
||||
OR step.history_tip_message_id > $2::bigint
|
||||
step.assistant_message_id > $3::bigint
|
||||
OR step.history_tip_message_id > $3::bigint
|
||||
)
|
||||
)
|
||||
DELETE FROM chat_debug_runs
|
||||
@@ -2926,12 +2930,17 @@ WHERE chat_id = $1::uuid
|
||||
`
|
||||
|
||||
type DeleteChatDebugDataAfterMessageIDParams struct {
|
||||
ChatID uuid.UUID `db:"chat_id" json:"chat_id"`
|
||||
MessageID int64 `db:"message_id" json:"message_id"`
|
||||
ChatID uuid.UUID `db:"chat_id" json:"chat_id"`
|
||||
StartedBefore time.Time `db:"started_before" json:"started_before"`
|
||||
MessageID int64 `db:"message_id" json:"message_id"`
|
||||
}
|
||||
|
||||
// Deletes debug runs (and their cascaded steps) whose message IDs
|
||||
// exceed the cutoff. The started_before bound prevents retried
|
||||
// cleanup from deleting runs created by a replacement turn that
|
||||
// raced ahead of the retry window.
|
||||
func (q *sqlQuerier) DeleteChatDebugDataAfterMessageID(ctx context.Context, arg DeleteChatDebugDataAfterMessageIDParams) (int64, error) {
|
||||
result, err := q.db.ExecContext(ctx, deleteChatDebugDataAfterMessageID, arg.ChatID, arg.MessageID)
|
||||
result, err := q.db.ExecContext(ctx, deleteChatDebugDataAfterMessageID, arg.ChatID, arg.StartedBefore, arg.MessageID)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
@@ -2941,10 +2950,20 @@ func (q *sqlQuerier) DeleteChatDebugDataAfterMessageID(ctx context.Context, arg
|
||||
const deleteChatDebugDataByChatID = `-- name: DeleteChatDebugDataByChatID :execrows
|
||||
DELETE FROM chat_debug_runs
|
||||
WHERE chat_id = $1::uuid
|
||||
AND started_at < $2::timestamptz
|
||||
`
|
||||
|
||||
func (q *sqlQuerier) DeleteChatDebugDataByChatID(ctx context.Context, chatID uuid.UUID) (int64, error) {
|
||||
result, err := q.db.ExecContext(ctx, deleteChatDebugDataByChatID, chatID)
|
||||
type DeleteChatDebugDataByChatIDParams struct {
|
||||
ChatID uuid.UUID `db:"chat_id" json:"chat_id"`
|
||||
StartedBefore time.Time `db:"started_before" json:"started_before"`
|
||||
}
|
||||
|
||||
// The started_before bound prevents retried cleanup from deleting
|
||||
// runs created by a replacement turn that races ahead of the retry
|
||||
// window (for example, after an unarchive races with a pending
|
||||
// archive-cleanup retry).
|
||||
func (q *sqlQuerier) DeleteChatDebugDataByChatID(ctx context.Context, arg DeleteChatDebugDataByChatIDParams) (int64, error) {
|
||||
result, err := q.db.ExecContext(ctx, deleteChatDebugDataByChatID, arg.ChatID, arg.StartedBefore)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
@@ -206,14 +206,24 @@ WHERE run_id = @run_id::uuid
|
||||
ORDER BY step_number ASC, started_at ASC;
|
||||
|
||||
-- name: DeleteChatDebugDataByChatID :execrows
|
||||
-- The started_before bound prevents retried cleanup from deleting
|
||||
-- runs created by a replacement turn that races ahead of the retry
|
||||
-- window (for example, after an unarchive races with a pending
|
||||
-- archive-cleanup retry).
|
||||
DELETE FROM chat_debug_runs
|
||||
WHERE chat_id = @chat_id::uuid;
|
||||
WHERE chat_id = @chat_id::uuid
|
||||
AND started_at < @started_before::timestamptz;
|
||||
|
||||
-- name: DeleteChatDebugDataAfterMessageID :execrows
|
||||
-- Deletes debug runs (and their cascaded steps) whose message IDs
|
||||
-- exceed the cutoff. The started_before bound prevents retried
|
||||
-- cleanup from deleting runs created by a replacement turn that
|
||||
-- raced ahead of the retry window.
|
||||
WITH affected_runs AS (
|
||||
SELECT DISTINCT run.id
|
||||
FROM chat_debug_runs run
|
||||
WHERE run.chat_id = @chat_id::uuid
|
||||
AND run.started_at < @started_before::timestamptz
|
||||
AND (
|
||||
run.history_tip_message_id > @message_id::bigint
|
||||
OR run.trigger_message_id > @message_id::bigint
|
||||
@@ -223,7 +233,10 @@ WITH affected_runs AS (
|
||||
|
||||
SELECT DISTINCT step.run_id AS id
|
||||
FROM chat_debug_steps step
|
||||
JOIN chat_debug_runs run ON run.id = step.run_id
|
||||
AND run.chat_id = step.chat_id
|
||||
WHERE step.chat_id = @chat_id::uuid
|
||||
AND run.started_at < @started_before::timestamptz
|
||||
AND (
|
||||
step.assistant_message_id > @message_id::bigint
|
||||
OR step.history_tip_message_id > @message_id::bigint
|
||||
|
||||
+667
-126
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,162 @@
|
||||
package chatd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog/v3"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatdebug"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
|
||||
)
|
||||
|
||||
const (
|
||||
debugCleanupRetryDelay = 500 * time.Millisecond
|
||||
debugCleanupAttempts = 3
|
||||
debugCleanupTimeout = 5 * time.Second
|
||||
// debugCreateRunTimeout caps how long a CreateRun insert can
|
||||
// block the caller's critical path. Debug persistence is
|
||||
// best-effort, so the turn proceeds without debug rows if the
|
||||
// DB is slow or locked. Matches the manual-title budget.
|
||||
debugCreateRunTimeout = 5 * time.Second
|
||||
// debugCleanupClockSkew gives cleanup cutoffs tolerance for cross-
|
||||
// replica clock drift. The cutoff is sampled from the DB
|
||||
// (updated_at returned by the status transition), and
|
||||
// chat_debug_runs.started_at is stamped by whatever replica
|
||||
// processes the replacement turn. If that replica's clock lags
|
||||
// the DB, its started_at can land behind a commit-time cutoff
|
||||
// even though the insert physically happened after commit.
|
||||
// Subtracting this buffer ensures the fast retry path cannot
|
||||
// delete replacement rows when clocks drift by up to this
|
||||
// amount; rows within the buffer survive the fast cleanup but
|
||||
// are still finalized (and eligible for stale-sweep cleanup) by
|
||||
// the existing FinalizeStale background loop.
|
||||
debugCleanupClockSkew = 30 * time.Second
|
||||
)
|
||||
|
||||
func (p *Server) debugService() *chatdebug.Service {
|
||||
if p == nil {
|
||||
return nil
|
||||
}
|
||||
if p.debugSvcFactory == nil {
|
||||
return p.debugSvc
|
||||
}
|
||||
p.debugSvcInit.Do(func() {
|
||||
p.debugSvc = p.debugSvcFactory()
|
||||
p.debugSvcReady.Store(p.debugSvc != nil)
|
||||
})
|
||||
return p.debugSvc
|
||||
}
|
||||
|
||||
func (p *Server) existingDebugService() *chatdebug.Service {
|
||||
if p == nil {
|
||||
return nil
|
||||
}
|
||||
if p.debugSvcFactory == nil {
|
||||
return p.debugSvc
|
||||
}
|
||||
if !p.debugSvcReady.Load() {
|
||||
return nil
|
||||
}
|
||||
return p.debugSvc
|
||||
}
|
||||
|
||||
func (p *Server) scheduleDebugCleanup(
|
||||
ctx context.Context,
|
||||
logMessage string,
|
||||
fields []slog.Field,
|
||||
cleanup func(context.Context, *chatdebug.Service) error,
|
||||
) {
|
||||
debugSvc := p.debugService()
|
||||
if debugSvc == nil {
|
||||
return
|
||||
}
|
||||
|
||||
// Acquire inflightMu around the positive Add so Close() cannot
|
||||
// call drainInflight concurrently when the counter is at zero.
|
||||
// See drainInflight for the WaitGroup contract this preserves.
|
||||
p.inflightMu.Lock()
|
||||
p.inflight.Add(1)
|
||||
p.inflightMu.Unlock()
|
||||
go func() {
|
||||
defer p.inflight.Done()
|
||||
|
||||
cleanupCtx := context.WithoutCancel(ctx)
|
||||
for attempt := 0; attempt < debugCleanupAttempts; attempt++ {
|
||||
if attempt > 0 {
|
||||
timer := p.clock.NewTimer(debugCleanupRetryDelay, "chatd", "debug_cleanup")
|
||||
<-timer.C
|
||||
}
|
||||
|
||||
passCtx, cancel := context.WithTimeout(cleanupCtx, debugCleanupTimeout)
|
||||
err := cleanup(passCtx, debugSvc)
|
||||
cancel()
|
||||
if err == nil {
|
||||
return
|
||||
}
|
||||
|
||||
logFields := append([]slog.Field{
|
||||
slog.F("attempt", attempt+1),
|
||||
slog.F("max_attempts", debugCleanupAttempts),
|
||||
}, fields...)
|
||||
logFields = append(logFields, slog.Error(err))
|
||||
p.logger.Warn(cleanupCtx, logMessage, logFields...)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
func (p *Server) newDebugAwareModelFromConfig(
|
||||
ctx context.Context,
|
||||
chat database.Chat,
|
||||
providerHint string,
|
||||
modelName string,
|
||||
providerKeys chatprovider.ProviderAPIKeys,
|
||||
userAgent string,
|
||||
extraHeaders map[string]string,
|
||||
) (fantasy.LanguageModel, bool, error) {
|
||||
provider, resolvedModel, err := chatprovider.ResolveModelWithProviderHint(modelName, providerHint)
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
|
||||
debugSvc := p.debugService()
|
||||
debugEnabled := debugSvc != nil && debugSvc.IsEnabled(ctx, chat.ID, chat.OwnerID)
|
||||
|
||||
var httpClient *http.Client
|
||||
if debugEnabled {
|
||||
httpClient = &http.Client{Transport: &chatdebug.RecordingTransport{}}
|
||||
}
|
||||
|
||||
model, err := chatprovider.ModelFromConfig(
|
||||
provider,
|
||||
resolvedModel,
|
||||
providerKeys,
|
||||
userAgent,
|
||||
extraHeaders,
|
||||
httpClient,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, debugEnabled, err
|
||||
}
|
||||
if model == nil {
|
||||
return nil, debugEnabled, xerrors.Errorf(
|
||||
"create model for %s/%s returned nil",
|
||||
provider,
|
||||
resolvedModel,
|
||||
)
|
||||
}
|
||||
if !debugEnabled {
|
||||
return model, false, nil
|
||||
}
|
||||
|
||||
return chatdebug.WrapModel(model, debugSvc, chatdebug.RecorderOptions{
|
||||
ChatID: chat.ID,
|
||||
OwnerID: chat.OwnerID,
|
||||
Provider: provider,
|
||||
Model: resolvedModel,
|
||||
}), true, nil
|
||||
}
|
||||
@@ -279,6 +279,14 @@ func TestStopAfterBehaviorTools(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
// TestWaitForActiveChatStop and TestWaitForActiveChatStop_WaitsForReplacementRun
|
||||
// were removed along with the process-local activeChats mechanism.
|
||||
// Debug cleanup is now best-effort; stale finalization handles orphaned rows.
|
||||
|
||||
// TestArchiveChatWaitsForActiveChatStop and
|
||||
// TestArchiveChatWaitsForEveryInterruptedChat were removed along with
|
||||
// the process-local activeChats mechanism. Archive cleanup is now
|
||||
// best-effort; stale finalization handles any orphaned rows.
|
||||
func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -2889,6 +2897,10 @@ func TestProcessChat_IgnoresStaleControlNotification(t *testing.T) {
|
||||
return database.Chat{ID: chatID, Status: params.Status}, nil
|
||||
},
|
||||
)
|
||||
db.EXPECT().GetChatByID(gomock.Any(), chatID).Return(
|
||||
database.Chat{ID: chatID, Status: database.ChatStatusError},
|
||||
nil,
|
||||
)
|
||||
|
||||
// resolveChatModel fails immediately — that's fine, we only
|
||||
// need processChat to get past initialization without being
|
||||
@@ -2920,6 +2932,69 @@ func TestProcessChat_IgnoresStaleControlNotification(t *testing.T) {
|
||||
"processChat should have reached runChat (error), not been interrupted (waiting)")
|
||||
}
|
||||
|
||||
func TestShouldPublishFinishedChatState(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
chatID := uuid.New()
|
||||
workerID := uuid.New()
|
||||
|
||||
server := &Server{db: db}
|
||||
updatedChat := database.Chat{
|
||||
ID: chatID,
|
||||
Status: database.ChatStatusWaiting,
|
||||
WorkerID: uuid.NullUUID{},
|
||||
}
|
||||
|
||||
db.EXPECT().GetChatByID(gomock.Any(), chatID).Return(database.Chat{
|
||||
ID: chatID,
|
||||
Status: database.ChatStatusWaiting,
|
||||
WorkerID: uuid.NullUUID{},
|
||||
}, nil)
|
||||
|
||||
require.True(t, server.shouldPublishFinishedChatState(ctx, logger, updatedChat))
|
||||
|
||||
db.EXPECT().GetChatByID(gomock.Any(), chatID).Return(database.Chat{
|
||||
ID: chatID,
|
||||
Status: database.ChatStatusRunning,
|
||||
WorkerID: uuid.NullUUID{UUID: workerID, Valid: true},
|
||||
}, nil)
|
||||
|
||||
require.False(t, server.shouldPublishFinishedChatState(ctx, logger, updatedChat))
|
||||
}
|
||||
|
||||
// TestShouldPublishFinishedChatState_DBErrorPublishes pins the
|
||||
// deliberate fail-open behavior when the re-read query errors: we
|
||||
// surface the finished state anyway so watchers don't get stuck
|
||||
// waiting for a status update that never arrives. The error path is
|
||||
// easy to regress into a fail-closed default otherwise.
|
||||
func TestShouldPublishFinishedChatState_DBErrorPublishes(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
chatID := uuid.New()
|
||||
|
||||
server := &Server{db: db}
|
||||
updatedChat := database.Chat{
|
||||
ID: chatID,
|
||||
Status: database.ChatStatusWaiting,
|
||||
WorkerID: uuid.NullUUID{},
|
||||
}
|
||||
|
||||
db.EXPECT().GetChatByID(gomock.Any(), chatID).Return(
|
||||
database.Chat{}, xerrors.New("boom"),
|
||||
)
|
||||
|
||||
require.True(t, server.shouldPublishFinishedChatState(ctx, logger, updatedChat),
|
||||
"fail-open: a re-read error must not swallow the status change")
|
||||
}
|
||||
|
||||
// TestHeartbeatTick_StolenChatIsInterrupted verifies that when the
|
||||
// batch heartbeat UPDATE does not return a registered chat's ID
|
||||
// (because another replica stole it or it was completed), the
|
||||
|
||||
@@ -2052,6 +2052,279 @@ func TestEditMessageRejectsNonUserMessage(t *testing.T) {
|
||||
require.True(t, errors.Is(err, chatd.ErrEditedMessageNotUser))
|
||||
}
|
||||
|
||||
// TestEditMessageDebugCleanupDeletesPreEditRuns verifies that
|
||||
// EditMessage schedules the chat debug cleanup goroutine when debug
|
||||
// logging is enabled and that it deletes debug runs tied to the
|
||||
// pre-edit conversation branch. This exercises the chatd wiring end
|
||||
// to end: lazy debugService init, editCutoff sampling from the DB,
|
||||
// and the scheduleDebugCleanup retry loop against a real Postgres
|
||||
// store.
|
||||
func TestEditMessageDebugCleanupDeletesPreEditRuns(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
replica := newDebugEnabledTestServer(t, db, ps, uuid.New())
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(ctx, t, db)
|
||||
|
||||
chat, err := replica.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "debug-edit-cleanup",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("first")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
msgs, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chat.ID, AfterID: 0,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, msgs, 1)
|
||||
editedMsgID := msgs[0].ID
|
||||
|
||||
// Stale debug run tied to the pre-edit message branch. Stamped
|
||||
// well outside the clock-skew buffer so the fast retry path
|
||||
// deletes it instead of deferring to the stale sweeper.
|
||||
staleStart := time.Now().Add(-time.Hour).UTC().Truncate(time.Microsecond)
|
||||
staleRun, err := db.InsertChatDebugRun(ctx, database.InsertChatDebugRunParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
TriggerMessageID: sql.NullInt64{Int64: editedMsgID, Valid: true},
|
||||
HistoryTipMessageID: sql.NullInt64{Int64: editedMsgID, Valid: true},
|
||||
Kind: "chat_turn",
|
||||
Status: "in_progress",
|
||||
Provider: sql.NullString{String: "openai", Valid: true},
|
||||
Model: sql.NullString{String: model.Model, Valid: true},
|
||||
StartedAt: sql.NullTime{Time: staleStart, Valid: true},
|
||||
UpdatedAt: sql.NullTime{Time: staleStart, Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Run tied to an earlier message branch that the message-id
|
||||
// filter should leave alone even though it predates the edit.
|
||||
unrelatedRun, err := db.InsertChatDebugRun(ctx, database.InsertChatDebugRunParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
TriggerMessageID: sql.NullInt64{Int64: editedMsgID - 1, Valid: true},
|
||||
HistoryTipMessageID: sql.NullInt64{Int64: editedMsgID - 1, Valid: true},
|
||||
Kind: "chat_turn",
|
||||
Status: "completed",
|
||||
Provider: sql.NullString{String: "openai", Valid: true},
|
||||
Model: sql.NullString{String: model.Model, Valid: true},
|
||||
StartedAt: sql.NullTime{Time: staleStart, Valid: true},
|
||||
UpdatedAt: sql.NullTime{Time: staleStart, Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = replica.EditMessage(ctx, chatd.EditMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
EditedMessageID: editedMsgID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
chatd.WaitUntilIdleForTest(replica)
|
||||
|
||||
// ErrNoRows on staleRun proves the fast-retry path DELETED the
|
||||
// row: FinalizeStale (the only other debug-row writer on the
|
||||
// server) only UPDATEs finished_at in place, it never deletes,
|
||||
// so the row can only disappear via DeleteAfterMessageID which
|
||||
// is reached solely from scheduleDebugCleanup.
|
||||
_, err = db.GetChatDebugRunByID(ctx, staleRun.ID)
|
||||
require.ErrorIs(t, err, sql.ErrNoRows,
|
||||
"pre-edit run matching the message-id filter should be deleted")
|
||||
|
||||
remaining, err := db.GetChatDebugRunByID(ctx, unrelatedRun.ID)
|
||||
require.NoError(t, err,
|
||||
"runs outside the edited message branch must survive cleanup")
|
||||
require.Equal(t, unrelatedRun.ID, remaining.ID)
|
||||
|
||||
// Count the seeded rows that survive so the delete count is
|
||||
// verified directly (not just by negative lookup). Scoped to
|
||||
// seeded IDs because the processor may start a new chat_turn
|
||||
// run in parallel when EditMessage transitions the chat back to
|
||||
// pending.
|
||||
remainingRuns, err := db.GetChatDebugRunsByChatID(ctx, database.GetChatDebugRunsByChatIDParams{
|
||||
ChatID: chat.ID, LimitVal: 100,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
seeded := map[uuid.UUID]bool{staleRun.ID: true, unrelatedRun.ID: true}
|
||||
survivors := 0
|
||||
for _, r := range remainingRuns {
|
||||
if seeded[r.ID] {
|
||||
survivors++
|
||||
}
|
||||
}
|
||||
require.Equal(t, 1, survivors,
|
||||
"exactly one of the two seeded runs should survive (the unrelated run)")
|
||||
}
|
||||
|
||||
// TestEditMessageDebugCleanupPreservesRecentRuns verifies that the
|
||||
// clock-skew buffer in the edit-cleanup cutoff prevents the fast
|
||||
// retry from deleting debug runs that started within the buffer
|
||||
// window. The stale sweep handles those leftovers later.
|
||||
func TestEditMessageDebugCleanupPreservesRecentRuns(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
replica := newDebugEnabledTestServer(t, db, ps, uuid.New())
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(ctx, t, db)
|
||||
|
||||
chat, err := replica.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "debug-edit-buffer",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("first")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
msgs, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chat.ID, AfterID: 0,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, msgs, 1)
|
||||
editedMsgID := msgs[0].ID
|
||||
|
||||
// Within the 30s skew buffer, so the fast retry must leave it
|
||||
// alone even though its message ID matches the delete filter.
|
||||
recentStart := time.Now().Add(-time.Second).UTC().Truncate(time.Microsecond)
|
||||
recentRun, err := db.InsertChatDebugRun(ctx, database.InsertChatDebugRunParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
TriggerMessageID: sql.NullInt64{Int64: editedMsgID, Valid: true},
|
||||
HistoryTipMessageID: sql.NullInt64{Int64: editedMsgID, Valid: true},
|
||||
Kind: "chat_turn",
|
||||
Status: "in_progress",
|
||||
Provider: sql.NullString{String: "openai", Valid: true},
|
||||
Model: sql.NullString{String: model.Model, Valid: true},
|
||||
StartedAt: sql.NullTime{Time: recentStart, Valid: true},
|
||||
UpdatedAt: sql.NullTime{Time: recentStart, Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = replica.EditMessage(ctx, chatd.EditMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
EditedMessageID: editedMsgID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
chatd.WaitUntilIdleForTest(replica)
|
||||
|
||||
remaining, err := db.GetChatDebugRunByID(ctx, recentRun.ID)
|
||||
require.NoError(t, err,
|
||||
"runs inside the clock-skew buffer must survive the fast retry")
|
||||
require.Equal(t, recentRun.ID, remaining.ID)
|
||||
|
||||
// If the clock-skew buffer were removed the fast retry would
|
||||
// have deleted recentRun. Verify the count of seeded survivors
|
||||
// directly, ignoring any new chat_turn run the processor may
|
||||
// create after the pending status transition.
|
||||
remainingRuns, err := db.GetChatDebugRunsByChatID(ctx, database.GetChatDebugRunsByChatIDParams{
|
||||
ChatID: chat.ID, LimitVal: 100,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
survivors := 0
|
||||
for _, r := range remainingRuns {
|
||||
if r.ID == recentRun.ID {
|
||||
survivors++
|
||||
}
|
||||
}
|
||||
require.Equal(t, 1, survivors,
|
||||
"the buffered run must survive the fast retry")
|
||||
}
|
||||
|
||||
// TestArchiveChatDebugCleanupDeletesPreArchiveRuns verifies that
|
||||
// ArchiveChat schedules cleanup that deletes pre-archive debug runs
|
||||
// for the archived chat. Covers the archiveCutoff sampled from
|
||||
// ArchiveChatByID's DB-stamped updated_at and the DeleteByChatID
|
||||
// delete path.
|
||||
func TestArchiveChatDebugCleanupDeletesPreArchiveRuns(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
replica := newDebugEnabledTestServer(t, db, ps, uuid.New())
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(ctx, t, db)
|
||||
|
||||
chat, err := replica.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "debug-archive-cleanup",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
staleStart := time.Now().Add(-time.Hour).UTC().Truncate(time.Microsecond)
|
||||
staleRun, err := db.InsertChatDebugRun(ctx, database.InsertChatDebugRunParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
Kind: "chat_turn",
|
||||
Status: "in_progress",
|
||||
Provider: sql.NullString{String: "openai", Valid: true},
|
||||
Model: sql.NullString{String: model.Model, Valid: true},
|
||||
StartedAt: sql.NullTime{Time: staleStart, Valid: true},
|
||||
UpdatedAt: sql.NullTime{Time: staleStart, Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Freshly-inserted run inside the skew buffer must survive the
|
||||
// fast retry for the same reason as the edit-cleanup buffer test.
|
||||
recentStart := time.Now().Add(-time.Second).UTC().Truncate(time.Microsecond)
|
||||
recentRun, err := db.InsertChatDebugRun(ctx, database.InsertChatDebugRunParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
Kind: "chat_turn",
|
||||
Status: "in_progress",
|
||||
Provider: sql.NullString{String: "openai", Valid: true},
|
||||
Model: sql.NullString{String: model.Model, Valid: true},
|
||||
StartedAt: sql.NullTime{Time: recentStart, Valid: true},
|
||||
UpdatedAt: sql.NullTime{Time: recentStart, Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = replica.ArchiveChat(ctx, chat)
|
||||
require.NoError(t, err)
|
||||
|
||||
chatd.WaitUntilIdleForTest(replica)
|
||||
|
||||
// ErrNoRows proves the fast-retry path DELETED the row:
|
||||
// FinalizeStale only UPDATEs in place, never deletes.
|
||||
_, err = db.GetChatDebugRunByID(ctx, staleRun.ID)
|
||||
require.ErrorIs(t, err, sql.ErrNoRows,
|
||||
"pre-archive run outside the buffer should be deleted")
|
||||
|
||||
remaining, err := db.GetChatDebugRunByID(ctx, recentRun.ID)
|
||||
require.NoError(t, err,
|
||||
"runs inside the clock-skew buffer must survive the fast retry")
|
||||
require.Equal(t, recentRun.ID, remaining.ID)
|
||||
|
||||
// Count the seeded survivors directly so the delete is verified
|
||||
// not just by absence of a specific row. Scoped to seeded IDs
|
||||
// because the archive transition may still race with other
|
||||
// background debug writes.
|
||||
remainingRuns, err := db.GetChatDebugRunsByChatID(ctx, database.GetChatDebugRunsByChatIDParams{
|
||||
ChatID: chat.ID, LimitVal: 100,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
seeded := map[uuid.UUID]bool{staleRun.ID: true, recentRun.ID: true}
|
||||
survivors := 0
|
||||
for _, r := range remainingRuns {
|
||||
if seeded[r.ID] {
|
||||
survivors++
|
||||
}
|
||||
}
|
||||
require.Equal(t, 1, survivors,
|
||||
"only the recent (buffered) seeded run should survive")
|
||||
}
|
||||
|
||||
func TestRecoverStaleChatsPeriodically(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -4138,6 +4411,34 @@ func newTestServer(
|
||||
return server
|
||||
}
|
||||
|
||||
// newDebugEnabledTestServer creates a passive test server with
|
||||
// AlwaysEnableDebugLogs=true so that IsEnabled(ctx, chatID, ownerID)
|
||||
// always returns true regardless of runtime admin config. This lets
|
||||
// chatd-level integration tests exercise the debug cleanup wiring
|
||||
// without seeding the admin/user opt-in settings tables.
|
||||
func newDebugEnabledTestServer(
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
ps dbpubsub.Pubsub,
|
||||
replicaID uuid.UUID,
|
||||
) *chatd.Server {
|
||||
t.Helper()
|
||||
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
server := chatd.New(chatd.Config{
|
||||
Logger: logger,
|
||||
Database: db,
|
||||
ReplicaID: replicaID,
|
||||
Pubsub: ps,
|
||||
PendingChatAcquireInterval: testutil.WaitLong,
|
||||
AlwaysEnableDebugLogs: true,
|
||||
})
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, server.Close())
|
||||
})
|
||||
return server
|
||||
}
|
||||
|
||||
// newActiveTestServer creates a chatd server that actively polls for
|
||||
// and processes pending chats. Use this instead of newTestServer when
|
||||
// the test needs the chat loop to actually run. Optional config
|
||||
|
||||
@@ -426,7 +426,7 @@ func (s *Service) CreateStep(
|
||||
}
|
||||
|
||||
return database.ChatDebugStep{}, xerrors.Errorf(
|
||||
"failed to create debug step after %d attempts (run_id=%s)",
|
||||
"chatdebug: failed to create step after %d retries (run %s)",
|
||||
maxCreateStepRetries, params.RunID,
|
||||
)
|
||||
}
|
||||
@@ -522,12 +522,24 @@ func (s *Service) TouchStep(
|
||||
})
|
||||
}
|
||||
|
||||
// DeleteByChatID deletes all debug data for a chat and emits a delete event.
|
||||
// DeleteByChatID deletes debug data for a chat and emits a delete event.
|
||||
// The startedBefore bound scopes deletion to runs created before that
|
||||
// instant so that retried cleanup does not remove runs created by a
|
||||
// replacement turn that raced ahead of the retry window (for example,
|
||||
// an unarchive that fires between the initial archive-cleanup attempt
|
||||
// and its retry).
|
||||
func (s *Service) DeleteByChatID(
|
||||
ctx context.Context,
|
||||
chatID uuid.UUID,
|
||||
startedBefore time.Time,
|
||||
) (int64, error) {
|
||||
deleted, err := s.db.DeleteChatDebugDataByChatID(chatdContext(ctx), chatID)
|
||||
deleted, err := s.db.DeleteChatDebugDataByChatID(
|
||||
chatdContext(ctx),
|
||||
database.DeleteChatDebugDataByChatIDParams{
|
||||
ChatID: chatID,
|
||||
StartedBefore: startedBefore,
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
@@ -537,16 +549,21 @@ func (s *Service) DeleteByChatID(
|
||||
}
|
||||
|
||||
// DeleteAfterMessageID deletes debug data newer than the given message.
|
||||
// The startedBefore bound scopes deletion to runs created before that
|
||||
// instant so that retried cleanup does not remove runs created by a
|
||||
// replacement turn that raced ahead of the retry window.
|
||||
func (s *Service) DeleteAfterMessageID(
|
||||
ctx context.Context,
|
||||
chatID uuid.UUID,
|
||||
messageID int64,
|
||||
startedBefore time.Time,
|
||||
) (int64, error) {
|
||||
deleted, err := s.db.DeleteChatDebugDataAfterMessageID(
|
||||
chatdContext(ctx),
|
||||
database.DeleteChatDebugDataAfterMessageIDParams{
|
||||
ChatID: chatID,
|
||||
MessageID: messageID,
|
||||
ChatID: chatID,
|
||||
MessageID: messageID,
|
||||
StartedBefore: startedBefore,
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
@@ -579,6 +596,79 @@ func (s *Service) FinalizeStale(
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// FinalizeRunParams bundles the arguments for FinalizeRun.
|
||||
type FinalizeRunParams struct {
|
||||
RunID uuid.UUID
|
||||
ChatID uuid.UUID
|
||||
Status Status
|
||||
SeedSummary map[string]any
|
||||
// Timeout for the aggregate + update calls. Zero defaults to 5s.
|
||||
Timeout time.Duration
|
||||
}
|
||||
|
||||
// FinalizeRun aggregates the run summary, updates the run status, and
|
||||
// cleans up the step counter. It detaches from the parent context's
|
||||
// cancellation so finalization succeeds even when the request context
|
||||
// is already done. Errors are returned but are always safe to ignore;
|
||||
// callers that treat debug instrumentation as best-effort can discard
|
||||
// them.
|
||||
func (s *Service) FinalizeRun(ctx context.Context, p FinalizeRunParams) error {
|
||||
timeout := p.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 5 * time.Second
|
||||
}
|
||||
|
||||
finalizeCtx, cancel := context.WithTimeout(
|
||||
context.WithoutCancel(ctx), timeout,
|
||||
)
|
||||
defer cancel()
|
||||
|
||||
finalSummary := p.SeedSummary
|
||||
if aggregated, aggErr := s.AggregateRunSummary(
|
||||
finalizeCtx,
|
||||
p.RunID,
|
||||
p.SeedSummary,
|
||||
); aggErr != nil {
|
||||
// Non-fatal: proceed with the seed summary.
|
||||
s.log.Warn(ctx, "failed to aggregate debug run summary",
|
||||
slog.F("chat_id", p.ChatID),
|
||||
slog.F("run_id", p.RunID),
|
||||
slog.Error(aggErr),
|
||||
)
|
||||
} else {
|
||||
finalSummary = aggregated
|
||||
}
|
||||
|
||||
if _, err := s.UpdateRun(finalizeCtx, UpdateRunParams{
|
||||
ID: p.RunID,
|
||||
ChatID: p.ChatID,
|
||||
Status: p.Status,
|
||||
Summary: finalSummary,
|
||||
FinishedAt: s.clock.Now(),
|
||||
}); err != nil {
|
||||
CleanupStepCounter(p.RunID)
|
||||
return xerrors.Errorf("update debug run: %w", err)
|
||||
}
|
||||
CleanupStepCounter(p.RunID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ClassifyError maps a run error to the appropriate debug status.
|
||||
// nil → StatusCompleted, context.Canceled → StatusInterrupted,
|
||||
// everything else → StatusError. Callers with additional
|
||||
// classification rules (e.g. ErrInterrupted, ErrDynamicToolCall)
|
||||
// should handle those before falling back to this helper.
|
||||
func ClassifyError(err error) Status {
|
||||
switch {
|
||||
case err == nil:
|
||||
return StatusCompleted
|
||||
case errors.Is(err, context.Canceled):
|
||||
return StatusInterrupted
|
||||
default:
|
||||
return StatusError
|
||||
}
|
||||
}
|
||||
|
||||
func nullUUID(id uuid.UUID) uuid.NullUUID {
|
||||
return uuid.NullUUID{UUID: id, Valid: id != uuid.Nil}
|
||||
}
|
||||
|
||||
@@ -581,7 +581,8 @@ func TestService_DeleteByChatID(t *testing.T) {
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
deleted, err := fixture.svc.DeleteByChatID(fixture.ctx, fixture.chat.ID)
|
||||
deleted, err := fixture.svc.DeleteByChatID(fixture.ctx, fixture.chat.ID,
|
||||
time.Now().Add(time.Minute))
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, deleted)
|
||||
|
||||
@@ -640,7 +641,7 @@ func TestService_DeleteAfterMessageID(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
deleted, err := fixture.svc.DeleteAfterMessageID(fixture.ctx, fixture.chat.ID,
|
||||
threshold.ID)
|
||||
threshold.ID, time.Now().Add(time.Minute))
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, deleted)
|
||||
|
||||
@@ -826,6 +827,192 @@ func TestService_FinalizeStale_NoChangesDoesNotBroadcast(t *testing.T) {
|
||||
_ = chat // keep seeded chat usage explicit for test readability.
|
||||
}
|
||||
|
||||
func TestClassifyError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
err error
|
||||
want chatdebug.Status
|
||||
}{
|
||||
{"nil", nil, chatdebug.StatusCompleted},
|
||||
{"context.Canceled", context.Canceled, chatdebug.StatusInterrupted},
|
||||
// Wrapped context.Canceled must still classify as interrupted so
|
||||
// callers that decorate cancellation errors do not flip to
|
||||
// StatusError.
|
||||
{
|
||||
"wrapped context.Canceled",
|
||||
xerrors.Errorf("canceled mid-stream: %w", context.Canceled),
|
||||
chatdebug.StatusInterrupted,
|
||||
},
|
||||
{"generic error", xerrors.New("boom"), chatdebug.StatusError},
|
||||
// context.DeadlineExceeded is not context.Canceled and is not
|
||||
// special-cased by ClassifyError, so it must fall through to
|
||||
// StatusError. This pins the priority ordering in the switch.
|
||||
{
|
||||
"context.DeadlineExceeded",
|
||||
context.DeadlineExceeded, chatdebug.StatusError,
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
require.Equal(t, tt.want, chatdebug.ClassifyError(tt.err))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_FinalizeRun_FallsBackToSeedSummary(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
svc := chatdebug.NewService(db, testutil.Logger(t), nil)
|
||||
|
||||
runID := uuid.New()
|
||||
chatID := uuid.New()
|
||||
seed := map[string]any{"first_message": "hello"}
|
||||
|
||||
// Force AggregateRunSummary to fail by returning an error from the
|
||||
// step fetch it depends on. FinalizeRun must log the warning and
|
||||
// continue with the caller-supplied SeedSummary.
|
||||
db.EXPECT().
|
||||
GetChatDebugStepsByRunID(gomock.Any(), runID).
|
||||
Return(nil, xerrors.New("boom"))
|
||||
|
||||
db.EXPECT().
|
||||
UpdateChatDebugRun(gomock.Any(), gomock.Any()).
|
||||
DoAndReturn(func(_ context.Context, arg database.UpdateChatDebugRunParams) (database.ChatDebugRun, error) {
|
||||
require.Equal(t, runID, arg.ID)
|
||||
require.Equal(t, chatID, arg.ChatID)
|
||||
require.True(t, arg.Summary.Valid)
|
||||
var got map[string]any
|
||||
require.NoError(t, json.Unmarshal(arg.Summary.RawMessage, &got))
|
||||
require.Equal(t, "hello", got["first_message"])
|
||||
return database.ChatDebugRun{
|
||||
ID: runID,
|
||||
ChatID: chatID,
|
||||
}, nil
|
||||
})
|
||||
|
||||
err := svc.FinalizeRun(context.Background(), chatdebug.FinalizeRunParams{
|
||||
RunID: runID,
|
||||
ChatID: chatID,
|
||||
Status: chatdebug.StatusCompleted,
|
||||
SeedSummary: seed,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestService_FinalizeRun_ReturnsWrappedUpdateError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
svc := chatdebug.NewService(db, testutil.Logger(t), nil)
|
||||
|
||||
runID := uuid.New()
|
||||
chatID := uuid.New()
|
||||
|
||||
db.EXPECT().
|
||||
GetChatDebugStepsByRunID(gomock.Any(), runID).
|
||||
Return(nil, nil)
|
||||
db.EXPECT().
|
||||
UpdateChatDebugRun(gomock.Any(), gomock.Any()).
|
||||
Return(database.ChatDebugRun{}, xerrors.New("update failed"))
|
||||
|
||||
err := svc.FinalizeRun(context.Background(), chatdebug.FinalizeRunParams{
|
||||
RunID: runID,
|
||||
ChatID: chatID,
|
||||
Status: chatdebug.StatusCompleted,
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "update debug run")
|
||||
require.Contains(t, err.Error(), "update failed")
|
||||
}
|
||||
|
||||
func TestService_FinalizeRun_CustomTimeoutAppliesToDBCalls(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
svc := chatdebug.NewService(db, testutil.Logger(t), nil)
|
||||
|
||||
runID := uuid.New()
|
||||
chatID := uuid.New()
|
||||
customTimeout := 123 * time.Millisecond
|
||||
// Allow for scheduling jitter but ensure the custom timeout is
|
||||
// honored rather than the 5s default. Both DB calls receive the
|
||||
// same timeout-bounded context.
|
||||
maxRemaining := customTimeout + 50*time.Millisecond
|
||||
|
||||
db.EXPECT().
|
||||
GetChatDebugStepsByRunID(gomock.Any(), runID).
|
||||
DoAndReturn(func(ctx context.Context, _ uuid.UUID) ([]database.ChatDebugStep, error) {
|
||||
deadline, ok := ctx.Deadline()
|
||||
require.True(t, ok, "FinalizeRun must apply its Timeout to aggregation context")
|
||||
require.LessOrEqual(t, time.Until(deadline), maxRemaining)
|
||||
return nil, nil
|
||||
})
|
||||
db.EXPECT().
|
||||
UpdateChatDebugRun(gomock.Any(), gomock.Any()).
|
||||
DoAndReturn(func(ctx context.Context, _ database.UpdateChatDebugRunParams) (database.ChatDebugRun, error) {
|
||||
deadline, ok := ctx.Deadline()
|
||||
require.True(t, ok, "FinalizeRun must apply its Timeout to update context")
|
||||
require.LessOrEqual(t, time.Until(deadline), maxRemaining)
|
||||
return database.ChatDebugRun{ID: runID, ChatID: chatID}, nil
|
||||
})
|
||||
|
||||
err := svc.FinalizeRun(context.Background(), chatdebug.FinalizeRunParams{
|
||||
RunID: runID,
|
||||
ChatID: chatID,
|
||||
Status: chatdebug.StatusCompleted,
|
||||
Timeout: customTimeout,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestService_FinalizeRun_DetachesFromParentCancellation(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
svc := chatdebug.NewService(db, testutil.Logger(t), nil)
|
||||
|
||||
runID := uuid.New()
|
||||
chatID := uuid.New()
|
||||
|
||||
// FinalizeRun uses context.WithoutCancel so a canceled parent must
|
||||
// not propagate to the DB calls. Verify both calls see a live
|
||||
// context with the FinalizeRun-owned deadline.
|
||||
parentCtx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
db.EXPECT().
|
||||
GetChatDebugStepsByRunID(gomock.Any(), runID).
|
||||
DoAndReturn(func(ctx context.Context, _ uuid.UUID) ([]database.ChatDebugStep, error) {
|
||||
require.NoError(t, ctx.Err(),
|
||||
"aggregation context must not inherit parent cancellation")
|
||||
_, ok := ctx.Deadline()
|
||||
require.True(t, ok)
|
||||
return nil, nil
|
||||
})
|
||||
db.EXPECT().
|
||||
UpdateChatDebugRun(gomock.Any(), gomock.Any()).
|
||||
DoAndReturn(func(ctx context.Context, _ database.UpdateChatDebugRunParams) (database.ChatDebugRun, error) {
|
||||
require.NoError(t, ctx.Err(),
|
||||
"update context must not inherit parent cancellation")
|
||||
return database.ChatDebugRun{ID: runID, ChatID: chatID}, nil
|
||||
})
|
||||
|
||||
err := svc.FinalizeRun(parentCtx, chatdebug.FinalizeRunParams{
|
||||
RunID: runID,
|
||||
ChatID: chatID,
|
||||
Status: chatdebug.StatusCompleted,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestService_PublishesEvents(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -15,6 +15,10 @@ import (
|
||||
stringutil "github.com/coder/coder/v2/coderd/util/strings"
|
||||
)
|
||||
|
||||
// MaxLabelLength is the maximum number of runes kept when building
|
||||
// first_message labels for debug run summaries.
|
||||
const MaxLabelLength = 200
|
||||
|
||||
// whitespaceRun matches one or more consecutive whitespace characters.
|
||||
var whitespaceRun = regexp.MustCompile(`\s+`)
|
||||
|
||||
|
||||
@@ -20,6 +20,7 @@ import (
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database/dbtime"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatdebug"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chaterror"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatretry"
|
||||
@@ -405,7 +406,8 @@ func Run(ctx context.Context, opts RunOptions) error {
|
||||
}
|
||||
|
||||
var result stepResult
|
||||
err := chatretry.Retry(ctx, func(retryCtx context.Context) error {
|
||||
stepCtx := chatdebug.ReuseStep(ctx)
|
||||
err := chatretry.Retry(stepCtx, func(retryCtx context.Context) error {
|
||||
attempt, streamErr := guardedStream(
|
||||
retryCtx,
|
||||
provider,
|
||||
|
||||
@@ -7,8 +7,10 @@ import (
|
||||
"time"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatdebug"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
@@ -17,6 +19,14 @@ const (
|
||||
minCompactionThresholdPercent = int32(0)
|
||||
maxCompactionThresholdPercent = int32(100)
|
||||
|
||||
// compactionDebugCreateRunTimeout caps the compaction debug
|
||||
// CreateRun budget so a slow or locked DB cannot consume the
|
||||
// compaction's configured Timeout and cause model.Generate to
|
||||
// fail with deadline exceeded. Debug instrumentation is
|
||||
// best-effort; running without the debug row is preferable to
|
||||
// failing the compaction.
|
||||
compactionDebugCreateRunTimeout = 5 * time.Second
|
||||
|
||||
defaultCompactionSummaryPrompt = "You are performing a context compaction. " +
|
||||
"Summarize the conversation so a new assistant can seamlessly " +
|
||||
"continue the work in progress.\n\n" +
|
||||
@@ -46,6 +56,9 @@ type CompactionOptions struct {
|
||||
SystemSummaryPrefix string
|
||||
Timeout time.Duration
|
||||
Persist func(context.Context, CompactionResult) error
|
||||
DebugSvc *chatdebug.Service
|
||||
ChatID uuid.UUID
|
||||
HistoryTipMessageID int64
|
||||
|
||||
// ToolCallID and ToolName identify the synthetic tool call
|
||||
// used to represent compaction in the message stream.
|
||||
@@ -269,6 +282,79 @@ func shouldCompact(contextTokens, contextLimit int64, thresholdPercent int32) (f
|
||||
return usagePercent, usagePercent >= float64(thresholdPercent)
|
||||
}
|
||||
|
||||
func startCompactionDebugRun(
|
||||
ctx context.Context,
|
||||
options CompactionOptions,
|
||||
) (context.Context, func(error)) {
|
||||
if options.DebugSvc == nil || options.ChatID == uuid.Nil {
|
||||
return ctx, func(error) {}
|
||||
}
|
||||
|
||||
parentRun, ok := chatdebug.RunFromContext(ctx)
|
||||
if !ok {
|
||||
return ctx, func(error) {}
|
||||
}
|
||||
|
||||
historyTipMessageID := options.HistoryTipMessageID
|
||||
if historyTipMessageID == 0 {
|
||||
historyTipMessageID = parentRun.HistoryTipMessageID
|
||||
}
|
||||
|
||||
// Use a separate short-lived context for the debug insert so a
|
||||
// slow or locked DB cannot consume the compaction timeout budget
|
||||
// and turn debug slowness into a compaction failure via
|
||||
// model.Generate hitting a deadline exceeded. Detached from the
|
||||
// parent so cancellation of the compaction run still lets the
|
||||
// insert reach a terminal state, matching the best-effort
|
||||
// contract of debug instrumentation.
|
||||
createRunCtx, createRunCancel := context.WithTimeout(
|
||||
context.WithoutCancel(ctx), compactionDebugCreateRunTimeout,
|
||||
)
|
||||
run, err := options.DebugSvc.CreateRun(createRunCtx, chatdebug.CreateRunParams{
|
||||
ChatID: options.ChatID,
|
||||
RootChatID: parentRun.RootChatID,
|
||||
ParentChatID: parentRun.ParentChatID,
|
||||
ModelConfigID: parentRun.ModelConfigID,
|
||||
TriggerMessageID: parentRun.TriggerMessageID,
|
||||
HistoryTipMessageID: historyTipMessageID,
|
||||
Kind: chatdebug.KindCompaction,
|
||||
Status: chatdebug.StatusInProgress,
|
||||
Provider: parentRun.Provider,
|
||||
Model: parentRun.Model,
|
||||
})
|
||||
createRunCancel()
|
||||
if err != nil {
|
||||
// Debug instrumentation must not surface as a compaction failure.
|
||||
return ctx, func(error) {}
|
||||
}
|
||||
|
||||
compactionCtx := chatdebug.ContextWithRun(ctx, &chatdebug.RunContext{
|
||||
RunID: run.ID,
|
||||
ChatID: options.ChatID,
|
||||
RootChatID: parentRun.RootChatID,
|
||||
ParentChatID: parentRun.ParentChatID,
|
||||
ModelConfigID: parentRun.ModelConfigID,
|
||||
TriggerMessageID: parentRun.TriggerMessageID,
|
||||
HistoryTipMessageID: historyTipMessageID,
|
||||
Kind: chatdebug.KindCompaction,
|
||||
Provider: parentRun.Provider,
|
||||
Model: parentRun.Model,
|
||||
})
|
||||
|
||||
return compactionCtx, func(runErr error) {
|
||||
status := chatdebug.ClassifyError(runErr)
|
||||
if runErr != nil && xerrors.Is(runErr, ErrInterrupted) {
|
||||
status = chatdebug.StatusInterrupted
|
||||
}
|
||||
// Debug instrumentation must not surface as a compaction failure.
|
||||
_ = options.DebugSvc.FinalizeRun(compactionCtx, chatdebug.FinalizeRunParams{
|
||||
RunID: run.ID,
|
||||
ChatID: options.ChatID,
|
||||
Status: status,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// generateCompactionSummary asks the model to summarize the
|
||||
// conversation so far. The provided messages should contain the
|
||||
// complete history (system prompt, user/assistant turns, tool
|
||||
@@ -279,7 +365,7 @@ func generateCompactionSummary(
|
||||
model fantasy.LanguageModel,
|
||||
messages []fantasy.Message,
|
||||
options CompactionOptions,
|
||||
) (string, error) {
|
||||
) (summary string, err error) {
|
||||
summaryPrompt := make([]fantasy.Message, 0, len(messages)+1)
|
||||
summaryPrompt = append(summaryPrompt, messages...)
|
||||
summaryPrompt = append(summaryPrompt, fantasy.Message{
|
||||
@@ -293,6 +379,22 @@ func generateCompactionSummary(
|
||||
summaryCtx, cancel := context.WithTimeout(ctx, options.Timeout)
|
||||
defer cancel()
|
||||
|
||||
summaryCtx, finishDebugRun := startCompactionDebugRun(summaryCtx, options)
|
||||
defer func() {
|
||||
// If model.Generate (or anything else below) panics, the
|
||||
// named err return is still nil at this point. Without the
|
||||
// recover hook we would finalize the debug run as Completed
|
||||
// in the exact crash path operators rely on to diagnose
|
||||
// failures. Finalize with the panic as an error status and
|
||||
// re-panic so the caller's recovery still observes the
|
||||
// original panic value.
|
||||
if r := recover(); r != nil {
|
||||
finishDebugRun(xerrors.Errorf("panic during compaction summary: %v", r))
|
||||
panic(r)
|
||||
}
|
||||
finishDebugRun(err)
|
||||
}()
|
||||
|
||||
response, err := model.Generate(summaryCtx, fantasy.Call{
|
||||
Prompt: summaryPrompt,
|
||||
ToolChoice: &toolChoice,
|
||||
|
||||
@@ -2,17 +2,240 @@ package chatloop //nolint:testpackage // Uses internal symbols.
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"github.com/google/uuid"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbmock"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatdebug"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
func TestStartCompactionDebugRun_DoesNotReportDebugErrors(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
newParentContext := func(chatID uuid.UUID) context.Context {
|
||||
return chatdebug.ContextWithRun(context.Background(), &chatdebug.RunContext{
|
||||
RunID: uuid.New(),
|
||||
ChatID: chatID,
|
||||
RootChatID: uuid.New(),
|
||||
ParentChatID: uuid.New(),
|
||||
ModelConfigID: uuid.New(),
|
||||
TriggerMessageID: 41,
|
||||
HistoryTipMessageID: 42,
|
||||
Kind: chatdebug.KindChatTurn,
|
||||
Provider: "fake-provider",
|
||||
Model: "fake-model",
|
||||
})
|
||||
}
|
||||
|
||||
t.Run("CreateRun", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
svc := chatdebug.NewService(db, testutil.Logger(t), nil)
|
||||
chatID := uuid.New()
|
||||
reportedErr := make(chan error, 1)
|
||||
|
||||
db.EXPECT().InsertChatDebugRun(
|
||||
gomock.Any(),
|
||||
gomock.AssignableToTypeOf(database.InsertChatDebugRunParams{}),
|
||||
).Return(database.ChatDebugRun{}, xerrors.New("insert compaction debug run"))
|
||||
|
||||
ctx := newParentContext(chatID)
|
||||
compactionCtx, finish := startCompactionDebugRun(ctx, CompactionOptions{
|
||||
DebugSvc: svc,
|
||||
ChatID: chatID,
|
||||
OnError: func(err error) {
|
||||
reportedErr <- err
|
||||
},
|
||||
})
|
||||
require.Same(t, ctx, compactionCtx)
|
||||
finish(nil)
|
||||
select {
|
||||
case err := <-reportedErr:
|
||||
t.Fatalf("unexpected OnError callback: %v", err)
|
||||
default:
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("FinalizeRunAggregatesSummary", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
svc := chatdebug.NewService(db, testutil.Logger(t), nil)
|
||||
chatID := uuid.New()
|
||||
runID := uuid.New()
|
||||
usageJSON, err := json.Marshal(fantasy.Usage{InputTokens: 7, OutputTokens: 3})
|
||||
require.NoError(t, err)
|
||||
attemptsJSON, err := json.Marshal([]chatdebug.Attempt{{
|
||||
Status: "completed",
|
||||
Method: "POST",
|
||||
Path: "/v1/messages",
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
|
||||
db.EXPECT().InsertChatDebugRun(
|
||||
gomock.Any(),
|
||||
gomock.AssignableToTypeOf(database.InsertChatDebugRunParams{}),
|
||||
).Return(database.ChatDebugRun{ //nolint:exhaustruct // Test only needs IDs.
|
||||
ID: runID,
|
||||
ChatID: chatID,
|
||||
}, nil)
|
||||
db.EXPECT().GetChatDebugStepsByRunID(gomock.Any(), runID).Return([]database.ChatDebugStep{{
|
||||
ID: uuid.New(),
|
||||
RunID: runID,
|
||||
ChatID: chatID,
|
||||
Status: string(chatdebug.StatusCompleted),
|
||||
Usage: pqtype.NullRawMessage{RawMessage: usageJSON, Valid: true},
|
||||
Attempts: attemptsJSON,
|
||||
}}, nil)
|
||||
db.EXPECT().UpdateChatDebugRun(
|
||||
gomock.Any(),
|
||||
gomock.AssignableToTypeOf(database.UpdateChatDebugRunParams{}),
|
||||
).DoAndReturn(func(_ context.Context, params database.UpdateChatDebugRunParams) (database.ChatDebugRun, error) {
|
||||
require.Equal(t, chatID, params.ChatID)
|
||||
require.Equal(t, runID, params.ID)
|
||||
require.True(t, params.Summary.Valid)
|
||||
require.JSONEq(t, `{"endpoint_label":"POST /v1/messages","step_count":1,"total_input_tokens":7,"total_output_tokens":3}`,
|
||||
string(params.Summary.RawMessage))
|
||||
return database.ChatDebugRun{ID: runID, ChatID: chatID}, nil
|
||||
})
|
||||
|
||||
ctx := newParentContext(chatID)
|
||||
compactionCtx, finish := startCompactionDebugRun(ctx, CompactionOptions{
|
||||
DebugSvc: svc,
|
||||
ChatID: chatID,
|
||||
})
|
||||
require.NotSame(t, ctx, compactionCtx)
|
||||
finish(nil)
|
||||
})
|
||||
|
||||
t.Run("FinalizeRun", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
svc := chatdebug.NewService(db, testutil.Logger(t), nil)
|
||||
chatID := uuid.New()
|
||||
reportedErr := make(chan error, 1)
|
||||
runID := uuid.New()
|
||||
|
||||
db.EXPECT().InsertChatDebugRun(
|
||||
gomock.Any(),
|
||||
gomock.AssignableToTypeOf(database.InsertChatDebugRunParams{}),
|
||||
).Return(database.ChatDebugRun{ //nolint:exhaustruct // Test only needs IDs.
|
||||
ID: runID,
|
||||
ChatID: chatID,
|
||||
}, nil)
|
||||
db.EXPECT().GetChatDebugStepsByRunID(gomock.Any(), runID).Return(nil, xerrors.New("aggregate compaction debug run"))
|
||||
db.EXPECT().UpdateChatDebugRun(
|
||||
gomock.Any(),
|
||||
gomock.AssignableToTypeOf(database.UpdateChatDebugRunParams{}),
|
||||
).Return(database.ChatDebugRun{}, xerrors.New("finalize compaction debug run"))
|
||||
|
||||
ctx := newParentContext(chatID)
|
||||
compactionCtx, finish := startCompactionDebugRun(ctx, CompactionOptions{
|
||||
DebugSvc: svc,
|
||||
ChatID: chatID,
|
||||
OnError: func(err error) {
|
||||
reportedErr <- err
|
||||
},
|
||||
})
|
||||
require.NotSame(t, ctx, compactionCtx)
|
||||
finish(nil)
|
||||
select {
|
||||
case err := <-reportedErr:
|
||||
t.Fatalf("unexpected OnError callback: %v", err)
|
||||
default:
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestGenerateCompactionSummary_PanicFinalizesAsError verifies that a
|
||||
// panic originating inside the model call during compaction is
|
||||
// captured by the deferred debug-run finalizer so the run is recorded
|
||||
// with StatusError rather than StatusCompleted. Without the recover
|
||||
// hook the named `err` return is still nil when the defer fires and
|
||||
// the row silently misclassifies the crash path.
|
||||
func TestGenerateCompactionSummary_PanicFinalizesAsError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
svc := chatdebug.NewService(db, testutil.Logger(t), nil)
|
||||
chatID := uuid.New()
|
||||
runID := uuid.New()
|
||||
|
||||
status := make(chan string, 1)
|
||||
|
||||
db.EXPECT().InsertChatDebugRun(
|
||||
gomock.Any(),
|
||||
gomock.AssignableToTypeOf(database.InsertChatDebugRunParams{}),
|
||||
).Return(database.ChatDebugRun{
|
||||
ID: runID,
|
||||
ChatID: chatID,
|
||||
}, nil)
|
||||
db.EXPECT().GetChatDebugStepsByRunID(gomock.Any(), runID).Return(nil, nil)
|
||||
db.EXPECT().UpdateChatDebugRun(
|
||||
gomock.Any(),
|
||||
gomock.AssignableToTypeOf(database.UpdateChatDebugRunParams{}),
|
||||
).DoAndReturn(func(_ context.Context, params database.UpdateChatDebugRunParams) (database.ChatDebugRun, error) {
|
||||
status <- params.Status.String
|
||||
return database.ChatDebugRun{ID: runID, ChatID: chatID}, nil
|
||||
})
|
||||
|
||||
model := &chattest.FakeModel{
|
||||
ProviderName: "fake",
|
||||
GenerateFn: func(_ context.Context, _ fantasy.Call) (*fantasy.Response, error) {
|
||||
panic("compaction model crash")
|
||||
},
|
||||
}
|
||||
|
||||
parentCtx := chatdebug.ContextWithRun(context.Background(), &chatdebug.RunContext{
|
||||
RunID: uuid.New(),
|
||||
ChatID: chatID,
|
||||
ModelConfigID: uuid.New(),
|
||||
TriggerMessageID: 1,
|
||||
HistoryTipMessageID: 2,
|
||||
Kind: chatdebug.KindChatTurn,
|
||||
Provider: "fake",
|
||||
Model: "fake-model",
|
||||
})
|
||||
|
||||
require.PanicsWithValue(t, "compaction model crash", func() {
|
||||
_, _ = generateCompactionSummary(parentCtx, model,
|
||||
[]fantasy.Message{textMessage(fantasy.MessageRoleUser, "hello")},
|
||||
CompactionOptions{
|
||||
DebugSvc: svc,
|
||||
ChatID: chatID,
|
||||
SummaryPrompt: "summarize",
|
||||
Timeout: time.Second,
|
||||
})
|
||||
})
|
||||
|
||||
select {
|
||||
case s := <-status:
|
||||
require.Equal(t, string(chatdebug.StatusError), s,
|
||||
"panic path must finalize the debug run with StatusError")
|
||||
case <-time.After(testutil.WaitShort):
|
||||
t.Fatal("FinalizeRun never reached UpdateChatDebugRun on panic")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRun_Compaction(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ package chatprovider
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
@@ -1115,13 +1116,15 @@ func CoderHeadersFromIDs(
|
||||
// language model client using the provided provider credentials. The
|
||||
// userAgent is sent as the User-Agent header on every outgoing LLM
|
||||
// API request. extraHeaders, when non-nil, are sent as additional
|
||||
// HTTP headers on every request.
|
||||
// HTTP headers on every request. httpClient, when non-nil, is used for
|
||||
// all provider HTTP requests.
|
||||
func ModelFromConfig(
|
||||
providerHint string,
|
||||
modelName string,
|
||||
providerKeys ProviderAPIKeys,
|
||||
userAgent string,
|
||||
extraHeaders map[string]string,
|
||||
httpClient *http.Client,
|
||||
) (fantasy.LanguageModel, error) {
|
||||
provider, modelID, err := ResolveModelWithProviderHint(modelName, providerHint)
|
||||
if err != nil {
|
||||
@@ -1147,6 +1150,9 @@ func ModelFromConfig(
|
||||
if baseURL != "" {
|
||||
options = append(options, fantasyanthropic.WithBaseURL(baseURL))
|
||||
}
|
||||
if httpClient != nil {
|
||||
options = append(options, fantasyanthropic.WithHTTPClient(httpClient))
|
||||
}
|
||||
providerClient, err = fantasyanthropic.New(options...)
|
||||
case fantasyazure.Name:
|
||||
if baseURL == "" {
|
||||
@@ -1161,6 +1167,9 @@ func ModelFromConfig(
|
||||
if len(extraHeaders) > 0 {
|
||||
azureOpts = append(azureOpts, fantasyazure.WithHeaders(extraHeaders))
|
||||
}
|
||||
if httpClient != nil {
|
||||
azureOpts = append(azureOpts, fantasyazure.WithHTTPClient(httpClient))
|
||||
}
|
||||
providerClient, err = fantasyazure.New(azureOpts...)
|
||||
case fantasybedrock.Name:
|
||||
bedrockOpts := []fantasybedrock.Option{
|
||||
@@ -1170,6 +1179,9 @@ func ModelFromConfig(
|
||||
if len(extraHeaders) > 0 {
|
||||
bedrockOpts = append(bedrockOpts, fantasybedrock.WithHeaders(extraHeaders))
|
||||
}
|
||||
if httpClient != nil {
|
||||
bedrockOpts = append(bedrockOpts, fantasybedrock.WithHTTPClient(httpClient))
|
||||
}
|
||||
providerClient, err = fantasybedrock.New(bedrockOpts...)
|
||||
case fantasygoogle.Name:
|
||||
options := []fantasygoogle.Option{
|
||||
@@ -1182,6 +1194,9 @@ func ModelFromConfig(
|
||||
if baseURL != "" {
|
||||
options = append(options, fantasygoogle.WithBaseURL(baseURL))
|
||||
}
|
||||
if httpClient != nil {
|
||||
options = append(options, fantasygoogle.WithHTTPClient(httpClient))
|
||||
}
|
||||
providerClient, err = fantasygoogle.New(options...)
|
||||
case fantasyopenai.Name:
|
||||
options := []fantasyopenai.Option{
|
||||
@@ -1195,6 +1210,9 @@ func ModelFromConfig(
|
||||
if baseURL != "" {
|
||||
options = append(options, fantasyopenai.WithBaseURL(baseURL))
|
||||
}
|
||||
if httpClient != nil {
|
||||
options = append(options, fantasyopenai.WithHTTPClient(httpClient))
|
||||
}
|
||||
providerClient, err = fantasyopenai.New(options...)
|
||||
case fantasyopenaicompat.Name:
|
||||
options := []fantasyopenaicompat.Option{
|
||||
@@ -1207,6 +1225,9 @@ func ModelFromConfig(
|
||||
if baseURL != "" {
|
||||
options = append(options, fantasyopenaicompat.WithBaseURL(baseURL))
|
||||
}
|
||||
if httpClient != nil {
|
||||
options = append(options, fantasyopenaicompat.WithHTTPClient(httpClient))
|
||||
}
|
||||
providerClient, err = fantasyopenaicompat.New(options...)
|
||||
case fantasyopenrouter.Name:
|
||||
routerOpts := []fantasyopenrouter.Option{
|
||||
@@ -1216,6 +1237,9 @@ func ModelFromConfig(
|
||||
if len(extraHeaders) > 0 {
|
||||
routerOpts = append(routerOpts, fantasyopenrouter.WithHeaders(extraHeaders))
|
||||
}
|
||||
if httpClient != nil {
|
||||
routerOpts = append(routerOpts, fantasyopenrouter.WithHTTPClient(httpClient))
|
||||
}
|
||||
providerClient, err = fantasyopenrouter.New(routerOpts...)
|
||||
case fantasyvercel.Name:
|
||||
options := []fantasyvercel.Option{
|
||||
@@ -1228,6 +1252,9 @@ func ModelFromConfig(
|
||||
if baseURL != "" {
|
||||
options = append(options, fantasyvercel.WithBaseURL(baseURL))
|
||||
}
|
||||
if httpClient != nil {
|
||||
options = append(options, fantasyvercel.WithHTTPClient(httpClient))
|
||||
}
|
||||
providerClient, err = fantasyvercel.New(options...)
|
||||
default:
|
||||
return nil, xerrors.Errorf("unsupported model provider %q", provider)
|
||||
|
||||
@@ -181,6 +181,12 @@ func TestResolveUserProviderKeys(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
type roundTripperFunc func(*http.Request) (*http.Response, error)
|
||||
|
||||
func (fn roundTripperFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
return fn(req)
|
||||
}
|
||||
|
||||
func TestReasoningEffortFromChat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -783,7 +789,7 @@ func TestModelFromConfig_ExtraHeaders(t *testing.T) {
|
||||
BaseURLByProvider: map[string]string{"openai": serverURL},
|
||||
}
|
||||
|
||||
model, err := chatprovider.ModelFromConfig("openai", "gpt-4", keys, chatprovider.UserAgent(), headers)
|
||||
model, err := chatprovider.ModelFromConfig("openai", "gpt-4", keys, chatprovider.UserAgent(), headers, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = model.Generate(ctx, fantasy.Call{
|
||||
@@ -814,7 +820,7 @@ func TestModelFromConfig_ExtraHeaders(t *testing.T) {
|
||||
BaseURLByProvider: map[string]string{"anthropic": serverURL},
|
||||
}
|
||||
|
||||
model, err := chatprovider.ModelFromConfig("anthropic", "claude-sonnet-4-20250514", keys, chatprovider.UserAgent(), headers)
|
||||
model, err := chatprovider.ModelFromConfig("anthropic", "claude-sonnet-4-20250514", keys, chatprovider.UserAgent(), headers, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = model.Generate(ctx, fantasy.Call{
|
||||
@@ -850,7 +856,7 @@ func TestModelFromConfig_NilExtraHeaders(t *testing.T) {
|
||||
BaseURLByProvider: map[string]string{"openai": serverURL},
|
||||
}
|
||||
|
||||
model, err := chatprovider.ModelFromConfig("openai", "gpt-4", keys, chatprovider.UserAgent(), nil)
|
||||
model, err := chatprovider.ModelFromConfig("openai", "gpt-4", keys, chatprovider.UserAgent(), nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = model.Generate(ctx, fantasy.Call{
|
||||
@@ -865,6 +871,48 @@ func TestModelFromConfig_NilExtraHeaders(t *testing.T) {
|
||||
_ = testutil.TryReceive(ctx, t, called)
|
||||
}
|
||||
|
||||
func TestModelFromConfig_HTTPClient(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
|
||||
called := make(chan struct{})
|
||||
serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
assert.Equal(t, "true", req.Header.Get("X-Test-Transport"))
|
||||
close(called)
|
||||
return chattest.OpenAINonStreamingResponse("hello")
|
||||
})
|
||||
|
||||
keys := chatprovider.ProviderAPIKeys{
|
||||
ByProvider: map[string]string{"openai": "test-key"},
|
||||
BaseURLByProvider: map[string]string{"openai": serverURL},
|
||||
}
|
||||
client := &http.Client{Transport: roundTripperFunc(func(req *http.Request) (*http.Response, error) {
|
||||
cloned := req.Clone(req.Context())
|
||||
cloned.Header = req.Header.Clone()
|
||||
cloned.Header.Set("X-Test-Transport", "true")
|
||||
return http.DefaultTransport.RoundTrip(cloned)
|
||||
})}
|
||||
|
||||
model, err := chatprovider.ModelFromConfig(
|
||||
"openai",
|
||||
"gpt-4",
|
||||
keys,
|
||||
chatprovider.UserAgent(),
|
||||
nil,
|
||||
client,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = model.Generate(ctx, fantasy.Call{
|
||||
Prompt: []fantasy.Message{{
|
||||
Role: fantasy.MessageRoleUser,
|
||||
Content: []fantasy.MessagePart{fantasy.TextPart{Text: "hello"}},
|
||||
}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
_ = testutil.TryReceive(ctx, t, called)
|
||||
}
|
||||
|
||||
func TestMergeMissingProviderOptions_OpenRouterNested(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -48,7 +48,7 @@ func TestModelFromConfig_UserAgent(t *testing.T) {
|
||||
BaseURLByProvider: map[string]string{"openai": serverURL},
|
||||
}
|
||||
|
||||
model, err := chatprovider.ModelFromConfig("openai", "gpt-4", keys, expectedUA, nil)
|
||||
model, err := chatprovider.ModelFromConfig("openai", "gpt-4", keys, expectedUA, nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Make a real call so Fantasy sends an HTTP request to the
|
||||
|
||||
+246
-11
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -21,6 +22,7 @@ import (
|
||||
|
||||
"cdr.dev/slog/v3"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatdebug"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatretry"
|
||||
@@ -64,6 +66,12 @@ var preferredTitleModels = []struct {
|
||||
{fantasyvercel.Name, "anthropic/claude-haiku-4.5"},
|
||||
}
|
||||
|
||||
type shortTextCandidate struct {
|
||||
provider string
|
||||
model string
|
||||
lm fantasy.LanguageModel
|
||||
}
|
||||
|
||||
func selectPreferredConfiguredShortTextModelConfig(
|
||||
configs []database.ChatModelConfig,
|
||||
) (database.ChatModelConfig, bool) {
|
||||
@@ -105,35 +113,88 @@ func (p *Server) maybeGenerateChatTitle(
|
||||
ctx context.Context,
|
||||
chat database.Chat,
|
||||
messages []database.ChatMessage,
|
||||
fallbackProvider string,
|
||||
fallbackModelName string,
|
||||
fallbackModel fantasy.LanguageModel,
|
||||
keys chatprovider.ProviderAPIKeys,
|
||||
generatedTitle *generatedChatTitle,
|
||||
logger slog.Logger,
|
||||
debugSvc *chatdebug.Service,
|
||||
) {
|
||||
input, ok := titleInput(chat, messages)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
debugEnabled := debugSvc != nil && debugSvc.IsEnabled(ctx, chat.ID, chat.OwnerID)
|
||||
|
||||
titleCtx, cancel := context.WithTimeout(ctx, 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// Build candidate list: preferred lightweight models first,
|
||||
// then the user's chat model as last resort.
|
||||
candidates := make([]fantasy.LanguageModel, 0, len(preferredTitleModels)+1)
|
||||
candidates := make([]shortTextCandidate, 0, len(preferredTitleModels)+1)
|
||||
for _, c := range preferredTitleModels {
|
||||
m, err := chatprovider.ModelFromConfig(
|
||||
c.provider, c.model, keys, chatprovider.UserAgent(),
|
||||
chatprovider.CoderHeaders(chat),
|
||||
nil,
|
||||
)
|
||||
if err == nil {
|
||||
candidates = append(candidates, m)
|
||||
candidates = append(candidates, shortTextCandidate{
|
||||
provider: c.provider,
|
||||
model: c.model,
|
||||
lm: m,
|
||||
})
|
||||
}
|
||||
}
|
||||
candidates = append(candidates, fallbackModel)
|
||||
candidates = append(candidates, shortTextCandidate{
|
||||
provider: fallbackProvider,
|
||||
model: fallbackModelName,
|
||||
lm: fallbackModel,
|
||||
})
|
||||
|
||||
var historyTipMessageID int64
|
||||
if len(messages) > 0 {
|
||||
historyTipMessageID = messages[len(messages)-1].ID
|
||||
}
|
||||
|
||||
var triggerMessageID int64
|
||||
for _, message := range messages {
|
||||
if message.Visibility == database.ChatMessageVisibilityModel {
|
||||
continue
|
||||
}
|
||||
if message.Role == database.ChatMessageRoleUser {
|
||||
triggerMessageID = message.ID
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
seedSummary := chatdebug.SeedSummary(
|
||||
chatdebug.TruncateLabel(input, chatdebug.MaxLabelLength),
|
||||
)
|
||||
|
||||
var lastErr error
|
||||
for _, model := range candidates {
|
||||
title, err := generateTitle(titleCtx, model, input)
|
||||
for _, candidate := range candidates {
|
||||
candidateCtx := titleCtx
|
||||
candidateModel := candidate.lm
|
||||
finishDebugRun := func(error) {}
|
||||
if debugEnabled {
|
||||
candidateCtx, candidateModel, finishDebugRun = prepareQuickgenDebugCandidate(
|
||||
titleCtx,
|
||||
chat,
|
||||
keys,
|
||||
debugSvc,
|
||||
candidate,
|
||||
chatdebug.KindTitleGeneration,
|
||||
triggerMessageID,
|
||||
historyTipMessageID,
|
||||
seedSummary,
|
||||
logger,
|
||||
)
|
||||
}
|
||||
|
||||
title, err := generateTitle(candidateCtx, candidateModel, input)
|
||||
finishDebugRun(err)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
logger.Debug(ctx, "title model candidate failed",
|
||||
@@ -171,6 +232,137 @@ func (p *Server) maybeGenerateChatTitle(
|
||||
}
|
||||
}
|
||||
|
||||
func newQuickgenDebugModel(
|
||||
chat database.Chat,
|
||||
keys chatprovider.ProviderAPIKeys,
|
||||
debugSvc *chatdebug.Service,
|
||||
provider string,
|
||||
model string,
|
||||
) (fantasy.LanguageModel, error) {
|
||||
httpClient := &http.Client{Transport: &chatdebug.RecordingTransport{}}
|
||||
debugModel, err := chatprovider.ModelFromConfig(
|
||||
provider,
|
||||
model,
|
||||
keys,
|
||||
chatprovider.UserAgent(),
|
||||
chatprovider.CoderHeaders(chat),
|
||||
httpClient,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if debugModel == nil {
|
||||
return nil, xerrors.Errorf(
|
||||
"create model for %s/%s returned nil",
|
||||
provider,
|
||||
model,
|
||||
)
|
||||
}
|
||||
|
||||
return chatdebug.WrapModel(debugModel, debugSvc, chatdebug.RecorderOptions{
|
||||
ChatID: chat.ID,
|
||||
OwnerID: chat.OwnerID,
|
||||
Provider: provider,
|
||||
Model: model,
|
||||
}), nil
|
||||
}
|
||||
|
||||
func prepareQuickgenDebugCandidate(
|
||||
ctx context.Context,
|
||||
chat database.Chat,
|
||||
keys chatprovider.ProviderAPIKeys,
|
||||
debugSvc *chatdebug.Service,
|
||||
candidate shortTextCandidate,
|
||||
kind chatdebug.RunKind,
|
||||
triggerMessageID int64,
|
||||
historyTipMessageID int64,
|
||||
seedSummary map[string]any,
|
||||
logger slog.Logger,
|
||||
) (context.Context, fantasy.LanguageModel, func(error)) {
|
||||
finishDebugRun := func(error) {}
|
||||
if debugSvc == nil {
|
||||
return ctx, candidate.lm, finishDebugRun
|
||||
}
|
||||
|
||||
debugModel, err := newQuickgenDebugModel(
|
||||
chat,
|
||||
keys,
|
||||
debugSvc,
|
||||
candidate.provider,
|
||||
candidate.model,
|
||||
)
|
||||
if err != nil {
|
||||
logger.Warn(ctx, "failed to build short-text debug model",
|
||||
slog.F("chat_id", chat.ID),
|
||||
slog.F("run_kind", kind),
|
||||
slog.F("provider", candidate.provider),
|
||||
slog.F("model", candidate.model),
|
||||
slog.Error(err),
|
||||
)
|
||||
return ctx, candidate.lm, finishDebugRun
|
||||
}
|
||||
|
||||
// Debug instrumentation must not eat into the quickgen budget
|
||||
// (30s titleCtx / summaryCtx on the caller). Detach and bound
|
||||
// the insert so a slow DB can't delay title generation or push
|
||||
// summaries, matching prepareManualTitleDebugRun,
|
||||
// prepareChatTurnDebugRun, and startCompactionDebugRun.
|
||||
createRunCtx, createRunCancel := context.WithTimeout(
|
||||
context.WithoutCancel(ctx), debugCreateRunTimeout,
|
||||
)
|
||||
run, err := debugSvc.CreateRun(createRunCtx, chatdebug.CreateRunParams{
|
||||
ChatID: chat.ID,
|
||||
TriggerMessageID: triggerMessageID,
|
||||
HistoryTipMessageID: historyTipMessageID,
|
||||
Kind: kind,
|
||||
Status: chatdebug.StatusInProgress,
|
||||
Provider: candidate.provider,
|
||||
Model: candidate.model,
|
||||
Summary: seedSummary,
|
||||
})
|
||||
createRunCancel()
|
||||
if err != nil {
|
||||
logger.Warn(ctx, "failed to create short-text debug run",
|
||||
slog.F("chat_id", chat.ID),
|
||||
slog.F("run_kind", kind),
|
||||
slog.F("provider", candidate.provider),
|
||||
slog.F("model", candidate.model),
|
||||
slog.Error(err),
|
||||
)
|
||||
return ctx, candidate.lm, finishDebugRun
|
||||
}
|
||||
|
||||
runCtx := chatdebug.ContextWithRun(
|
||||
ctx,
|
||||
&chatdebug.RunContext{
|
||||
RunID: run.ID,
|
||||
ChatID: chat.ID,
|
||||
TriggerMessageID: triggerMessageID,
|
||||
HistoryTipMessageID: historyTipMessageID,
|
||||
Kind: kind,
|
||||
Provider: candidate.provider,
|
||||
Model: candidate.model,
|
||||
},
|
||||
)
|
||||
finishDebugRun = func(runErr error) {
|
||||
if finalizeErr := debugSvc.FinalizeRun(ctx, chatdebug.FinalizeRunParams{
|
||||
RunID: run.ID,
|
||||
ChatID: chat.ID,
|
||||
Status: chatdebug.ClassifyError(runErr),
|
||||
SeedSummary: seedSummary,
|
||||
Timeout: 10 * time.Second,
|
||||
}); finalizeErr != nil {
|
||||
logger.Warn(ctx, "failed to finalize short-text debug run",
|
||||
slog.F("chat_id", chat.ID),
|
||||
slog.F("run_kind", kind),
|
||||
slog.F("run_id", run.ID),
|
||||
slog.Error(finalizeErr),
|
||||
)
|
||||
}
|
||||
}
|
||||
return runCtx, debugModel, finishDebugRun
|
||||
}
|
||||
|
||||
// generateTitle calls the model with a title-generation system prompt
|
||||
// and returns the normalized result. It retries transient LLM errors
|
||||
// (rate limits, overloaded, etc.) with exponential backoff.
|
||||
@@ -571,30 +763,72 @@ func generatePushSummary(
|
||||
ctx context.Context,
|
||||
chat database.Chat,
|
||||
assistantText string,
|
||||
fallbackProvider string,
|
||||
fallbackModelName string,
|
||||
fallbackModel fantasy.LanguageModel,
|
||||
keys chatprovider.ProviderAPIKeys,
|
||||
logger slog.Logger,
|
||||
debugSvc *chatdebug.Service,
|
||||
triggerMessageID int64,
|
||||
historyTipMessageID int64,
|
||||
) string {
|
||||
debugEnabled := debugSvc != nil && debugSvc.IsEnabled(ctx, chat.ID, chat.OwnerID)
|
||||
|
||||
summaryCtx, cancel := context.WithTimeout(ctx, 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
assistantText = truncateRunes(assistantText, maxConversationContextRunes)
|
||||
input := "Chat title: " + chat.Title + "\n\nAgent's last message:\n" + assistantText
|
||||
|
||||
candidates := make([]fantasy.LanguageModel, 0, len(preferredTitleModels)+1)
|
||||
candidates := make([]shortTextCandidate, 0, len(preferredTitleModels)+1)
|
||||
for _, c := range preferredTitleModels {
|
||||
m, err := chatprovider.ModelFromConfig(
|
||||
c.provider, c.model, keys, chatprovider.UserAgent(),
|
||||
chatprovider.CoderHeaders(chat),
|
||||
nil,
|
||||
)
|
||||
if err == nil {
|
||||
candidates = append(candidates, m)
|
||||
candidates = append(candidates, shortTextCandidate{
|
||||
provider: c.provider,
|
||||
model: c.model,
|
||||
lm: m,
|
||||
})
|
||||
}
|
||||
}
|
||||
candidates = append(candidates, fallbackModel)
|
||||
candidates = append(candidates, shortTextCandidate{
|
||||
provider: fallbackProvider,
|
||||
model: fallbackModelName,
|
||||
lm: fallbackModel,
|
||||
})
|
||||
|
||||
for _, model := range candidates {
|
||||
summary, err := generateShortText(summaryCtx, model, pushSummaryPrompt, input)
|
||||
pushSeedSummary := chatdebug.SeedSummary("Push summary")
|
||||
|
||||
for _, candidate := range candidates {
|
||||
candidateCtx := summaryCtx
|
||||
candidateModel := candidate.lm
|
||||
finishDebugRun := func(error) {}
|
||||
if debugEnabled {
|
||||
candidateCtx, candidateModel, finishDebugRun = prepareQuickgenDebugCandidate(
|
||||
summaryCtx,
|
||||
chat,
|
||||
keys,
|
||||
debugSvc,
|
||||
candidate,
|
||||
chatdebug.KindQuickgen,
|
||||
triggerMessageID,
|
||||
historyTipMessageID,
|
||||
pushSeedSummary,
|
||||
logger,
|
||||
)
|
||||
}
|
||||
|
||||
summary, err := generateShortText(
|
||||
candidateCtx,
|
||||
candidateModel,
|
||||
pushSummaryPrompt,
|
||||
input,
|
||||
)
|
||||
finishDebugRun(err)
|
||||
if err != nil {
|
||||
logger.Debug(ctx, "push summary model candidate failed",
|
||||
slog.Error(err),
|
||||
@@ -610,7 +844,8 @@ func generatePushSummary(
|
||||
|
||||
// generateShortText calls a model with a system prompt and user
|
||||
// input, returning a cleaned-up short text response. It reuses the
|
||||
// same retry logic as title generation.
|
||||
// same retry logic as title generation. Retries can therefore
|
||||
// produce multiple debug steps for a single quickgen run.
|
||||
func generateShortText(
|
||||
ctx context.Context,
|
||||
model fantasy.LanguageModel,
|
||||
|
||||
Reference in New Issue
Block a user