feat(coderd/chatd): unify chat storage on SDK parts and fix file-reference rendering (#22958)

File-reference parts in user messages were flattened to `TextContent` at
write time because fantasy has no file-reference content type. The
frontend never saw them as structured parts.

This moves all write paths (user, assistant, tool) from fantasy envelope
format to `codersdk.ChatMessagePart`. The streaming layer (`chatloop`)
is untouched, the conversion happens at the serialization boundary in
`persistStep`.

Old rows are still readable. `ParseContent` uses a structural heuristic
(`isFantasyEnvelopeFormat`) to distinguish legacy envelopes from SDK
parts. We chose this over try/fallback because fantasy envelopes
partially unmarshal into `ChatMessagePart` (the `type` field matches)
while silently losing content. A guard test enforces that no SDK part
can produce the envelope shape.

This is forward-only: new rows are unreadable by old code. Chat is
behind a feature flag so rollback risk is contained.

Also adds a typed `ChatMessageRole` to replace raw strings and
`fantasy.MessageRole*` casts at the persistence boundary. The type
covers `ChatMessage.Role`, `ChatStreamMessagePart.Role`, the
`PublishMessagePart` callback chain, and all DB write sites.
`fantasy.MessageRole*` remains only where we build `fantasy.Message`
structs for LLM dispatch.

Separately, `ProviderMetadata` was leaking to SSE clients via
`publishMessagePart`. `StripInternal` now runs on both the SSE and REST
paths, covering this.

Other cleanup:

- Old `db2sdk.contentBlockToPart` silently dropped metadata on
text/reasoning/tool-call content. New code preserves it.
- `providerMetadataToOptions` now logs warnings instead of silently
returning nil.
- `db2sdk` shrinks from ~250 lines of parallel conversion to ~15 lines
delegating to `chatprompt.ParseContent()`, removing the `fantasy` import
entirely.

Refs #22821
This commit is contained in:
Mathias Fredriksson
2026-03-13 17:53:26 +02:00
committed by GitHub
parent 870583224d
commit bdbcd3428b
19 changed files with 1734 additions and 912 deletions
+42 -48
View File
@@ -179,10 +179,7 @@ type CreateOptions struct {
Title string
ModelConfigID uuid.UUID
SystemPrompt string
InitialUserContent []fantasy.Content
// ContentFileIDs maps content block indices to their chat_files IDs
// so the file_id can be preserved in the stored message JSON.
ContentFileIDs map[int]uuid.UUID
InitialUserContent []codersdk.ChatMessagePart
}
// SendMessageBusyBehavior controls what happens when a chat is already active.
@@ -200,12 +197,11 @@ const (
// SendMessageOptions controls user message insertion with busy-state behavior.
type SendMessageOptions struct {
ChatID uuid.UUID
CreatedBy uuid.UUID
Content []fantasy.Content
ContentFileIDs map[int]uuid.UUID
ModelConfigID *uuid.UUID
BusyBehavior SendMessageBusyBehavior
ChatID uuid.UUID
CreatedBy uuid.UUID
Content []codersdk.ChatMessagePart
ModelConfigID *uuid.UUID
BusyBehavior SendMessageBusyBehavior
}
// SendMessageResult contains the outcome of user message processing.
@@ -221,8 +217,7 @@ type EditMessageOptions struct {
ChatID uuid.UUID
CreatedBy uuid.UUID
EditedMessageID int64
Content []fantasy.Content
ContentFileIDs map[int]uuid.UUID
Content []codersdk.ChatMessagePart
}
// EditMessageResult contains the updated user message and chat status.
@@ -284,7 +279,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
UUID: opts.ModelConfigID,
Valid: true,
},
Role: "system",
Role: string(codersdk.ChatMessageRoleSystem),
Content: pqtype.NullRawMessage{
RawMessage: systemContent,
Valid: len(systemContent) > 0,
@@ -304,7 +299,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
}
}
userContent, err := chatprompt.MarshalContent(opts.InitialUserContent, opts.ContentFileIDs)
userContent, err := chatprompt.MarshalParts(opts.InitialUserContent)
if err != nil {
return xerrors.Errorf("marshal initial user content: %w", err)
}
@@ -314,7 +309,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
UUID: opts.ModelConfigID,
Valid: true,
},
Role: "user",
Role: string(codersdk.ChatMessageRoleUser),
Content: userContent,
CreatedBy: uuid.NullUUID{UUID: opts.OwnerID, Valid: opts.OwnerID != uuid.Nil},
Visibility: database.ChatMessageVisibilityBoth,
@@ -372,7 +367,7 @@ func (p *Server) SendMessage(
return SendMessageResult{}, xerrors.Errorf("invalid busy behavior %q", opts.BusyBehavior)
}
content, err := chatprompt.MarshalContent(opts.Content, opts.ContentFileIDs)
content, err := chatprompt.MarshalParts(opts.Content)
if err != nil {
return SendMessageResult{}, xerrors.Errorf("marshal message content: %w", err)
}
@@ -506,7 +501,7 @@ func (p *Server) EditMessage(
return EditMessageResult{}, xerrors.New("content is required")
}
content, err := chatprompt.MarshalContent(opts.Content, opts.ContentFileIDs)
content, err := chatprompt.MarshalParts(opts.Content)
if err != nil {
return EditMessageResult{}, xerrors.Errorf("marshal message content: %w", err)
}
@@ -902,7 +897,7 @@ func insertUserMessageAndSetPending(
message, err := insertChatMessageWithStore(ctx, store, database.InsertChatMessageParams{
ChatID: lockedChat.ID,
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
Role: "user",
Role: string(codersdk.ChatMessageRoleUser),
Content: content,
CreatedBy: uuid.NullUUID{UUID: createdBy, Valid: createdBy != uuid.Nil},
Visibility: database.ChatMessageVisibilityBoth,
@@ -1740,10 +1735,13 @@ func (p *Server) publishEditedMessage(chatID uuid.UUID, message database.ChatMes
})
}
func (p *Server) publishMessagePart(chatID uuid.UUID, role string, part codersdk.ChatMessagePart) {
func (p *Server) publishMessagePart(chatID uuid.UUID, role codersdk.ChatMessageRole, part codersdk.ChatMessagePart) {
if part.Type == "" {
return
}
// Strip internal-only fields before client delivery.
// Mirrors db2sdk.chatMessageParts stripping for REST.
part.StripInternal()
p.publishEvent(chatID, codersdk.ChatStreamEvent{
Type: codersdk.ChatStreamEventTypeMessagePart,
MessagePart: &codersdk.ChatStreamMessagePart{
@@ -1954,7 +1952,7 @@ func (p *Server) processChat(ctx context.Context, chat database.Chat) {
msg, insertErr := tx.InsertChatMessage(cleanupCtx, database.InsertChatMessageParams{
ChatID: chat.ID,
ModelConfigID: uuid.NullUUID{UUID: latestChat.LastModelConfigID, Valid: true},
Role: "user",
Role: string(codersdk.ChatMessageRoleUser),
Content: pqtype.NullRawMessage{
RawMessage: nextQueued.Content,
Valid: len(nextQueued.Content) > 0,
@@ -2336,7 +2334,11 @@ func (p *Server) runChat(
}
if len(assistantBlocks) > 0 {
assistantContent, marshalErr := chatprompt.MarshalContent(assistantBlocks, nil)
sdkParts := make([]codersdk.ChatMessagePart, 0, len(assistantBlocks))
for _, block := range assistantBlocks {
sdkParts = append(sdkParts, chatprompt.PartFromContent(block))
}
assistantContent, marshalErr := chatprompt.MarshalParts(sdkParts)
if marshalErr != nil {
return marshalErr
}
@@ -2346,7 +2348,7 @@ func (p *Server) runChat(
ChatID: chat.ID,
CreatedBy: uuid.NullUUID{},
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
Role: string(fantasy.MessageRoleAssistant),
Role: string(codersdk.ChatMessageRoleAssistant),
Content: assistantContent,
Visibility: database.ChatMessageVisibilityBoth,
InputTokens: usageNullInt64(step.Usage.InputTokens, hasUsage),
@@ -2371,7 +2373,8 @@ func (p *Server) runChat(
}
for _, tr := range toolResults {
resultContent, marshalErr := chatprompt.MarshalToolResultContent(tr)
trPart := chatprompt.PartFromContent(tr)
resultContent, marshalErr := chatprompt.MarshalParts([]codersdk.ChatMessagePart{trPart})
if marshalErr != nil {
return marshalErr
}
@@ -2380,7 +2383,7 @@ func (p *Server) runChat(
ChatID: chat.ID,
CreatedBy: uuid.NullUUID{},
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
Role: string(fantasy.MessageRoleTool),
Role: string(codersdk.ChatMessageRoleTool),
Content: resultContent,
Visibility: database.ChatMessageVisibilityBoth,
InputTokens: sql.NullInt64{},
@@ -2461,8 +2464,8 @@ func (p *Server) runChat(
},
ToolCallID: compactionToolCallID,
ToolName: "chat_summarized",
PublishMessagePart: func(role fantasy.MessageRole, part codersdk.ChatMessagePart) {
p.publishMessagePart(chat.ID, string(role), part)
PublishMessagePart: func(role codersdk.ChatMessageRole, part codersdk.ChatMessagePart) {
p.publishMessagePart(chat.ID, role, part)
},
OnError: func(err error) {
logger.Warn(ctx, "failed to compact chat context", slog.Error(err))
@@ -2550,10 +2553,10 @@ func (p *Server) runChat(
PersistStep: persistStep,
PublishMessagePart: func(
role fantasy.MessageRole,
role codersdk.ChatMessageRole,
part codersdk.ChatMessagePart,
) {
p.publishMessagePart(chat.ID, string(role), part)
p.publishMessagePart(chat.ID, role, part)
},
Compaction: compactionOptions,
ReloadMessages: func(reloadCtx context.Context) ([]fantasy.Message, error) {
@@ -2686,13 +2689,9 @@ func (p *Server) persistChatContextSummary(
return xerrors.Errorf("encode summary tool args: %w", err)
}
assistantContent, err := chatprompt.MarshalContent([]fantasy.Content{
fantasy.ToolCallContent{
ToolCallID: toolCallID,
ToolName: "chat_summarized",
Input: string(args),
},
}, nil)
assistantContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
codersdk.ChatMessageToolCall(toolCallID, "chat_summarized", args),
})
if err != nil {
return xerrors.Errorf("encode summary tool call: %w", err)
}
@@ -2708,14 +2707,9 @@ func (p *Server) persistChatContextSummary(
if err != nil {
return xerrors.Errorf("encode summary result payload: %w", err)
}
toolResult, err := chatprompt.MarshalToolResult(
toolCallID,
"chat_summarized",
summaryResult,
false,
false,
nil,
)
toolResult, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
codersdk.ChatMessageToolResult(toolCallID, "chat_summarized", summaryResult, false),
})
if err != nil {
return xerrors.Errorf("encode summary tool result: %w", err)
}
@@ -2727,7 +2721,7 @@ func (p *Server) persistChatContextSummary(
ChatID: chatID,
CreatedBy: uuid.NullUUID{},
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
Role: string(fantasy.MessageRoleUser),
Role: string(codersdk.ChatMessageRoleUser),
Content: pqtype.NullRawMessage{
RawMessage: systemContent,
Valid: len(systemContent) > 0,
@@ -2750,7 +2744,7 @@ func (p *Server) persistChatContextSummary(
ChatID: chatID,
CreatedBy: uuid.NullUUID{},
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
Role: string(fantasy.MessageRoleAssistant),
Role: string(codersdk.ChatMessageRoleAssistant),
Content: assistantContent,
Visibility: database.ChatMessageVisibilityUser,
Compressed: sql.NullBool{
@@ -2774,7 +2768,7 @@ func (p *Server) persistChatContextSummary(
ChatID: chatID,
CreatedBy: uuid.NullUUID{},
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
Role: string(fantasy.MessageRoleTool),
Role: string(codersdk.ChatMessageRoleTool),
Content: toolResult,
Visibility: database.ChatMessageVisibilityBoth,
Compressed: sql.NullBool{
@@ -3140,10 +3134,10 @@ func (p *Server) maybeSendPushNotification(
msg, err := p.db.GetLastChatMessageByRole(pushCtx, database.GetLastChatMessageByRoleParams{
ChatID: chat.ID,
Role: "assistant",
Role: string(codersdk.ChatMessageRoleAssistant),
})
if err == nil {
content, parseErr := chatprompt.ParseContent(msg.Role, msg.Content)
content, parseErr := chatprompt.ParseContent(codersdk.ChatMessageRoleAssistant, msg.Content)
if parseErr == nil {
assistantText := strings.TrimSpace(contentBlocksToText(content))
if assistantText != "" {
+25 -26
View File
@@ -12,7 +12,6 @@ import (
"testing"
"time"
"charm.land/fantasy"
"github.com/google/uuid"
"github.com/sqlc-dev/pqtype"
"github.com/stretchr/testify/require"
@@ -48,7 +47,7 @@ func TestInterruptChatBroadcastsStatusAcrossInstances(t *testing.T) {
OwnerID: user.ID,
Title: "interrupt-me",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
@@ -251,7 +250,7 @@ func TestInterruptChatClearsWorkerInDatabase(t *testing.T) {
OwnerID: user.ID,
Title: "db-transition",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
@@ -287,7 +286,7 @@ func TestUpdateChatHeartbeatRequiresOwnership(t *testing.T) {
OwnerID: user.ID,
Title: "heartbeat-ownership",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
@@ -329,7 +328,7 @@ func TestSendMessageQueueBehaviorQueuesWhenBusy(t *testing.T) {
OwnerID: user.ID,
Title: "queue-when-busy",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
@@ -345,7 +344,7 @@ func TestSendMessageQueueBehaviorQueuesWhenBusy(t *testing.T) {
result, err := replica.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: chat.ID,
Content: []fantasy.Content{fantasy.TextContent{Text: "queued"}},
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued")},
BusyBehavior: chatd.SendMessageBusyBehaviorQueue,
})
require.NoError(t, err)
@@ -380,7 +379,7 @@ func TestSendMessageInterruptBehaviorQueuesAndInterruptsWhenBusy(t *testing.T) {
OwnerID: user.ID,
Title: "interrupt-when-busy",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
@@ -395,7 +394,7 @@ func TestSendMessageInterruptBehaviorQueuesAndInterruptsWhenBusy(t *testing.T) {
result, err := replica.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: chat.ID,
Content: []fantasy.Content{fantasy.TextContent{Text: "interrupt"}},
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("interrupt")},
BusyBehavior: chatd.SendMessageBusyBehaviorInterrupt,
})
require.NoError(t, err)
@@ -439,7 +438,7 @@ func TestEditMessageUpdatesAndTruncatesAndClearsQueue(t *testing.T) {
OwnerID: user.ID,
Title: "edit-message",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "original"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("original")},
})
require.NoError(t, err)
@@ -453,13 +452,13 @@ func TestEditMessageUpdatesAndTruncatesAndClearsQueue(t *testing.T) {
_, err = replica.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: chat.ID,
Content: []fantasy.Content{fantasy.TextContent{Text: "follow-up"}},
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("follow-up")},
BusyBehavior: chatd.SendMessageBusyBehaviorInterrupt,
})
require.NoError(t, err)
_, err = replica.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: chat.ID,
Content: []fantasy.Content{fantasy.TextContent{Text: "another"}},
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("another")},
BusyBehavior: chatd.SendMessageBusyBehaviorInterrupt,
})
require.NoError(t, err)
@@ -482,7 +481,7 @@ func TestEditMessageUpdatesAndTruncatesAndClearsQueue(t *testing.T) {
editResult, err := replica.EditMessage(ctx, chatd.EditMessageOptions{
ChatID: chat.ID,
EditedMessageID: editedMessageID,
Content: []fantasy.Content{fantasy.TextContent{Text: "edited"}},
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")},
})
require.NoError(t, err)
require.Equal(t, editedMessageID, editResult.Message.ID)
@@ -527,14 +526,14 @@ func TestEditMessageRejectsMissingMessage(t *testing.T) {
OwnerID: user.ID,
Title: "missing-edited-message",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
_, err = replica.EditMessage(ctx, chatd.EditMessageOptions{
ChatID: chat.ID,
EditedMessageID: 999999,
Content: []fantasy.Content{fantasy.TextContent{Text: "edited"}},
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")},
})
require.Error(t, err)
require.True(t, errors.Is(err, chatd.ErrEditedMessageNotFound))
@@ -553,7 +552,7 @@ func TestEditMessageRejectsNonUserMessage(t *testing.T) {
OwnerID: user.ID,
Title: "non-user-edited-message",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
@@ -580,7 +579,7 @@ func TestEditMessageRejectsNonUserMessage(t *testing.T) {
_, err = replica.EditMessage(ctx, chatd.EditMessageOptions{
ChatID: chat.ID,
EditedMessageID: assistantMessage.ID,
Content: []fantasy.Content{fantasy.TextContent{Text: "edited"}},
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")},
})
require.Error(t, err)
require.True(t, errors.Is(err, chatd.ErrEditedMessageNotUser))
@@ -827,7 +826,7 @@ func TestSubscribeSnapshotIncludesStatusEvent(t *testing.T) {
OwnerID: user.ID,
Title: "status-snapshot",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
@@ -856,7 +855,7 @@ func TestSubscribeNoPubsubNoDuplicateMessageParts(t *testing.T) {
OwnerID: user.ID,
Title: "no-dup-parts",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
@@ -896,7 +895,7 @@ func TestSubscribeAfterMessageID(t *testing.T) {
OwnerID: user.ID,
Title: "after-id-test",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "first"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("first")},
})
require.NoError(t, err)
@@ -952,7 +951,7 @@ func TestSubscribeAfterMessageID(t *testing.T) {
partialMessages := filterMessageEvents(partialSnapshot)
require.Len(t, partialMessages, 1, "afterMessageID=msg2.ID should return only messages after msg2")
require.Equal(t, "user", partialMessages[0].Message.Role)
require.Equal(t, codersdk.ChatMessageRoleUser, partialMessages[0].Message.Role)
}
// filterMessageEvents returns only the Message-type events from a
@@ -1084,7 +1083,7 @@ func TestCreateWorkspaceTool_EndToEnd(t *testing.T) {
var foundCreateWorkspaceResult bool
for _, message := range chatMsgs.Messages {
if message.Role != "tool" {
if message.Role != codersdk.ChatMessageRoleTool {
continue
}
for _, part := range message.Content {
@@ -1256,7 +1255,7 @@ func TestStartWorkspaceTool_EndToEnd(t *testing.T) {
// Verify start_workspace tool result exists in the chat messages.
var foundStartWorkspaceResult bool
for _, message := range chatMsgs.Messages {
if message.Role != "tool" {
if message.Role != codersdk.ChatMessageRoleTool {
continue
}
for _, part := range message.Content {
@@ -1431,7 +1430,7 @@ func TestInterruptChatDoesNotSendWebPushNotification(t *testing.T) {
OwnerID: user.ID,
Title: "interrupt-no-push",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
@@ -1542,7 +1541,7 @@ func TestSuccessfulChatSendsWebPushWithNavigationData(t *testing.T) {
OwnerID: user.ID,
Title: "push-nav-test",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
@@ -1626,7 +1625,7 @@ func TestCloseDuringShutdownContextCanceledShouldRetryOnNewReplica(t *testing.T)
OwnerID: user.ID,
Title: "shutdown-retry",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
@@ -1732,7 +1731,7 @@ func TestSuccessfulChatSendsWebPushWithSummary(t *testing.T) {
OwnerID: user.ID,
Title: "summary-push-test",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "do the thing"}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("do the thing")},
})
require.NoError(t, err)
+10 -16
View File
@@ -71,7 +71,7 @@ type RunOptions struct {
PersistStep func(context.Context, PersistedStep) error
PublishMessagePart func(
role fantasy.MessageRole,
role codersdk.ChatMessageRole,
part codersdk.ChatMessagePart,
)
Compaction *CompactionOptions
@@ -216,7 +216,7 @@ func Run(ctx context.Context, opts RunOptions) error {
opts.MaxSteps = 1
}
publishMessagePart := func(role fantasy.MessageRole, part codersdk.ChatMessagePart) {
publishMessagePart := func(role codersdk.ChatMessageRole, part codersdk.ChatMessagePart) {
if opts.PublishMessagePart == nil {
return
}
@@ -317,7 +317,7 @@ func Run(ctx context.Context, opts RunOptions) error {
toolResults = executeTools(ctx, opts.Tools, result.toolCalls, func(tr fantasy.ToolResultContent) {
publishMessagePart(
fantasy.MessageRoleTool,
codersdk.ChatMessageRoleTool,
chatprompt.PartFromContent(tr),
)
})
@@ -455,7 +455,7 @@ func Run(ctx context.Context, opts RunOptions) error {
func processStepStream(
ctx context.Context,
stream fantasy.StreamResponse,
publishMessagePart func(fantasy.MessageRole, codersdk.ChatMessagePart),
publishMessagePart func(codersdk.ChatMessageRole, codersdk.ChatMessagePart),
) (stepResult, error) {
var result stepResult
@@ -474,10 +474,7 @@ func processStepStream(
if _, exists := activeTextContent[part.ID]; exists {
activeTextContent[part.ID] += part.Delta
}
publishMessagePart(fantasy.MessageRoleAssistant, codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeText,
Text: part.Delta,
})
publishMessagePart(codersdk.ChatMessageRoleAssistant, codersdk.ChatMessageText(part.Delta))
case fantasy.StreamPartTypeTextEnd:
if text, exists := activeTextContent[part.ID]; exists {
@@ -500,10 +497,7 @@ func processStepStream(
active.options = part.ProviderMetadata
activeReasoningContent[part.ID] = active
}
publishMessagePart(fantasy.MessageRoleAssistant, codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeReasoning,
Text: part.Delta,
})
publishMessagePart(codersdk.ChatMessageRoleAssistant, codersdk.ChatMessageReasoning(part.Delta))
case fantasy.StreamPartTypeReasoningEnd:
if active, exists := activeReasoningContent[part.ID]; exists {
@@ -535,7 +529,7 @@ func processStepStream(
providerExecuted = toolCall.ProviderExecuted
}
toolName := toolNames[part.ID]
publishMessagePart(fantasy.MessageRoleAssistant, codersdk.ChatMessagePart{
publishMessagePart(codersdk.ChatMessageRoleAssistant, codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeToolCall,
ToolCallID: part.ID,
ToolName: toolName,
@@ -563,7 +557,7 @@ func processStepStream(
delete(activeToolCalls, part.ID)
publishMessagePart(
fantasy.MessageRoleAssistant,
codersdk.ChatMessageRoleAssistant,
chatprompt.PartFromContent(tc),
)
@@ -577,7 +571,7 @@ func processStepStream(
}
result.content = append(result.content, sourceContent)
publishMessagePart(
fantasy.MessageRoleAssistant,
codersdk.ChatMessageRoleAssistant,
chatprompt.PartFromContent(sourceContent),
)
@@ -595,7 +589,7 @@ func processStepStream(
}
result.content = append(result.content, tr)
publishMessagePart(
fantasy.MessageRoleTool,
codersdk.ChatMessageRoleTool,
chatprompt.PartFromContent(tr),
)
}
+7 -22
View File
@@ -55,7 +55,7 @@ type CompactionOptions struct {
// PublishMessagePart publishes streaming parts to connected
// clients so they see "Summarizing..." / "Summarized" UI
// transitions during compaction.
PublishMessagePart func(fantasy.MessageRole, codersdk.ChatMessagePart)
PublishMessagePart func(codersdk.ChatMessageRole, codersdk.ChatMessagePart)
OnError func(error)
}
@@ -110,12 +110,8 @@ func tryCompact(
// connected clients see activity during summary generation.
if config.PublishMessagePart != nil && config.ToolCallID != "" {
config.PublishMessagePart(
fantasy.MessageRoleAssistant,
codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeToolCall,
ToolCallID: config.ToolCallID,
ToolName: config.ToolName,
},
codersdk.ChatMessageRoleAssistant,
codersdk.ChatMessageToolCall(config.ToolCallID, config.ToolName, nil),
)
}
@@ -163,13 +159,8 @@ func tryCompact(
"context_limit_tokens": contextLimit,
})
config.PublishMessagePart(
fantasy.MessageRoleTool,
codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeToolResult,
ToolCallID: config.ToolCallID,
ToolName: config.ToolName,
Result: resultJSON,
},
codersdk.ChatMessageRoleTool,
codersdk.ChatMessageToolResult(config.ToolCallID, config.ToolName, resultJSON, false),
)
}
@@ -186,14 +177,8 @@ func publishCompactionError(config CompactionOptions, msg string) {
"error": msg,
})
config.PublishMessagePart(
fantasy.MessageRoleTool,
codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeToolResult,
ToolCallID: config.ToolCallID,
ToolName: config.ToolName,
Result: errJSON,
IsError: true,
},
codersdk.ChatMessageRoleTool,
codersdk.ChatMessageToolResult(config.ToolCallID, config.ToolName, errJSON, true),
)
}
+2 -2
View File
@@ -149,7 +149,7 @@ func TestRun_Compaction(t *testing.T) {
SummaryPrompt: "summarize now",
ToolCallID: "test-tool-call-id",
ToolName: "chat_summarized",
PublishMessagePart: func(role fantasy.MessageRole, part codersdk.ChatMessagePart) {
PublishMessagePart: func(role codersdk.ChatMessageRole, part codersdk.ChatMessagePart) {
switch part.Type {
case codersdk.ChatMessagePartTypeToolCall:
callOrder = append(callOrder, "publish_tool_call")
@@ -218,7 +218,7 @@ func TestRun_Compaction(t *testing.T) {
ThresholdPercent: 70,
ToolCallID: "test-tool-call-id",
ToolName: "chat_summarized",
PublishMessagePart: func(_ fantasy.MessageRole, _ codersdk.ChatMessagePart) {
PublishMessagePart: func(_ codersdk.ChatMessageRole, _ codersdk.ChatMessagePart) {
publishCalled = true
},
Persist: func(_ context.Context, _ CompactionResult) error {
File diff suppressed because it is too large Load Diff
+771 -7
View File
@@ -1,18 +1,23 @@
package chatprompt_test
import (
"bytes"
"context"
"encoding/json"
"testing"
"charm.land/fantasy"
fantasyanthropic "charm.land/fantasy/providers/anthropic"
"github.com/google/uuid"
"github.com/sqlc-dev/pqtype"
"github.com/stretchr/testify/assert"
"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/database"
"github.com/coder/coder/v2/coderd/database/db2sdk"
"github.com/coder/coder/v2/codersdk"
)
func TestConvertMessages_NormalizesAssistantToolCallInput(t *testing.T) {
@@ -46,7 +51,6 @@ func TestConvertMessages_NormalizesAssistantToolCallInput(t *testing.T) {
}
for _, tc := range testCases {
tc := tc
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
@@ -153,22 +157,37 @@ func TestConvertMessagesWithFiles_ResolvesFileData(t *testing.T) {
func TestConvertMessagesWithFiles_BackwardCompat(t *testing.T) {
t.Parallel()
// A message with inline data and a file_id should use the
// inline data even when the resolver returns nothing.
// A legacy message with inline data and a file_id: ParseContent
// extracts the file_id and clears inline data (resolved at LLM
// dispatch time). When a resolver provides data, the file part
// in the LLM prompt should contain the resolved data.
fileID := uuid.New()
inlineData := []byte("inline-image-data")
resolvedData := []byte("resolved-image-data")
rawContent := mustJSON(t, []json.RawMessage{
mustJSON(t, map[string]any{
"type": "file",
"data": map[string]any{
"media_type": "image/png",
"data": inlineData,
"data": []byte("inline-image-data"),
"file_id": fileID.String(),
},
}),
})
resolver := func(_ context.Context, ids []uuid.UUID) (map[uuid.UUID]chatprompt.FileData, error) {
result := make(map[uuid.UUID]chatprompt.FileData)
for _, id := range ids {
if id == fileID {
result[id] = chatprompt.FileData{
Data: resolvedData,
MediaType: "image/png",
}
}
}
return result, nil
}
prompt, err := chatprompt.ConvertMessagesWithFiles(
context.Background(),
[]database.ChatMessage{
@@ -178,7 +197,7 @@ func TestConvertMessagesWithFiles_BackwardCompat(t *testing.T) {
Content: pqtype.NullRawMessage{RawMessage: rawContent, Valid: true},
},
},
nil, // No resolver.
resolver,
slogtest.Make(t, nil),
)
require.NoError(t, err)
@@ -187,7 +206,8 @@ func TestConvertMessagesWithFiles_BackwardCompat(t *testing.T) {
filePart, ok := fantasy.AsMessagePart[fantasy.FilePart](prompt[0].Content[0])
require.True(t, ok, "expected FilePart")
require.Equal(t, inlineData, filePart.Data)
require.Equal(t, resolvedData, filePart.Data)
require.Equal(t, "image/png", filePart.MediaType)
}
func TestInjectFileID_StripsInlineData(t *testing.T) {
@@ -587,6 +607,750 @@ func TestProviderExecutedResult_LegacyToolRow(t *testing.T) {
require.Equal(t, []string{"toolu_exec"}, toolIDs)
}
// TestSDKPartsNeverProduceFantasyEnvelopeShape guards the structural
// invariant that isFantasyEnvelopeFormat relies on: no SDK part type
// serializes with a top-level "data" field containing a JSON object
// (starting with '{'). Fantasy envelopes always have
// "data":{object}, while ChatMessagePart.Data is []byte which
// serializes to a base64 string or is omitted. If this test fails,
// the format discriminator can no longer distinguish legacy fantasy
// content from SDK parts, and parseAssistantRole / parseUserRole
// would silently lose data on legacy rows.
func TestSDKPartsNeverProduceFantasyEnvelopeShape(t *testing.T) {
t.Parallel()
parts := []codersdk.ChatMessagePart{
{Type: codersdk.ChatMessagePartTypeText, Text: "hello"},
{Type: codersdk.ChatMessagePartTypeFile, FileID: uuid.NullUUID{UUID: uuid.New(), Valid: true}, MediaType: "image/png"},
{Type: codersdk.ChatMessagePartTypeFile, MediaType: "image/png", Data: []byte("fake-image-data")},
{Type: codersdk.ChatMessagePartTypeFileReference, FileName: "main.go", StartLine: 1, EndLine: 10, Content: "func main() {}"},
{Type: codersdk.ChatMessagePartTypeReasoning, Text: "thinking..."},
{Type: codersdk.ChatMessagePartTypeToolCall, ToolCallID: "abc", ToolName: "read_file", Args: json.RawMessage(`{"path":"main.go"}`)},
{Type: codersdk.ChatMessagePartTypeToolResult, ToolCallID: "abc", ToolName: "read_file", Result: json.RawMessage(`{"output":"code"}`)},
{Type: codersdk.ChatMessagePartTypeSource, SourceID: "s1", URL: "https://example.com", Title: "Example"},
}
for _, part := range parts {
raw, err := json.Marshal(part)
require.NoError(t, err)
var fields map[string]json.RawMessage
require.NoError(t, json.Unmarshal(raw, &fields))
if data, ok := fields["data"]; ok {
trimmed := bytes.TrimSpace(data)
require.NotEmpty(t, trimmed)
assert.NotEqual(t, byte('{'), trimmed[0],
"SDK part type %q serializes with data field starting with '{', "+
"would be misidentified as fantasy envelope by isFantasyEnvelopeFormat",
part.Type)
}
}
}
// nullRaw wraps raw JSON bytes in a NullRawMessage for test input.
func nullRaw(data json.RawMessage) pqtype.NullRawMessage {
return pqtype.NullRawMessage{RawMessage: data, Valid: true}
}
func TestParseContent_BackwardCompat(t *testing.T) {
t.Parallel()
fileID := uuid.New()
// Build legacy fantasy assistant content using MarshalContent.
legacyAssistantReasoning, err := chatprompt.MarshalContent([]fantasy.Content{
fantasy.ReasoningContent{
Text: "let me think...",
ProviderMetadata: fantasy.ProviderMetadata{
"anthropic": &fantasyanthropic.ProviderCacheControlOptions{
CacheControl: fantasyanthropic.CacheControl{Type: "ephemeral"},
},
},
},
}, nil)
require.NoError(t, err)
legacyAssistantSource, err := chatprompt.MarshalContent([]fantasy.Content{
fantasy.SourceContent{
ID: "src_001",
URL: "https://example.com/doc",
Title: "Example Doc",
},
}, nil)
require.NoError(t, err)
legacyAssistantToolCall, err := chatprompt.MarshalContent([]fantasy.Content{
fantasy.ToolCallContent{
ToolCallID: "call_123",
ToolName: "read_file",
Input: `{"path":"main.go"}`,
},
}, nil)
require.NoError(t, err)
// Build new SDK format using MarshalParts.
sdkMetadata := json.RawMessage(`{"anthropic":{"type":"anthropic.cache_control_options","data":{"cache_control":{"type":"ephemeral"}}}}`)
newAssistantWithMeta, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{{
Type: codersdk.ChatMessagePartTypeText,
Text: "here is my answer",
ProviderMetadata: sdkMetadata,
}})
require.NoError(t, err)
newAssistantToolCall, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{{
Type: codersdk.ChatMessagePartTypeToolCall,
ToolCallID: "call_456",
ToolName: "execute",
Args: json.RawMessage(`{"cmd":"ls"}`),
}})
require.NoError(t, err)
newToolResult, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{{
Type: codersdk.ChatMessagePartTypeToolResult,
ToolCallID: "call_456",
ToolName: "execute",
Result: json.RawMessage(`{"output":"file1.go"}`),
}})
require.NoError(t, err)
tests := []struct {
name string
role codersdk.ChatMessageRole
raw pqtype.NullRawMessage
check func(t *testing.T, parts []codersdk.ChatMessagePart)
}{
{
name: "system/plain_string",
role: codersdk.ChatMessageRoleSystem,
raw: nullRaw(mustJSON(t, "You are helpful.")),
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeText, parts[0].Type)
assert.Equal(t, "You are helpful.", parts[0].Text)
},
},
{
name: "user/fantasy_text",
role: codersdk.ChatMessageRoleUser,
raw: nullRaw(mustJSON(t, []json.RawMessage{
mustJSON(t, map[string]any{
"type": "text",
"data": map[string]any{"text": "hello from user"},
}),
})),
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeText, parts[0].Type)
assert.Equal(t, "hello from user", parts[0].Text)
},
},
{
name: "assistant/fantasy_text",
role: codersdk.ChatMessageRoleAssistant,
raw: nullRaw(mustJSON(t, []json.RawMessage{
mustJSON(t, map[string]any{
"type": "text",
"data": map[string]any{"text": "hello from assistant"},
}),
})),
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeText, parts[0].Type)
assert.Equal(t, "hello from assistant", parts[0].Text)
},
},
{
name: "user/plain_string",
role: codersdk.ChatMessageRoleUser,
raw: nullRaw(mustJSON(t, "just a plain string")),
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeText, parts[0].Type)
assert.Equal(t, "just a plain string", parts[0].Text)
},
},
{
name: "user/fantasy_file_with_file_id",
role: codersdk.ChatMessageRoleUser,
raw: nullRaw(mustJSON(t, []json.RawMessage{
mustJSON(t, map[string]any{
"type": "file",
"data": map[string]any{
"media_type": "image/png",
"file_id": fileID.String(),
},
}),
})),
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeFile, parts[0].Type)
assert.Equal(t, "image/png", parts[0].MediaType)
assert.True(t, parts[0].FileID.Valid)
assert.Equal(t, fileID, parts[0].FileID.UUID)
assert.Nil(t, parts[0].Data, "inline data cleared when file_id present")
},
},
{
name: "assistant/fantasy_reasoning_with_metadata",
role: codersdk.ChatMessageRoleAssistant,
raw: legacyAssistantReasoning,
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeReasoning, parts[0].Type)
assert.Equal(t, "let me think...", parts[0].Text)
require.NotNil(t, parts[0].ProviderMetadata, "ProviderMetadata must be preserved")
assert.Contains(t, string(parts[0].ProviderMetadata), "anthropic")
},
},
{
name: "assistant/fantasy_source",
role: codersdk.ChatMessageRoleAssistant,
raw: legacyAssistantSource,
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeSource, parts[0].Type)
assert.Equal(t, "src_001", parts[0].SourceID)
assert.Equal(t, "https://example.com/doc", parts[0].URL)
assert.Equal(t, "Example Doc", parts[0].Title)
},
},
{
name: "assistant/fantasy_tool_call",
role: codersdk.ChatMessageRoleAssistant,
raw: legacyAssistantToolCall,
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeToolCall, parts[0].Type)
assert.Equal(t, "call_123", parts[0].ToolCallID)
assert.Equal(t, "read_file", parts[0].ToolName)
assert.JSONEq(t, `{"path":"main.go"}`, string(parts[0].Args))
},
},
{
name: "tool/legacy_result_row",
role: codersdk.ChatMessageRoleTool,
raw: nullRaw(mustJSON(t, []map[string]any{{
"tool_call_id": "call_123",
"tool_name": "read_file",
"result": json.RawMessage(`{"output":"package main"}`),
}})),
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeToolResult, parts[0].Type)
assert.Equal(t, "call_123", parts[0].ToolCallID)
assert.Equal(t, "read_file", parts[0].ToolName)
assert.JSONEq(t, `{"output":"package main"}`, string(parts[0].Result))
},
},
{
name: "user/sdk_text",
role: codersdk.ChatMessageRoleUser,
raw: nullRaw(mustJSON(t, []codersdk.ChatMessagePart{
{Type: codersdk.ChatMessagePartTypeText, Text: "hello sdk"},
})),
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeText, parts[0].Type)
assert.Equal(t, "hello sdk", parts[0].Text)
},
},
{
name: "user/sdk_file_reference",
role: codersdk.ChatMessageRoleUser,
raw: nullRaw(mustJSON(t, []codersdk.ChatMessagePart{
{Type: codersdk.ChatMessagePartTypeFileReference, FileName: "main.go", StartLine: 1, EndLine: 10, Content: "func main() {}"},
})),
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeFileReference, parts[0].Type)
assert.Equal(t, "main.go", parts[0].FileName)
assert.Equal(t, 1, parts[0].StartLine)
assert.Equal(t, 10, parts[0].EndLine)
assert.Equal(t, "func main() {}", parts[0].Content)
},
},
{
name: "user/sdk_file",
role: codersdk.ChatMessageRoleUser,
raw: nullRaw(mustJSON(t, []codersdk.ChatMessagePart{
{Type: codersdk.ChatMessagePartTypeFile, FileID: uuid.NullUUID{UUID: fileID, Valid: true}, MediaType: "image/png"},
})),
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeFile, parts[0].Type)
assert.True(t, parts[0].FileID.Valid)
assert.Equal(t, fileID, parts[0].FileID.UUID)
assert.Equal(t, "image/png", parts[0].MediaType)
},
},
{
name: "assistant/sdk_text_with_metadata",
role: codersdk.ChatMessageRoleAssistant,
raw: newAssistantWithMeta,
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeText, parts[0].Type)
assert.Equal(t, "here is my answer", parts[0].Text)
assert.JSONEq(t, string(sdkMetadata), string(parts[0].ProviderMetadata))
},
},
{
name: "assistant/sdk_tool_call",
role: codersdk.ChatMessageRoleAssistant,
raw: newAssistantToolCall,
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeToolCall, parts[0].Type)
assert.Equal(t, "call_456", parts[0].ToolCallID)
assert.Equal(t, "execute", parts[0].ToolName)
assert.JSONEq(t, `{"cmd":"ls"}`, string(parts[0].Args))
},
},
{
name: "tool/sdk_tool_result",
role: codersdk.ChatMessageRoleTool,
raw: newToolResult,
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeToolResult, parts[0].Type)
assert.Equal(t, "call_456", parts[0].ToolCallID)
assert.Equal(t, "execute", parts[0].ToolName)
assert.JSONEq(t, `{"output":"file1.go"}`, string(parts[0].Result))
},
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
parts, err := chatprompt.ParseContent(tc.role, tc.raw)
require.NoError(t, err)
tc.check(t, parts)
})
}
}
// TestProviderMetadataRoundTrip verifies that Anthropic cache
// control hints survive the full path: legacy fantasy DB row →
// ParseContent → SDK part (ProviderMetadata) → partsToMessageParts
// → fantasy.MessagePart (ProviderOptions).
func TestProviderMetadataRoundTrip(t *testing.T) {
t.Parallel()
legacyContent, err := chatprompt.MarshalContent([]fantasy.Content{
fantasy.TextContent{
Text: "cached response",
ProviderMetadata: fantasy.ProviderMetadata{
"anthropic": &fantasyanthropic.ProviderCacheControlOptions{
CacheControl: fantasyanthropic.CacheControl{Type: "ephemeral"},
},
},
},
}, nil)
require.NoError(t, err)
// Step 1: ParseContent preserves metadata on the SDK part.
parts, err := chatprompt.ParseContent(codersdk.ChatMessageRoleAssistant, legacyContent)
require.NoError(t, err)
require.Len(t, parts, 1)
require.NotNil(t, parts[0].ProviderMetadata,
"ProviderMetadata must survive ParseContent")
// Step 2: ConvertMessagesWithFiles reconstructs typed
// ProviderOptions on the fantasy part.
prompt, err := chatprompt.ConvertMessagesWithFiles(
context.Background(),
[]database.ChatMessage{{
Role: "assistant",
Visibility: database.ChatMessageVisibilityBoth,
Content: legacyContent,
}},
nil,
slogtest.Make(t, nil),
)
require.NoError(t, err)
require.Len(t, prompt, 1)
require.Len(t, prompt[0].Content, 1)
textPart, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[0].Content[0])
require.True(t, ok, "expected TextPart")
require.Equal(t, "cached response", textPart.Text)
cc := fantasyanthropic.GetCacheControl(textPart.ProviderOptions)
require.NotNil(t, cc, "Anthropic cache control must survive round-trip")
require.Equal(t, "ephemeral", cc.Type)
}
// TestFileReferencePreservation verifies file-reference parts
// survive the storage round-trip and convert to text for LLMs.
func TestFileReferencePreservation(t *testing.T) {
t.Parallel()
raw, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{{
Type: codersdk.ChatMessagePartTypeFileReference,
FileName: "main.go",
StartLine: 10,
EndLine: 20,
Content: "func main() {}",
}})
require.NoError(t, err)
// Storage round-trip: all fields intact.
parts, err := chatprompt.ParseContent(codersdk.ChatMessageRoleUser, raw)
require.NoError(t, err)
require.Len(t, parts, 1)
assert.Equal(t, codersdk.ChatMessagePartTypeFileReference, parts[0].Type)
assert.Equal(t, "main.go", parts[0].FileName)
assert.Equal(t, 10, parts[0].StartLine)
assert.Equal(t, 20, parts[0].EndLine)
assert.Equal(t, "func main() {}", parts[0].Content)
// LLM dispatch: file-reference becomes a TextPart.
prompt, err := chatprompt.ConvertMessagesWithFiles(
context.Background(),
[]database.ChatMessage{{
Role: "user",
Visibility: database.ChatMessageVisibilityBoth,
Content: raw,
}},
nil,
slogtest.Make(t, nil),
)
require.NoError(t, err)
require.Len(t, prompt, 1)
require.Len(t, prompt[0].Content, 1)
textPart, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[0].Content[0])
require.True(t, ok, "file-reference should become TextPart for LLM")
assert.Contains(t, textPart.Text, "[file-reference]")
assert.Contains(t, textPart.Text, "main.go")
assert.Contains(t, textPart.Text, "10-20")
assert.Contains(t, textPart.Text, "func main() {}")
}
// TestAssistantWriteRoundTrip verifies the Stage 4 write path:
// fantasy.Content (with ProviderMetadata) → PartFromContent →
// MarshalParts → DB → ParseContent (SDK path) →
// ConvertMessagesWithFiles → fantasy part with ProviderOptions.
func TestAssistantWriteRoundTrip(t *testing.T) {
t.Parallel()
original := fantasy.TextContent{
Text: "response with cache hints",
ProviderMetadata: fantasy.ProviderMetadata{
"anthropic": &fantasyanthropic.ProviderCacheControlOptions{
CacheControl: fantasyanthropic.CacheControl{Type: "ephemeral"},
},
},
}
// Simulate persistStep: PartFromContent → MarshalParts.
sdkPart := chatprompt.PartFromContent(original)
require.Equal(t, codersdk.ChatMessagePartTypeText, sdkPart.Type)
require.NotNil(t, sdkPart.ProviderMetadata)
raw, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{sdkPart})
require.NoError(t, err)
// Read back via ParseContent (takes the new SDK path, not
// the legacy fallback, because the stored format is flat).
parts, err := chatprompt.ParseContent(codersdk.ChatMessageRoleAssistant, raw)
require.NoError(t, err)
require.Len(t, parts, 1)
assert.Equal(t, "response with cache hints", parts[0].Text)
assert.JSONEq(t, string(sdkPart.ProviderMetadata), string(parts[0].ProviderMetadata))
// Full LLM dispatch: metadata reconstructed as typed options.
prompt, err := chatprompt.ConvertMessagesWithFiles(
context.Background(),
[]database.ChatMessage{{
Role: "assistant",
Visibility: database.ChatMessageVisibilityBoth,
Content: raw,
}},
nil,
slogtest.Make(t, nil),
)
require.NoError(t, err)
require.Len(t, prompt, 1)
require.Len(t, prompt[0].Content, 1)
textPart, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[0].Content[0])
require.True(t, ok)
require.Equal(t, "response with cache hints", textPart.Text)
cc := fantasyanthropic.GetCacheControl(textPart.ProviderOptions)
require.NotNil(t, cc, "cache control must survive new write → new read round-trip")
require.Equal(t, "ephemeral", cc.Type)
}
// TestMixedFormatConversation verifies ConvertMessagesWithFiles
// handles a realistic post-deploy conversation where legacy and new
// storage formats coexist.
func TestMixedFormatConversation(t *testing.T) {
t.Parallel()
fileID := uuid.New()
resolvedFileData := []byte("resolved-png-bytes")
resolver := func(_ context.Context, ids []uuid.UUID) (map[uuid.UUID]chatprompt.FileData, error) {
out := make(map[uuid.UUID]chatprompt.FileData)
for _, id := range ids {
if id == fileID {
out[id] = chatprompt.FileData{Data: resolvedFileData, MediaType: "image/png"}
}
}
return out, nil
}
// 1. System (JSON string).
systemRaw, err := json.Marshal("You are helpful.")
require.NoError(t, err)
// 2. Old user (fantasy envelope: text + file with file_id).
oldUserRaw := mustJSON(t, []json.RawMessage{
mustJSON(t, map[string]any{
"type": "text",
"data": map[string]any{"text": "Look at this image."},
}),
mustJSON(t, map[string]any{
"type": "file",
"data": map[string]any{
"media_type": "image/png",
"file_id": fileID.String(),
},
}),
})
// 3. Old assistant (fantasy envelope: tool-call).
oldAssistantRaw, err := chatprompt.MarshalContent([]fantasy.Content{
fantasy.ToolCallContent{
ToolCallID: "call_1",
ToolName: "analyze_image",
Input: `{"detail":"high"}`,
},
}, nil)
require.NoError(t, err)
// 4. Old tool (legacy result rows).
oldToolRaw, err := chatprompt.MarshalToolResult(
"call_1", "analyze_image",
json.RawMessage(`{"description":"a cat"}`), false,
false, nil,
)
require.NoError(t, err)
// 5. New user (SDK parts: text + file-reference).
newUserRaw, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
{Type: codersdk.ChatMessagePartTypeText, Text: "Check this diff."},
{Type: codersdk.ChatMessagePartTypeFileReference, FileName: "main.go", StartLine: 5, EndLine: 15, Content: "func main() {}"},
})
require.NoError(t, err)
// 6. New assistant (SDK parts: text with metadata).
newAssistantMeta := json.RawMessage(`{"anthropic":{"type":"anthropic.cache_control_options","data":{"cache_control":{"type":"ephemeral"}}}}`)
newAssistantRaw, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
{Type: codersdk.ChatMessagePartTypeText, Text: "Here is my analysis.", ProviderMetadata: newAssistantMeta},
})
require.NoError(t, err)
messages := []database.ChatMessage{
{Role: "system", Visibility: database.ChatMessageVisibilityModel, Content: pqtype.NullRawMessage{RawMessage: systemRaw, Valid: true}},
{Role: "user", Visibility: database.ChatMessageVisibilityBoth, Content: pqtype.NullRawMessage{RawMessage: oldUserRaw, Valid: true}},
{Role: "assistant", Visibility: database.ChatMessageVisibilityBoth, Content: oldAssistantRaw},
{Role: "tool", Visibility: database.ChatMessageVisibilityBoth, Content: oldToolRaw},
{Role: "user", Visibility: database.ChatMessageVisibilityBoth, Content: newUserRaw},
{Role: "assistant", Visibility: database.ChatMessageVisibilityBoth, Content: newAssistantRaw},
}
prompt, err := chatprompt.ConvertMessagesWithFiles(
context.Background(), messages, resolver, slogtest.Make(t, nil),
)
require.NoError(t, err)
require.Len(t, prompt, 6, "all 6 messages should produce prompt entries")
// 1. System.
require.Equal(t, fantasy.MessageRoleSystem, prompt[0].Role)
systemText, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[0].Content[0])
require.True(t, ok)
assert.Equal(t, "You are helpful.", systemText.Text)
// 2. Old user: text + file with resolved data.
require.Equal(t, fantasy.MessageRoleUser, prompt[1].Role)
require.Len(t, prompt[1].Content, 2)
userText, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[1].Content[0])
require.True(t, ok)
assert.Equal(t, "Look at this image.", userText.Text)
filePart, ok := fantasy.AsMessagePart[fantasy.FilePart](prompt[1].Content[1])
require.True(t, ok)
assert.Equal(t, resolvedFileData, filePart.Data)
assert.Equal(t, "image/png", filePart.MediaType)
// 3. Old assistant: tool-call with normalized input.
require.Equal(t, fantasy.MessageRoleAssistant, prompt[2].Role)
toolCalls := chatprompt.ExtractToolCalls(prompt[2].Content)
require.Len(t, toolCalls, 1)
assert.Equal(t, "call_1", toolCalls[0].ToolCallID)
assert.Equal(t, "analyze_image", toolCalls[0].ToolName)
assert.JSONEq(t, `{"detail":"high"}`, toolCalls[0].Input)
// 4. Old tool: result paired with call_1.
require.Equal(t, fantasy.MessageRoleTool, prompt[3].Role)
require.Len(t, prompt[3].Content, 1)
toolResult, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[3].Content[0])
require.True(t, ok)
assert.Equal(t, "call_1", toolResult.ToolCallID)
// 5. New user: text + file-reference (converted to TextPart).
require.Equal(t, fantasy.MessageRoleUser, prompt[4].Role)
require.Len(t, prompt[4].Content, 2)
newUserText, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[4].Content[0])
require.True(t, ok)
assert.Equal(t, "Check this diff.", newUserText.Text)
refText, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[4].Content[1])
require.True(t, ok)
assert.Contains(t, refText.Text, "[file-reference]")
assert.Contains(t, refText.Text, "main.go")
// 6. New assistant: text with ProviderMetadata → ProviderOptions.
require.Equal(t, fantasy.MessageRoleAssistant, prompt[5].Role)
require.Len(t, prompt[5].Content, 1)
newAssistantText, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[5].Content[0])
require.True(t, ok)
assert.Equal(t, "Here is my analysis.", newAssistantText.Text)
cc := fantasyanthropic.GetCacheControl(newAssistantText.ProviderOptions)
require.NotNil(t, cc, "ProviderMetadata must survive on new-format assistant messages")
assert.Equal(t, "ephemeral", cc.Type)
}
// TestQueuedMessageRoundTrip verifies that a user message with
// file-reference parts survives the queue → promote cycle. The
// queued path stores MarshalParts output as raw JSON in
// chat_queued_messages, db2sdk.ChatQueuedMessage parses it for
// display while queued, then PromoteQueued copies the same raw
// bytes into chat_messages where ParseContent reads them.
func TestQueuedMessageRoundTrip(t *testing.T) {
t.Parallel()
// Simulate the write path: user sends a message with text +
// file-reference, which gets queued.
parts := []codersdk.ChatMessagePart{
{Type: codersdk.ChatMessagePartTypeText, Text: "Review this change."},
{Type: codersdk.ChatMessagePartTypeFileReference, FileName: "api.go", StartLine: 42, EndLine: 58, Content: "func handleRequest() {}"},
}
raw, err := chatprompt.MarshalParts(parts)
require.NoError(t, err)
// Step 1: While queued, db2sdk.ChatQueuedMessage parses the
// content for display. Verify it produces correct parts
// (with internal fields stripped).
queuedMsg := db2sdk.ChatQueuedMessage(database.ChatQueuedMessage{
ID: 1,
ChatID: uuid.New(),
Content: raw.RawMessage,
})
require.Len(t, queuedMsg.Content, 2)
assert.Equal(t, codersdk.ChatMessagePartTypeText, queuedMsg.Content[0].Type)
assert.Equal(t, "Review this change.", queuedMsg.Content[0].Text)
assert.Equal(t, codersdk.ChatMessagePartTypeFileReference, queuedMsg.Content[1].Type)
assert.Equal(t, "api.go", queuedMsg.Content[1].FileName)
assert.Equal(t, 42, queuedMsg.Content[1].StartLine)
assert.Equal(t, 58, queuedMsg.Content[1].EndLine)
assert.Equal(t, "func handleRequest() {}", queuedMsg.Content[1].Content)
// Step 2: PromoteQueued copies the raw bytes into
// chat_messages. ParseContent must handle them identically.
promoted, err := chatprompt.ParseContent(codersdk.ChatMessageRoleUser, pqtype.NullRawMessage{
RawMessage: raw.RawMessage,
Valid: true,
})
require.NoError(t, err)
require.Len(t, promoted, 2)
assert.Equal(t, codersdk.ChatMessagePartTypeText, promoted[0].Type)
assert.Equal(t, "Review this change.", promoted[0].Text)
assert.Equal(t, codersdk.ChatMessagePartTypeFileReference, promoted[1].Type)
assert.Equal(t, "api.go", promoted[1].FileName)
assert.Equal(t, 42, promoted[1].StartLine)
assert.Equal(t, 58, promoted[1].EndLine)
assert.Equal(t, "func handleRequest() {}", promoted[1].Content)
// Step 3: The promoted message is used for LLM dispatch.
// File-reference becomes a TextPart.
prompt, err := chatprompt.ConvertMessagesWithFiles(
context.Background(),
[]database.ChatMessage{{
Role: "user",
Visibility: database.ChatMessageVisibilityBoth,
Content: pqtype.NullRawMessage{RawMessage: raw.RawMessage, Valid: true},
}},
nil,
slogtest.Make(t, nil),
)
require.NoError(t, err)
require.Len(t, prompt, 1)
require.Len(t, prompt[0].Content, 2)
textPart, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[0].Content[0])
require.True(t, ok)
assert.Equal(t, "Review this change.", textPart.Text)
refPart, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[0].Content[1])
require.True(t, ok)
assert.Contains(t, refPart.Text, "[file-reference]")
assert.Contains(t, refPart.Text, "api.go")
}
func TestParseContent_ErrorPaths(t *testing.T) {
t.Parallel()
t.Run("null_content_returns_nil", func(t *testing.T) {
t.Parallel()
parts, err := chatprompt.ParseContent(codersdk.ChatMessageRoleUser, pqtype.NullRawMessage{})
require.NoError(t, err)
assert.Nil(t, parts)
})
t.Run("empty_content_returns_nil", func(t *testing.T) {
t.Parallel()
parts, err := chatprompt.ParseContent(codersdk.ChatMessageRoleAssistant, pqtype.NullRawMessage{
RawMessage: []byte{},
Valid: true,
})
require.NoError(t, err)
assert.Nil(t, parts)
})
t.Run("unknown_role", func(t *testing.T) {
t.Parallel()
_, err := chatprompt.ParseContent(codersdk.ChatMessageRole("banana"), nullRaw(json.RawMessage(`"hello"`)))
require.Error(t, err)
assert.Contains(t, err.Error(), "unsupported chat message role")
})
t.Run("system/malformed_json", func(t *testing.T) {
t.Parallel()
_, err := chatprompt.ParseContent(codersdk.ChatMessageRoleSystem, nullRaw(json.RawMessage(`not json`)))
require.Error(t, err)
assert.Contains(t, err.Error(), "parse system content")
})
t.Run("user/malformed_json", func(t *testing.T) {
t.Parallel()
_, err := chatprompt.ParseContent(codersdk.ChatMessageRoleUser, nullRaw(json.RawMessage(`{not json`)))
require.Error(t, err)
})
t.Run("assistant/malformed_json", func(t *testing.T) {
t.Parallel()
_, err := chatprompt.ParseContent(codersdk.ChatMessageRoleAssistant, nullRaw(json.RawMessage(`{not json`)))
require.Error(t, err)
})
t.Run("tool/malformed_json", func(t *testing.T) {
t.Parallel()
_, err := chatprompt.ParseContent(codersdk.ChatMessageRoleTool, nullRaw(json.RawMessage(`{not json`)))
require.Error(t, err)
})
}
func mustJSON(t *testing.T, v any) json.RawMessage {
t.Helper()
data, err := json.Marshal(v)
+20 -14
View File
@@ -21,6 +21,7 @@ import (
"github.com/coder/coder/v2/coderd/chatd/chatretry"
"github.com/coder/coder/v2/coderd/database"
coderdpubsub "github.com/coder/coder/v2/coderd/pubsub"
"github.com/coder/coder/v2/codersdk"
)
const titleGenerationPrompt = "You are a title generator. Your ONLY job is to output a short title (2-8 words) " +
@@ -158,13 +159,13 @@ func titleInput(
}
switch message.Role {
case string(fantasy.MessageRoleAssistant), string(fantasy.MessageRoleTool):
case string(codersdk.ChatMessageRoleAssistant), string(codersdk.ChatMessageRoleTool):
return "", false
case string(fantasy.MessageRoleUser):
case string(codersdk.ChatMessageRoleUser):
userCount++
if firstUserText == "" {
parsed, err := chatprompt.ParseContent(
string(fantasy.MessageRoleUser), message.Content,
codersdk.ChatMessageRoleUser, message.Content,
)
if err != nil {
return "", false
@@ -226,22 +227,21 @@ func fallbackChatTitle(message string) string {
return truncateRunes(title, maxRunes)
}
// contentBlocksToText concatenates the text parts of content blocks
// into a single space-separated string.
func contentBlocksToText(content []fantasy.Content) string {
parts := make([]string, 0, len(content))
for _, block := range content {
textBlock, ok := fantasy.AsContentType[fantasy.TextContent](block)
if !ok {
// contentBlocksToText concatenates the text parts of SDK chat
// message parts into a single space-separated string.
func contentBlocksToText(parts []codersdk.ChatMessagePart) string {
texts := make([]string, 0, len(parts))
for _, part := range parts {
if part.Type != codersdk.ChatMessagePartTypeText {
continue
}
text := strings.TrimSpace(textBlock.Text)
text := strings.TrimSpace(part.Text)
if text == "" {
continue
}
parts = append(parts, text)
texts = append(texts, text)
}
return strings.Join(parts, " ")
return strings.Join(texts, " ")
}
func truncateRunes(value string, maxLen int) string {
@@ -343,7 +343,13 @@ func generateShortText(
return "", xerrors.Errorf("generate short text: %w", err)
}
text := strings.TrimSpace(contentBlocksToText(response.Content))
responseParts := make([]codersdk.ChatMessagePart, 0, len(response.Content))
for _, block := range response.Content {
if p := chatprompt.PartFromContent(block); p.Type != "" {
responseParts = append(responseParts, p)
}
}
text := strings.TrimSpace(contentBlocksToText(responseParts))
text = strings.Trim(text, "\"'`")
return text, nil
}
+5 -4
View File
@@ -15,6 +15,7 @@ import (
"github.com/coder/coder/v2/coderd/chatd/chatprompt"
"github.com/coder/coder/v2/coderd/database"
coderdpubsub "github.com/coder/coder/v2/coderd/pubsub"
"github.com/coder/coder/v2/codersdk"
)
var ErrSubagentNotDescendant = xerrors.New("target chat is not a descendant of current chat")
@@ -263,7 +264,7 @@ func (p *Server) createChildSubagentChat(
},
ModelConfigID: parent.LastModelConfigID,
Title: title,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: prompt}},
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText(prompt)},
})
if err != nil {
return database.Chat{}, xerrors.Errorf("create child chat: %w", err)
@@ -301,7 +302,7 @@ func (p *Server) sendSubagentMessage(
sendResult, err := p.SendMessage(ctx, SendMessageOptions{
ChatID: targetChatID,
CreatedBy: targetChat.OwnerID,
Content: []fantasy.Content{fantasy.TextContent{Text: message}},
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText(message)},
BusyBehavior: busyBehavior,
})
if err != nil {
@@ -481,12 +482,12 @@ func latestSubagentAssistantMessage(
for i := len(messages) - 1; i >= 0; i-- {
message := messages[i]
if message.Role != string(fantasy.MessageRoleAssistant) ||
if message.Role != string(codersdk.ChatMessageRoleAssistant) ||
message.Visibility == database.ChatMessageVisibilityModel {
continue
}
content, parseErr := chatprompt.ParseContent(message.Role, message.Content)
content, parseErr := chatprompt.ParseContent(codersdk.ChatMessageRole(message.Role), message.Content)
if parseErr != nil {
continue
}