mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix: use cursor-based query for chat stream notifications (#22510)
## Problem
The pubsub notification handler in `chatd` re-fetched **all** messages
from the DB on every new message notification, then filtered in Go with
`msg.ID > lastMessageID`. This grows linearly with conversation length —
every new message triggers a full table scan of that chat's history.
The `AfterMessageID` field in the pubsub notification payload was
clearly designed for cursor-based fetching, but no matching query
existed.
## Fix
- Add `GetChatMessagesByChatIDAfter` SQL query with `WHERE id >
@after_id`, so the database does the filtering instead of Go.
- Use it in the pubsub notification handler in `chatd.go`, passing
`lastMessageID` as the cursor.
- Implement the dbauthz wrapper (was a `panic("not implemented")` stub
from codegen) with the same read-check-on-parent-chat pattern as
adjacent methods.
- Add dbauthz test coverage for the new method.
**Not changed:** The initial snapshot in `Subscribe()` still loads all
messages — that's correct, since a newly-connecting client needs the
full conversation state. The waste was only in the ongoing notification
path.
This commit is contained in:
+24
-16
@@ -988,6 +988,7 @@ func (p *Server) Subscribe(
|
||||
ctx context.Context,
|
||||
chatID uuid.UUID,
|
||||
requestHeader http.Header,
|
||||
afterMessageID int64,
|
||||
) (
|
||||
[]codersdk.ChatStreamEvent,
|
||||
<-chan codersdk.ChatStreamEvent,
|
||||
@@ -1013,8 +1014,14 @@ func (p *Server) Subscribe(
|
||||
}
|
||||
}
|
||||
|
||||
// Load initial messages from DB
|
||||
messages, err := p.db.GetChatMessagesByChatID(ctx, chatID)
|
||||
// Load initial messages from DB. When afterMessageID > 0 the
|
||||
// caller already has messages up to that ID (e.g. from the REST
|
||||
// endpoint), so we only fetch newer ones to avoid sending
|
||||
// duplicate data.
|
||||
messages, err := p.db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chatID,
|
||||
AfterID: afterMessageID,
|
||||
})
|
||||
if err == nil {
|
||||
for _, msg := range messages {
|
||||
sdkMsg := db2sdk.ChatMessage(msg)
|
||||
@@ -1191,23 +1198,24 @@ func (p *Server) Subscribe(
|
||||
case notify := <-notifications:
|
||||
// Handle different notification types
|
||||
if notify.AfterMessageID > 0 {
|
||||
// Read new messages from DB
|
||||
messages, err := p.db.GetChatMessagesByChatID(mergedCtx, chatID)
|
||||
// Read only new messages from DB.
|
||||
messages, err := p.db.GetChatMessagesByChatID(mergedCtx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chatID,
|
||||
AfterID: lastMessageID,
|
||||
})
|
||||
if err == nil {
|
||||
for _, msg := range messages {
|
||||
if msg.ID > lastMessageID {
|
||||
sdkMsg := db2sdk.ChatMessage(msg)
|
||||
select {
|
||||
case <-mergedCtx.Done():
|
||||
return
|
||||
case mergedEvents <- codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeMessage,
|
||||
ChatID: chatID,
|
||||
Message: &sdkMsg,
|
||||
}:
|
||||
}
|
||||
lastMessageID = msg.ID
|
||||
sdkMsg := db2sdk.ChatMessage(msg)
|
||||
select {
|
||||
case <-mergedCtx.Done():
|
||||
return
|
||||
case mergedEvents <- codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeMessage,
|
||||
ChatID: chatID,
|
||||
Message: &sdkMsg,
|
||||
}:
|
||||
}
|
||||
lastMessageID = msg.ID
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+101
-7
@@ -27,6 +27,7 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/util/slice"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/provisioner/echo"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
@@ -60,7 +61,7 @@ func TestInterruptChatBroadcastsStatusAcrossInstances(t *testing.T) {
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, events, cancel, ok := replicaB.Subscribe(ctx, chat.ID, nil)
|
||||
_, events, cancel, ok := replicaB.Subscribe(ctx, chat.ID, nil, 0)
|
||||
require.True(t, ok)
|
||||
t.Cleanup(cancel)
|
||||
|
||||
@@ -202,7 +203,10 @@ func TestSendMessageQueueBehaviorQueuesWhenBusy(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
require.Len(t, queued, 1)
|
||||
|
||||
messages, err := db.GetChatMessagesByChatID(ctx, chat.ID)
|
||||
messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chat.ID,
|
||||
AfterID: 0,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, messages, 1)
|
||||
}
|
||||
@@ -252,7 +256,10 @@ func TestSendMessageInterruptBehaviorSendsImmediatelyWhenBusy(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
require.Len(t, queued, 0)
|
||||
|
||||
messages, err := db.GetChatMessagesByChatID(ctx, chat.ID)
|
||||
messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chat.ID,
|
||||
AfterID: 0,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, messages, 2)
|
||||
require.Equal(t, messages[len(messages)-1].ID, result.Message.ID)
|
||||
@@ -275,7 +282,10 @@ func TestEditMessageUpdatesAndTruncatesAndClearsQueue(t *testing.T) {
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
initialMessages, err := db.GetChatMessagesByChatID(ctx, chat.ID)
|
||||
initialMessages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chat.ID,
|
||||
AfterID: 0,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, initialMessages, 1)
|
||||
editedMessageID := initialMessages[0].ID
|
||||
@@ -322,7 +332,10 @@ func TestEditMessageUpdatesAndTruncatesAndClearsQueue(t *testing.T) {
|
||||
require.Len(t, editedSDK.Content, 1)
|
||||
require.Equal(t, "edited", editedSDK.Content[0].Text)
|
||||
|
||||
messages, err := db.GetChatMessagesByChatID(ctx, chat.ID)
|
||||
messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chat.ID,
|
||||
AfterID: 0,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, messages, 1)
|
||||
require.Equal(t, editedMessageID, messages[0].ID)
|
||||
@@ -657,7 +670,7 @@ func TestSubscribeSnapshotIncludesStatusEvent(t *testing.T) {
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
snapshot, _, cancel, ok := replica.Subscribe(ctx, chat.ID, nil)
|
||||
snapshot, _, cancel, ok := replica.Subscribe(ctx, chat.ID, nil, 0)
|
||||
require.True(t, ok)
|
||||
t.Cleanup(cancel)
|
||||
|
||||
@@ -686,7 +699,7 @@ func TestSubscribeNoPubsubNoDuplicateMessageParts(t *testing.T) {
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
snapshot, events, cancel, ok := replica.Subscribe(ctx, chat.ID, nil)
|
||||
snapshot, events, cancel, ok := replica.Subscribe(ctx, chat.ID, nil, 0)
|
||||
require.True(t, ok)
|
||||
t.Cleanup(cancel)
|
||||
|
||||
@@ -708,6 +721,87 @@ func TestSubscribeNoPubsubNoDuplicateMessageParts(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSubscribeAfterMessageID(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
replica := newTestServer(t, db, ps, uuid.New())
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, model := seedChatDependencies(ctx, t, db)
|
||||
|
||||
// Create a chat — this inserts one initial "user" message.
|
||||
chat, err := replica.CreateChat(ctx, chatd.CreateOptions{
|
||||
OwnerID: user.ID,
|
||||
Title: "after-id-test",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "first"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Insert two more messages so we have three total visible
|
||||
// messages (the initial user message plus these two).
|
||||
msg2, err := db.InsertChatMessage(ctx, database.InsertChatMessageParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
Role: "assistant",
|
||||
Content: pqtype.NullRawMessage{RawMessage: json.RawMessage(`"second"`), Valid: true},
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
InputTokens: sql.NullInt64{},
|
||||
OutputTokens: sql.NullInt64{},
|
||||
TotalTokens: sql.NullInt64{},
|
||||
ReasoningTokens: sql.NullInt64{},
|
||||
CacheCreationTokens: sql.NullInt64{},
|
||||
CacheReadTokens: sql.NullInt64{},
|
||||
ContextLimit: sql.NullInt64{},
|
||||
Compressed: sql.NullBool{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = db.InsertChatMessage(ctx, database.InsertChatMessageParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
Role: "user",
|
||||
Content: pqtype.NullRawMessage{RawMessage: json.RawMessage(`"third"`), Valid: true},
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
InputTokens: sql.NullInt64{},
|
||||
OutputTokens: sql.NullInt64{},
|
||||
TotalTokens: sql.NullInt64{},
|
||||
ReasoningTokens: sql.NullInt64{},
|
||||
CacheCreationTokens: sql.NullInt64{},
|
||||
CacheReadTokens: sql.NullInt64{},
|
||||
ContextLimit: sql.NullInt64{},
|
||||
Compressed: sql.NullBool{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Control: Subscribe with afterMessageID=0 returns ALL messages.
|
||||
allSnapshot, _, cancelAll, ok := replica.Subscribe(ctx, chat.ID, nil, 0)
|
||||
require.True(t, ok)
|
||||
cancelAll()
|
||||
|
||||
allMessages := filterMessageEvents(allSnapshot)
|
||||
require.Len(t, allMessages, 3, "afterMessageID=0 should return all three messages")
|
||||
|
||||
// Subscribe with afterMessageID set to the second message's ID.
|
||||
// Only the third message (inserted after msg2) should appear.
|
||||
partialSnapshot, _, cancelPartial, ok := replica.Subscribe(ctx, chat.ID, nil, msg2.ID)
|
||||
require.True(t, ok)
|
||||
cancelPartial()
|
||||
|
||||
partialMessages := filterMessageEvents(partialSnapshot)
|
||||
require.Len(t, partialMessages, 1, "afterMessageID=msg2.ID should return only messages after msg2")
|
||||
require.Equal(t, "user", partialMessages[0].Message.Role)
|
||||
}
|
||||
|
||||
// filterMessageEvents returns only the Message-type events from a
|
||||
// snapshot slice, which is useful for ignoring status / queue events.
|
||||
func filterMessageEvents(events []codersdk.ChatStreamEvent) []codersdk.ChatStreamEvent {
|
||||
return slice.Filter(events, func(e codersdk.ChatStreamEvent) bool {
|
||||
return e.Type == codersdk.ChatStreamEventTypeMessage
|
||||
})
|
||||
}
|
||||
|
||||
func TestCreateWorkspaceTool_EndToEnd(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -397,7 +397,10 @@ func latestSubagentAssistantMessage(
|
||||
store database.Store,
|
||||
chatID uuid.UUID,
|
||||
) (string, error) {
|
||||
messages, err := store.GetChatMessagesByChatID(ctx, chatID)
|
||||
messages, err := store.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chatID,
|
||||
AfterID: 0,
|
||||
})
|
||||
if err != nil {
|
||||
return "", xerrors.Errorf("get chat messages: %w", err)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user