Thanks to visit codestin.com
Credit goes to github.com

Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion coderd/database/dbauthz/dbauthz.go
Original file line number Diff line number Diff line change
Expand Up @@ -6086,7 +6086,7 @@ func (q *querier) InsertChatFile(ctx context.Context, arg database.InsertChatFil
return insert(q.log, q.auth, rbac.ResourceChat.WithOwner(arg.OwnerID.String()).InOrg(arg.OrganizationID), q.db.InsertChatFile)(ctx, arg)
}

func (q *querier) InsertChatMessages(ctx context.Context, arg database.InsertChatMessagesParams) ([]database.ChatMessage, error) {
func (q *querier) InsertChatMessages(ctx context.Context, arg database.InsertChatMessagesParams) ([]database.InsertChatMessagesRow, error) {
// Authorize create on the parent chat (using update permission).
chat, err := q.db.GetChatByID(ctx, arg.ChatID)
if err != nil {
Expand Down
2 changes: 1 addition & 1 deletion coderd/database/dbauthz/dbauthz_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1281,7 +1281,7 @@ func (s *MethodTestSuite) TestChats() {
s.Run("InsertChatMessages", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
chat := testutil.Fake(s.T(), faker, database.Chat{})
arg := testutil.Fake(s.T(), faker, database.InsertChatMessagesParams{ChatID: chat.ID})
msgs := []database.ChatMessage{testutil.Fake(s.T(), faker, database.ChatMessage{ChatID: chat.ID})}
msgs := []database.InsertChatMessagesRow{testutil.Fake(s.T(), faker, database.InsertChatMessagesRow{ChatID: chat.ID})}
dbm.EXPECT().GetChatByID(gomock.Any(), chat.ID).Return(chat, nil).AnyTimes()
dbm.EXPECT().InsertChatMessages(gomock.Any(), arg).Return(msgs, nil).AnyTimes()
check.Args(arg).Asserts(chat, policy.ActionUpdate).Returns(msgs)
Expand Down
2 changes: 1 addition & 1 deletion coderd/database/dbgen/dbgen.go
Original file line number Diff line number Diff line change
Expand Up @@ -145,7 +145,7 @@ func ChatMessage(t testing.TB, db database.Store, seed database.ChatMessage) dat
})
require.NoError(t, err, "insert chat message")
require.Len(t, msgs, 1)
return msgs[0]
return database.ChatMessage(msgs[0])
}

const (
Expand Down
2 changes: 1 addition & 1 deletion coderd/database/dbmetrics/querymetrics.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

4 changes: 2 additions & 2 deletions coderd/database/dbmock/dbmock.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

2 changes: 2 additions & 0 deletions coderd/database/dump.sql

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
DROP INDEX IF EXISTS idx_chat_messages_chat_role_id;
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
-- Serves GetLastChatMessageByRole. It orders by id, so the existing
-- (chat_id, created_at) index cannot supply the LIMIT 1 row in index order.
CREATE INDEX idx_chat_messages_chat_role_id
ON chat_messages (chat_id, role, id DESC)
WHERE deleted = false;
22 changes: 22 additions & 0 deletions coderd/database/modelqueries_internal_test.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package database

import (
"reflect"
"regexp"
"slices"
"strings"
Expand Down Expand Up @@ -168,6 +169,27 @@ func TestFinalizeStaleChatDebugRows_TerminalStatusAlignment(t *testing.T) {
}
}

// TestInsertChatMessagesOrderContract guards the input-order guarantee that
// callers rely on when indexing the returned slice. A behavior test cannot:
// Postgres evaluates the id default in row order anyway, so a batch still looks
// ordered once the guarantee is removed.
func TestInsertChatMessagesOrderContract(t *testing.T) {
t.Parallel()

require.Contains(t, insertChatMessages, "nextval('chat_messages_id_seq')",
"ids must be allocated explicitly so they can be correlated to input array position")
require.Contains(t, insertChatMessages, "ROW_NUMBER() OVER (ORDER BY id)",
"the k-th smallest allocated id must be assigned to input index k")
require.Regexp(t, `(?s)ORDER BY id\s*\z`, strings.TrimSpace(insertChatMessages),
"returned rows must be explicitly ordered by id rather than relying on RETURNING order")

// Every parallel input array must be read at the allocated ordinal. A column
// left on UNNEST would be positioned by the executor instead.
subscripted := regexp.MustCompile(`\)\[allocated\.ord\]`).FindAllString(insertChatMessages, -1)
require.Len(t, subscripted, reflect.TypeOf(InsertChatMessagesParams{}).NumField()-1,
"each InsertChatMessagesParams array field, all but ChatID, must be subscripted by allocated.ord")
}

// extractWhereClause extracts the WHERE clause from a SQL query string
func extractWhereClause(query string) string {
// Find WHERE and get everything after it
Expand Down
11 changes: 10 additions & 1 deletion coderd/database/querier.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

133 changes: 131 additions & 2 deletions coderd/database/querier_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -12212,6 +12212,135 @@ func TestInsertChatMessages(t *testing.T) {
})
}

// The returned ids are in insert order, which the inverted created_at values
// deliberately contradict.
func insertChatMessagesInvertedTimestamps(t *testing.T, db database.Store, sqlDB *sql.DB, roles []database.ChatMessageRole) (database.Chat, []int64) {
t.Helper()

ctx := context.Background()
org := dbgen.Organization(t, db, database.Organization{})
owner := dbgen.User(t, db, database.User{})
modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
})
chat := dbgen.Chat(t, db, database.Chat{
OrganizationID: org.ID,
OwnerID: owner.ID,
LastModelConfigID: modelCfg.ID,
})

count := len(roles)
inserted, err := db.InsertChatMessages(ctx, database.InsertChatMessagesParams{
ChatID: chat.ID,
CreatedBy: slices.Repeat([]uuid.UUID{owner.ID}, count),
ModelConfigID: slices.Repeat([]uuid.UUID{modelCfg.ID}, count),
Role: roles,
ContentVersion: slices.Repeat([]int16{chatprompt.CurrentContentVersion}, count),
Visibility: slices.Repeat([]database.ChatMessageVisibility{database.ChatMessageVisibilityBoth}, count),
Content: slices.Repeat([]string{`"message"`}, 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),
TotalCostMicros: make([]int64, count),
RuntimeMs: make([]int64, count),
})
require.NoError(t, err)
require.Len(t, inserted, count)

insertedIDs := make([]int64, count)
for i, message := range inserted {
insertedIDs[i] = message.ID
_, err := sqlDB.ExecContext(ctx,
"UPDATE chat_messages SET created_at = $1 WHERE id = $2",
message.CreatedAt.Add(time.Duration(count-i)*time.Minute), message.ID)
require.NoError(t, err)
}

return chat, insertedIDs
}

func chatMessageIDs(messages []database.ChatMessage) []int64 {
ids := make([]int64, len(messages))
for i, message := range messages {
ids[i] = message.ID
}
return ids
}

func TestGetChatMessagesByChatIDOrdersByID(t *testing.T) {
t.Parallel()

db, _, sqlDB := dbtestutil.NewDBWithSQLDB(t)
ctx := context.Background()

chat, insertedIDs := insertChatMessagesInvertedTimestamps(t, db, sqlDB,
slices.Repeat([]database.ChatMessageRole{database.ChatMessageRoleUser}, 3))

messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
ChatID: chat.ID,
AfterID: 0,
})
require.NoError(t, err)
require.Equal(t, insertedIDs, chatMessageIDs(messages))
}

func TestGetChatMessagesByRevisionForStreamOrdersByID(t *testing.T) {
t.Parallel()

db, _, sqlDB := dbtestutil.NewDBWithSQLDB(t)
ctx := context.Background()

chat, insertedIDs := insertChatMessagesInvertedTimestamps(t, db, sqlDB,
slices.Repeat([]database.ChatMessageRole{database.ChatMessageRoleUser}, 3))

messages, err := db.GetChatMessagesByRevisionForStream(ctx, database.GetChatMessagesByRevisionForStreamParams{
ChatID: chat.ID,
AfterRevision: 0,
})
require.NoError(t, err)
require.Equal(t, insertedIDs, chatMessageIDs(messages))
}

func TestGetLastChatMessageByRoleOrdersByID(t *testing.T) {
t.Parallel()

db, _, sqlDB := dbtestutil.NewDBWithSQLDB(t)
ctx := context.Background()

chat, insertedIDs := insertChatMessagesInvertedTimestamps(t, db, sqlDB,
slices.Repeat([]database.ChatMessageRole{database.ChatMessageRoleAssistant}, 3))

last, err := db.GetLastChatMessageByRole(ctx, database.GetLastChatMessageByRoleParams{
ChatID: chat.ID,
Role: database.ChatMessageRoleAssistant,
})
require.NoError(t, err)
require.Equal(t, insertedIDs[len(insertedIDs)-1], last.ID)
}

// Sequence cache blocks are handed out per session, so above cache 1 a backend
// holding stale cached values can take the chat row lock second and still commit
// lower ids. Bumping a sequence cache is an ordinary throughput tweak.
func TestChatMessagesSequenceCacheIsOne(t *testing.T) {
t.Parallel()

_, _, sqlDB := dbtestutil.NewDBWithSQLDB(t)
ctx := context.Background()

var cacheSize int64
err := sqlDB.QueryRowContext(ctx,
"SELECT cache_size FROM pg_sequences WHERE sequencename = 'chat_messages_id_seq'").
Scan(&cacheSize)
require.NoError(t, err)
require.Equal(t, int64(1), cacheSize, "chat_messages_id_seq must use cache 1")
}

func TestGetChatMessagesForPromptByChatID(t *testing.T) {
t.Parallel()

Expand Down Expand Up @@ -12294,7 +12423,7 @@ func TestGetChatMessagesForPromptByChatID(t *testing.T) {
RuntimeMs: []int64{0},
})
require.NoError(t, err)
return results[0]
return database.ChatMessage(results[0])
}

msgIDs := func(msgs []database.ChatMessage) []int64 {
Expand Down Expand Up @@ -17173,7 +17302,7 @@ func TestGetChatsSearch(t *testing.T) {
})
require.NoError(t, err)
require.Len(t, msgs, 1)
return msgs[0]
return database.ChatMessage(msgs[0])
}

linkPR := func(chatID uuid.UUID, url, state, prTitle string, prNumber int32, gitRemoteOrigin string) {
Expand Down
Loading
Loading