fix: persist per-turn model on chats and queued messages (#24688)

Previously, `chats.last_model_config_id` was not updated when a user
sent a mid-chat message with a different model, and queued messages did
not store their own per-turn model, so promotion ran against whatever
the chat row said at promote time. Chat watch events also did not merge
`last_model_config_id` into the site's root, child, and per-chat
caches, so sidebar labels stayed stale after direct sends and queued
promotions.

- Add nullable `chat_queued_messages.model_config_id`, backfilled from
  `chats.last_model_config_id`. Queued inserts round-trip the effective
  model id at enqueue time.
- In `coderd/x/chatd`, direct sends update `chats.last_model_config_id`
  inside the same transaction that inserts the admitted user message.
  Manual promotion and auto-promotion use the queued row's stored
  `model_config_id`, with a fallback to `chats.last_model_config_id`
for legacy NULL rows during rollout.
`PromoteQueuedOptions.ModelConfigID`
  is now ignored.
- On the site, extract `mergeWatchedChatSummary` and
  `mergeWatchedChatIntoCaches` in `site/src/api/queries/chats.ts` so
  status-change watch events merge `last_model_config_id` into the
  root infinite chat list, the parent-embedded child entry, and the
  per-chat `chatKey(chatId)` cache. `updated_at` guards against stale
  watch payloads clobbering newer cached state, while diff status
  events still merge their PR metadata because they are timestamped
  outside the chat row. Watch timestamps are compared as instants so
  variable fractional precision does not make fresh events look stale.
- Queued promotion validates stored model config IDs before admission.
  Invalid legacy queued IDs fall back to the chat's current model config
  instead of dropping the queued message during auto-promotion.
- Backend and frontend regression coverage added for admission, queue
  promotion (including FIFO across mixed models, legacy NULL fallback,
  and invalid queued model IDs), and chat watch cache merging.

> Mux is acting on Mike's behalf.
This commit is contained in:
Michael Suchacz
2026-04-24 15:36:08 +02:00
committed by GitHub
parent a876287d36
commit c7cac9debe
16 changed files with 1580 additions and 182 deletions
+147 -17
View File
@@ -871,6 +871,8 @@ func (c *streamStateCollector) Collect(ch chan<- prometheus.Metric) {
const MaxQueueSize = 20
var (
// ErrInvalidModelConfigID indicates the requested model config does not exist.
ErrInvalidModelConfigID = xerrors.New("invalid model config ID")
// ErrMessageQueueFull indicates the per-chat queue limit was reached.
ErrMessageQueueFull = xerrors.New("chat message queue is full")
// ErrEditedMessageNotFound indicates the edited message does not exist
@@ -950,7 +952,7 @@ type SendMessageOptions struct {
ChatID uuid.UUID
CreatedBy uuid.UUID
Content []codersdk.ChatMessagePart
ModelConfigID *uuid.UUID
ModelConfigID uuid.UUID
BusyBehavior SendMessageBusyBehavior
PlanMode *database.NullChatPlanMode
MCPServerIDs *[]uuid.UUID
@@ -983,7 +985,6 @@ type PromoteQueuedOptions struct {
ChatID uuid.UUID
CreatedBy uuid.UUID
QueuedMessageID int64
ModelConfigID *uuid.UUID
}
// PromoteQueuedResult contains post-promotion message metadata.
@@ -1217,9 +1218,14 @@ func (p *Server) SendMessage(
}
}
modelConfigID := lockedChat.LastModelConfigID
if opts.ModelConfigID != nil {
modelConfigID = *opts.ModelConfigID
modelConfigID, err := resolveSendMessageModelConfigID(
ctx,
tx,
lockedChat,
opts.ModelConfigID,
)
if err != nil {
return err
}
// Update MCP server IDs on the chat when explicitly provided.
@@ -1264,6 +1270,10 @@ func (p *Server) SendMessage(
queued, err := tx.InsertChatQueuedMessage(ctx, database.InsertChatQueuedMessageParams{
ChatID: opts.ChatID,
Content: content.RawMessage,
ModelConfigID: uuid.NullUUID{
UUID: modelConfigID,
Valid: modelConfigID != uuid.Nil,
},
})
if err != nil {
return xerrors.Errorf("insert queued message: %w", err)
@@ -1368,6 +1378,90 @@ func (p *Server) checkUsageLimit(ctx context.Context, store database.Store, owne
return nil
}
func chatdModelConfigLookupContext(ctx context.Context) context.Context {
//nolint:gocritic // Chat message admission needs daemon-scoped
// deployment-config reads for model config validation.
return dbauthz.AsChatd(ctx)
}
func resolveSendMessageModelConfigID(
ctx context.Context,
store database.Store,
chat database.Chat,
requested uuid.UUID,
) (uuid.UUID, error) {
if requested == uuid.Nil {
return resolveFallbackModelConfigID(ctx, store, chat.LastModelConfigID)
}
chatdCtx := chatdModelConfigLookupContext(ctx)
if _, err := store.GetChatModelConfigByID(chatdCtx, requested); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return uuid.Nil, xerrors.Errorf(
"%w: %s",
ErrInvalidModelConfigID,
requested,
)
}
return uuid.Nil, xerrors.Errorf(
"get requested model config %s: %w",
requested,
err,
)
}
return requested, nil
}
func resolveQueuedMessageModelConfigID(
ctx context.Context,
store database.Store,
chat database.Chat,
queuedModelConfigID uuid.NullUUID,
) (uuid.UUID, error) {
chatdCtx := chatdModelConfigLookupContext(ctx)
if queuedModelConfigID.Valid && queuedModelConfigID.UUID != uuid.Nil {
if _, err := store.GetChatModelConfigByID(chatdCtx, queuedModelConfigID.UUID); err == nil {
return queuedModelConfigID.UUID, nil
} else if !errors.Is(err, sql.ErrNoRows) {
return uuid.Nil, xerrors.Errorf(
"get queued model config %s: %w",
queuedModelConfigID.UUID,
err,
)
}
}
return resolveFallbackModelConfigID(ctx, store, chat.LastModelConfigID)
}
func resolveFallbackModelConfigID(
ctx context.Context,
store database.Store,
modelConfigID uuid.UUID,
) (uuid.UUID, error) {
chatdCtx := chatdModelConfigLookupContext(ctx)
if modelConfigID != uuid.Nil {
if _, err := store.GetChatModelConfigByID(chatdCtx, modelConfigID); err == nil {
return modelConfigID, nil
} else if !errors.Is(err, sql.ErrNoRows) {
return uuid.Nil, xerrors.Errorf(
"get chat model config %s: %w",
modelConfigID,
err,
)
}
}
defaultConfig, err := store.GetDefaultChatModelConfig(chatdCtx)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return uuid.Nil, xerrors.New("no default chat model config is available")
}
return uuid.Nil, xerrors.Errorf("get default chat model config: %w", err)
}
return defaultConfig.ID, nil
}
// EditMessage marks the old user message as deleted, soft-deletes all
// following messages, inserts a new message with the updated content,
// clears queued messages, and moves the chat into pending status.
@@ -1768,23 +1862,20 @@ func (p *Server) PromoteQueued(
return ErrChatArchived
}
modelConfigID := lockedChat.LastModelConfigID
if opts.ModelConfigID != nil {
modelConfigID = *opts.ModelConfigID
}
queuedMessages, err := tx.GetChatQueuedMessages(ctx, opts.ChatID)
if err != nil {
return xerrors.Errorf("get queued messages: %w", err)
}
var (
targetContent json.RawMessage
found bool
targetContent json.RawMessage
targetModelConfigID uuid.NullUUID
found bool
)
for _, qm := range queuedMessages {
if qm.ID == opts.QueuedMessageID {
targetContent = qm.Content
targetModelConfigID = qm.ModelConfigID
found = true
break
}
@@ -1793,6 +1884,16 @@ func (p *Server) PromoteQueued(
return xerrors.New("queued message not found")
}
effectiveModelConfigID, err := resolveQueuedMessageModelConfigID(
ctx,
tx,
lockedChat,
targetModelConfigID,
)
if err != nil {
return err
}
err = tx.DeleteChatQueuedMessage(ctx, database.DeleteChatQueuedMessageParams{
ID: opts.QueuedMessageID,
ChatID: opts.ChatID,
@@ -1805,7 +1906,7 @@ func (p *Server) PromoteQueued(
ctx,
tx,
lockedChat,
modelConfigID,
effectiveModelConfigID,
pqtype.NullRawMessage{
RawMessage: targetContent,
Valid: len(targetContent) > 0,
@@ -3313,6 +3414,8 @@ func BuildSingleChatMessageInsertParams(
return params
}
// insertUserMessageAndSetPending inserts a user message, transitions the
// chat to pending when needed, and returns the refreshed chat row.
func insertUserMessageAndSetPending(
ctx context.Context,
store database.Store,
@@ -3338,7 +3441,16 @@ func insertUserMessageAndSetPending(
message := messages[0]
if lockedChat.Status == database.ChatStatusPending {
return message, lockedChat, nil
if modelConfigID == uuid.Nil || lockedChat.LastModelConfigID == modelConfigID {
return message, lockedChat, nil
}
// The InsertChatMessages CTE updates chats.last_model_config_id when
// the message's model config differs. Reload to surface that change.
updatedChat, err := store.GetChatByID(ctx, lockedChat.ID)
if err != nil {
return database.ChatMessage{}, database.Chat{}, xerrors.Errorf("get chat after model config update: %w", err)
}
return message, updatedChat, nil
}
updatedChat, err := store.UpdateChatStatus(ctx, database.UpdateChatStatusParams{
@@ -4752,13 +4864,31 @@ func (p *Server) tryAutoPromoteQueuedMessage(
) (*database.ChatMessage, []database.ChatQueuedMessage, bool, error) {
logger := p.logger.With(slog.F("chat_id", chat.ID))
nextQueued, err := tx.PopNextQueuedMessage(ctx, chat.ID)
if errors.Is(err, sql.ErrNoRows) {
queuedMessages, err := tx.GetChatQueuedMessages(ctx, chat.ID)
if err != nil {
return nil, nil, false, xerrors.Errorf("get queued messages: %w", err)
}
if len(queuedMessages) == 0 {
return nil, nil, false, nil
}
nextQueued := queuedMessages[0]
effectiveModelConfigID, err := resolveQueuedMessageModelConfigID(
ctx,
tx,
chat,
nextQueued.ModelConfigID,
)
if err != nil {
return nil, nil, false, err
}
poppedQueued, err := tx.PopNextQueuedMessage(ctx, chat.ID)
if err != nil {
return nil, nil, false, xerrors.Errorf("pop next queued message: %w", err)
}
if poppedQueued.ID != nextQueued.ID {
return nil, nil, false, xerrors.New("popped queued message out of order")
}
msgParams := database.InsertChatMessagesParams{ //nolint:exhaustruct // Fields populated by appendChatMessage.
ChatID: chat.ID,
@@ -4770,7 +4900,7 @@ func (p *Server) tryAutoPromoteQueuedMessage(
Valid: len(nextQueued.Content) > 0,
},
database.ChatMessageVisibilityBoth,
chat.LastModelConfigID,
effectiveModelConfigID,
chatprompt.CurrentContentVersion,
).withCreatedBy(chat.OwnerID))
msgs, err := insertChatMessageWithStore(ctx, tx, msgParams)
+490
View File
@@ -39,6 +39,7 @@ import (
"github.com/coder/coder/v2/coderd/database/dbtestutil"
"github.com/coder/coder/v2/coderd/database/dbtime"
dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub"
coderdpubsub "github.com/coder/coder/v2/coderd/pubsub"
"github.com/coder/coder/v2/coderd/rbac"
"github.com/coder/coder/v2/coderd/util/slice"
"github.com/coder/coder/v2/coderd/workspacestats"
@@ -2045,6 +2046,38 @@ func TestSendMessageQueuesWhenWaitingWithQueuedBacklog(t *testing.T) {
require.Len(t, messages, 1)
}
func TestSendMessageRejectsInvalidQueuedModelConfigID(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
replica := newTestServer(t, db, ps, uuid.New())
ctx := testutil.Context(t, testutil.WaitLong)
user, org, modelConfig := seedChatDependencies(ctx, t, db)
chat, err := db.InsertChat(ctx, database.InsertChatParams{
OrganizationID: org.ID,
Status: database.ChatStatusPending,
ClientType: database.ChatClientTypeUi,
OwnerID: user.ID,
LastModelConfigID: modelConfig.ID,
Title: "reject invalid queued model config",
})
require.NoError(t, err)
invalidModelConfigID := uuid.New()
_, err = replica.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: chat.ID,
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued")},
ModelConfigID: invalidModelConfigID,
})
require.ErrorIs(t, err, chatd.ErrInvalidModelConfigID)
queued, err := db.GetChatQueuedMessages(ctx, chat.ID)
require.NoError(t, err)
require.Empty(t, queued)
}
func TestSendMessageInterruptBehaviorQueuesAndInterruptsWhenBusy(t *testing.T) {
t.Parallel()
@@ -2501,6 +2534,463 @@ func TestPromoteQueuedAllowsAlreadyQueuedMessageWhenUsageLimitReached(t *testing
require.Equal(t, database.ChatMessageRoleUser, messages[3].Role)
}
func TestPromoteQueuedMessageUsesQueuedModelConfigID(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
replica := newTestServer(t, db, ps, uuid.New())
ctx := testutil.Context(t, testutil.WaitLong)
user, org, modelConfigA := seedChatDependencies(ctx, t, db)
modelConfigB := insertChatModelConfigWithCallConfig(
ctx,
t,
db,
user.ID,
"openai",
"gpt-4o-mini-promote-"+uuid.NewString(),
codersdk.ChatModelCallConfig{},
)
chat, err := db.InsertChat(ctx, database.InsertChatParams{
OrganizationID: org.ID,
Status: database.ChatStatusWaiting,
ClientType: database.ChatClientTypeUi,
OwnerID: user.ID,
LastModelConfigID: modelConfigA.ID,
Title: "promote queued uses stored model",
})
require.NoError(t, err)
queuedContent, err := json.Marshal([]codersdk.ChatMessagePart{codersdk.ChatMessageText("queued with model b")})
require.NoError(t, err)
queuedMessage, err := db.InsertChatQueuedMessage(ctx, database.InsertChatQueuedMessageParams{
ChatID: chat.ID,
Content: queuedContent,
ModelConfigID: uuid.NullUUID{
UUID: modelConfigB.ID,
Valid: true,
},
})
require.NoError(t, err)
result, err := replica.PromoteQueued(ctx, chatd.PromoteQueuedOptions{
ChatID: chat.ID,
QueuedMessageID: queuedMessage.ID,
CreatedBy: user.ID,
})
require.NoError(t, err)
require.True(t, result.PromotedMessage.ModelConfigID.Valid)
require.Equal(t, modelConfigB.ID, result.PromotedMessage.ModelConfigID.UUID)
storedChat, err := db.GetChatByID(ctx, chat.ID)
require.NoError(t, err)
require.Equal(t, modelConfigB.ID, storedChat.LastModelConfigID)
require.Equal(t, database.ChatStatusPending, storedChat.Status)
}
func TestPromoteQueuedMessageReloadsChatWhenModelConfigChangesDuringPending(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
replica := newTestServer(t, db, ps, uuid.New())
ctx := testutil.Context(t, testutil.WaitLong)
user, org, modelConfigA := seedChatDependencies(ctx, t, db)
modelConfigB := insertChatModelConfigWithCallConfig(
ctx,
t,
db,
user.ID,
"openai",
"gpt-4o-mini-promote-pending-"+uuid.NewString(),
codersdk.ChatModelCallConfig{},
)
watchEvents := make(chan struct {
payload codersdk.ChatWatchEvent
err error
}, 1)
cancelWatch, err := ps.SubscribeWithErr(
coderdpubsub.ChatWatchEventChannel(user.ID),
coderdpubsub.HandleChatWatchEvent(func(_ context.Context, payload codersdk.ChatWatchEvent, err error) {
select {
case watchEvents <- struct {
payload codersdk.ChatWatchEvent
err error
}{payload: payload, err: err}:
default:
}
}),
)
require.NoError(t, err)
defer cancelWatch()
chat, err := db.InsertChat(ctx, database.InsertChatParams{
OrganizationID: org.ID,
Status: database.ChatStatusPending,
ClientType: database.ChatClientTypeUi,
OwnerID: user.ID,
LastModelConfigID: modelConfigA.ID,
Title: "promote queued reloads pending chat",
})
require.NoError(t, err)
queuedContent, err := json.Marshal([]codersdk.ChatMessagePart{codersdk.ChatMessageText("queued with new model")})
require.NoError(t, err)
queuedMessage, err := db.InsertChatQueuedMessage(ctx, database.InsertChatQueuedMessageParams{
ChatID: chat.ID,
Content: queuedContent,
ModelConfigID: uuid.NullUUID{
UUID: modelConfigB.ID,
Valid: true,
},
})
require.NoError(t, err)
result, err := replica.PromoteQueued(ctx, chatd.PromoteQueuedOptions{
ChatID: chat.ID,
QueuedMessageID: queuedMessage.ID,
CreatedBy: user.ID,
})
require.NoError(t, err)
require.True(t, result.PromotedMessage.ModelConfigID.Valid)
require.Equal(t, modelConfigB.ID, result.PromotedMessage.ModelConfigID.UUID)
storedChat, err := db.GetChatByID(ctx, chat.ID)
require.NoError(t, err)
require.Equal(t, database.ChatStatusPending, storedChat.Status)
require.Equal(t, modelConfigB.ID, storedChat.LastModelConfigID)
select {
case event := <-watchEvents:
require.NoError(t, event.err)
require.Equal(t, codersdk.ChatWatchEventKindStatusChange, event.payload.Kind)
require.Equal(t, chat.ID, event.payload.Chat.ID)
require.Equal(t, codersdk.ChatStatusPending, event.payload.Chat.Status)
require.Equal(t, modelConfigB.ID, event.payload.Chat.LastModelConfigID)
case <-ctx.Done():
t.Fatal("timed out waiting for status change watch event")
}
}
func TestAutoPromoteQueuedMessagesPreservesPerTurnModelOrder(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitLong)
firstRunStarted := make(chan struct{})
allowFirstRunFinish := make(chan struct{})
var requestCount atomic.Int32
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
if !req.Stream {
return chattest.OpenAINonStreamingResponse("title")
}
switch requestCount.Add(1) {
case 1:
chunks := make(chan chattest.OpenAIChunk, 1)
go func() {
defer close(chunks)
chunks <- chattest.OpenAITextChunks("first run partial")[0]
select {
case <-firstRunStarted:
default:
close(firstRunStarted)
}
<-allowFirstRunFinish
}()
return chattest.OpenAIResponse{StreamingChunks: chunks}
case 2:
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("second run done")...)
case 3:
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("third run done")...)
default:
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("extra run done")...)
}
})
server := newActiveTestServer(t, db, ps)
user, org, modelConfigA := seedChatDependenciesWithProvider(ctx, t, db, "openai-compat", openAIURL)
modelConfigB := insertChatModelConfigWithCallConfig(
ctx,
t,
db,
user.ID,
"openai-compat",
"gpt-4o-mini-queue-b-"+uuid.NewString(),
codersdk.ChatModelCallConfig{},
)
modelConfigC := insertChatModelConfigWithCallConfig(
ctx,
t,
db,
user.ID,
"openai-compat",
"gpt-4o-mini-queue-c-"+uuid.NewString(),
codersdk.ChatModelCallConfig{},
)
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
OwnerID: user.ID,
Title: "auto-promote per-turn model order",
ModelConfigID: modelConfigA.ID,
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
testutil.TryReceive(ctx, t, firstRunStarted)
queuedB, err := server.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: chat.ID,
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued b")},
ModelConfigID: modelConfigB.ID,
BusyBehavior: chatd.SendMessageBusyBehaviorQueue,
})
require.NoError(t, err)
require.True(t, queuedB.Queued)
queuedC, err := server.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: chat.ID,
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued c")},
ModelConfigID: modelConfigC.ID,
BusyBehavior: chatd.SendMessageBusyBehaviorQueue,
})
require.NoError(t, err)
require.True(t, queuedC.Queued)
close(allowFirstRunFinish)
require.Eventually(t, func() bool {
return requestCount.Load() >= 3
}, testutil.WaitLong, testutil.IntervalFast)
chatd.WaitUntilIdleForTest(server)
queuedMessages, err := db.GetChatQueuedMessages(ctx, chat.ID)
require.NoError(t, err)
require.Empty(t, queuedMessages)
storedChat, err := db.GetChatByID(ctx, chat.ID)
require.NoError(t, err)
require.Equal(t, database.ChatStatusWaiting, storedChat.Status)
require.Equal(t, modelConfigC.ID, storedChat.LastModelConfigID)
messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
ChatID: chat.ID,
AfterID: 0,
})
require.NoError(t, err)
var userTexts []string
var userModelConfigIDs []uuid.UUID
for _, message := range messages {
if message.Role != database.ChatMessageRoleUser {
continue
}
sdkMessage := db2sdk.ChatMessage(message)
require.Len(t, sdkMessage.Content, 1)
userTexts = append(userTexts, sdkMessage.Content[0].Text)
require.True(t, message.ModelConfigID.Valid)
userModelConfigIDs = append(userModelConfigIDs, message.ModelConfigID.UUID)
}
require.Equal(t, []string{"hello", "queued b", "queued c"}, userTexts)
require.Equal(t, []uuid.UUID{modelConfigA.ID, modelConfigB.ID, modelConfigC.ID}, userModelConfigIDs)
}
func TestAutoPromoteQueuedMessageFallsBackForLegacyQueuedRows(t *testing.T) {
t.Parallel()
testAutoPromoteQueuedMessageFallback(t, uuid.NullUUID{})
}
func TestAutoPromoteQueuedMessageFallsBackForInvalidQueuedModelConfigID(t *testing.T) {
t.Parallel()
testAutoPromoteQueuedMessageFallback(t, uuid.NullUUID{
UUID: uuid.New(),
Valid: true,
})
}
func testAutoPromoteQueuedMessageFallback(t *testing.T, queuedModelConfigID uuid.NullUUID) {
db, ps := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitLong)
firstRunStarted := make(chan struct{})
allowFirstRunFinish := make(chan struct{})
var requestCount atomic.Int32
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
if !req.Stream {
return chattest.OpenAINonStreamingResponse("title")
}
switch requestCount.Add(1) {
case 1:
chunks := make(chan chattest.OpenAIChunk, 1)
go func() {
defer close(chunks)
chunks <- chattest.OpenAITextChunks("first run partial")[0]
select {
case <-firstRunStarted:
default:
close(firstRunStarted)
}
<-allowFirstRunFinish
}()
return chattest.OpenAIResponse{StreamingChunks: chunks}
default:
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("fallback run done")...)
}
})
server := newActiveTestServer(t, db, ps)
user, org, modelConfig := seedChatDependenciesWithProvider(ctx, t, db, "openai-compat", openAIURL)
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
OwnerID: user.ID,
Title: "auto-promote queued fallback",
ModelConfigID: modelConfig.ID,
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
testutil.TryReceive(ctx, t, firstRunStarted)
queuedContent, err := json.Marshal([]codersdk.ChatMessagePart{codersdk.ChatMessageText("legacy queued row")})
require.NoError(t, err)
_, err = db.InsertChatQueuedMessage(ctx, database.InsertChatQueuedMessageParams{
ChatID: chat.ID,
Content: queuedContent,
ModelConfigID: queuedModelConfigID,
})
require.NoError(t, err)
close(allowFirstRunFinish)
require.Eventually(t, func() bool {
return requestCount.Load() >= 2
}, testutil.WaitLong, testutil.IntervalFast)
chatd.WaitUntilIdleForTest(server)
queuedMessages, err := db.GetChatQueuedMessages(ctx, chat.ID)
require.NoError(t, err)
require.Empty(t, queuedMessages)
storedChat, err := db.GetChatByID(ctx, chat.ID)
require.NoError(t, err)
require.Equal(t, database.ChatStatusWaiting, storedChat.Status)
require.Equal(t, modelConfig.ID, storedChat.LastModelConfigID)
messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
ChatID: chat.ID,
AfterID: 0,
})
require.NoError(t, err)
var found bool
for _, message := range messages {
if message.Role != database.ChatMessageRoleUser {
continue
}
sdkMessage := db2sdk.ChatMessage(message)
require.Len(t, sdkMessage.Content, 1)
if sdkMessage.Content[0].Text != "legacy queued row" {
continue
}
require.True(t, message.ModelConfigID.Valid)
require.Equal(t, modelConfig.ID, message.ModelConfigID.UUID)
found = true
}
require.True(t, found)
}
func TestPromoteQueuedMessageFallsBackForLegacyQueuedRows(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
replica := newTestServer(t, db, ps, uuid.New())
ctx := testutil.Context(t, testutil.WaitLong)
user, org, modelConfigA := seedChatDependencies(ctx, t, db)
chat, err := db.InsertChat(ctx, database.InsertChatParams{
OrganizationID: org.ID,
Status: database.ChatStatusWaiting,
ClientType: database.ChatClientTypeUi,
OwnerID: user.ID,
LastModelConfigID: modelConfigA.ID,
Title: "promote queued legacy fallback",
})
require.NoError(t, err)
queuedContent, err := json.Marshal([]codersdk.ChatMessagePart{codersdk.ChatMessageText("legacy queued row")})
require.NoError(t, err)
queuedMessage, err := db.InsertChatQueuedMessage(ctx, database.InsertChatQueuedMessageParams{
ChatID: chat.ID,
Content: queuedContent,
})
require.NoError(t, err)
result, err := replica.PromoteQueued(ctx, chatd.PromoteQueuedOptions{
ChatID: chat.ID,
QueuedMessageID: queuedMessage.ID,
CreatedBy: user.ID,
})
require.NoError(t, err)
require.True(t, result.PromotedMessage.ModelConfigID.Valid)
require.Equal(t, modelConfigA.ID, result.PromotedMessage.ModelConfigID.UUID)
storedChat, err := db.GetChatByID(ctx, chat.ID)
require.NoError(t, err)
require.Equal(t, modelConfigA.ID, storedChat.LastModelConfigID)
}
func TestPromoteQueuedMessageFallsBackForInvalidQueuedModelConfigID(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
replica := newTestServer(t, db, ps, uuid.New())
ctx := testutil.Context(t, testutil.WaitLong)
user, org, modelConfig := seedChatDependencies(ctx, t, db)
chat, err := db.InsertChat(ctx, database.InsertChatParams{
OrganizationID: org.ID,
Status: database.ChatStatusWaiting,
ClientType: database.ChatClientTypeUi,
OwnerID: user.ID,
LastModelConfigID: modelConfig.ID,
Title: "promote queued invalid fallback",
})
require.NoError(t, err)
queuedContent, err := json.Marshal([]codersdk.ChatMessagePart{codersdk.ChatMessageText("invalid queued model")})
require.NoError(t, err)
queuedMessage, err := db.InsertChatQueuedMessage(ctx, database.InsertChatQueuedMessageParams{
ChatID: chat.ID,
Content: queuedContent,
ModelConfigID: uuid.NullUUID{
UUID: uuid.New(),
Valid: true,
},
})
require.NoError(t, err)
result, err := replica.PromoteQueued(ctx, chatd.PromoteQueuedOptions{
ChatID: chat.ID,
QueuedMessageID: queuedMessage.ID,
CreatedBy: user.ID,
})
require.NoError(t, err)
require.True(t, result.PromotedMessage.ModelConfigID.Valid)
require.Equal(t, modelConfig.ID, result.PromotedMessage.ModelConfigID.UUID)
storedChat, err := db.GetChatByID(ctx, chat.ID)
require.NoError(t, err)
require.Equal(t, modelConfig.ID, storedChat.LastModelConfigID)
}
func TestInterruptAutoPromotionIgnoresLaterUsageLimitIncrease(t *testing.T) {
t.Parallel()