feat: add after_id pagination for chat messages (#24531)

This commit is contained in:
david-fraley
2026-04-28 08:31:33 -05:00
committed by GitHub
parent 8fe11e9b14
commit 5222db86c7
8 changed files with 511 additions and 13 deletions
+12 -2
View File
@@ -6545,22 +6545,32 @@ WHERE
WHEN $2::bigint > 0 THEN id < $2::bigint
ELSE true
END
AND CASE
WHEN $3::bigint > 0 THEN id > $3::bigint
ELSE true
END
AND visibility IN ('user', 'both')
AND deleted = false
ORDER BY
id DESC
LIMIT
COALESCE(NULLIF($3::int, 0), 50)
COALESCE(NULLIF($4::int, 0), 50)
`
type GetChatMessagesByChatIDDescPaginatedParams struct {
ChatID uuid.UUID `db:"chat_id" json:"chat_id"`
BeforeID int64 `db:"before_id" json:"before_id"`
AfterID int64 `db:"after_id" json:"after_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)
rows, err := q.db.QueryContext(ctx, getChatMessagesByChatIDDescPaginated,
arg.ChatID,
arg.BeforeID,
arg.AfterID,
arg.LimitVal,
)
if err != nil {
return nil, err
}
+4
View File
@@ -266,6 +266,10 @@ WHERE
WHEN @before_id::bigint > 0 THEN id < @before_id::bigint
ELSE true
END
AND CASE
WHEN @after_id::bigint > 0 THEN id > @after_id::bigint
ELSE true
END
AND visibility IN ('user', 'both')
AND deleted = false
ORDER BY
+35 -8
View File
@@ -1702,6 +1702,7 @@ func (api *API) getChatMessages(rw http.ResponseWriter, r *http.Request) {
queryParams := r.URL.Query()
parser := httpapi.NewQueryParamParser()
beforeID := parser.PositiveInt64(queryParams, 0, "before_id")
afterID := parser.PositiveInt64(queryParams, 0, "after_id")
limit := parser.PositiveInt32(queryParams, 50, "limit")
if len(parser.Errors) > 0 {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
@@ -1716,12 +1717,36 @@ func (api *API) getChatMessages(rw http.ResponseWriter, r *http.Request) {
})
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,
})
// Reject transposed or equal cursors so an empty open range is loud,
// not silently indistinguishable from "no messages in this range."
if beforeID > 0 && afterID > 0 && afterID >= beforeID {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "after_id must be less than before_id.",
})
return
}
// Polling with only after_id uses ASC so the cursor advances
// monotonically; a DESC limit would drop rows when a burst larger
// than `limit` lands between polls. Fetch limit+1 in both paths to
// detect whether more pages exist.
var messages []database.ChatMessage
var err error
switch {
case afterID > 0 && beforeID == 0:
messages, err = api.Database.GetChatMessagesByChatIDAscPaginated(ctx, database.GetChatMessagesByChatIDAscPaginatedParams{
ChatID: chatID,
AfterID: afterID,
LimitVal: limit + 1,
})
default:
messages, err = api.Database.GetChatMessagesByChatIDDescPaginated(ctx, database.GetChatMessagesByChatIDDescPaginatedParams{
ChatID: chatID,
BeforeID: beforeID,
AfterID: afterID,
LimitVal: limit + 1,
})
}
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to get chat messages.",
@@ -1735,9 +1760,11 @@ func (api *API) getChatMessages(rw http.ResponseWriter, r *http.Request) {
messages = messages[:limit]
}
// Only fetch queued messages on the first page (no cursor).
// Queued messages are only meaningful for the initial top-of-history
// load. Suppress them whenever any cursor is set so polling callers do
// not receive the snapshot on every page fetch.
var queuedMessages []database.ChatQueuedMessage
if beforeID == 0 {
if beforeID == 0 && afterID == 0 {
queuedMessages, err = api.Database.GetChatQueuedMessages(ctx, chatID)
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
+417
View File
@@ -10,6 +10,7 @@ import (
"mime"
"net/http"
"regexp"
"slices"
"strings"
"sync/atomic"
"testing"
@@ -12186,6 +12187,422 @@ func TestPostChats_DynamicToolValidation(t *testing.T) {
})
}
func TestGetChatMessages_Pagination(t *testing.T) {
t.Parallel()
// seedChat creates a chat and inserts `count` user messages, returning
// the chat and the inserted message IDs in the order they were
// persisted (ascending). Callers use these IDs as cursor values.
seedChat := func(
ctx context.Context,
t *testing.T,
db database.Store,
ownerID uuid.UUID,
organizationID uuid.UUID,
modelConfigID uuid.UUID,
count int,
) (database.Chat, []int64) {
t.Helper()
chat, err := db.InsertChat(dbauthz.AsSystemRestricted(ctx), database.InsertChatParams{
OrganizationID: organizationID,
Status: database.ChatStatusWaiting,
ClientType: database.ChatClientTypeUi,
OwnerID: ownerID,
LastModelConfigID: modelConfigID,
Title: "pagination-test",
})
require.NoError(t, err)
createdBy := make([]uuid.UUID, count)
modelIDs := make([]uuid.UUID, count)
roles := make([]database.ChatMessageRole, count)
contents := make([]string, count)
contentVersions := make([]int16, count)
visibility := make([]database.ChatMessageVisibility, count)
inputTokens := make([]int64, count)
outputTokens := make([]int64, count)
totalTokens := make([]int64, count)
reasoningTokens := make([]int64, count)
cacheCreationTokens := make([]int64, count)
cacheReadTokens := make([]int64, count)
contextLimit := make([]int64, count)
compressed := make([]bool, count)
totalCost := make([]int64, count)
runtime := make([]int64, count)
for i := range count {
part, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
codersdk.ChatMessageText(fmt.Sprintf("msg %d", i)),
})
require.NoError(t, err)
createdBy[i] = ownerID
modelIDs[i] = modelConfigID
roles[i] = database.ChatMessageRoleUser
contents[i] = string(part.RawMessage)
contentVersions[i] = chatprompt.CurrentContentVersion
visibility[i] = database.ChatMessageVisibilityBoth
}
results, err := db.InsertChatMessages(dbauthz.AsSystemRestricted(ctx), database.InsertChatMessagesParams{
ChatID: chat.ID,
CreatedBy: createdBy,
ModelConfigID: modelIDs,
Role: roles,
Content: contents,
ContentVersion: contentVersions,
Visibility: visibility,
InputTokens: inputTokens,
OutputTokens: outputTokens,
TotalTokens: totalTokens,
ReasoningTokens: reasoningTokens,
CacheCreationTokens: cacheCreationTokens,
CacheReadTokens: cacheReadTokens,
ContextLimit: contextLimit,
Compressed: compressed,
TotalCostMicros: totalCost,
RuntimeMs: runtime,
})
require.NoError(t, err)
require.Len(t, results, count)
ids := make([]int64, count)
for i, m := range results {
ids[i] = m.ID
}
return chat, ids
}
seedQueuedMessage := func(
ctx context.Context,
t *testing.T,
db database.Store,
chatID uuid.UUID,
) {
t.Helper()
content, err := json.Marshal([]codersdk.ChatMessagePart{
codersdk.ChatMessageText("queued"),
})
require.NoError(t, err)
_, err = db.InsertChatQueuedMessage(
dbauthz.AsSystemRestricted(ctx),
database.InsertChatQueuedMessageParams{
ChatID: chatID,
Content: content,
},
)
require.NoError(t, err)
}
t.Run("NoCursorReturnsAllDESCPlusQueued", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
user := coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
chat, ids := seedChat(ctx, t, db, user.UserID, user.OrganizationID, modelConfig.ID, 5)
seedQueuedMessage(ctx, t, db, chat.ID)
resp, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
require.Len(t, resp.Messages, 5)
require.False(t, resp.HasMore)
require.Len(t, resp.QueuedMessages, 1)
want := []int64{ids[4], ids[3], ids[2], ids[1], ids[0]}
got := make([]int64, len(resp.Messages))
for i, m := range resp.Messages {
got[i] = m.ID
}
require.Equal(t, want, got)
})
t.Run("BeforeIDReturnsOlderAndSuppressesQueued", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
user := coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
chat, ids := seedChat(ctx, t, db, user.UserID, user.OrganizationID, modelConfig.ID, 5)
seedQueuedMessage(ctx, t, db, chat.ID)
resp, err := client.GetChatMessages(ctx, chat.ID, &codersdk.ChatMessagesPaginationOptions{
BeforeID: ids[2],
})
require.NoError(t, err)
require.False(t, resp.HasMore)
require.Empty(t, resp.QueuedMessages)
want := []int64{ids[1], ids[0]}
got := make([]int64, len(resp.Messages))
for i, m := range resp.Messages {
got[i] = m.ID
}
require.Equal(t, want, got)
})
t.Run("AfterIDReturnsNewerInASCOrderForMonotonicPolling", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
user := coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
chat, ids := seedChat(ctx, t, db, user.UserID, user.OrganizationID, modelConfig.ID, 5)
seedQueuedMessage(ctx, t, db, chat.ID)
resp, err := client.GetChatMessages(ctx, chat.ID, &codersdk.ChatMessagesPaginationOptions{
AfterID: ids[1],
})
require.NoError(t, err)
require.False(t, resp.HasMore)
require.Empty(t, resp.QueuedMessages)
// ASC order so a polling caller can advance its cursor to
// max(returned_ids) without gaps.
want := []int64{ids[2], ids[3], ids[4]}
got := make([]int64, len(resp.Messages))
for i, m := range resp.Messages {
got[i] = m.ID
}
require.Equal(t, want, got)
})
t.Run("AfterAndBeforeIDReturnsOpenRange", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
user := coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
chat, ids := seedChat(ctx, t, db, user.UserID, user.OrganizationID, modelConfig.ID, 5)
seedQueuedMessage(ctx, t, db, chat.ID)
resp, err := client.GetChatMessages(ctx, chat.ID, &codersdk.ChatMessagesPaginationOptions{
AfterID: ids[0],
BeforeID: ids[4],
})
require.NoError(t, err)
require.False(t, resp.HasMore)
require.Empty(t, resp.QueuedMessages)
want := []int64{ids[3], ids[2], ids[1]}
got := make([]int64, len(resp.Messages))
for i, m := range resp.Messages {
got[i] = m.ID
}
require.Equal(t, want, got)
})
t.Run("LimitCapsAfterIDPageToOldestAndSetsHasMore", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
user := coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
chat, ids := seedChat(ctx, t, db, user.UserID, user.OrganizationID, modelConfig.ID, 5)
// Seed a queued message so the Empty assertion below verifies
// the cursor suppresses queued rows, not just that none exist.
seedQueuedMessage(ctx, t, db, chat.ID)
resp, err := client.GetChatMessages(ctx, chat.ID, &codersdk.ChatMessagesPaginationOptions{
AfterID: ids[0],
Limit: 2,
})
require.NoError(t, err)
require.True(t, resp.HasMore)
require.Empty(t, resp.QueuedMessages)
// The ASC polling path returns the OLDEST unseen messages
// first. A burst larger than `limit` would otherwise silently
// drop the oldest rows between polls on the DESC path.
want := []int64{ids[1], ids[2]}
got := make([]int64, len(resp.Messages))
for i, m := range resp.Messages {
got[i] = m.ID
}
require.Equal(t, want, got)
})
t.Run("NegativeAfterIDReturns400", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
user := coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
chat, _ := seedChat(ctx, t, db, user.UserID, user.OrganizationID, modelConfig.ID, 1)
res, err := client.Request(
ctx,
http.MethodGet,
fmt.Sprintf("/api/experimental/chats/%s/messages?after_id=-1", chat.ID),
nil,
)
require.NoError(t, err)
defer res.Body.Close()
require.Equal(t, http.StatusBadRequest, res.StatusCode)
var sdkResp codersdk.Response
require.NoError(t, json.NewDecoder(res.Body).Decode(&sdkResp))
require.Equal(t, "Query parameters have invalid values.", sdkResp.Message)
require.True(t,
slices.ContainsFunc(sdkResp.Validations, func(v codersdk.ValidationError) bool {
return v.Field == "after_id"
}),
"expected validation error for after_id field",
)
})
t.Run("NonNumericAfterIDReturns400", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
user := coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
chat, _ := seedChat(ctx, t, db, user.UserID, user.OrganizationID, modelConfig.ID, 1)
res, err := client.Request(
ctx,
http.MethodGet,
fmt.Sprintf("/api/experimental/chats/%s/messages?after_id=abc", chat.ID),
nil,
)
require.NoError(t, err)
defer res.Body.Close()
require.Equal(t, http.StatusBadRequest, res.StatusCode)
var sdkResp codersdk.Response
require.NoError(t, json.NewDecoder(res.Body).Decode(&sdkResp))
require.Equal(t, "Query parameters have invalid values.", sdkResp.Message)
require.True(t,
slices.ContainsFunc(sdkResp.Validations, func(v codersdk.ValidationError) bool {
return v.Field == "after_id"
}),
"expected validation error for after_id field",
)
})
t.Run("AfterIDAtOrAboveMaxReturnsEmpty", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
user := coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
chat, ids := seedChat(ctx, t, db, user.UserID, user.OrganizationID, modelConfig.ID, 3)
// Seed a queued message to prove the cursor path suppresses
// it even when nothing else comes back.
seedQueuedMessage(ctx, t, db, chat.ID)
// The steady-state polling case: the caller already has every
// message, so after_id equals the largest seen id. The server
// must return an empty page, not the last row again.
resp, err := client.GetChatMessages(ctx, chat.ID, &codersdk.ChatMessagesPaginationOptions{
AfterID: ids[len(ids)-1],
})
require.NoError(t, err)
require.Empty(t, resp.Messages)
require.False(t, resp.HasMore)
require.Empty(t, resp.QueuedMessages)
})
t.Run("AfterIDGreaterThanOrEqualBeforeIDReturns400", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
user := coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
chat, ids := seedChat(ctx, t, db, user.UserID, user.OrganizationID, modelConfig.ID, 3)
// Transposed cursors: after >= before. Fail loudly rather
// than return an empty page indistinguishable from
// "no messages in this range."
for _, tc := range []struct {
name string
after int64
before int64
}{
{"Transposed", ids[2], ids[0]},
{"Equal", ids[1], ids[1]},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitShort)
_, err := client.GetChatMessages(ctx, chat.ID, &codersdk.ChatMessagesPaginationOptions{
AfterID: tc.after,
BeforeID: tc.before,
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "after_id must be less than before_id.", sdkErr.Message)
})
}
})
t.Run("AfterIDPollingWalksBurstWithoutGaps", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, db := newChatClientWithDatabase(t)
user := coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
// Simulate a polling client that has already acknowledged the
// first message (cursor = ids[0]) when a burst of
// `burstSize` new messages arrives. With `limit=pageSize` and
// `burstSize > pageSize`, the naive DESC-ordered path would
// silently drop the oldest rows between polls. The ASC
// dispatch lets the client walk the whole burst by advancing
// after_id to max(returned_ids) on each tick.
const burstSize = 60
const pageSize = 25
// Seed burstSize+1 rows; ids[0] is the "already acknowledged"
// message the client saw before the burst.
chat, ids := seedChat(ctx, t, db, user.UserID, user.OrganizationID, modelConfig.ID, burstSize+1)
var seen []int64
cursor := ids[0]
maxPages := (burstSize / pageSize) + 2
for range maxPages {
resp, err := client.GetChatMessages(ctx, chat.ID, &codersdk.ChatMessagesPaginationOptions{
AfterID: cursor,
Limit: pageSize,
})
require.NoError(t, err)
if len(resp.Messages) == 0 {
require.False(t, resp.HasMore)
break
}
for _, m := range resp.Messages {
seen = append(seen, m.ID)
}
// Advance to max(returned). On the ASC path this is the
// last element of the returned slice.
cursor = resp.Messages[len(resp.Messages)-1].ID
if !resp.HasMore {
break
}
}
require.Equal(t, ids[1:], seen,
"polling walk must return every burst row exactly once in ascending order")
})
}
func requireSDKError(t *testing.T, err error, expectedStatus int) *codersdk.Error {
t.Helper()
+11 -1
View File
@@ -2571,7 +2571,14 @@ func (c *ExperimentalClient) GetChat(ctx context.Context, chatID uuid.UUID) (Cha
// GetChatMessages.
type ChatMessagesPaginationOptions struct {
BeforeID int64
Limit int
// AfterID, when > 0, restricts results to messages with id strictly
// greater than AfterID. When set without BeforeID, results come back
// in ASCENDING id order so a polling caller can advance its cursor
// to max(returned_ids) without gaps. When combined with BeforeID,
// results come back in DESC order over the open range
// (AfterID, BeforeID).
AfterID int64
Limit int
}
// GetChatMessages returns the messages and queued messages for a chat.
@@ -2583,6 +2590,9 @@ func (c *ExperimentalClient) GetChatMessages(ctx context.Context, chatID uuid.UU
if opts.BeforeID > 0 {
q.Set("before_id", strconv.FormatInt(opts.BeforeID, 10))
}
if opts.AfterID > 0 {
q.Set("after_id", strconv.FormatInt(opts.AfterID, 10))
}
if opts.Limit > 0 {
q.Set("limit", strconv.Itoa(opts.Limit))
}
+19 -1
View File
@@ -241,7 +241,25 @@ absent when all files are linked successfully.
`GET /api/experimental/chats/{chat}/messages`
Returns the messages and queued messages for a chat.
Returns messages for a chat in descending ID order (newest first).
| Query parameter | Type | Required | Description |
|-----------------|---------|----------|---------------------------------------------|
| `before_id` | `int64` | no | Only return messages with `id < before_id`. |
| `after_id` | `int64` | no | Only return messages with `id > after_id`. |
| `limit` | `int32` | no | Page size, 1 to 200. Defaults to 50. |
Results are returned in descending ID order (newest first), except when
only `after_id` is set: that shape is intended for polling and returns
ASCENDING ID order so a client can advance its cursor to the largest
returned ID without gaps. When both cursors are set they must satisfy
`after_id < before_id`; otherwise the server returns `400 Bad Request`.
`queued_messages` is only populated on the initial load (no
cursor). Pass either cursor to page through history or to poll
for new messages without receiving the queued snapshot on every
request. The `has_more` flag indicates more rows exist beyond
this page in the same direction.
### List models
+4 -1
View File
@@ -3115,12 +3115,15 @@ class ExperimentalApiMethods {
};
getChatMessages = async (
chatId: string,
opts?: { before_id?: number; limit?: number },
opts?: { before_id?: number; after_id?: number; limit?: number },
): Promise<TypesGen.ChatMessagesResponse> => {
const params = new URLSearchParams();
if (opts?.before_id) {
params.set("before_id", opts.before_id.toString());
}
if (opts?.after_id) {
params.set("after_id", opts.after_id.toString());
}
if (opts?.limit) {
params.set("limit", opts.limit.toString());
}
+9
View File
@@ -1857,6 +1857,15 @@ export interface ChatMessageUsage {
*/
export interface ChatMessagesPaginationOptions {
readonly BeforeID: number;
/**
* AfterID, when > 0, restricts results to messages with id strictly
* greater than AfterID. When set without BeforeID, results come back
* in ASCENDING id order so a polling caller can advance its cursor
* to max(returned_ids) without gaps. When combined with BeforeID,
* results come back in DESC order over the open range
* (AfterID, BeforeID).
*/
readonly AfterID: number;
readonly Limit: number;
}