mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: replace InsertChatMessage with batch InsertChatMessages (#23220)
Replaces the singular `InsertChatMessage` query with `InsertChatMessages` that uses PostgreSQL's `unnest()` for batch inserts. This reduces the number of database round-trips when inserting multiple messages in a single transaction. ## Changes - **SQL**: New `InsertChatMessages :many` query using `unnest()` arrays following the existing codebase pattern (e.g., `InsertWorkspaceAgentStats`). Preserves the CTE that updates `chats.last_model_config_id` using the last non-null model config from the batch. Uses `NULLIF` for UUID columns to handle NULL foreign keys. - **Go layers**: Updated `querier.go`, `dbauthz.go`, `dbmetrics/querymetrics.go`, `dbmock/dbmock.go`, and `queries.sql.go` to use the new batch signature (`[]ChatMessage` return type, array params). - **chatd.go**: All call sites converted to batch inserts: - **CreateChat**: System prompt + user message batched into one call - **persistStep**: Assistant message + tool messages batched into one call - **persistSummary**: Hidden summary + assistant + tool messages batched into one call - Single-message sites use the same API with single-element arrays - **Helper**: New `appendChatMessage` function simplifies building batch params at each call site. - **Tests**: All test files updated to use the new API. Builds on top of #23213.
This commit is contained in:
+115
-81
@@ -4841,16 +4841,34 @@ func (q *sqlQuerier) InsertChat(ctx context.Context, arg InsertChatParams) (Chat
|
||||
return i, err
|
||||
}
|
||||
|
||||
const insertChatMessage = `-- name: InsertChatMessage :one
|
||||
const insertChatMessages = `-- name: InsertChatMessages :many
|
||||
WITH updated_chat AS (
|
||||
UPDATE
|
||||
chats
|
||||
SET
|
||||
last_model_config_id = $3::uuid
|
||||
last_model_config_id = (
|
||||
SELECT val
|
||||
FROM unnest($3::uuid[])
|
||||
WITH ORDINALITY AS t(val, ord)
|
||||
WHERE val != '00000000-0000-0000-0000-000000000000'::uuid
|
||||
ORDER BY ord DESC
|
||||
LIMIT 1
|
||||
)
|
||||
WHERE
|
||||
id = $1::uuid
|
||||
AND $3::uuid IS NOT NULL
|
||||
AND chats.last_model_config_id IS DISTINCT FROM $3::uuid
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM unnest($3::uuid[])
|
||||
WHERE unnest != '00000000-0000-0000-0000-000000000000'::uuid
|
||||
)
|
||||
AND chats.last_model_config_id IS DISTINCT FROM (
|
||||
SELECT val
|
||||
FROM unnest($3::uuid[])
|
||||
WITH ORDINALITY AS t(val, ord)
|
||||
WHERE val != '00000000-0000-0000-0000-000000000000'::uuid
|
||||
ORDER BY ord DESC
|
||||
LIMIT 1
|
||||
)
|
||||
)
|
||||
INSERT INTO chat_messages (
|
||||
chat_id,
|
||||
@@ -4870,92 +4888,108 @@ INSERT INTO chat_messages (
|
||||
compressed,
|
||||
total_cost_micros,
|
||||
runtime_ms
|
||||
) VALUES (
|
||||
$1::uuid,
|
||||
$2::uuid,
|
||||
$3::uuid,
|
||||
$4::chat_message_role,
|
||||
$5::jsonb,
|
||||
$6::smallint,
|
||||
$7::chat_message_visibility,
|
||||
$8::bigint,
|
||||
$9::bigint,
|
||||
$10::bigint,
|
||||
$11::bigint,
|
||||
$12::bigint,
|
||||
$13::bigint,
|
||||
$14::bigint,
|
||||
COALESCE($15::boolean, FALSE),
|
||||
$16::bigint,
|
||||
$17::bigint
|
||||
)
|
||||
SELECT
|
||||
$1::uuid,
|
||||
NULLIF(unnest($2::uuid[]), '00000000-0000-0000-0000-000000000000'::uuid),
|
||||
NULLIF(unnest($3::uuid[]), '00000000-0000-0000-0000-000000000000'::uuid),
|
||||
unnest($4::chat_message_role[]),
|
||||
unnest($5::text[])::jsonb,
|
||||
unnest($6::smallint[]),
|
||||
unnest($7::chat_message_visibility[]),
|
||||
NULLIF(unnest($8::bigint[]), 0),
|
||||
NULLIF(unnest($9::bigint[]), 0),
|
||||
NULLIF(unnest($10::bigint[]), 0),
|
||||
NULLIF(unnest($11::bigint[]), 0),
|
||||
NULLIF(unnest($12::bigint[]), 0),
|
||||
NULLIF(unnest($13::bigint[]), 0),
|
||||
NULLIF(unnest($14::bigint[]), 0),
|
||||
unnest($15::boolean[]),
|
||||
NULLIF(unnest($16::bigint[]), 0),
|
||||
NULLIF(unnest($17::bigint[]), 0)
|
||||
RETURNING
|
||||
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, runtime_ms
|
||||
`
|
||||
|
||||
type InsertChatMessageParams struct {
|
||||
ChatID uuid.UUID `db:"chat_id" json:"chat_id"`
|
||||
CreatedBy uuid.NullUUID `db:"created_by" json:"created_by"`
|
||||
ModelConfigID uuid.NullUUID `db:"model_config_id" json:"model_config_id"`
|
||||
Role ChatMessageRole `db:"role" json:"role"`
|
||||
Content pqtype.NullRawMessage `db:"content" json:"content"`
|
||||
ContentVersion int16 `db:"content_version" json:"content_version"`
|
||||
Visibility ChatMessageVisibility `db:"visibility" json:"visibility"`
|
||||
InputTokens sql.NullInt64 `db:"input_tokens" json:"input_tokens"`
|
||||
OutputTokens sql.NullInt64 `db:"output_tokens" json:"output_tokens"`
|
||||
TotalTokens sql.NullInt64 `db:"total_tokens" json:"total_tokens"`
|
||||
ReasoningTokens sql.NullInt64 `db:"reasoning_tokens" json:"reasoning_tokens"`
|
||||
CacheCreationTokens sql.NullInt64 `db:"cache_creation_tokens" json:"cache_creation_tokens"`
|
||||
CacheReadTokens sql.NullInt64 `db:"cache_read_tokens" json:"cache_read_tokens"`
|
||||
ContextLimit sql.NullInt64 `db:"context_limit" json:"context_limit"`
|
||||
Compressed sql.NullBool `db:"compressed" json:"compressed"`
|
||||
TotalCostMicros sql.NullInt64 `db:"total_cost_micros" json:"total_cost_micros"`
|
||||
RuntimeMs sql.NullInt64 `db:"runtime_ms" json:"runtime_ms"`
|
||||
type InsertChatMessagesParams struct {
|
||||
ChatID uuid.UUID `db:"chat_id" json:"chat_id"`
|
||||
CreatedBy []uuid.UUID `db:"created_by" json:"created_by"`
|
||||
ModelConfigID []uuid.UUID `db:"model_config_id" json:"model_config_id"`
|
||||
Role []ChatMessageRole `db:"role" json:"role"`
|
||||
Content []string `db:"content" json:"content"`
|
||||
ContentVersion []int16 `db:"content_version" json:"content_version"`
|
||||
Visibility []ChatMessageVisibility `db:"visibility" json:"visibility"`
|
||||
InputTokens []int64 `db:"input_tokens" json:"input_tokens"`
|
||||
OutputTokens []int64 `db:"output_tokens" json:"output_tokens"`
|
||||
TotalTokens []int64 `db:"total_tokens" json:"total_tokens"`
|
||||
ReasoningTokens []int64 `db:"reasoning_tokens" json:"reasoning_tokens"`
|
||||
CacheCreationTokens []int64 `db:"cache_creation_tokens" json:"cache_creation_tokens"`
|
||||
CacheReadTokens []int64 `db:"cache_read_tokens" json:"cache_read_tokens"`
|
||||
ContextLimit []int64 `db:"context_limit" json:"context_limit"`
|
||||
Compressed []bool `db:"compressed" json:"compressed"`
|
||||
TotalCostMicros []int64 `db:"total_cost_micros" json:"total_cost_micros"`
|
||||
RuntimeMs []int64 `db:"runtime_ms" json:"runtime_ms"`
|
||||
}
|
||||
|
||||
func (q *sqlQuerier) InsertChatMessage(ctx context.Context, arg InsertChatMessageParams) (ChatMessage, error) {
|
||||
row := q.db.QueryRowContext(ctx, insertChatMessage,
|
||||
func (q *sqlQuerier) InsertChatMessages(ctx context.Context, arg InsertChatMessagesParams) ([]ChatMessage, error) {
|
||||
rows, err := q.db.QueryContext(ctx, insertChatMessages,
|
||||
arg.ChatID,
|
||||
arg.CreatedBy,
|
||||
arg.ModelConfigID,
|
||||
arg.Role,
|
||||
arg.Content,
|
||||
arg.ContentVersion,
|
||||
arg.Visibility,
|
||||
arg.InputTokens,
|
||||
arg.OutputTokens,
|
||||
arg.TotalTokens,
|
||||
arg.ReasoningTokens,
|
||||
arg.CacheCreationTokens,
|
||||
arg.CacheReadTokens,
|
||||
arg.ContextLimit,
|
||||
arg.Compressed,
|
||||
arg.TotalCostMicros,
|
||||
arg.RuntimeMs,
|
||||
pq.Array(arg.CreatedBy),
|
||||
pq.Array(arg.ModelConfigID),
|
||||
pq.Array(arg.Role),
|
||||
pq.Array(arg.Content),
|
||||
pq.Array(arg.ContentVersion),
|
||||
pq.Array(arg.Visibility),
|
||||
pq.Array(arg.InputTokens),
|
||||
pq.Array(arg.OutputTokens),
|
||||
pq.Array(arg.TotalTokens),
|
||||
pq.Array(arg.ReasoningTokens),
|
||||
pq.Array(arg.CacheCreationTokens),
|
||||
pq.Array(arg.CacheReadTokens),
|
||||
pq.Array(arg.ContextLimit),
|
||||
pq.Array(arg.Compressed),
|
||||
pq.Array(arg.TotalCostMicros),
|
||||
pq.Array(arg.RuntimeMs),
|
||||
)
|
||||
var i ChatMessage
|
||||
err := row.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,
|
||||
&i.RuntimeMs,
|
||||
)
|
||||
return i, err
|
||||
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,
|
||||
&i.RuntimeMs,
|
||||
); 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 insertChatQueuedMessage = `-- name: InsertChatQueuedMessage :one
|
||||
|
||||
Reference in New Issue
Block a user