mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
refactor: add chat_message_role enum and content_version column (#23042)
Migration 000434 converts chat_messages.role from text to a Postgres enum, rebuilds the partial index, and adds content_version smallint. The column is backfilled with DEFAULT 0, then the default is dropped so future inserts must set it explicitly. Version 0 uses the role-aware heuristic from #22958. Version 1 (all new inserts) stores []ChatMessagePart JSON for all roles, including system messages. ParseContent takes database.ChatMessage directly and dispatches on version internally. Unknown versions error. All string(codersdk.ChatMessageRole*) casts at DB write sites are replaced with database.ChatMessageRole* constants from sqlc. Refs #22958
This commit is contained in:
@@ -1071,7 +1071,7 @@ func ChatMessage(m database.ChatMessage) codersdk.ChatMessage {
|
||||
Role: codersdk.ChatMessageRole(m.Role),
|
||||
}
|
||||
if m.Content.Valid {
|
||||
parts, err := chatMessageParts(codersdk.ChatMessageRole(m.Role), m.Content)
|
||||
parts, err := chatMessageParts(m)
|
||||
if err == nil {
|
||||
msg.Content = parts
|
||||
}
|
||||
@@ -1113,9 +1113,15 @@ func chatMessageUsage(m database.ChatMessage) *codersdk.ChatMessageUsage {
|
||||
|
||||
// ChatQueuedMessage converts a queued message to its SDK representation.
|
||||
func ChatQueuedMessage(message database.ChatQueuedMessage) codersdk.ChatQueuedMessage {
|
||||
parts, err := chatMessageParts(codersdk.ChatMessageRoleUser, pqtype.NullRawMessage{
|
||||
RawMessage: message.Content,
|
||||
Valid: len(message.Content) > 0,
|
||||
// Queued messages are always written by current code via
|
||||
// MarshalParts, so they are always current content version.
|
||||
parts, err := chatMessageParts(database.ChatMessage{
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: pqtype.NullRawMessage{
|
||||
RawMessage: message.Content,
|
||||
Valid: len(message.Content) > 0,
|
||||
},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
})
|
||||
if err != nil {
|
||||
parts = nil
|
||||
@@ -1139,8 +1145,8 @@ func ChatQueuedMessages(messages []database.ChatQueuedMessage) []codersdk.ChatQu
|
||||
return out
|
||||
}
|
||||
|
||||
func chatMessageParts(role codersdk.ChatMessageRole, raw pqtype.NullRawMessage) ([]codersdk.ChatMessagePart, error) {
|
||||
parts, err := chatprompt.ParseContent(role, raw)
|
||||
func chatMessageParts(m database.ChatMessage) ([]codersdk.ChatMessagePart, error) {
|
||||
parts, err := chatprompt.ParseContent(m)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -467,7 +467,7 @@ func TestChatMessage_PreservesProviderExecutedOnToolResults(t *testing.T) {
|
||||
dbMsg := database.ChatMessage{
|
||||
ID: 1,
|
||||
ChatID: uuid.New(),
|
||||
Role: string(codersdk.ChatMessageRoleAssistant),
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
Content: pqtype.NullRawMessage{
|
||||
RawMessage: rawContent,
|
||||
Valid: true,
|
||||
@@ -495,8 +495,9 @@ func TestChatMessage_PreservesProviderExecutedOnToolResults(t *testing.T) {
|
||||
func TestChatQueuedMessage_ParsesUserContentParts(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
rawContent, err := json.Marshal([]fantasy.Content{
|
||||
fantasy.TextContent{Text: "queued text"},
|
||||
// Queued messages are always written via MarshalParts (SDK format).
|
||||
rawContent, err := json.Marshal([]codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("queued text"),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -512,35 +513,15 @@ func TestChatQueuedMessage_ParsesUserContentParts(t *testing.T) {
|
||||
require.Equal(t, "queued text", queued.Content[0].Text)
|
||||
}
|
||||
|
||||
func TestChatQueuedMessage_FallsBackToTextForLegacyContent(t *testing.T) {
|
||||
func TestChatQueuedMessage_MalformedContent(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("legacy_string", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
queued := db2sdk.ChatQueuedMessage(database.ChatQueuedMessage{
|
||||
ID: 1,
|
||||
ChatID: uuid.New(),
|
||||
Content: json.RawMessage(`"legacy queued text"`),
|
||||
CreatedAt: time.Now(),
|
||||
})
|
||||
|
||||
require.Len(t, queued.Content, 1)
|
||||
require.Equal(t, codersdk.ChatMessagePartTypeText, queued.Content[0].Type)
|
||||
require.Equal(t, "legacy queued text", queued.Content[0].Text)
|
||||
queued := db2sdk.ChatQueuedMessage(database.ChatQueuedMessage{
|
||||
ID: 1,
|
||||
ChatID: uuid.New(),
|
||||
Content: json.RawMessage(`{"unexpected":"shape"}`),
|
||||
CreatedAt: time.Now(),
|
||||
})
|
||||
|
||||
t.Run("malformed_payload", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
raw := json.RawMessage(`{"unexpected":"shape"}`)
|
||||
queued := db2sdk.ChatQueuedMessage(database.ChatQueuedMessage{
|
||||
ID: 1,
|
||||
ChatID: uuid.New(),
|
||||
Content: raw,
|
||||
CreatedAt: time.Now(),
|
||||
})
|
||||
|
||||
require.Empty(t, queued.Content)
|
||||
})
|
||||
require.Empty(t, queued.Content)
|
||||
}
|
||||
|
||||
@@ -486,7 +486,7 @@ func (s *MethodTestSuite) TestChats() {
|
||||
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})
|
||||
arg := database.GetLastChatMessageByRoleParams{ChatID: chat.ID, Role: "assistant"}
|
||||
arg := database.GetLastChatMessageByRoleParams{ChatID: chat.ID, Role: database.ChatMessageRoleAssistant}
|
||||
dbm.EXPECT().GetChatByID(gomock.Any(), chat.ID).Return(chat, nil).AnyTimes()
|
||||
dbm.EXPECT().GetLastChatMessageByRole(gomock.Any(), arg).Return(msg, nil).AnyTimes()
|
||||
check.Args(arg).Asserts(chat, policy.ActionRead).Returns(msg)
|
||||
|
||||
Generated
+11
-3
@@ -265,6 +265,13 @@ CREATE TYPE build_reason AS ENUM (
|
||||
'task_resume'
|
||||
);
|
||||
|
||||
CREATE TYPE chat_message_role AS ENUM (
|
||||
'system',
|
||||
'user',
|
||||
'assistant',
|
||||
'tool'
|
||||
);
|
||||
|
||||
CREATE TYPE chat_message_visibility AS ENUM (
|
||||
'user',
|
||||
'model',
|
||||
@@ -1207,7 +1214,7 @@ CREATE TABLE chat_messages (
|
||||
chat_id uuid NOT NULL,
|
||||
model_config_id uuid,
|
||||
created_at timestamp with time zone DEFAULT now() NOT NULL,
|
||||
role text NOT NULL,
|
||||
role chat_message_role NOT NULL,
|
||||
content jsonb,
|
||||
visibility chat_message_visibility DEFAULT 'both'::chat_message_visibility NOT NULL,
|
||||
input_tokens bigint,
|
||||
@@ -1218,7 +1225,8 @@ CREATE TABLE chat_messages (
|
||||
cache_read_tokens bigint,
|
||||
context_limit bigint,
|
||||
compressed boolean DEFAULT false NOT NULL,
|
||||
created_by uuid
|
||||
created_by uuid,
|
||||
content_version smallint NOT NULL
|
||||
);
|
||||
|
||||
CREATE SEQUENCE chat_messages_id_seq
|
||||
@@ -3524,7 +3532,7 @@ CREATE INDEX idx_chat_messages_chat ON chat_messages USING btree (chat_id);
|
||||
|
||||
CREATE INDEX idx_chat_messages_chat_created ON chat_messages USING btree (chat_id, created_at);
|
||||
|
||||
CREATE INDEX idx_chat_messages_compressed_summary_boundary ON chat_messages USING btree (chat_id, created_at DESC, id DESC) WHERE ((compressed = true) AND (role = 'system'::text) AND (visibility = ANY (ARRAY['model'::chat_message_visibility, 'both'::chat_message_visibility])));
|
||||
CREATE INDEX idx_chat_messages_compressed_summary_boundary ON chat_messages USING btree (chat_id, created_at DESC, id DESC) WHERE ((compressed = true) AND (role = 'system'::chat_message_role) AND (visibility = ANY (ARRAY['model'::chat_message_visibility, 'both'::chat_message_visibility])));
|
||||
|
||||
CREATE INDEX idx_chat_model_configs_enabled ON chat_model_configs USING btree (enabled);
|
||||
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
ALTER TABLE chat_messages DROP COLUMN content_version;
|
||||
|
||||
DROP INDEX idx_chat_messages_compressed_summary_boundary;
|
||||
|
||||
ALTER TABLE chat_messages
|
||||
ALTER COLUMN role TYPE text
|
||||
USING (role::text);
|
||||
|
||||
CREATE INDEX idx_chat_messages_compressed_summary_boundary
|
||||
ON chat_messages(chat_id, created_at DESC, id DESC)
|
||||
WHERE compressed = TRUE
|
||||
AND role = 'system'
|
||||
AND visibility IN ('model', 'both');
|
||||
|
||||
DROP TYPE chat_message_role;
|
||||
@@ -0,0 +1,32 @@
|
||||
-- Add chat_message_role enum.
|
||||
CREATE TYPE chat_message_role AS ENUM (
|
||||
'system',
|
||||
'user',
|
||||
'assistant',
|
||||
'tool'
|
||||
);
|
||||
|
||||
-- Drop the partial index that references role as text before
|
||||
-- converting the column type.
|
||||
DROP INDEX idx_chat_messages_compressed_summary_boundary;
|
||||
|
||||
-- Convert role column from text to enum.
|
||||
ALTER TABLE chat_messages
|
||||
ALTER COLUMN role TYPE chat_message_role
|
||||
USING (role::chat_message_role);
|
||||
|
||||
-- Recreate the partial index with enum-typed comparison.
|
||||
CREATE INDEX idx_chat_messages_compressed_summary_boundary
|
||||
ON chat_messages(chat_id, created_at DESC, id DESC)
|
||||
WHERE compressed = TRUE
|
||||
AND role = 'system'
|
||||
AND visibility IN ('model', 'both');
|
||||
|
||||
-- Add content_version column. Default 0 backfills existing rows.
|
||||
-- The default is then dropped so future inserts must specify the
|
||||
-- version explicitly.
|
||||
ALTER TABLE chat_messages
|
||||
ADD COLUMN content_version smallint NOT NULL DEFAULT 0;
|
||||
|
||||
ALTER TABLE chat_messages
|
||||
ALTER COLUMN content_version DROP DEFAULT;
|
||||
@@ -1049,6 +1049,70 @@ func AllBuildReasonValues() []BuildReason {
|
||||
}
|
||||
}
|
||||
|
||||
type ChatMessageRole string
|
||||
|
||||
const (
|
||||
ChatMessageRoleSystem ChatMessageRole = "system"
|
||||
ChatMessageRoleUser ChatMessageRole = "user"
|
||||
ChatMessageRoleAssistant ChatMessageRole = "assistant"
|
||||
ChatMessageRoleTool ChatMessageRole = "tool"
|
||||
)
|
||||
|
||||
func (e *ChatMessageRole) Scan(src interface{}) error {
|
||||
switch s := src.(type) {
|
||||
case []byte:
|
||||
*e = ChatMessageRole(s)
|
||||
case string:
|
||||
*e = ChatMessageRole(s)
|
||||
default:
|
||||
return fmt.Errorf("unsupported scan type for ChatMessageRole: %T", src)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type NullChatMessageRole struct {
|
||||
ChatMessageRole ChatMessageRole `json:"chat_message_role"`
|
||||
Valid bool `json:"valid"` // Valid is true if ChatMessageRole is not NULL
|
||||
}
|
||||
|
||||
// Scan implements the Scanner interface.
|
||||
func (ns *NullChatMessageRole) Scan(value interface{}) error {
|
||||
if value == nil {
|
||||
ns.ChatMessageRole, ns.Valid = "", false
|
||||
return nil
|
||||
}
|
||||
ns.Valid = true
|
||||
return ns.ChatMessageRole.Scan(value)
|
||||
}
|
||||
|
||||
// Value implements the driver Valuer interface.
|
||||
func (ns NullChatMessageRole) Value() (driver.Value, error) {
|
||||
if !ns.Valid {
|
||||
return nil, nil
|
||||
}
|
||||
return string(ns.ChatMessageRole), nil
|
||||
}
|
||||
|
||||
func (e ChatMessageRole) Valid() bool {
|
||||
switch e {
|
||||
case ChatMessageRoleSystem,
|
||||
ChatMessageRoleUser,
|
||||
ChatMessageRoleAssistant,
|
||||
ChatMessageRoleTool:
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func AllChatMessageRoleValues() []ChatMessageRole {
|
||||
return []ChatMessageRole{
|
||||
ChatMessageRoleSystem,
|
||||
ChatMessageRoleUser,
|
||||
ChatMessageRoleAssistant,
|
||||
ChatMessageRoleTool,
|
||||
}
|
||||
}
|
||||
|
||||
type ChatMessageVisibility string
|
||||
|
||||
const (
|
||||
@@ -3943,7 +4007,7 @@ type ChatMessage struct {
|
||||
ChatID uuid.UUID `db:"chat_id" json:"chat_id"`
|
||||
ModelConfigID uuid.NullUUID `db:"model_config_id" json:"model_config_id"`
|
||||
CreatedAt time.Time `db:"created_at" json:"created_at"`
|
||||
Role string `db:"role" json:"role"`
|
||||
Role ChatMessageRole `db:"role" json:"role"`
|
||||
Content pqtype.NullRawMessage `db:"content" json:"content"`
|
||||
Visibility ChatMessageVisibility `db:"visibility" json:"visibility"`
|
||||
InputTokens sql.NullInt64 `db:"input_tokens" json:"input_tokens"`
|
||||
@@ -3955,6 +4019,7 @@ type ChatMessage struct {
|
||||
ContextLimit sql.NullInt64 `db:"context_limit" json:"context_limit"`
|
||||
Compressed bool `db:"compressed" json:"compressed"`
|
||||
CreatedBy uuid.NullUUID `db:"created_by" json:"created_by"`
|
||||
ContentVersion int16 `db:"content_version" json:"content_version"`
|
||||
}
|
||||
|
||||
type ChatModelConfig struct {
|
||||
|
||||
@@ -21,6 +21,7 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog/v3/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/coderd/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/coderdtest"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbauthz"
|
||||
@@ -9044,17 +9045,18 @@ func TestGetChatMessagesForPromptByChatID(t *testing.T) {
|
||||
insertMsg := func(
|
||||
t *testing.T,
|
||||
chatID uuid.UUID,
|
||||
role string,
|
||||
role database.ChatMessageRole,
|
||||
vis database.ChatMessageVisibility,
|
||||
compressed bool,
|
||||
content string,
|
||||
) database.ChatMessage {
|
||||
t.Helper()
|
||||
msg, err := db.InsertChatMessage(ctx, database.InsertChatMessageParams{
|
||||
ChatID: chatID,
|
||||
Role: role,
|
||||
Visibility: vis,
|
||||
Compressed: sql.NullBool{Bool: compressed, Valid: true},
|
||||
ChatID: chatID,
|
||||
Role: role,
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
Visibility: vis,
|
||||
Compressed: sql.NullBool{Bool: compressed, Valid: true},
|
||||
Content: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`"` + content + `"`),
|
||||
Valid: true,
|
||||
@@ -9076,9 +9078,9 @@ func TestGetChatMessagesForPromptByChatID(t *testing.T) {
|
||||
t.Parallel()
|
||||
chat := newChat(t)
|
||||
|
||||
sys := insertMsg(t, chat.ID, "system", database.ChatMessageVisibilityModel, false, "system prompt")
|
||||
usr := insertMsg(t, chat.ID, "user", database.ChatMessageVisibilityBoth, false, "hello")
|
||||
ast := insertMsg(t, chat.ID, "assistant", database.ChatMessageVisibilityBoth, false, "hi there")
|
||||
sys := insertMsg(t, chat.ID, database.ChatMessageRoleSystem, database.ChatMessageVisibilityModel, false, "system prompt")
|
||||
usr := insertMsg(t, chat.ID, database.ChatMessageRoleUser, database.ChatMessageVisibilityBoth, false, "hello")
|
||||
ast := insertMsg(t, chat.ID, database.ChatMessageRoleAssistant, database.ChatMessageVisibilityBoth, false, "hi there")
|
||||
|
||||
got, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
@@ -9091,9 +9093,9 @@ func TestGetChatMessagesForPromptByChatID(t *testing.T) {
|
||||
|
||||
// Messages with visibility=user should NOT appear in the
|
||||
// prompt (they are only for the UI).
|
||||
insertMsg(t, chat.ID, "system", database.ChatMessageVisibilityModel, false, "system prompt")
|
||||
insertMsg(t, chat.ID, "user", database.ChatMessageVisibilityUser, false, "user-only msg")
|
||||
usr := insertMsg(t, chat.ID, "user", database.ChatMessageVisibilityBoth, false, "hello")
|
||||
insertMsg(t, chat.ID, database.ChatMessageRoleSystem, database.ChatMessageVisibilityModel, false, "system prompt")
|
||||
insertMsg(t, chat.ID, database.ChatMessageRoleUser, database.ChatMessageVisibilityUser, false, "user-only msg")
|
||||
usr := insertMsg(t, chat.ID, database.ChatMessageRoleUser, database.ChatMessageVisibilityBoth, false, "hello")
|
||||
|
||||
got, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
@@ -9109,21 +9111,21 @@ func TestGetChatMessagesForPromptByChatID(t *testing.T) {
|
||||
chat := newChat(t)
|
||||
|
||||
// Pre-compaction conversation.
|
||||
sys := insertMsg(t, chat.ID, "system", database.ChatMessageVisibilityModel, false, "system prompt")
|
||||
preUser := insertMsg(t, chat.ID, "user", database.ChatMessageVisibilityBoth, false, "old question")
|
||||
preAsst := insertMsg(t, chat.ID, "assistant", database.ChatMessageVisibilityBoth, false, "old answer")
|
||||
sys := insertMsg(t, chat.ID, database.ChatMessageRoleSystem, database.ChatMessageVisibilityModel, false, "system prompt")
|
||||
preUser := insertMsg(t, chat.ID, database.ChatMessageRoleUser, database.ChatMessageVisibilityBoth, false, "old question")
|
||||
preAsst := insertMsg(t, chat.ID, database.ChatMessageRoleAssistant, database.ChatMessageVisibilityBoth, false, "old answer")
|
||||
|
||||
// Compaction messages:
|
||||
// 1. Summary (role=user, visibility=model, compressed=true).
|
||||
summary := insertMsg(t, chat.ID, "user", database.ChatMessageVisibilityModel, true, "compaction summary")
|
||||
summary := insertMsg(t, chat.ID, database.ChatMessageRoleUser, database.ChatMessageVisibilityModel, true, "compaction summary")
|
||||
// 2. Compressed assistant tool-call (visibility=user).
|
||||
insertMsg(t, chat.ID, "assistant", database.ChatMessageVisibilityUser, true, "tool call")
|
||||
insertMsg(t, chat.ID, database.ChatMessageRoleAssistant, database.ChatMessageVisibilityUser, true, "tool call")
|
||||
// 3. Compressed tool result (visibility=both).
|
||||
insertMsg(t, chat.ID, "tool", database.ChatMessageVisibilityBoth, true, "tool result")
|
||||
insertMsg(t, chat.ID, database.ChatMessageRoleTool, database.ChatMessageVisibilityBoth, true, "tool result")
|
||||
|
||||
// Post-compaction messages.
|
||||
postUser := insertMsg(t, chat.ID, "user", database.ChatMessageVisibilityBoth, false, "new question")
|
||||
postAsst := insertMsg(t, chat.ID, "assistant", database.ChatMessageVisibilityBoth, false, "new answer")
|
||||
postUser := insertMsg(t, chat.ID, database.ChatMessageRoleUser, database.ChatMessageVisibilityBoth, false, "new question")
|
||||
postAsst := insertMsg(t, chat.ID, database.ChatMessageRoleAssistant, database.ChatMessageVisibilityBoth, false, "new answer")
|
||||
|
||||
got, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
@@ -9151,9 +9153,9 @@ func TestGetChatMessagesForPromptByChatID(t *testing.T) {
|
||||
// After compaction the summary must appear as role=user so
|
||||
// that LLM APIs (e.g. Anthropic) see at least one
|
||||
// non-system message in the prompt.
|
||||
insertMsg(t, chat.ID, "system", database.ChatMessageVisibilityModel, false, "system prompt")
|
||||
summary := insertMsg(t, chat.ID, "user", database.ChatMessageVisibilityModel, true, "summary text")
|
||||
newUsr := insertMsg(t, chat.ID, "user", database.ChatMessageVisibilityBoth, false, "new question")
|
||||
insertMsg(t, chat.ID, database.ChatMessageRoleSystem, database.ChatMessageVisibilityModel, false, "system prompt")
|
||||
summary := insertMsg(t, chat.ID, database.ChatMessageRoleUser, database.ChatMessageVisibilityModel, true, "summary text")
|
||||
newUsr := insertMsg(t, chat.ID, database.ChatMessageRoleUser, database.ChatMessageVisibilityBoth, false, "new question")
|
||||
|
||||
got, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
@@ -9179,10 +9181,10 @@ func TestGetChatMessagesForPromptByChatID(t *testing.T) {
|
||||
// used IN ('model','both'), the compressed tool result
|
||||
// (visibility=both) would be picked as the "summary"
|
||||
// instead of the actual summary.
|
||||
insertMsg(t, chat.ID, "system", database.ChatMessageVisibilityModel, false, "system prompt")
|
||||
summary := insertMsg(t, chat.ID, "user", database.ChatMessageVisibilityModel, true, "real summary")
|
||||
compressedTool := insertMsg(t, chat.ID, "tool", database.ChatMessageVisibilityBoth, true, "tool result")
|
||||
postUser := insertMsg(t, chat.ID, "user", database.ChatMessageVisibilityBoth, false, "follow-up")
|
||||
insertMsg(t, chat.ID, database.ChatMessageRoleSystem, database.ChatMessageVisibilityModel, false, "system prompt")
|
||||
summary := insertMsg(t, chat.ID, database.ChatMessageRoleUser, database.ChatMessageVisibilityModel, true, "real summary")
|
||||
compressedTool := insertMsg(t, chat.ID, database.ChatMessageRoleTool, database.ChatMessageVisibilityBoth, true, "tool result")
|
||||
postUser := insertMsg(t, chat.ID, database.ChatMessageRoleUser, database.ChatMessageVisibilityBoth, false, "follow-up")
|
||||
|
||||
got, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -3311,7 +3311,7 @@ func (q *sqlQuerier) GetChatDiffStatusesByChatIDs(ctx context.Context, chatIds [
|
||||
|
||||
const getChatMessageByID = `-- name: GetChatMessageByID :one
|
||||
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
|
||||
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
|
||||
FROM
|
||||
chat_messages
|
||||
WHERE
|
||||
@@ -3338,13 +3338,14 @@ func (q *sqlQuerier) GetChatMessageByID(ctx context.Context, id int64) (ChatMess
|
||||
&i.ContextLimit,
|
||||
&i.Compressed,
|
||||
&i.CreatedBy,
|
||||
&i.ContentVersion,
|
||||
)
|
||||
return i, err
|
||||
}
|
||||
|
||||
const getChatMessagesByChatID = `-- name: GetChatMessagesByChatID :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
|
||||
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
|
||||
FROM
|
||||
chat_messages
|
||||
WHERE
|
||||
@@ -3386,6 +3387,7 @@ func (q *sqlQuerier) GetChatMessagesByChatID(ctx context.Context, arg GetChatMes
|
||||
&i.ContextLimit,
|
||||
&i.Compressed,
|
||||
&i.CreatedBy,
|
||||
&i.ContentVersion,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -3417,7 +3419,7 @@ WITH latest_compressed_summary AS (
|
||||
1
|
||||
)
|
||||
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
|
||||
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
|
||||
FROM
|
||||
chat_messages
|
||||
WHERE
|
||||
@@ -3483,6 +3485,7 @@ func (q *sqlQuerier) GetChatMessagesForPromptByChatID(ctx context.Context, chatI
|
||||
&i.ContextLimit,
|
||||
&i.Compressed,
|
||||
&i.CreatedBy,
|
||||
&i.ContentVersion,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -3626,12 +3629,12 @@ func (q *sqlQuerier) GetChatsByOwnerID(ctx context.Context, arg GetChatsByOwnerI
|
||||
|
||||
const getLastChatMessageByRole = `-- name: GetLastChatMessageByRole :one
|
||||
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
|
||||
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
|
||||
FROM
|
||||
chat_messages
|
||||
WHERE
|
||||
chat_id = $1::uuid
|
||||
AND role = $2::text
|
||||
AND role = $2::chat_message_role
|
||||
ORDER BY
|
||||
created_at DESC, id DESC
|
||||
LIMIT
|
||||
@@ -3639,8 +3642,8 @@ LIMIT
|
||||
`
|
||||
|
||||
type GetLastChatMessageByRoleParams struct {
|
||||
ChatID uuid.UUID `db:"chat_id" json:"chat_id"`
|
||||
Role string `db:"role" json:"role"`
|
||||
ChatID uuid.UUID `db:"chat_id" json:"chat_id"`
|
||||
Role ChatMessageRole `db:"role" json:"role"`
|
||||
}
|
||||
|
||||
func (q *sqlQuerier) GetLastChatMessageByRole(ctx context.Context, arg GetLastChatMessageByRoleParams) (ChatMessage, error) {
|
||||
@@ -3663,6 +3666,7 @@ func (q *sqlQuerier) GetLastChatMessageByRole(ctx context.Context, arg GetLastCh
|
||||
&i.ContextLimit,
|
||||
&i.Compressed,
|
||||
&i.CreatedBy,
|
||||
&i.ContentVersion,
|
||||
)
|
||||
return i, err
|
||||
}
|
||||
@@ -3793,6 +3797,7 @@ INSERT INTO chat_messages (
|
||||
model_config_id,
|
||||
role,
|
||||
content,
|
||||
content_version,
|
||||
visibility,
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
@@ -3806,28 +3811,30 @@ INSERT INTO chat_messages (
|
||||
$1::uuid,
|
||||
$2::uuid,
|
||||
$3::uuid,
|
||||
$4::text,
|
||||
$4::chat_message_role,
|
||||
$5::jsonb,
|
||||
$6::chat_message_visibility,
|
||||
$7::bigint,
|
||||
$6::smallint,
|
||||
$7::chat_message_visibility,
|
||||
$8::bigint,
|
||||
$9::bigint,
|
||||
$10::bigint,
|
||||
$11::bigint,
|
||||
$12::bigint,
|
||||
$13::bigint,
|
||||
COALESCE($14::boolean, FALSE)
|
||||
$14::bigint,
|
||||
COALESCE($15::boolean, FALSE)
|
||||
)
|
||||
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
|
||||
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
|
||||
`
|
||||
|
||||
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 string `db:"role" json:"role"`
|
||||
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"`
|
||||
@@ -3846,6 +3853,7 @@ func (q *sqlQuerier) InsertChatMessage(ctx context.Context, arg InsertChatMessag
|
||||
arg.ModelConfigID,
|
||||
arg.Role,
|
||||
arg.Content,
|
||||
arg.ContentVersion,
|
||||
arg.Visibility,
|
||||
arg.InputTokens,
|
||||
arg.OutputTokens,
|
||||
@@ -3874,6 +3882,7 @@ func (q *sqlQuerier) InsertChatMessage(ctx context.Context, arg InsertChatMessag
|
||||
&i.ContextLimit,
|
||||
&i.Compressed,
|
||||
&i.CreatedBy,
|
||||
&i.ContentVersion,
|
||||
)
|
||||
return i, err
|
||||
}
|
||||
@@ -4008,7 +4017,7 @@ SET
|
||||
WHERE
|
||||
id = $3::bigint
|
||||
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
|
||||
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
|
||||
`
|
||||
|
||||
type UpdateChatMessageByIDParams struct {
|
||||
@@ -4037,6 +4046,7 @@ func (q *sqlQuerier) UpdateChatMessageByID(ctx context.Context, arg UpdateChatMe
|
||||
&i.ContextLimit,
|
||||
&i.Compressed,
|
||||
&i.CreatedBy,
|
||||
&i.ContentVersion,
|
||||
)
|
||||
return i, err
|
||||
}
|
||||
|
||||
@@ -170,6 +170,7 @@ INSERT INTO chat_messages (
|
||||
model_config_id,
|
||||
role,
|
||||
content,
|
||||
content_version,
|
||||
visibility,
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
@@ -183,8 +184,9 @@ INSERT INTO chat_messages (
|
||||
@chat_id::uuid,
|
||||
sqlc.narg('created_by')::uuid,
|
||||
sqlc.narg('model_config_id')::uuid,
|
||||
@role::text,
|
||||
@role::chat_message_role,
|
||||
sqlc.narg('content')::jsonb,
|
||||
@content_version::smallint,
|
||||
@visibility::chat_message_visibility,
|
||||
sqlc.narg('input_tokens')::bigint,
|
||||
sqlc.narg('output_tokens')::bigint,
|
||||
@@ -422,7 +424,7 @@ FROM
|
||||
chat_messages
|
||||
WHERE
|
||||
chat_id = @chat_id::uuid
|
||||
AND role = @role::text
|
||||
AND role = @role::chat_message_role
|
||||
ORDER BY
|
||||
created_at DESC, id DESC
|
||||
LIMIT
|
||||
|
||||
Reference in New Issue
Block a user