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
}
+25 -34
View File
@@ -18,7 +18,6 @@ import (
"sync"
"time"
"charm.land/fantasy"
"github.com/go-chi/chi/v5"
"github.com/google/uuid"
"golang.org/x/xerrors"
@@ -222,7 +221,7 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
return
}
contentBlocks, contentFileIDs, titleSource, inputError := createChatInputFromRequest(ctx, api.Database, req)
contentBlocks, titleSource, inputError := createChatInputFromRequest(ctx, api.Database, req)
if inputError != nil {
httpapi.Write(ctx, rw, http.StatusBadRequest, *inputError)
return
@@ -257,7 +256,6 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
ModelConfigID: modelConfigID,
SystemPrompt: api.resolvedChatSystemPrompt(ctx),
InitialUserContent: contentBlocks,
ContentFileIDs: contentFileIDs,
})
if err != nil {
if database.IsForeignKeyViolation(
@@ -630,7 +628,7 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
return
}
contentBlocks, contentFileIDs, _, inputError := createChatInputFromParts(ctx, api.Database, req.Content, "content")
contentBlocks, _, inputError := createChatInputFromParts(ctx, api.Database, req.Content, "content")
if inputError != nil {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: inputError.Message,
@@ -642,12 +640,11 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
sendResult, sendErr := api.chatDaemon.SendMessage(
ctx,
chatd.SendMessageOptions{
ChatID: chatID,
CreatedBy: apiKey.UserID,
Content: contentBlocks,
ContentFileIDs: contentFileIDs,
ModelConfigID: req.ModelConfigID,
BusyBehavior: chatd.SendMessageBusyBehaviorQueue,
ChatID: chatID,
CreatedBy: apiKey.UserID,
Content: contentBlocks,
ModelConfigID: req.ModelConfigID,
BusyBehavior: chatd.SendMessageBusyBehaviorQueue,
},
)
if sendErr != nil {
@@ -707,7 +704,7 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) {
return
}
contentBlocks, contentFileIDs, _, inputError := createChatInputFromParts(ctx, api.Database, req.Content, "content")
contentBlocks, _, inputError := createChatInputFromParts(ctx, api.Database, req.Content, "content")
if inputError != nil {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: inputError.Message,
@@ -721,7 +718,6 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) {
CreatedBy: apiKey.UserID,
EditedMessageID: messageID,
Content: contentBlocks,
ContentFileIDs: contentFileIDs,
})
if editErr != nil {
switch {
@@ -1985,8 +1981,7 @@ func (api *API) chatFileByID(rw http.ResponseWriter, r *http.Request) {
}
func createChatInputFromRequest(ctx context.Context, db database.Store, req codersdk.CreateChatRequest) (
[]fantasy.Content,
map[int]uuid.UUID,
[]codersdk.ChatMessagePart,
string,
*codersdk.Response,
) {
@@ -1998,32 +1993,31 @@ func createChatInputFromParts(
db database.Store,
parts []codersdk.ChatInputPart,
fieldName string,
) ([]fantasy.Content, map[int]uuid.UUID, string, *codersdk.Response) {
) ([]codersdk.ChatMessagePart, string, *codersdk.Response) {
if len(parts) == 0 {
return nil, nil, "", &codersdk.Response{
return nil, "", &codersdk.Response{
Message: "Content is required.",
Detail: "Content cannot be empty.",
}
}
content := make([]fantasy.Content, 0, len(parts))
fileIDs := make(map[int]uuid.UUID)
content := make([]codersdk.ChatMessagePart, 0, len(parts))
textParts := make([]string, 0, len(parts))
for i, part := range parts {
switch strings.ToLower(strings.TrimSpace(string(part.Type))) {
case string(codersdk.ChatInputPartTypeText):
text := strings.TrimSpace(part.Text)
if text == "" {
return nil, nil, "", &codersdk.Response{
return nil, "", &codersdk.Response{
Message: "Invalid input part.",
Detail: fmt.Sprintf("%s[%d].text cannot be empty.", fieldName, i),
}
}
content = append(content, fantasy.TextContent{Text: text})
content = append(content, codersdk.ChatMessageText(text))
textParts = append(textParts, text)
case string(codersdk.ChatInputPartTypeFile):
if part.FileID == uuid.Nil {
return nil, nil, "", &codersdk.Response{
return nil, "", &codersdk.Response{
Message: "Invalid input part.",
Detail: fmt.Sprintf("%s[%d].file_id is required for file parts.", fieldName, i),
}
@@ -2034,27 +2028,26 @@ func createChatInputFromParts(
chatFile, err := db.GetChatFileByID(ctx, part.FileID)
if err != nil {
if httpapi.Is404Error(err) {
return nil, nil, "", &codersdk.Response{
return nil, "", &codersdk.Response{
Message: "Invalid input part.",
Detail: fmt.Sprintf("%s[%d].file_id references a file that does not exist.", fieldName, i),
}
}
return nil, nil, "", &codersdk.Response{
return nil, "", &codersdk.Response{
Message: "Internal error.",
Detail: fmt.Sprintf("Failed to retrieve file for %s[%d].", fieldName, i),
}
}
content = append(content, fantasy.FileContent{
MediaType: chatFile.Mimetype,
})
fileIDs[len(content)-1] = part.FileID
content = append(content, codersdk.ChatMessageFile(part.FileID, chatFile.Mimetype))
case string(codersdk.ChatInputPartTypeFileReference):
if part.FileName == "" {
return nil, nil, "", &codersdk.Response{
return nil, "", &codersdk.Response{
Message: "Invalid input part.",
Detail: fmt.Sprintf("%s[%d].file_name cannot be empty for file-reference.", fieldName, i),
}
}
content = append(content, codersdk.ChatMessageFileReference(part.FileName, part.StartLine, part.EndLine, part.Content))
// Build text representation for title generation.
lineRange := fmt.Sprintf("%d", part.StartLine)
if part.StartLine != part.EndLine {
lineRange = fmt.Sprintf("%d-%d", part.StartLine, part.EndLine)
@@ -2064,11 +2057,9 @@ func createChatInputFromParts(
if strings.TrimSpace(part.Content) != "" {
_, _ = fmt.Fprintf(&sb, "\n```%s\n%s\n```", part.FileName, strings.TrimSpace(part.Content))
}
text := sb.String()
content = append(content, fantasy.TextContent{Text: text})
textParts = append(textParts, text)
textParts = append(textParts, sb.String())
default:
return nil, nil, "", &codersdk.Response{
return nil, "", &codersdk.Response{
Message: "Invalid input part.",
Detail: fmt.Sprintf(
"%s[%d].type %q is not supported.",
@@ -2083,13 +2074,13 @@ func createChatInputFromParts(
// Allow file-only messages. The titleSource may be empty
// when only file parts are provided, callers handle this.
if len(content) == 0 {
return nil, nil, "", &codersdk.Response{
return nil, "", &codersdk.Response{
Message: "Content is required.",
Detail: fmt.Sprintf("%s must include at least one text or file part.", fieldName),
}
}
titleSource := strings.TrimSpace(strings.Join(textParts, " "))
return content, fileIDs, titleSource, nil
return content, titleSource, nil
}
func chatTitleFromMessage(message string) string {
+92 -58
View File
@@ -93,7 +93,7 @@ func TestPostChats(t *testing.T) {
foundUserMessage := false
for _, message := range messagesResult.Messages {
if message.Role != "user" {
if message.Role != codersdk.ChatMessageRoleUser {
continue
}
for _, part := range message.Content {
@@ -128,7 +128,7 @@ func TestPostChats(t *testing.T) {
messagesResult, err := client.GetChatMessages(ctx, chat.ID)
require.NoError(t, err)
for _, message := range messagesResult.Messages {
require.NotEqual(t, "system", message.Role)
require.NotEqual(t, codersdk.ChatMessageRoleSystem, message.Role)
}
})
@@ -1340,9 +1340,9 @@ func TestGetChat(t *testing.T) {
foundUserMessage := false
for _, message := range messagesResult.Messages {
require.Equal(t, createdChat.ID, message.ChatID)
require.NotEqual(t, "system", message.Role)
require.NotEqual(t, codersdk.ChatMessageRoleSystem, message.Role)
for _, part := range message.Content {
if message.Role == "user" &&
if message.Role == codersdk.ChatMessageRoleUser &&
part.Type == codersdk.ChatMessagePartTypeText &&
part.Text == "get chat route payload" {
foundUserMessage = true
@@ -1663,7 +1663,7 @@ func TestPostChatMessages(t *testing.T) {
}
}
for _, message := range messagesResult.Messages {
if message.Role == "user" && hasTextPart(message.Content, messageText) {
if message.Role == codersdk.ChatMessageRoleUser && hasTextPart(message.Content, messageText) {
return true
}
}
@@ -1673,7 +1673,7 @@ func TestPostChatMessages(t *testing.T) {
require.Nil(t, created.QueuedMessage)
require.NotNil(t, created.Message)
require.Equal(t, chat.ID, created.Message.ChatID)
require.Equal(t, "user", created.Message.Role)
require.Equal(t, codersdk.ChatMessageRoleUser, created.Message.Role)
require.NotZero(t, created.Message.ID)
require.True(t, hasTextPart(created.Message.Content, messageText))
@@ -1684,7 +1684,7 @@ func TestPostChatMessages(t *testing.T) {
}
for _, message := range messagesResult.Messages {
if message.ID == created.Message.ID &&
message.Role == "user" &&
message.Role == codersdk.ChatMessageRoleUser &&
hasTextPart(message.Content, messageText) {
return true
}
@@ -1782,9 +1782,14 @@ func TestChatMessageWithFileReferences(t *testing.T) {
})
require.NoError(t, err)
// The file-reference is stored as a formatted text block.
wantText := "[file-reference] main.go:10-15\n" +
"```main.go\nfunc broken() {}\n```"
// File-reference parts are stored as structured parts.
checkFileRef := func(part codersdk.ChatMessagePart) bool {
return part.Type == codersdk.ChatMessagePartTypeFileReference &&
part.FileName == "main.go" &&
part.StartLine == 10 &&
part.EndLine == 15 &&
part.Content == "func broken() {}"
}
var found bool
require.Eventually(t, func() bool {
@@ -1793,12 +1798,11 @@ func TestChatMessageWithFileReferences(t *testing.T) {
return false
}
for _, message := range messagesResult.Messages {
if message.Role != "user" {
if message.Role != codersdk.ChatMessageRoleUser {
continue
}
for _, part := range message.Content {
if part.Type == codersdk.ChatMessagePartTypeText &&
part.Text == wantText {
if checkFileRef(part) {
found = true
return true
}
@@ -1808,8 +1812,7 @@ func TestChatMessageWithFileReferences(t *testing.T) {
if created.Queued && created.QueuedMessage != nil {
for _, queued := range messagesResult.QueuedMessages {
for _, part := range queued.Content {
if part.Type == codersdk.ChatMessagePartTypeText &&
part.Text == wantText {
if checkFileRef(part) {
found = true
return true
}
@@ -1818,7 +1821,7 @@ func TestChatMessageWithFileReferences(t *testing.T) {
}
return false
}, testutil.WaitLong, testutil.IntervalFast)
require.True(t, found, "expected to find file-reference text in stored message")
require.True(t, found, "expected to find file-reference part in stored message")
})
t.Run("FileReferenceSingleLine", func(t *testing.T) {
@@ -1841,9 +1844,13 @@ func TestChatMessageWithFileReferences(t *testing.T) {
})
require.NoError(t, err)
// Single-line range should use "42" not "42-42".
wantText := "[file-reference] lib/utils.ts:42\n" +
"```lib/utils.ts\nconst x = 1;\n```"
checkFileRef := func(part codersdk.ChatMessagePart) bool {
return part.Type == codersdk.ChatMessagePartTypeFileReference &&
part.FileName == "lib/utils.ts" &&
part.StartLine == 42 &&
part.EndLine == 42 &&
part.Content == "const x = 1;"
}
require.Eventually(t, func() bool {
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID)
@@ -1852,7 +1859,7 @@ func TestChatMessageWithFileReferences(t *testing.T) {
}
for _, msg := range messagesResult.Messages {
for _, part := range msg.Content {
if part.Type == codersdk.ChatMessagePartTypeText && part.Text == wantText {
if checkFileRef(part) {
return true
}
}
@@ -1860,7 +1867,7 @@ func TestChatMessageWithFileReferences(t *testing.T) {
if created.Queued && created.QueuedMessage != nil {
for _, queued := range messagesResult.QueuedMessages {
for _, part := range queued.Content {
if part.Type == codersdk.ChatMessagePartTypeText && part.Text == wantText {
if checkFileRef(part) {
return true
}
}
@@ -1890,8 +1897,14 @@ func TestChatMessageWithFileReferences(t *testing.T) {
})
require.NoError(t, err)
// No fenced code block when content is empty.
wantText := "[file-reference] README.md:1"
checkFileRef := func(part codersdk.ChatMessagePart) bool {
return part.Type == codersdk.ChatMessagePartTypeFileReference &&
part.FileName == "README.md" &&
part.StartLine == 1 &&
part.EndLine == 1 &&
part.Content == ""
}
require.Eventually(t, func() bool {
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID)
if getErr != nil {
@@ -1899,7 +1912,7 @@ func TestChatMessageWithFileReferences(t *testing.T) {
}
for _, msg := range messagesResult.Messages {
for _, part := range msg.Content {
if part.Type == codersdk.ChatMessagePartTypeText && part.Text == wantText {
if checkFileRef(part) {
return true
}
}
@@ -1907,7 +1920,7 @@ func TestChatMessageWithFileReferences(t *testing.T) {
if created.Queued && created.QueuedMessage != nil {
for _, queued := range messagesResult.QueuedMessages {
for _, part := range queued.Content {
if part.Type == codersdk.ChatMessagePartTypeText && part.Text == wantText {
if checkFileRef(part) {
return true
}
}
@@ -1937,8 +1950,13 @@ func TestChatMessageWithFileReferences(t *testing.T) {
})
require.NoError(t, err)
wantText := "[file-reference] server.go:5-8\n" +
"```server.go\nfunc main() {\n\tfmt.Println()\n}\n```"
checkFileRef := func(part codersdk.ChatMessagePart) bool {
return part.Type == codersdk.ChatMessagePartTypeFileReference &&
part.FileName == "server.go" &&
part.StartLine == 5 &&
part.EndLine == 8 &&
part.Content == "func main() {\n\tfmt.Println()\n}"
}
require.Eventually(t, func() bool {
messagesResult, getErr := client.GetChatMessages(ctx, chat.ID)
@@ -1947,7 +1965,7 @@ func TestChatMessageWithFileReferences(t *testing.T) {
}
for _, msg := range messagesResult.Messages {
for _, part := range msg.Content {
if part.Type == codersdk.ChatMessagePartTypeText && part.Text == wantText {
if checkFileRef(part) {
return true
}
}
@@ -1955,7 +1973,7 @@ func TestChatMessageWithFileReferences(t *testing.T) {
if created.Queued && created.QueuedMessage != nil {
for _, queued := range messagesResult.QueuedMessages {
for _, part := range queued.Content {
if part.Type == codersdk.ChatMessagePartTypeText && part.Text == wantText {
if checkFileRef(part) {
return true
}
}
@@ -2010,14 +2028,24 @@ func TestChatMessageWithFileReferences(t *testing.T) {
})
require.NoError(t, err)
// Verify that all six parts are stored in order.
wantTexts := []string{
"Please review these two issues:",
"[file-reference] a.go:1-3\n```a.go\nline1\nline2\nline3\n```",
"first issue",
"and also:",
"[file-reference] b.go:10\n```b.go\nreturn nil\n```",
"second issue",
// Verify that all six parts are stored in order with
// correct types: text, file-reference, text, text,
// file-reference, text.
type wantPart struct {
typ codersdk.ChatMessagePartType
text string
fileName string
startLine int
endLine int
content string
}
want := []wantPart{
{typ: codersdk.ChatMessagePartTypeText, text: "Please review these two issues:"},
{typ: codersdk.ChatMessagePartTypeFileReference, fileName: "a.go", startLine: 1, endLine: 3, content: "line1\nline2\nline3"},
{typ: codersdk.ChatMessagePartTypeText, text: "first issue"},
{typ: codersdk.ChatMessagePartTypeText, text: "and also:"},
{typ: codersdk.ChatMessagePartTypeFileReference, fileName: "b.go", startLine: 10, endLine: 10, content: "return nil"},
{typ: codersdk.ChatMessagePartTypeText, text: "second issue"},
}
require.Eventually(t, func() bool {
@@ -2026,28 +2054,34 @@ func TestChatMessageWithFileReferences(t *testing.T) {
return false
}
// Check messages and queued messages for the
// interleaved parts in order.
checkParts := func(parts []codersdk.ChatMessagePart) bool {
textParts := make([]string, 0, len(parts))
for _, part := range parts {
if part.Type == codersdk.ChatMessagePartTypeText {
textParts = append(textParts, part.Text)
}
}
if len(textParts) != len(wantTexts) {
if len(parts) != len(want) {
return false
}
for i, want := range wantTexts {
if textParts[i] != want {
for i, w := range want {
p := parts[i]
if p.Type != w.typ {
return false
}
switch w.typ {
case codersdk.ChatMessagePartTypeText:
if p.Text != w.text {
return false
}
case codersdk.ChatMessagePartTypeFileReference:
if p.FileName != w.fileName ||
p.StartLine != w.startLine ||
p.EndLine != w.endLine ||
p.Content != w.content {
return false
}
}
}
return true
}
for _, msg := range messagesResult.Messages {
if msg.Role == "user" && checkParts(msg.Content) {
if msg.Role == codersdk.ChatMessageRoleUser && checkParts(msg.Content) {
return true
}
}
@@ -2154,7 +2188,7 @@ func TestChatMessageWithFiles(t *testing.T) {
require.NotNil(t, resp.QueuedMessage)
} else {
require.NotNil(t, resp.Message)
require.Equal(t, "user", resp.Message.Role)
require.Equal(t, codersdk.ChatMessageRoleUser, resp.Message.Role)
}
})
@@ -2201,7 +2235,7 @@ func TestChatMessageWithFiles(t *testing.T) {
require.NotNil(t, resp.QueuedMessage)
} else {
require.NotNil(t, resp.Message)
require.Equal(t, "user", resp.Message.Role)
require.Equal(t, codersdk.ChatMessageRoleUser, resp.Message.Role)
}
// Verify file parts omit inline data in the API response.
@@ -2306,7 +2340,7 @@ func TestPatchChatMessage(t *testing.T) {
var userMessageID int64
for _, message := range messagesResult.Messages {
if message.Role == "user" {
if message.Role == codersdk.ChatMessageRoleUser {
userMessageID = message.ID
break
}
@@ -2323,7 +2357,7 @@ func TestPatchChatMessage(t *testing.T) {
})
require.NoError(t, err)
require.Equal(t, userMessageID, edited.ID)
require.Equal(t, "user", edited.Role)
require.Equal(t, codersdk.ChatMessageRoleUser, edited.Role)
foundEditedText := false
for _, part := range edited.Content {
@@ -2338,7 +2372,7 @@ func TestPatchChatMessage(t *testing.T) {
foundEditedInChat := false
foundOriginalInChat := false
for _, message := range messagesResult.Messages {
if message.Role != "user" {
if message.Role != codersdk.ChatMessageRoleUser {
continue
}
for _, part := range message.Content {
@@ -2391,7 +2425,7 @@ func TestPatchChatMessage(t *testing.T) {
var userMessageID int64
for _, message := range messagesResult.Messages {
if message.Role == "user" {
if message.Role == codersdk.ChatMessageRoleUser {
userMessageID = message.ID
break
}
@@ -2434,7 +2468,7 @@ func TestPatchChatMessage(t *testing.T) {
var foundTextInChat, foundFileInChat bool
for _, message := range messagesResult.Messages {
if message.Role != "user" {
if message.Role != codersdk.ChatMessageRoleUser {
continue
}
for _, part := range message.Content {
@@ -2568,7 +2602,7 @@ func TestStreamChat(t *testing.T) {
if event.Type == codersdk.ChatStreamEventTypeMessage &&
event.Message != nil &&
event.Message.Role == "user" &&
event.Message.Role == codersdk.ChatMessageRoleUser &&
hasTextPart(event.Message.Content, initialMessage) {
foundInitialUserMessage = true
}
@@ -3128,7 +3162,7 @@ func TestPromoteChatQueuedMessage(t *testing.T) {
require.NoError(t, err)
require.NotZero(t, promoted.ID)
require.Equal(t, chat.ID, promoted.ChatID)
require.Equal(t, "user", promoted.Role)
require.Equal(t, codersdk.ChatMessageRoleUser, promoted.Role)
foundPromotedText := false
for _, part := range promoted.Content {
+11 -250
View File
@@ -12,7 +12,6 @@ import (
"strings"
"time"
"charm.land/fantasy"
"github.com/google/uuid"
"github.com/hashicorp/hcl/v2"
"github.com/sqlc-dev/pqtype"
@@ -1069,10 +1068,10 @@ func ChatMessage(m database.ChatMessage) codersdk.ChatMessage {
CreatedBy: createdBy,
ModelConfigID: modelConfigID,
CreatedAt: m.CreatedAt,
Role: m.Role,
Role: codersdk.ChatMessageRole(m.Role),
}
if m.Content.Valid {
parts, err := chatMessageParts(m.Role, m.Content)
parts, err := chatMessageParts(codersdk.ChatMessageRole(m.Role), m.Content)
if err == nil {
msg.Content = parts
}
@@ -1114,7 +1113,7 @@ 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(string(fantasy.MessageRoleUser), pqtype.NullRawMessage{
parts, err := chatMessageParts(codersdk.ChatMessageRoleUser, pqtype.NullRawMessage{
RawMessage: message.Content,
Valid: len(message.Content) > 0,
})
@@ -1140,254 +1139,16 @@ func ChatQueuedMessages(messages []database.ChatQueuedMessage) []codersdk.ChatQu
return out
}
func chatMessageParts(role string, raw pqtype.NullRawMessage) ([]codersdk.ChatMessagePart, error) {
switch role {
case string(fantasy.MessageRoleSystem):
content, err := parseSystemContent(raw)
if err != nil {
return nil, err
}
if strings.TrimSpace(content) == "" {
return nil, nil
}
return []codersdk.ChatMessagePart{{
Type: codersdk.ChatMessagePartTypeText,
Text: content,
}}, nil
case string(fantasy.MessageRoleUser), string(fantasy.MessageRoleAssistant):
content, err := parseContentBlocks(role, raw)
if err != nil {
return nil, err
}
var rawBlocks []json.RawMessage
_ = json.Unmarshal(raw.RawMessage, &rawBlocks)
parts := make([]codersdk.ChatMessagePart, 0, len(content))
for i, block := range content {
part := contentBlockToPart(block)
if part.Type == "" {
continue
}
if i < len(rawBlocks) {
if part.Type == codersdk.ChatMessagePartTypeFile {
if fid, err := chatprompt.ExtractFileID(rawBlocks[i]); err == nil {
part.FileID = uuid.NullUUID{UUID: fid, Valid: true}
}
// When a file_id is present, omit inline data
// from the response. Clients fetch content via
// the GET /chats/files/{id} endpoint instead.
if part.FileID.Valid {
part.Data = nil
}
}
}
parts = append(parts, part)
}
return parts, nil
case string(fantasy.MessageRoleTool):
results, err := parseToolResults(raw)
if err != nil {
return nil, err
}
parts := make([]codersdk.ChatMessagePart, 0, len(results))
for _, result := range results {
parts = append(parts, codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeToolResult,
ToolCallID: result.ToolCallID,
ToolName: result.ToolName,
Result: result.Result,
IsError: result.IsError,
ProviderExecuted: result.ProviderExecuted,
})
}
return parts, nil
default:
return nil, nil
func chatMessageParts(role codersdk.ChatMessageRole, raw pqtype.NullRawMessage) ([]codersdk.ChatMessagePart, error) {
parts, err := chatprompt.ParseContent(role, raw)
if err != nil {
return nil, err
}
}
func parseSystemContent(raw pqtype.NullRawMessage) (string, error) {
if !raw.Valid || len(raw.RawMessage) == 0 {
return "", nil
// Strip internal-only fields before API responses.
for i := range parts {
parts[i].StripInternal()
}
var content string
if err := json.Unmarshal(raw.RawMessage, &content); err != nil {
return "", xerrors.Errorf("parse system content: %w", err)
}
return content, nil
}
func parseContentBlocks(role string, raw pqtype.NullRawMessage) ([]fantasy.Content, error) {
if !raw.Valid || len(raw.RawMessage) == 0 {
return nil, nil
}
if role == string(fantasy.MessageRoleUser) {
var text string
if err := json.Unmarshal(raw.RawMessage, &text); err == nil {
return []fantasy.Content{
fantasy.TextContent{Text: text},
}, nil
}
}
var blocks []json.RawMessage
if err := json.Unmarshal(raw.RawMessage, &blocks); err != nil {
return nil, xerrors.Errorf("parse content blocks: %w", err)
}
content := make([]fantasy.Content, 0, len(blocks))
for _, block := range blocks {
decoded, err := fantasy.UnmarshalContent(block)
if err != nil {
return nil, xerrors.Errorf("parse content block: %w", err)
}
content = append(content, decoded)
}
return content, nil
}
// toolResultRow is used only for extracting top-level fields from
// persisted tool result JSON. The result payload is kept as raw JSON.
type toolResultRow struct {
ToolCallID string `json:"tool_call_id"`
ToolName string `json:"tool_name"`
Result json.RawMessage `json:"result"`
IsError bool `json:"is_error,omitempty"`
ProviderExecuted bool `json:"provider_executed,omitempty"`
}
func parseToolResults(raw pqtype.NullRawMessage) ([]toolResultRow, error) {
if !raw.Valid || len(raw.RawMessage) == 0 {
return nil, nil
}
var results []toolResultRow
if err := json.Unmarshal(raw.RawMessage, &results); err != nil {
return nil, xerrors.Errorf("parse tool results: %w", err)
}
return results, nil
}
func contentBlockToPart(block fantasy.Content) codersdk.ChatMessagePart {
switch value := block.(type) {
case fantasy.TextContent:
return codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeText,
Text: value.Text,
}
case *fantasy.TextContent:
return codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeText,
Text: value.Text,
}
case fantasy.ReasoningContent:
return codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeReasoning,
Text: value.Text,
}
case *fantasy.ReasoningContent:
return codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeReasoning,
Text: value.Text,
}
case fantasy.ToolCallContent:
return codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeToolCall,
ToolCallID: value.ToolCallID,
ToolName: value.ToolName,
Args: []byte(value.Input),
ProviderExecuted: value.ProviderExecuted,
}
case *fantasy.ToolCallContent:
return codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeToolCall,
ToolCallID: value.ToolCallID,
ToolName: value.ToolName,
Args: []byte(value.Input),
ProviderExecuted: value.ProviderExecuted,
}
case fantasy.SourceContent:
return codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeSource,
SourceID: value.ID,
URL: value.URL,
Title: value.Title,
}
case *fantasy.SourceContent:
return codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeSource,
SourceID: value.ID,
URL: value.URL,
Title: value.Title,
}
case fantasy.FileContent:
return codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeFile,
MediaType: value.MediaType,
Data: value.Data,
}
case *fantasy.FileContent:
return codersdk.ChatMessagePart{
Type: codersdk.ChatMessagePartTypeFile,
MediaType: value.MediaType,
Data: value.Data,
}
case fantasy.ToolResultContent:
part := chatprompt.ToolResultToPart(
value.ToolCallID,
value.ToolName,
toolResultOutputToRawJSON(value.Result),
toolResultOutputIsError(value.Result),
)
part.ProviderExecuted = value.ProviderExecuted
return part
case *fantasy.ToolResultContent:
part := chatprompt.ToolResultToPart(
value.ToolCallID,
value.ToolName,
toolResultOutputToRawJSON(value.Result),
toolResultOutputIsError(value.Result),
)
part.ProviderExecuted = value.ProviderExecuted
return part
default:
return codersdk.ChatMessagePart{}
}
}
func toolResultOutputToRawJSON(output fantasy.ToolResultOutputContent) json.RawMessage {
switch v := output.(type) {
case fantasy.ToolResultOutputContentError:
if v.Error != nil {
data, _ := json.Marshal(map[string]any{"error": v.Error.Error()})
return data
}
return json.RawMessage(`{"error":""}`)
case fantasy.ToolResultOutputContentText:
raw := json.RawMessage(v.Text)
if json.Valid(raw) {
return raw
}
data, _ := json.Marshal(map[string]any{"output": v.Text})
return data
case fantasy.ToolResultOutputContentMedia:
data, _ := json.Marshal(map[string]any{
"data": v.Data,
"mime_type": v.MediaType,
"text": v.Text,
})
return data
default:
return json.RawMessage(`{}`)
}
}
func toolResultOutputIsError(output fantasy.ToolResultOutputContent) bool {
_, ok := output.(fantasy.ToolResultOutputContentError)
return ok
return parts, nil
}
func nullInt64Ptr(v sql.NullInt64) *int64 {
+1 -1
View File
@@ -467,7 +467,7 @@ func TestChatMessage_PreservesProviderExecutedOnToolResults(t *testing.T) {
dbMsg := database.ChatMessage{
ID: 1,
ChatID: uuid.New(),
Role: string(fantasy.MessageRoleAssistant),
Role: string(codersdk.ChatMessageRoleAssistant),
Content: pqtype.NullRawMessage{
RawMessage: rawContent,
Valid: true,