feat: paginate chat messages endpoint with cursor-based infinite scroll (#23083)

Adds cursor-based pagination to the chat messages endpoint.

## Backend

- New `GetChatMessagesByChatIDPaginated` SQL query: returns messages in
`id DESC` order with a `before_id` keyset cursor and configurable
`limit`
- Handler parses `?before_id=N&limit=N` query params, uses the `LIMIT
N+1` trick to set `has_more` without a separate COUNT query
- Queued messages only returned on the first page (no cursor) since
they're always the most recent
- SDK client updated with `ChatMessagesPaginationOptions`
- Fully backward compatible: omitting params returns the 50 newest
messages

## Frontend

- Switches `getChatMessages` from `useQuery` to `useInfiniteQuery` with
cursor chaining via `getNextPageParam`
- Pages flattened and sorted by `id` ascending for chronological display
- `MessagesPaginationSentinel` component uses `IntersectionObserver`
(200px rootMargin prefetch) inside the existing `flex-col-reverse`
scroll container
- `flex-col-reverse` handles scroll anchoring natively when older
messages are prepended — no manual `scrollTop` adjustment needed (same
pattern as coder/blink)

## Why cursor-based instead of offset/limit

Offset-based pagination breaks when new messages arrive while paginating
backward (offsets shift, causing duplicates or missed messages). The
`before_id` cursor is stable regardless of inserts — each page is
deterministic.
This commit is contained in:
Kyle Carberry
2026-03-16 16:40:59 +00:00
committed by GitHub
parent 32a894d4a7
commit 741af057dc
20 changed files with 437 additions and 82 deletions
+2 -2
View File
@@ -1105,7 +1105,7 @@ func TestCreateWorkspaceTool_EndToEnd(t *testing.T) {
require.NoError(t, err)
require.Equal(t, workspaceName, workspace.Name)
chatMsgs, err := client.GetChatMessages(ctx, chat.ID)
chatMsgs, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
var foundCreateWorkspaceResult bool
@@ -1276,7 +1276,7 @@ func TestStartWorkspaceTool_EndToEnd(t *testing.T) {
require.NoError(t, err)
require.Equal(t, codersdk.WorkspaceTransitionStart, updatedWorkspace.LatestBuild.Transition)
chatMsgs, err := client.GetChatMessages(ctx, chat.ID)
chatMsgs, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
// Verify start_workspace tool result exists in the chat messages.
+2 -2
View File
@@ -92,7 +92,7 @@ func TestAnthropicWebSearchRoundTrip(t *testing.T) {
// Verify the chat completed and messages were persisted.
chatData, err := client.GetChat(ctx, chat.ID)
require.NoError(t, err)
chatMsgs, err := client.GetChatMessages(ctx, chat.ID)
chatMsgs, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
t.Logf("Chat status after step 1: %s, messages: %d",
chatData.Status, len(chatMsgs.Messages))
@@ -154,7 +154,7 @@ func TestAnthropicWebSearchRoundTrip(t *testing.T) {
// Verify the follow-up completed and produced content.
chatData2, err := client.GetChat(ctx, chat.ID)
require.NoError(t, err)
chatMsgs2, err := client.GetChatMessages(ctx, chat.ID)
chatMsgs2, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
t.Logf("Chat status after step 2: %s, messages: %d",
chatData2.Status, len(chatMsgs2.Messages))
+40 -10
View File
@@ -573,9 +573,29 @@ func (api *API) getChatMessages(rw http.ResponseWriter, r *http.Request) {
chat := httpmw.ChatParam(r)
chatID := chat.ID
messages, err := api.Database.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
ChatID: chatID,
AfterID: 0,
// Parse optional cursor-based pagination parameters.
queryParams := r.URL.Query()
parser := httpapi.NewQueryParamParser()
beforeID := parser.PositiveInt64(queryParams, 0, "before_id")
limit := parser.PositiveInt32(queryParams, 50, "limit")
if len(parser.Errors) > 0 {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "Query parameters have invalid values.",
Validations: parser.Errors,
})
return
}
if limit < 1 || limit > 200 {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "Invalid limit parameter (1-200).",
})
return
}
// Fetch limit+1 rows to detect whether more pages exist.
messages, err := api.Database.GetChatMessagesByChatIDDescPaginated(ctx, database.GetChatMessagesByChatIDDescPaginatedParams{
ChatID: chatID,
BeforeID: beforeID,
LimitVal: limit + 1,
})
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
@@ -585,18 +605,28 @@ func (api *API) getChatMessages(rw http.ResponseWriter, r *http.Request) {
return
}
queuedMessages, err := api.Database.GetChatQueuedMessages(ctx, chatID)
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to get queued messages.",
Detail: err.Error(),
})
return
hasMore := len(messages) > int(limit)
if hasMore {
messages = messages[:limit]
}
// Only fetch queued messages on the first page (no cursor).
var queuedMessages []database.ChatQueuedMessage
if beforeID == 0 {
queuedMessages, err = api.Database.GetChatQueuedMessages(ctx, chatID)
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to get queued messages.",
Detail: err.Error(),
})
return
}
}
httpapi.Write(ctx, rw, http.StatusOK, codersdk.ChatMessagesResponse{
Messages: convertChatMessages(messages),
QueuedMessages: convertChatQueuedMessages(queuedMessages),
HasMore: hasMore,
})
}
+17 -17
View File
@@ -89,7 +89,7 @@ func TestPostChats(t *testing.T) {
chatResult, err := client.GetChat(ctx, chat.ID)
require.NoError(t, err)
messagesResult, err := client.GetChatMessages(ctx, chat.ID)
messagesResult, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
require.Equal(t, chat.ID, chatResult.ID)
@@ -127,7 +127,7 @@ func TestPostChats(t *testing.T) {
})
require.NoError(t, err)
messagesResult, err := client.GetChatMessages(ctx, chat.ID)
messagesResult, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
for _, message := range messagesResult.Messages {
require.NotEqual(t, codersdk.ChatMessageRoleSystem, message.Role)
@@ -1484,7 +1484,7 @@ func TestGetChat(t *testing.T) {
chatResult, err := client.GetChat(ctx, createdChat.ID)
require.NoError(t, err)
messagesResult, err := client.GetChatMessages(ctx, createdChat.ID)
messagesResult, err := client.GetChatMessages(ctx, createdChat.ID, nil)
require.NoError(t, err)
require.Equal(t, createdChat.ID, chatResult.ID)
require.Equal(t, firstUser.UserID, chatResult.OwnerID)
@@ -1808,7 +1808,7 @@ func TestPostChatMessages(t *testing.T) {
require.True(t, hasTextPart(created.QueuedMessage.Content, messageText))
require.Eventually(t, func() bool {
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID)
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID, nil)
if getErr != nil {
return false
}
@@ -1836,7 +1836,7 @@ func TestPostChatMessages(t *testing.T) {
require.True(t, hasTextPart(created.Message.Content, messageText))
require.Eventually(t, func() bool {
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID)
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID, nil)
if getErr != nil {
return false
}
@@ -1951,7 +1951,7 @@ func TestChatMessageWithFileReferences(t *testing.T) {
var found bool
require.Eventually(t, func() bool {
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID)
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID, nil)
if getErr != nil {
return false
}
@@ -2011,7 +2011,7 @@ func TestChatMessageWithFileReferences(t *testing.T) {
}
require.Eventually(t, func() bool {
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID)
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID, nil)
if getErr != nil {
return false
}
@@ -2064,7 +2064,7 @@ func TestChatMessageWithFileReferences(t *testing.T) {
}
require.Eventually(t, func() bool {
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID)
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID, nil)
if getErr != nil {
return false
}
@@ -2117,7 +2117,7 @@ func TestChatMessageWithFileReferences(t *testing.T) {
}
require.Eventually(t, func() bool {
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID)
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID, nil)
if getErr != nil {
return false
}
@@ -2207,7 +2207,7 @@ func TestChatMessageWithFileReferences(t *testing.T) {
}
require.Eventually(t, func() bool {
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID)
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID, nil)
if getErr != nil {
return false
}
@@ -2397,7 +2397,7 @@ func TestChatMessageWithFiles(t *testing.T) {
}
// Verify file parts omit inline data in the API response.
messagesResult, err := client.GetChatMessages(ctx, chat.ID)
messagesResult, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
for _, msg := range messagesResult.Messages {
for _, part := range msg.Content {
@@ -2493,7 +2493,7 @@ func TestPatchChatMessage(t *testing.T) {
})
require.NoError(t, err)
messagesResult, err := client.GetChatMessages(ctx, chat.ID)
messagesResult, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
var userMessageID int64
@@ -2525,7 +2525,7 @@ func TestPatchChatMessage(t *testing.T) {
}
require.True(t, foundEditedText)
messagesResult, err = client.GetChatMessages(ctx, chat.ID)
messagesResult, err = client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
foundEditedInChat := false
foundOriginalInChat := false
@@ -2578,7 +2578,7 @@ func TestPatchChatMessage(t *testing.T) {
require.NoError(t, err)
// Find the user message ID.
messagesResult, err := client.GetChatMessages(ctx, chat.ID)
messagesResult, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
var userMessageID int64
@@ -2621,7 +2621,7 @@ func TestPatchChatMessage(t *testing.T) {
require.True(t, foundFile, "edited message should preserve file_id")
// GET the chat messages and verify the file_id persists.
messagesResult, err = client.GetChatMessages(ctx, chat.ID)
messagesResult, err = client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
var foundTextInChat, foundFileInChat bool
@@ -3236,7 +3236,7 @@ func TestDeleteChatQueuedMessage(t *testing.T) {
res.Body.Close()
require.Equal(t, http.StatusNoContent, res.StatusCode)
messagesResult, err := client.GetChatMessages(ctx, chat.ID)
messagesResult, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
for _, queued := range messagesResult.QueuedMessages {
require.NotEqual(t, queuedMessage.ID, queued.ID)
@@ -3339,7 +3339,7 @@ func TestPromoteChatQueuedMessage(t *testing.T) {
}
require.True(t, foundPromotedText)
messagesResult, err := client.GetChatMessages(ctx, chat.ID)
messagesResult, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
for _, queued := range messagesResult.QueuedMessages {
require.NotEqual(t, queuedMessage.ID, queued.ID)
+8
View File
@@ -2532,6 +2532,14 @@ func (q *querier) GetChatMessagesByChatID(ctx context.Context, arg database.GetC
return q.db.GetChatMessagesByChatID(ctx, arg)
}
func (q *querier) GetChatMessagesByChatIDDescPaginated(ctx context.Context, arg database.GetChatMessagesByChatIDDescPaginatedParams) ([]database.ChatMessage, error) {
_, err := q.GetChatByID(ctx, arg.ChatID)
if err != nil {
return nil, err
}
return q.db.GetChatMessagesByChatIDDescPaginated(ctx, arg)
}
func (q *querier) GetChatMessagesForPromptByChatID(ctx context.Context, chatID uuid.UUID) ([]database.ChatMessage, error) {
// Authorize read on the parent chat.
_, err := q.GetChatByID(ctx, chatID)
+8
View File
@@ -558,6 +558,14 @@ func (s *MethodTestSuite) TestChats() {
dbm.EXPECT().GetChatMessagesByChatID(gomock.Any(), arg).Return(msgs, nil).AnyTimes()
check.Args(arg).Asserts(chat, policy.ActionRead).Returns(msgs)
}))
s.Run("GetChatMessagesByChatIDDescPaginated", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
chat := testutil.Fake(s.T(), faker, database.Chat{})
msgs := []database.ChatMessage{testutil.Fake(s.T(), faker, database.ChatMessage{ChatID: chat.ID})}
arg := database.GetChatMessagesByChatIDDescPaginatedParams{ChatID: chat.ID, BeforeID: 0, LimitVal: 50}
dbm.EXPECT().GetChatByID(gomock.Any(), chat.ID).Return(chat, nil).AnyTimes()
dbm.EXPECT().GetChatMessagesByChatIDDescPaginated(gomock.Any(), arg).Return(msgs, nil).AnyTimes()
check.Args(arg).Asserts(chat, policy.ActionRead).Returns(msgs)
}))
s.Run("GetLastChatMessageByRole", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
chat := testutil.Fake(s.T(), faker, database.Chat{})
msg := testutil.Fake(s.T(), faker, database.ChatMessage{ChatID: chat.ID})
@@ -1063,6 +1063,14 @@ func (m queryMetricsStore) GetChatMessagesByChatID(ctx context.Context, chatID d
return r0, r1
}
func (m queryMetricsStore) GetChatMessagesByChatIDDescPaginated(ctx context.Context, arg database.GetChatMessagesByChatIDDescPaginatedParams) ([]database.ChatMessage, error) {
start := time.Now()
r0, r1 := m.s.GetChatMessagesByChatIDDescPaginated(ctx, arg)
m.queryLatencies.WithLabelValues("GetChatMessagesByChatIDDescPaginated").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatMessagesByChatIDDescPaginated").Inc()
return r0, r1
}
func (m queryMetricsStore) GetChatMessagesForPromptByChatID(ctx context.Context, chatID uuid.UUID) ([]database.ChatMessage, error) {
start := time.Now()
r0, r1 := m.s.GetChatMessagesForPromptByChatID(ctx, chatID)
+15
View File
@@ -1928,6 +1928,21 @@ func (mr *MockStoreMockRecorder) GetChatMessagesByChatID(ctx, arg any) *gomock.C
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatMessagesByChatID", reflect.TypeOf((*MockStore)(nil).GetChatMessagesByChatID), ctx, arg)
}
// GetChatMessagesByChatIDDescPaginated mocks base method.
func (m *MockStore) GetChatMessagesByChatIDDescPaginated(ctx context.Context, arg database.GetChatMessagesByChatIDDescPaginatedParams) ([]database.ChatMessage, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetChatMessagesByChatIDDescPaginated", ctx, arg)
ret0, _ := ret[0].([]database.ChatMessage)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetChatMessagesByChatIDDescPaginated indicates an expected call of GetChatMessagesByChatIDDescPaginated.
func (mr *MockStoreMockRecorder) GetChatMessagesByChatIDDescPaginated(ctx, arg any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatMessagesByChatIDDescPaginated", reflect.TypeOf((*MockStore)(nil).GetChatMessagesByChatIDDescPaginated), ctx, arg)
}
// GetChatMessagesForPromptByChatID mocks base method.
func (m *MockStore) GetChatMessagesForPromptByChatID(ctx context.Context, chatID uuid.UUID) ([]database.ChatMessage, error) {
m.ctrl.T.Helper()
+1
View File
@@ -234,6 +234,7 @@ type sqlcQuerier interface {
GetChatFilesByIDs(ctx context.Context, ids []uuid.UUID) ([]ChatFile, error)
GetChatMessageByID(ctx context.Context, id int64) (ChatMessage, error)
GetChatMessagesByChatID(ctx context.Context, arg GetChatMessagesByChatIDParams) ([]ChatMessage, error)
GetChatMessagesByChatIDDescPaginated(ctx context.Context, arg GetChatMessagesByChatIDDescPaginatedParams) ([]ChatMessage, error)
GetChatMessagesForPromptByChatID(ctx context.Context, chatID uuid.UUID) ([]ChatMessage, error)
GetChatModelConfigByID(ctx context.Context, id uuid.UUID) (ChatModelConfig, error)
GetChatModelConfigs(ctx context.Context) ([]ChatModelConfig, error)
+66
View File
@@ -3814,6 +3814,72 @@ func (q *sqlQuerier) GetChatMessagesByChatID(ctx context.Context, arg GetChatMes
return items, nil
}
const getChatMessagesByChatIDDescPaginated = `-- name: GetChatMessagesByChatIDDescPaginated :many
SELECT
id, chat_id, model_config_id, created_at, role, content, visibility, input_tokens, output_tokens, total_tokens, reasoning_tokens, cache_creation_tokens, cache_read_tokens, context_limit, compressed, created_by, content_version, total_cost_micros
FROM
chat_messages
WHERE
chat_id = $1::uuid
AND CASE
WHEN $2::bigint > 0 THEN id < $2::bigint
ELSE true
END
AND visibility IN ('user', 'both')
ORDER BY
id DESC
LIMIT
COALESCE(NULLIF($3::int, 0), 50)
`
type GetChatMessagesByChatIDDescPaginatedParams struct {
ChatID uuid.UUID `db:"chat_id" json:"chat_id"`
BeforeID int64 `db:"before_id" json:"before_id"`
LimitVal int32 `db:"limit_val" json:"limit_val"`
}
func (q *sqlQuerier) GetChatMessagesByChatIDDescPaginated(ctx context.Context, arg GetChatMessagesByChatIDDescPaginatedParams) ([]ChatMessage, error) {
rows, err := q.db.QueryContext(ctx, getChatMessagesByChatIDDescPaginated, arg.ChatID, arg.BeforeID, arg.LimitVal)
if err != nil {
return nil, err
}
defer rows.Close()
var items []ChatMessage
for rows.Next() {
var i ChatMessage
if err := rows.Scan(
&i.ID,
&i.ChatID,
&i.ModelConfigID,
&i.CreatedAt,
&i.Role,
&i.Content,
&i.Visibility,
&i.InputTokens,
&i.OutputTokens,
&i.TotalTokens,
&i.ReasoningTokens,
&i.CacheCreationTokens,
&i.CacheReadTokens,
&i.ContextLimit,
&i.Compressed,
&i.CreatedBy,
&i.ContentVersion,
&i.TotalCostMicros,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const getChatMessagesForPromptByChatID = `-- name: GetChatMessagesForPromptByChatID :many
WITH latest_compressed_summary AS (
SELECT
+17
View File
@@ -40,6 +40,23 @@ WHERE
ORDER BY
created_at ASC;
-- name: GetChatMessagesByChatIDDescPaginated :many
SELECT
*
FROM
chat_messages
WHERE
chat_id = @chat_id::uuid
AND CASE
WHEN @before_id::bigint > 0 THEN id < @before_id::bigint
ELSE true
END
AND visibility IN ('user', 'both')
ORDER BY
id DESC
LIMIT
COALESCE(NULLIF(@limit_val::int, 0), 50);
-- name: GetChatMessagesForPromptByChatID :many
WITH latest_compressed_summary AS (
SELECT
+24 -2
View File
@@ -272,6 +272,7 @@ type UploadChatFileResponse struct {
type ChatMessagesResponse struct {
Messages []ChatMessage `json:"messages"`
QueuedMessages []ChatQueuedMessage `json:"queued_messages"`
HasMore bool `json:"has_more"`
}
// ChatModelProviderUnavailableReason explains why a provider cannot be used.
@@ -1243,8 +1244,29 @@ func (c *Client) GetChat(ctx context.Context, chatID uuid.UUID) (Chat, error) {
}
// GetChatMessages returns the messages and queued messages for a chat.
func (c *Client) GetChatMessages(ctx context.Context, chatID uuid.UUID) (ChatMessagesResponse, error) {
res, err := c.Request(ctx, http.MethodGet, fmt.Sprintf("/api/experimental/chats/%s/messages", chatID), nil)
// ChatMessagesPaginationOptions are optional pagination params for
// GetChatMessages.
type ChatMessagesPaginationOptions struct {
BeforeID int64
Limit int
}
// GetChatMessages returns the messages and queued messages for a chat.
func (c *Client) GetChatMessages(ctx context.Context, chatID uuid.UUID, opts *ChatMessagesPaginationOptions) (ChatMessagesResponse, error) {
reqOpts := []RequestOption{}
if opts != nil {
reqOpts = append(reqOpts, func(r *http.Request) {
q := r.URL.Query()
if opts.BeforeID > 0 {
q.Set("before_id", strconv.FormatInt(opts.BeforeID, 10))
}
if opts.Limit > 0 {
q.Set("limit", strconv.Itoa(opts.Limit))
}
r.URL.RawQuery = q.Encode()
})
}
res, err := c.Request(ctx, http.MethodGet, fmt.Sprintf("/api/experimental/chats/%s/messages", chatID), nil, reqOpts...)
if err != nil {
return ChatMessagesResponse{}, err
}
+11 -3
View File
@@ -2960,10 +2960,18 @@ class ApiMethods {
};
getChatMessages = async (
chatId: string,
opts?: { before_id?: number; limit?: number },
): Promise<TypesGen.ChatMessagesResponse> => {
const response = await this.axios.get<TypesGen.ChatMessagesResponse>(
`/api/experimental/chats/${chatId}/messages`,
);
const params = new URLSearchParams();
if (opts?.before_id) {
params.set("before_id", opts.before_id.toString());
}
if (opts?.limit) {
params.set("limit", opts.limit.toString());
}
const query = params.toString();
const url = `/api/experimental/chats/${chatId}/messages${query ? `?${query}` : ""}`;
const response = await this.axios.get<TypesGen.ChatMessagesResponse>(url);
return response.data;
};
+18 -2
View File
@@ -150,9 +150,25 @@ export const chat = (chatId: string) => ({
queryFn: () => API.getChat(chatId),
});
export const chatMessages = (chatId: string) => ({
const MESSAGES_PAGE_SIZE = 50;
export const chatMessagesForInfiniteScroll = (chatId: string) => ({
queryKey: chatMessagesKey(chatId),
queryFn: () => API.getChatMessages(chatId),
initialPageParam: undefined as number | undefined,
queryFn: ({ pageParam }: { pageParam: number | undefined }) =>
API.getChatMessages(chatId, {
before_id: pageParam,
limit: MESSAGES_PAGE_SIZE,
}),
getNextPageParam: (lastPage: TypesGen.ChatMessagesResponse) => {
if (!lastPage.has_more || lastPage.messages.length === 0) {
return undefined;
}
// The API returns messages in DESC order (newest first).
// The last item in the array is the oldest in this page.
// Use its ID as the cursor for the next (older) page.
return lastPage.messages[lastPage.messages.length - 1].id;
},
});
export const archiveChat = (queryClient: QueryClient) => ({
+12
View File
@@ -1359,6 +1359,17 @@ export interface ChatMessageUsage {
readonly context_limit?: number;
}
// From codersdk/chats.go
/**
* GetChatMessages returns the messages and queued messages for a chat.
* ChatMessagesPaginationOptions are optional pagination params for
* GetChatMessages.
*/
export interface ChatMessagesPaginationOptions {
readonly BeforeID: number;
readonly Limit: number;
}
// From codersdk/chats.go
/**
* ChatMessagesResponse contains the messages and queued messages for a chat.
@@ -1366,6 +1377,7 @@ export interface ChatMessageUsage {
export interface ChatMessagesResponse {
readonly messages: readonly ChatMessage[];
readonly queued_messages: readonly ChatQueuedMessage[];
readonly has_more: boolean;
}
// From codersdk/chats.go
@@ -153,7 +153,10 @@ const buildQueries = (
};
return [
{ key: chatKey(CHAT_ID), data: chatWithDiffStatus },
{ key: chatMessagesKey(CHAT_ID), data: messagesData },
{
key: chatMessagesKey(CHAT_ID),
data: { pages: [messagesData], pageParams: [undefined] },
},
{ key: chatsKey, data: [chatWithDiffStatus] },
{
key: chatDiffContentsKey(CHAT_ID),
@@ -532,6 +535,7 @@ export const WithMessageHistory: Story = {
},
],
queued_messages: [],
has_more: false,
},
{ diffUrl: undefined },
),
@@ -555,7 +559,7 @@ export const CompletedWithDiffPanel: Story = {
title: "Build a feature",
status: "completed",
},
{ messages: [], queued_messages: [] },
{ messages: [], queued_messages: [], has_more: false },
{ diffUrl: "https://github.com/coder/coder/pull/123" },
),
},
@@ -590,7 +594,7 @@ export const NoDiffUrl: Story = {
title: "No diff yet",
status: "completed",
},
{ messages: [], queued_messages: [] },
{ messages: [], queued_messages: [], has_more: false },
{ diffUrl: undefined },
),
},
@@ -634,6 +638,7 @@ export const WithSubagentCards: Story = {
},
],
queued_messages: [],
has_more: false,
},
{ diffUrl: undefined },
),
@@ -675,6 +680,7 @@ export const WithReasoningCollapsed: Story = {
},
],
queued_messages: [],
has_more: false,
},
{ diffUrl: undefined },
),
@@ -710,7 +716,7 @@ export const StreamedSubagentTitle: Story = {
title: "Streaming title",
status: "running",
},
{ messages: [], queued_messages: [] },
{ messages: [], queued_messages: [], has_more: false },
{ diffUrl: undefined },
),
webSocket: {
@@ -762,7 +768,7 @@ export const SidebarWithPRAndRepos: Story = {
title: "Full sidebar demo",
status: "completed",
},
{ messages: [], queued_messages: [] },
{ messages: [], queued_messages: [], has_more: false },
{ diffUrl: "https://github.com/coder/coder/pull/456" },
),
webSocket: {
@@ -943,7 +949,7 @@ export const SidebarWithSingleRepo: Story = {
title: "Single repo sidebar",
status: "completed",
},
{ messages: [], queued_messages: [] },
{ messages: [], queued_messages: [], has_more: false },
{ diffUrl: undefined },
),
webSocket: {
@@ -1005,7 +1011,7 @@ export const StreamedReasoningCollapsed: Story = {
title: "Streaming reasoning title",
status: "running",
},
{ messages: [], queued_messages: [] },
{ messages: [], queued_messages: [], has_more: false },
{ diffUrl: undefined },
),
webSocket: {
+40 -8
View File
@@ -1,7 +1,7 @@
import { API, watchWorkspace } from "api/api";
import {
chat,
chatMessages,
chatMessagesForInfiniteScroll,
chatModelConfigs,
chatModels,
chats,
@@ -30,7 +30,12 @@ import {
useRef,
useState,
} from "react";
import { useMutation, useQuery, useQueryClient } from "react-query";
import {
useInfiniteQuery,
useMutation,
useQuery,
useQueryClient,
} from "react-query";
import { useNavigate, useOutletContext, useParams } from "react-router";
import { toast } from "sonner";
import type { UrlTransform } from "streamdown";
@@ -603,8 +608,8 @@ const AgentDetail: FC = () => {
...chat(agentId ?? ""),
enabled: Boolean(agentId),
});
const chatMessagesQuery = useQuery({
...chatMessages(agentId ?? ""),
const chatMessagesQuery = useInfiniteQuery({
...chatMessagesForInfiniteScroll(agentId ?? ""),
enabled: Boolean(agentId),
});
const chatsQuery = useQuery(chats());
@@ -670,10 +675,34 @@ const AgentDetail: FC = () => {
);
const chatRecord = chatQuery.data;
const chatMessagesData = chatMessagesQuery.data;
// Flatten paginated messages into chronological order.
// Pages arrive newest-first per page, and pages[0] is the
// most recent page.
const chatMessagesList = useMemo(() => {
const pages = chatMessagesQuery.data?.pages;
if (!pages || pages.length === 0) return undefined;
// Collect all messages, then sort chronologically by ID.
const all = pages.flatMap((p) => p.messages);
// Sort ascending by ID for chronological order.
all.sort((a, b) => a.id - b.id);
return all;
}, [chatMessagesQuery.data]);
// Queued messages are only in the first page (most recent).
const chatQueuedMessages = chatMessagesQuery.data?.pages[0]?.queued_messages;
// Build a synthetic ChatMessagesResponse from the flattened
// data for backward compat with useChatStore.
const chatMessagesData: TypesGen.ChatMessagesResponse | undefined =
useMemo(() => {
if (!chatMessagesList) return undefined;
return {
messages: chatMessagesList,
queued_messages: chatQueuedMessages ?? [],
has_more: chatMessagesQuery.data?.pages.at(-1)?.has_more ?? false,
};
}, [chatMessagesList, chatQueuedMessages, chatMessagesQuery.data]);
const isArchived = chatRecord?.archived ?? false;
const chatMessagesList = chatMessagesData?.messages;
const chatQueuedMessages = chatMessagesData?.queued_messages;
const chatLastModelConfigID = chatRecord?.last_model_config_id;
const modelOptions = useMemo(
@@ -1093,7 +1122,7 @@ const AgentDetail: FC = () => {
);
}
if (!chatQuery.data || !chatMessagesQuery.data || !agentId) {
if (!chatQuery.data || !chatMessagesQuery.data?.pages?.length || !agentId) {
return (
<AgentDetailNotFoundView
titleElement={titleElement}
@@ -1150,6 +1179,9 @@ const AgentDetail: FC = () => {
}
urlTransform={urlTransform}
scrollContainerRef={scrollContainerRef}
hasMoreMessages={chatMessagesQuery.hasNextPage ?? false}
isFetchingMoreMessages={chatMessagesQuery.isFetchingNextPage}
onFetchMoreMessages={chatMessagesQuery.fetchNextPage}
desktopChatId={agentId}
/>
);
@@ -253,6 +253,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -333,6 +334,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -407,6 +409,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -500,6 +503,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -573,6 +577,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -647,6 +652,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -734,6 +740,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [queuedMessage],
has_more: false,
},
chatQueuedMessages: [queuedMessage],
setChatErrorReason,
@@ -777,6 +784,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [queuedMessage],
has_more: false,
},
chatQueuedMessages: [queuedMessage],
});
@@ -810,6 +818,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [queuedMessage],
has_more: false,
},
chatQueuedMessages: [queuedMessage],
setChatErrorReason,
@@ -844,6 +853,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
});
@@ -875,8 +885,14 @@ describe("useChatStore", () => {
const initialChatMessagesData: TypesGen.ChatMessagesResponse = {
messages: [existingMessage],
queued_messages: [queuedMessage],
has_more: false,
};
queryClient.setQueryData(chatMessagesKey(chatID), initialChatMessagesData);
// The cache is InfiniteData<ChatMessagesResponse> after the
// migration to useInfiniteQuery for chat messages.
queryClient.setQueryData(chatMessagesKey(chatID), {
pages: [initialChatMessagesData],
pageParams: [undefined],
});
const wrapper = ({ children }: PropsWithChildren) => (
<QueryClientProvider client={queryClient}>{children}</QueryClientProvider>
@@ -917,11 +933,11 @@ describe("useChatStore", () => {
await waitFor(() => {
expect(result.current.queuedMessages).toEqual([]);
});
expect(
queryClient.getQueryData<TypesGen.ChatMessagesResponse | undefined>(
chatMessagesKey(chatID),
)?.queued_messages,
).toEqual([]);
const cachedData = queryClient.getQueryData<{
pages: TypesGen.ChatMessagesResponse[];
pageParams: unknown[];
}>(chatMessagesKey(chatID));
expect(cachedData?.pages[0]?.queued_messages).toEqual([]);
});
it("closes old WebSocket and resets state when chatID changes", async () => {
@@ -956,6 +972,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [msg1],
queued_messages: [] as TypesGen.ChatQueuedMessage[],
has_more: false,
},
chatQueuedMessages: [] as TypesGen.ChatQueuedMessage[],
setChatErrorReason,
@@ -1004,6 +1021,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [msg2],
queued_messages: [],
has_more: false,
},
});
@@ -1041,6 +1059,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [queuedMessage],
has_more: false,
},
chatQueuedMessages: [queuedMessage],
setChatErrorReason,
@@ -1096,6 +1115,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -1197,6 +1217,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -1292,6 +1313,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [msg1],
queued_messages: [] as TypesGen.ChatQueuedMessage[],
has_more: false,
},
chatQueuedMessages: [] as TypesGen.ChatQueuedMessage[],
setChatErrorReason,
@@ -1340,6 +1362,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [msg2],
queued_messages: [],
has_more: false,
},
});
@@ -1381,6 +1404,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [msg1],
queued_messages: [queuedMsg],
has_more: false,
},
chatQueuedMessages: [queuedMsg],
setChatErrorReason,
@@ -1416,6 +1440,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
});
@@ -1451,6 +1476,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -1520,6 +1546,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -1580,6 +1607,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -1632,6 +1660,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -1692,6 +1721,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -1766,6 +1796,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -1827,6 +1858,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -1898,6 +1930,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -1963,6 +1996,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason: vi.fn(),
@@ -2015,6 +2049,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [msg],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason: vi.fn(),
@@ -2079,6 +2114,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [existingMessage],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -2202,6 +2238,7 @@ describe("useChatStore", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -2270,6 +2307,7 @@ describe("updateSidebarChat via stream events", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -2333,6 +2371,7 @@ describe("updateSidebarChat via stream events", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -2405,6 +2444,7 @@ describe("updateSidebarChat via stream events", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -2470,6 +2510,7 @@ describe("updateSidebarChat via stream events", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -2542,6 +2583,7 @@ describe("updateSidebarChat via stream events", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -2612,6 +2654,7 @@ describe("updateSidebarChat via stream events", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -2677,6 +2720,7 @@ describe("updateSidebarChat via stream events", () => {
chatMessagesData: {
messages: [],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
@@ -9,7 +9,7 @@ import {
useRef,
useSyncExternalStore,
} from "react";
import { useQueryClient } from "react-query";
import { type InfiniteData, useQueryClient } from "react-query";
import type { OneWayMessageEvent } from "utils/OneWayWebSocket";
import { createReconnectingWebSocket } from "utils/reconnectingWebSocket";
import { applyMessagePartToStreamState } from "./streamState";
@@ -519,26 +519,29 @@ export const useChatStore = (
return;
}
const nextQueuedMessages = queuedMessages ?? [];
queryClient.setQueryData<TypesGen.ChatMessagesResponse | undefined>(
chatMessagesKey(chatID),
(currentData) => {
if (!currentData) {
return currentData;
}
if (
chatQueuedMessagesEqualByID(
currentData.queued_messages,
nextQueuedMessages,
)
) {
return currentData;
}
return {
...currentData,
queued_messages: nextQueuedMessages,
};
},
);
queryClient.setQueryData<
InfiniteData<TypesGen.ChatMessagesResponse> | undefined
>(chatMessagesKey(chatID), (currentData) => {
if (!currentData?.pages?.length) {
return currentData;
}
const firstPage = currentData.pages[0];
if (
chatQueuedMessagesEqualByID(
firstPage.queued_messages,
nextQueuedMessages,
)
) {
return currentData;
}
return {
...currentData,
pages: [
{ ...firstPage, queued_messages: nextQueuedMessages },
...currentData.pages.slice(1),
],
};
});
},
[chatID, queryClient],
);
@@ -551,7 +554,15 @@ export const useChatStore = (
prevChatIDRef.current = chatID;
store.replaceMessages([]);
}
store.replaceMessages(chatMessages);
// Merge REST-fetched messages into the store one-by-one instead
// of replacing the entire map. This preserves any messages the
// WebSocket delivered via upsertDurableMessage that haven't
// appeared in a REST page yet.
if (chatMessages) {
for (const message of chatMessages) {
store.upsertDurableMessage(message);
}
}
}, [chatID, chatMessages, store]);
useEffect(() => {
+52 -1
View File
@@ -2,7 +2,7 @@ import type * as TypesGen from "api/typesGenerated";
import type { ChatDiffStatus } from "api/typesGenerated";
import type { ModelSelectorOption } from "components/ai-elements";
import { ArchiveIcon } from "lucide-react";
import { type FC, type RefObject, useState } from "react";
import { type FC, type RefObject, useEffect, useRef, useState } from "react";
import type { UrlTransform } from "streamdown";
import { cn } from "utils/cn";
import { pageTitle } from "utils/page";
@@ -120,6 +120,11 @@ interface AgentDetailViewProps {
// Scroll container ref.
scrollContainerRef: RefObject<HTMLDivElement | null>;
// Pagination for loading older messages.
hasMoreMessages: boolean;
isFetchingMoreMessages: boolean;
onFetchMoreMessages: () => void;
urlTransform?: UrlTransform;
// Desktop chat ID (optional).
@@ -170,6 +175,9 @@ export const AgentDetailView: FC<AgentDetailViewProps> = ({
handleUnarchiveAgentAction,
handleArchiveAndDeleteWorkspaceAction,
scrollContainerRef,
hasMoreMessages,
isFetchingMoreMessages,
onFetchMoreMessages,
urlTransform,
desktopChatId,
}) => {
@@ -263,6 +271,13 @@ export const AgentDetailView: FC<AgentDetailViewProps> = ({
urlTransform={urlTransform}
/>
</div>
{hasMoreMessages && (
<MessagesPaginationSentinel
containerRef={scrollContainerRef}
isFetching={isFetchingMoreMessages}
onLoadMore={onFetchMoreMessages}
/>
)}
</div>
<div className="shrink-0 overflow-y-auto px-4 [scrollbar-gutter:stable] [scrollbar-width:thin]">
<AgentDetailInput
@@ -477,3 +492,39 @@ export const AgentDetailNotFoundView: FC<AgentDetailNotFoundViewProps> = ({
</div>
);
};
/**
* Invisible sentinel that triggers loading older messages when it
* scrolls into view. Placed at the visual top of the flex-col-reverse
* container (which is the DOM bottom).
*/
const MessagesPaginationSentinel: FC<{
containerRef: RefObject<HTMLDivElement | null>;
isFetching: boolean;
onLoadMore: () => void;
}> = ({ containerRef, isFetching, onLoadMore }) => {
const sentinelRef = useRef<HTMLDivElement>(null);
useEffect(() => {
const sentinel = sentinelRef.current;
const container = containerRef.current;
if (!sentinel || !container) return;
const observer = new IntersectionObserver(
([entry]) => {
if (entry.isIntersecting && !isFetching) {
onLoadMore();
}
},
{
root: container,
rootMargin: "200px 0px 0px 0px",
threshold: 0.01,
},
);
observer.observe(sentinel);
return () => observer.disconnect();
}, [containerRef, isFetching, onLoadMore]);
return <div ref={sentinelRef} className="h-px shrink-0" />;
};