mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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.
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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;
|
||||
};
|
||||
|
||||
|
||||
@@ -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) => ({
|
||||
|
||||
Generated
+12
@@ -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: {
|
||||
|
||||
@@ -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(() => {
|
||||
|
||||
@@ -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" />;
|
||||
};
|
||||
|
||||
Reference in New Issue
Block a user