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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user