mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add after_id pagination for chat messages (#24531)
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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{
|
||||
|
||||
@@ -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
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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
@@ -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());
|
||||
}
|
||||
|
||||
Generated
+9
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user