From bdbcd3428bbb9ecd11890efbdde9a983f1b9b8d6 Mon Sep 17 00:00:00 2001 From: Mathias Fredriksson Date: Fri, 13 Mar 2026 17:53:26 +0200 Subject: [PATCH] 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 --- coderd/chatd/chatd.go | 90 +- coderd/chatd/chatd_test.go | 51 +- coderd/chatd/chatloop/chatloop.go | 26 +- coderd/chatd/chatloop/compaction.go | 29 +- coderd/chatd/chatloop/compaction_test.go | 4 +- coderd/chatd/chatprompt/chatprompt.go | 873 ++++++++++-------- coderd/chatd/chatprompt/chatprompt_test.go | 778 +++++++++++++++- coderd/chatd/quickgen.go | 34 +- coderd/chatd/subagent.go | 9 +- coderd/chats.go | 59 +- coderd/chats_test.go | 150 +-- coderd/database/db2sdk/db2sdk.go | 261 +----- coderd/database/db2sdk/db2sdk_test.go | 2 +- codersdk/chats.go | 136 ++- codersdk/chats_test.go | 58 ++ enterprise/coderd/chatd/chatd_test.go | 39 +- enterprise/coderd/chats_test.go | 18 +- site/src/api/typesGenerated.ts | 27 +- .../AgentDetail/ChatContext.test.tsx | 2 +- 19 files changed, 1734 insertions(+), 912 deletions(-) diff --git a/coderd/chatd/chatd.go b/coderd/chatd/chatd.go index 267d9b2130..227eea9c4f 100644 --- a/coderd/chatd/chatd.go +++ b/coderd/chatd/chatd.go @@ -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 != "" { diff --git a/coderd/chatd/chatd_test.go b/coderd/chatd/chatd_test.go index b76bfac772..ac9d5bd750 100644 --- a/coderd/chatd/chatd_test.go +++ b/coderd/chatd/chatd_test.go @@ -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) diff --git a/coderd/chatd/chatloop/chatloop.go b/coderd/chatd/chatloop/chatloop.go index f493e38175..0431fad6d9 100644 --- a/coderd/chatd/chatloop/chatloop.go +++ b/coderd/chatd/chatloop/chatloop.go @@ -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), ) } diff --git a/coderd/chatd/chatloop/compaction.go b/coderd/chatd/chatloop/compaction.go index 696698b5dd..e6280ab7c2 100644 --- a/coderd/chatd/chatloop/compaction.go +++ b/coderd/chatd/chatloop/compaction.go @@ -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), ) } diff --git a/coderd/chatd/chatloop/compaction_test.go b/coderd/chatd/chatloop/compaction_test.go index 254dc8b57e..5c0f501126 100644 --- a/coderd/chatd/chatloop/compaction_test.go +++ b/coderd/chatd/chatloop/compaction_test.go @@ -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 { diff --git a/coderd/chatd/chatprompt/chatprompt.go b/coderd/chatd/chatprompt/chatprompt.go index 72a35b3a10..0085254ad6 100644 --- a/coderd/chatd/chatprompt/chatprompt.go +++ b/coderd/chatd/chatprompt/chatprompt.go @@ -1,8 +1,10 @@ package chatprompt import ( + "bytes" "context" "encoding/json" + "fmt" "regexp" "strings" @@ -49,62 +51,6 @@ func ExtractFileID(raw json.RawMessage) (uuid.UUID, error) { return uuid.Parse(envelope.Data.FileID) } -// extractFileIDs scans raw message content for file_id references. -// Returns a map of block index to file ID. Returns nil for -// non-array content or content with no file references. -func extractFileIDs(raw pqtype.NullRawMessage) map[int]uuid.UUID { - if !raw.Valid || len(raw.RawMessage) == 0 { - return nil - } - var rawBlocks []json.RawMessage - if err := json.Unmarshal(raw.RawMessage, &rawBlocks); err != nil { - return nil - } - var result map[int]uuid.UUID - for i, block := range rawBlocks { - fid, err := ExtractFileID(block) - if err == nil { - if result == nil { - result = make(map[int]uuid.UUID) - } - result[i] = fid - } - } - return result -} - -// patchFileContent fills in empty Data on FileContent blocks from -// resolved file data. Blocks that already have inline data (backward -// compat) or have no resolved data are left unchanged. -func patchFileContent( - content []fantasy.Content, - fileIDs map[int]uuid.UUID, - resolved map[uuid.UUID]FileData, -) { - for blockIdx, fid := range fileIDs { - if blockIdx >= len(content) { - continue - } - switch fc := content[blockIdx].(type) { - case fantasy.FileContent: - if len(fc.Data) > 0 { - continue - } - if data, found := resolved[fid]; found { - fc.Data = data.Data - content[blockIdx] = fc - } - case *fantasy.FileContent: - if len(fc.Data) > 0 { - continue - } - if data, found := resolved[fid]; found { - fc.Data = data.Data - } - } - } -} - // ConvertMessages converts persisted chat messages into LLM prompt // messages without resolving file references from storage. Inline // file data is preserved when present (backward compat). @@ -124,31 +70,41 @@ func ConvertMessagesWithFiles( resolver FileResolver, logger slog.Logger, ) ([]fantasy.Message, error) { - // Phase 1: Pre-scan user messages for file_id references. + // Phase 1: Parse all messages via ParseContent (→ SDK parts) + // and collect file_id references from user messages for batch + // resolution. + type parsedMessage struct { + role codersdk.ChatMessageRole + parts []codersdk.ChatMessagePart + } + parsed := make([]parsedMessage, len(messages)) var allFileIDs []uuid.UUID seenFileIDs := make(map[uuid.UUID]struct{}) - fileIDsByMsg := make(map[int]map[int]uuid.UUID) - if resolver != nil { - for i, msg := range messages { - visibility := msg.Visibility - if visibility == "" { - visibility = database.ChatMessageVisibilityBoth - } - if visibility != database.ChatMessageVisibilityModel && - visibility != database.ChatMessageVisibilityBoth { - continue - } - if msg.Role != string(fantasy.MessageRoleUser) { - continue - } - fids := extractFileIDs(msg.Content) - if len(fids) > 0 { - fileIDsByMsg[i] = fids - for _, fid := range fids { - if _, seen := seenFileIDs[fid]; !seen { - seenFileIDs[fid] = struct{}{} - allFileIDs = append(allFileIDs, fid) + for i, msg := range messages { + visibility := msg.Visibility + if visibility == "" { + visibility = database.ChatMessageVisibilityBoth + } + if visibility != database.ChatMessageVisibilityModel && + visibility != database.ChatMessageVisibilityBoth { + continue + } + + role := codersdk.ChatMessageRole(msg.Role) + parts, err := ParseContent(role, msg.Content) + if err != nil { + return nil, err + } + parsed[i] = parsedMessage{role: role, parts: parts} + + // Collect file IDs from user messages for resolution. + if resolver != nil && msg.Role == string(codersdk.ChatMessageRoleUser) { + for _, part := range parts { + if part.Type == codersdk.ChatMessagePartTypeFile && part.FileID.Valid { + if _, seen := seenFileIDs[part.FileID.UUID]; !seen { + seenFileIDs[part.FileID.UUID] = struct{}{} + allFileIDs = append(allFileIDs, part.FileID.UUID) } } } @@ -165,53 +121,34 @@ func ConvertMessagesWithFiles( } } - // Phase 3: Convert messages, patching file content as needed. + // Phase 3: Build fantasy messages from SDK parts via + // partsToMessageParts. Track tool names for injection. prompt := make([]fantasy.Message, 0, len(messages)) toolNameByCallID := make(map[string]string) - for i, message := range messages { - visibility := message.Visibility - if visibility == "" { - visibility = database.ChatMessageVisibilityBoth - } - if visibility != database.ChatMessageVisibilityModel && - visibility != database.ChatMessageVisibilityBoth { + for _, pm := range parsed { + if len(pm.parts) == 0 { continue } - switch message.Role { - case string(fantasy.MessageRoleSystem): - content, err := parseSystemContent(message.Content) - if err != nil { - return nil, err - } - if strings.TrimSpace(content) == "" { - continue - } + switch pm.role { + case codersdk.ChatMessageRoleSystem: + // System parts are always a single text part. prompt = append(prompt, fantasy.Message{ Role: fantasy.MessageRoleSystem, Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: content}, + fantasy.TextPart{Text: pm.parts[0].Text}, }, }) - case string(fantasy.MessageRoleUser): - content, err := ParseContent(string(fantasy.MessageRoleUser), message.Content) - if err != nil { - return nil, err - } - if fids, ok := fileIDsByMsg[i]; ok { - patchFileContent(content, fids, resolved) - } + case codersdk.ChatMessageRoleUser: prompt = append(prompt, fantasy.Message{ Role: fantasy.MessageRoleUser, - Content: ToMessageParts(content), + Content: partsToMessageParts(logger, pm.parts, resolved), }) - case string(fantasy.MessageRoleAssistant): - content, err := ParseContent(string(fantasy.MessageRoleAssistant), message.Content) - if err != nil { - return nil, err - } - parts := normalizeAssistantToolCallInputs(ToMessageParts(content)) - for _, toolCall := range ExtractToolCalls(parts) { + case codersdk.ChatMessageRoleAssistant: + fantasyParts := normalizeAssistantToolCallInputs( + partsToMessageParts(logger, pm.parts, resolved), + ) + for _, toolCall := range ExtractToolCalls(fantasyParts) { if toolCall.ToolCallID == "" || strings.TrimSpace(toolCall.ToolName) == "" { continue } @@ -219,26 +156,21 @@ func ConvertMessagesWithFiles( } prompt = append(prompt, fantasy.Message{ Role: fantasy.MessageRoleAssistant, - Content: parts, + Content: fantasyParts, }) - case string(fantasy.MessageRoleTool): - rows, err := parseToolResultRows(message.Content) - if err != nil { - return nil, err - } - parts := make([]fantasy.MessagePart, 0, len(rows)) - for _, row := range rows { - if row.ToolCallID != "" && row.ToolName != "" { - toolNameByCallID[sanitizeToolCallID(row.ToolCallID)] = row.ToolName + case codersdk.ChatMessageRoleTool: + // Track tool names from SDK parts before conversion. + for _, part := range pm.parts { + if part.Type == codersdk.ChatMessagePartTypeToolResult { + if part.ToolCallID != "" && part.ToolName != "" { + toolNameByCallID[sanitizeToolCallID(part.ToolCallID)] = part.ToolName + } } - parts = append(parts, row.toToolResultPart(logger)) } prompt = append(prompt, fantasy.Message{ Role: fantasy.MessageRoleTool, - Content: parts, + Content: partsToMessageParts(logger, pm.parts, resolved), }) - default: - return nil, xerrors.Errorf("unsupported chat message role %q", message.Role) } } prompt = injectMissingToolResults(prompt) @@ -330,31 +262,199 @@ func AppendUser(prompt []fantasy.Message, instruction string) []fantasy.Message return out } -// ParseContent decodes persisted chat message content blocks. -func ParseContent(role string, raw pqtype.NullRawMessage) ([]fantasy.Content, error) { +// ParseContent decodes persisted chat message content blocks into +// SDK parts. Role-aware: system messages are JSON strings, +// assistant and user messages use a structural heuristic +// (isFantasyEnvelopeFormat) to distinguish legacy fantasy envelope +// from SDK parts, and tool messages use try/fallback. +func ParseContent(role codersdk.ChatMessageRole, raw pqtype.NullRawMessage) ([]codersdk.ChatMessagePart, error) { if !raw.Valid || len(raw.RawMessage) == 0 { return nil, nil } + switch role { + case codersdk.ChatMessageRoleSystem: + return parseSystemRole(raw) + case codersdk.ChatMessageRoleAssistant: + return parseAssistantRole(raw) + case codersdk.ChatMessageRoleTool: + return parseToolRole(raw) + case codersdk.ChatMessageRoleUser: + return parseUserRole(raw) + default: + return nil, xerrors.Errorf("unsupported chat message role %q", role) + } +} + +// parseSystemRole decodes a system message (JSON string) into a +// single text part. +func parseSystemRole(raw pqtype.NullRawMessage) ([]codersdk.ChatMessagePart, error) { var text string - if err := json.Unmarshal(raw.RawMessage, &text); err == nil { - return []fantasy.Content{fantasy.TextContent{Text: text}}, nil + if err := json.Unmarshal(raw.RawMessage, &text); err != nil { + return nil, xerrors.Errorf("parse system content: %w", err) + } + if strings.TrimSpace(text) == "" { + return nil, nil + } + return []codersdk.ChatMessagePart{codersdk.ChatMessageText(text)}, nil +} + +// parseAssistantRole uses the structural heuristic to distinguish +// legacy fantasy envelope from new SDK parts. We don't use +// try/fallback here because json.Unmarshal of a fantasy envelope +// into []ChatMessagePart can partially succeed (Type gets set from +// the envelope's "type" field) while silently losing content. The +// only thing preventing that today is that Data ([]byte) rejects +// the envelope's "data" JSON object, but that's a brittle +// invariant tied to Go's json decoder behavior for []byte. +func parseAssistantRole(raw pqtype.NullRawMessage) ([]codersdk.ChatMessagePart, error) { + if isFantasyEnvelopeFormat(raw.RawMessage) { + return parseLegacyFantasyBlocks(string(codersdk.ChatMessageRoleAssistant), raw) } + // New SDK format. + var parts []codersdk.ChatMessagePart + if err := json.Unmarshal(raw.RawMessage, &parts); err != nil { + return nil, xerrors.Errorf("parse assistant content: %w", err) + } + if !hasNonEmptyType(parts) { + return nil, nil + } + return parts, nil +} + +// parseToolRole tries SDK parts first, then falls back to legacy +// tool result rows. Unlike assistant/user roles, tool messages +// don't need the isFantasyEnvelopeFormat heuristic: legacy tool +// result rows have no "type" field (just tool_call_id, tool_name, +// result), so hasToolResultType reliably rejects them. +func parseToolRole(raw pqtype.NullRawMessage) ([]codersdk.ChatMessagePart, error) { + // Try SDK parts. + var parts []codersdk.ChatMessagePart + if err := json.Unmarshal(raw.RawMessage, &parts); err == nil && hasToolResultType(parts) { + return parts, nil + } + + // Fall back to legacy tool result rows. + rows, err := parseToolResultRows(raw) + if err != nil { + return nil, err + } + parts = make([]codersdk.ChatMessagePart, 0, len(rows)) + for _, row := range rows { + part := codersdk.ChatMessageToolResult(row.ToolCallID, row.ToolName, row.Result, row.IsError) + part.ProviderExecuted = row.ProviderExecuted + part.ProviderMetadata = row.ProviderMetadata + parts = append(parts, part) + } + return parts, nil +} + +// parseUserRole uses a structural heuristic to distinguish legacy +// fantasy envelope from new SDK parts. +func parseUserRole(raw pqtype.NullRawMessage) ([]codersdk.ChatMessagePart, error) { + // Legacy: plain JSON string (very old format). + var text string + if err := json.Unmarshal(raw.RawMessage, &text); err == nil { + if strings.TrimSpace(text) == "" { + return nil, nil + } + return []codersdk.ChatMessagePart{codersdk.ChatMessageText(text)}, nil + } + + if isFantasyEnvelopeFormat(raw.RawMessage) { + return parseLegacyUserBlocks(raw) + } + + // New SDK format. + var parts []codersdk.ChatMessagePart + if err := json.Unmarshal(raw.RawMessage, &parts); err != nil { + return nil, xerrors.Errorf("parse user content: %w", err) + } + if !hasNonEmptyType(parts) { + return nil, nil + } + return parts, nil +} + +// parseLegacyUserBlocks decodes a user message stored in fantasy +// envelope format, extracting file_id references from the raw +// envelope for file-type blocks. +func parseLegacyUserBlocks(raw pqtype.NullRawMessage) ([]codersdk.ChatMessagePart, error) { + var rawBlocks []json.RawMessage + if err := json.Unmarshal(raw.RawMessage, &rawBlocks); err != nil { + return nil, xerrors.Errorf("parse user content: %w", err) + } + + parts := make([]codersdk.ChatMessagePart, 0, len(rawBlocks)) + for i, rawBlock := range rawBlocks { + block, err := fantasy.UnmarshalContent(rawBlock) + if err != nil { + return nil, xerrors.Errorf("parse user content block %d: %w", i, err) + } + part := PartFromContent(block) + if part.Type == "" { + continue + } + // For file-type blocks, extract file_id from the raw + // envelope's data sub-object. + if part.Type == codersdk.ChatMessagePartTypeFile { + if fid, err := ExtractFileID(rawBlock); err == nil { + part.FileID = uuid.NullUUID{UUID: fid, Valid: true} + // Clear inline data when file_id is present; + // resolved at LLM dispatch time. + part.Data = nil + } + } + parts = append(parts, part) + } + return parts, nil +} + +// parseLegacyFantasyBlocks decodes an assistant message stored in +// fantasy envelope format, converting each block via PartFromContent +// which preserves ProviderMetadata. +func parseLegacyFantasyBlocks(role string, raw pqtype.NullRawMessage) ([]codersdk.ChatMessagePart, error) { var rawBlocks []json.RawMessage if err := json.Unmarshal(raw.RawMessage, &rawBlocks); err != nil { return nil, xerrors.Errorf("parse %s content: %w", role, err) } - content := make([]fantasy.Content, 0, len(rawBlocks)) + parts := make([]codersdk.ChatMessagePart, 0, len(rawBlocks)) for i, rawBlock := range rawBlocks { block, err := fantasy.UnmarshalContent(rawBlock) if err != nil { return nil, xerrors.Errorf("parse %s content block %d: %w", role, i, err) } - content = append(content, block) + part := PartFromContent(block) + if part.Type == "" { + continue + } + parts = append(parts, part) } - return content, nil + return parts, nil +} + +// hasNonEmptyType returns true if at least one part has a non-empty +// Type field, indicating a valid SDK parts array. +func hasNonEmptyType(parts []codersdk.ChatMessagePart) bool { + for _, p := range parts { + if p.Type != "" { + return true + } + } + return false +} + +// hasToolResultType returns true if at least one part has Type == +// ToolResult, indicating a valid SDK tool-result array. +func hasToolResultType(parts []codersdk.ChatMessagePart) bool { + for _, p := range parts { + if p.Type == codersdk.ChatMessagePartTypeToolResult { + return true + } + } + return false } // toolResultRaw is an untyped representation of a persisted tool @@ -382,66 +482,6 @@ func parseToolResultRows(raw pqtype.NullRawMessage) ([]toolResultRaw, error) { return rows, nil } -func (r toolResultRaw) toToolResultPart(logger slog.Logger) fantasy.ToolResultPart { - toolCallID := sanitizeToolCallID(r.ToolCallID) - resultText := string(r.Result) - if resultText == "" || resultText == "null" { - resultText = "{}" - } - - if r.IsError { - message := strings.TrimSpace(resultText) - if extracted := extractErrorString(r.Result); extracted != "" { - message = extracted - } - return fantasy.ToolResultPart{ - ToolCallID: toolCallID, - ProviderExecuted: r.ProviderExecuted, - ProviderOptions: r.providerOptions(logger), - Output: fantasy.ToolResultOutputContentError{ - Error: xerrors.New(message), - }, - } - } - - return fantasy.ToolResultPart{ - ToolCallID: toolCallID, - ProviderExecuted: r.ProviderExecuted, - ProviderOptions: r.providerOptions(logger), - Output: fantasy.ToolResultOutputContentText{ - Text: resultText, - }, - } -} - -// providerOptions deserializes the stored provider metadata -// JSON into a ProviderOptions map using the fantasy type -// registry. Returns nil when no metadata is stored. -func (r toolResultRaw) providerOptions(logger slog.Logger) fantasy.ProviderOptions { - if len(r.ProviderMetadata) == 0 { - return nil - } - var raw map[string]json.RawMessage - if err := json.Unmarshal(r.ProviderMetadata, &raw); err != nil { - logger.Warn(context.Background(), - "failed to unmarshal provider metadata JSON", - slog.F("tool_call_id", r.ToolCallID), - slog.Error(err), - ) - return nil - } - opts, err := fantasy.UnmarshalProviderOptions(raw) - if err != nil { - logger.Warn(context.Background(), - "failed to deserialize provider metadata", - slog.F("tool_call_id", r.ToolCallID), - slog.Error(err), - ) - return nil - } - return opts -} - // extractErrorString pulls the "error" field from a JSON object if // present, returning it as a string. Returns "" if the field is // missing or the input is not an object. @@ -461,78 +501,6 @@ func extractErrorString(raw json.RawMessage) string { return strings.TrimSpace(s) } -// ToMessageParts converts fantasy content blocks into message parts. -func ToMessageParts(content []fantasy.Content) []fantasy.MessagePart { - parts := make([]fantasy.MessagePart, 0, len(content)) - for _, block := range content { - switch value := block.(type) { - case fantasy.TextContent: - parts = append(parts, fantasy.TextPart{ - Text: value.Text, - ProviderOptions: fantasy.ProviderOptions(value.ProviderMetadata), - }) - case *fantasy.TextContent: - parts = append(parts, fantasy.TextPart{ - Text: value.Text, - ProviderOptions: fantasy.ProviderOptions(value.ProviderMetadata), - }) - case fantasy.ReasoningContent: - parts = append(parts, fantasy.ReasoningPart{ - Text: value.Text, - ProviderOptions: fantasy.ProviderOptions(value.ProviderMetadata), - }) - case *fantasy.ReasoningContent: - parts = append(parts, fantasy.ReasoningPart{ - Text: value.Text, - ProviderOptions: fantasy.ProviderOptions(value.ProviderMetadata), - }) - case fantasy.ToolCallContent: - parts = append(parts, fantasy.ToolCallPart{ - ToolCallID: sanitizeToolCallID(value.ToolCallID), - ToolName: value.ToolName, - Input: value.Input, - ProviderExecuted: value.ProviderExecuted, - ProviderOptions: fantasy.ProviderOptions(value.ProviderMetadata), - }) - case *fantasy.ToolCallContent: - parts = append(parts, fantasy.ToolCallPart{ - ToolCallID: sanitizeToolCallID(value.ToolCallID), - ToolName: value.ToolName, - Input: value.Input, - ProviderExecuted: value.ProviderExecuted, - ProviderOptions: fantasy.ProviderOptions(value.ProviderMetadata), - }) - case fantasy.FileContent: - parts = append(parts, fantasy.FilePart{ - Data: value.Data, - MediaType: value.MediaType, - ProviderOptions: fantasy.ProviderOptions(value.ProviderMetadata), - }) - case *fantasy.FileContent: - parts = append(parts, fantasy.FilePart{ - Data: value.Data, - MediaType: value.MediaType, - ProviderOptions: fantasy.ProviderOptions(value.ProviderMetadata), - }) - case fantasy.ToolResultContent: - parts = append(parts, fantasy.ToolResultPart{ - ToolCallID: sanitizeToolCallID(value.ToolCallID), - ProviderExecuted: value.ProviderExecuted, - Output: value.Result, - ProviderOptions: fantasy.ProviderOptions(value.ProviderMetadata), - }) - case *fantasy.ToolResultContent: - parts = append(parts, fantasy.ToolResultPart{ - ToolCallID: sanitizeToolCallID(value.ToolCallID), - ProviderExecuted: value.ProviderExecuted, - Output: value.Result, - ProviderOptions: fantasy.ProviderOptions(value.ProviderMetadata), - }) - } - } - return parts -} - func normalizeAssistantToolCallInputs( parts []fantasy.MessagePart, ) []fantasy.MessagePart { @@ -585,10 +553,10 @@ func ExtractToolCalls(parts []fantasy.MessagePart) []fantasy.ToolCallContent { return toolCalls } -// MarshalContent encodes message content blocks for persistence. -// fileIDs optionally maps block indices to chat_files IDs, which -// are injected into the JSON envelope for file-type blocks so -// the reference survives round-trips through storage. +// MarshalContent encodes message content blocks in legacy fantasy +// envelope format. Retained for backward-compatible test fixtures +// that create legacy-format DB rows. Production write paths use +// MarshalParts instead. func MarshalContent(blocks []fantasy.Content, fileIDs map[int]uuid.UUID) (pqtype.NullRawMessage, error) { if len(blocks) == 0 { return pqtype.NullRawMessage{}, nil @@ -596,7 +564,7 @@ func MarshalContent(blocks []fantasy.Content, fileIDs map[int]uuid.UUID) (pqtype encodedBlocks := make([]json.RawMessage, 0, len(blocks)) for i, block := range blocks { - encoded, err := marshalContentBlock(block) + encoded, err := json.Marshal(block) if err != nil { return pqtype.NullRawMessage{}, xerrors.Errorf( "encode content block %d: %w", @@ -605,13 +573,23 @@ func MarshalContent(blocks []fantasy.Content, fileIDs map[int]uuid.UUID) (pqtype ) } if fid, ok := fileIDs[i]; ok { - encoded, err = injectFileID(encoded, fid) - if err != nil { - return pqtype.NullRawMessage{}, xerrors.Errorf( - "inject file_id into content block %d: %w", - i, - err, - ) + // Inline file_id injection into the fantasy envelope's + // data sub-object, stripping inline data. + var envelope struct { + Type string `json:"type"` + Data struct { + MediaType string `json:"media_type"` + Data json.RawMessage `json:"data,omitempty"` + FileID string `json:"file_id,omitempty"` + ProviderMetadata *json.RawMessage `json:"provider_metadata,omitempty"` + } `json:"data"` + } + if err := json.Unmarshal(encoded, &envelope); err == nil { + envelope.Data.FileID = fid.String() + envelope.Data.Data = nil + if patched, err := json.Marshal(envelope); err == nil { + encoded = patched + } } } encodedBlocks = append(encodedBlocks, encoded) @@ -624,28 +602,10 @@ func MarshalContent(blocks []fantasy.Content, fileIDs map[int]uuid.UUID) (pqtype return pqtype.NullRawMessage{RawMessage: data, Valid: true}, nil } -// injectFileID adds a file_id field into the data sub-object of a -// serialized content block envelope. -func injectFileID(encoded json.RawMessage, fileID uuid.UUID) (json.RawMessage, error) { - var envelope struct { - Type string `json:"type"` - Data struct { - MediaType string `json:"media_type"` - Data json.RawMessage `json:"data,omitempty"` - FileID string `json:"file_id,omitempty"` - ProviderMetadata *json.RawMessage `json:"provider_metadata,omitempty"` - } `json:"data"` - } - if err := json.Unmarshal(encoded, &envelope); err != nil { - return encoded, err - } - envelope.Data.FileID = fileID.String() - envelope.Data.Data = nil // Strip inline data; resolved at LLM dispatch time. - return json.Marshal(envelope) -} - -// MarshalToolResult encodes a single tool result for persistence as -// an opaque JSON blob. The stored shape is +// MarshalToolResult encodes a single tool result in the legacy +// tool-row format. Retained for test fixtures that create +// legacy-format DB rows. Production write paths use MarshalParts. +// The stored shape is // [{"tool_call_id":…,"tool_name":…,"result":…,"is_error":…}]. func MarshalToolResult(toolCallID, toolName string, result json.RawMessage, isError bool, providerExecuted bool, providerMetadata fantasy.ProviderMetadata) (pqtype.NullRawMessage, error) { var metaJSON json.RawMessage @@ -671,103 +631,81 @@ func MarshalToolResult(toolCallID, toolName string, result json.RawMessage, isEr return pqtype.NullRawMessage{RawMessage: data, Valid: true}, nil } -// MarshalToolResultContent encodes a fantasy tool result content -// block for persistence. It extracts the raw fields and delegates -// to MarshalToolResult. -func MarshalToolResultContent(content fantasy.ToolResultContent) (pqtype.NullRawMessage, error) { - var result json.RawMessage - var isError bool - - switch output := content.Result.(type) { - case fantasy.ToolResultOutputContentError: - isError = true - if output.Error != nil { - result, _ = json.Marshal(map[string]any{"error": output.Error.Error()}) - } else { - result = []byte(`{"error":""}`) - } - case fantasy.ToolResultOutputContentText: - result = json.RawMessage(output.Text) - if !json.Valid(result) { - result, _ = json.Marshal(map[string]any{"output": output.Text}) - } - case fantasy.ToolResultOutputContentMedia: - result, _ = json.Marshal(map[string]any{ - "data": output.Data, - "mime_type": output.MediaType, - "text": output.Text, - }) - default: - result = []byte(`{}`) - } - - return MarshalToolResult(content.ToolCallID, content.ToolName, result, isError, content.ProviderExecuted, content.ProviderMetadata) -} - -// PartFromContent converts fantasy content into a SDK chat message part. +// PartFromContent converts fantasy content into a SDK chat message +// part, preserving ProviderMetadata and ProviderExecuted fields. func PartFromContent(block fantasy.Content) codersdk.ChatMessagePart { switch value := block.(type) { case fantasy.TextContent: return codersdk.ChatMessagePart{ - Type: codersdk.ChatMessagePartTypeText, - Text: value.Text, + Type: codersdk.ChatMessagePartTypeText, + Text: value.Text, + ProviderMetadata: marshalProviderMetadata(value.ProviderMetadata), } case *fantasy.TextContent: return codersdk.ChatMessagePart{ - Type: codersdk.ChatMessagePartTypeText, - Text: value.Text, + Type: codersdk.ChatMessagePartTypeText, + Text: value.Text, + ProviderMetadata: marshalProviderMetadata(value.ProviderMetadata), } case fantasy.ReasoningContent: return codersdk.ChatMessagePart{ - Type: codersdk.ChatMessagePartTypeReasoning, - Text: value.Text, + Type: codersdk.ChatMessagePartTypeReasoning, + Text: value.Text, + ProviderMetadata: marshalProviderMetadata(value.ProviderMetadata), } case *fantasy.ReasoningContent: return codersdk.ChatMessagePart{ - Type: codersdk.ChatMessagePartTypeReasoning, - Text: value.Text, + Type: codersdk.ChatMessagePartTypeReasoning, + Text: value.Text, + ProviderMetadata: marshalProviderMetadata(value.ProviderMetadata), } case fantasy.ToolCallContent: return codersdk.ChatMessagePart{ Type: codersdk.ChatMessagePartTypeToolCall, ToolCallID: value.ToolCallID, ToolName: value.ToolName, - Args: []byte(value.Input), + Args: safeToolCallArgs(value.Input), ProviderExecuted: value.ProviderExecuted, + ProviderMetadata: marshalProviderMetadata(value.ProviderMetadata), } case *fantasy.ToolCallContent: return codersdk.ChatMessagePart{ Type: codersdk.ChatMessagePartTypeToolCall, ToolCallID: value.ToolCallID, ToolName: value.ToolName, - Args: []byte(value.Input), + Args: safeToolCallArgs(value.Input), ProviderExecuted: value.ProviderExecuted, + ProviderMetadata: marshalProviderMetadata(value.ProviderMetadata), } case fantasy.SourceContent: return codersdk.ChatMessagePart{ - Type: codersdk.ChatMessagePartTypeSource, - SourceID: value.ID, - URL: value.URL, - Title: value.Title, + Type: codersdk.ChatMessagePartTypeSource, + SourceID: value.ID, + URL: value.URL, + Title: value.Title, + ProviderMetadata: marshalProviderMetadata(value.ProviderMetadata), } case *fantasy.SourceContent: return codersdk.ChatMessagePart{ - Type: codersdk.ChatMessagePartTypeSource, - SourceID: value.ID, - URL: value.URL, - Title: value.Title, + Type: codersdk.ChatMessagePartTypeSource, + SourceID: value.ID, + URL: value.URL, + Title: value.Title, + ProviderMetadata: marshalProviderMetadata(value.ProviderMetadata), } case fantasy.FileContent: return codersdk.ChatMessagePart{ - Type: codersdk.ChatMessagePartTypeFile, - MediaType: value.MediaType, - Data: value.Data, + Type: codersdk.ChatMessagePartTypeFile, + MediaType: value.MediaType, + Data: value.Data, + ProviderMetadata: marshalProviderMetadata(value.ProviderMetadata), } case *fantasy.FileContent: return codersdk.ChatMessagePart{ - Type: codersdk.ChatMessagePartTypeFile, - MediaType: value.MediaType, - Data: value.Data, + Type: codersdk.ChatMessagePartTypeFile, + MediaType: value.MediaType, + Data: value.Data, + ProviderMetadata: marshalProviderMetadata(value.ProviderMetadata), } case fantasy.ToolResultContent: return toolResultContentToPart(value) @@ -782,13 +720,7 @@ func PartFromContent(block fantasy.Content) codersdk.ChatMessagePart { // flag into a ChatMessagePart. This is the minimal conversion used // both during streaming and when reading from the database. func ToolResultToPart(toolCallID, toolName string, result json.RawMessage, isError bool) codersdk.ChatMessagePart { - return codersdk.ChatMessagePart{ - Type: codersdk.ChatMessagePartTypeToolResult, - ToolCallID: toolCallID, - ToolName: toolName, - Result: result, - IsError: isError, - } + return codersdk.ChatMessageToolResult(toolCallID, toolName, result, isError) } // toolResultContentToPart converts a fantasy ToolResultContent @@ -823,6 +755,7 @@ func toolResultContentToPart(content fantasy.ToolResultContent) codersdk.ChatMes part := ToolResultToPart(content.ToolCallID, content.ToolName, result, isError) part.ProviderExecuted = content.ProviderExecuted + part.ProviderMetadata = marshalProviderMetadata(content.ProviderMetadata) return part } @@ -1039,18 +972,6 @@ func syntheticToolUseMessage( } } -func parseSystemContent(raw pqtype.NullRawMessage) (string, error) { - if !raw.Valid || len(raw.RawMessage) == 0 { - return "", nil - } - - var content string - if err := json.Unmarshal(raw.RawMessage, &content); err != nil { - return "", xerrors.Errorf("parse system message content: %w", err) - } - return content, nil -} - func sanitizeToolCallID(id string) string { if id == "" { return "" @@ -1058,6 +979,202 @@ func sanitizeToolCallID(id string) string { return toolCallIDSanitizer.ReplaceAllString(id, "_") } -func marshalContentBlock(block fantasy.Content) (json.RawMessage, error) { - return json.Marshal(block) +// MarshalParts encodes SDK chat message parts for persistence. +func MarshalParts(parts []codersdk.ChatMessagePart) (pqtype.NullRawMessage, error) { + if len(parts) == 0 { + return pqtype.NullRawMessage{}, nil + } + data, err := json.Marshal(parts) + if err != nil { + return pqtype.NullRawMessage{}, xerrors.Errorf("encode chat message parts: %w", err) + } + return pqtype.NullRawMessage{RawMessage: data, Valid: true}, nil +} + +// isFantasyEnvelopeFormat checks whether raw message content uses +// the fantasy envelope format (legacy) vs SDK parts (new). It +// examines the first array element for a "data" field containing a +// JSON object (starts with '{'). Fantasy always serializes Data +// from json.Marshal(struct{...}), producing a JSON object. +// ChatMessagePart.Data is []byte, which serializes to a base64 +// string or is omitted via omitempty. This structural invariant +// means a "data" field starting with '{' can only come from +// fantasy. +func isFantasyEnvelopeFormat(raw json.RawMessage) bool { + var arr []json.RawMessage + if err := json.Unmarshal(raw, &arr); err != nil || len(arr) == 0 { + return false + } + var fields map[string]json.RawMessage + if err := json.Unmarshal(arr[0], &fields); err != nil { + return false + } + data, ok := fields["data"] + if !ok { + return false + } + trimmed := bytes.TrimSpace(data) + return len(trimmed) > 0 && trimmed[0] == '{' +} + +// marshalProviderMetadata converts fantasy provider metadata to raw +// JSON for storage in SDK parts. +func marshalProviderMetadata(metadata fantasy.ProviderMetadata) json.RawMessage { + if len(metadata) == 0 { + return nil + } + data, err := json.Marshal(metadata) + if err != nil { + return nil + } + return data +} + +// providerMetadataToOptions reconstructs fantasy ProviderOptions +// from raw JSON stored in an SDK part's ProviderMetadata field. +// Uses fantasy.UnmarshalProviderOptions to restore registered +// provider-specific types. Returns nil on failure. +func providerMetadataToOptions(logger slog.Logger, raw json.RawMessage) fantasy.ProviderOptions { + if len(raw) == 0 { + return nil + } + var intermediate map[string]json.RawMessage + if err := json.Unmarshal(raw, &intermediate); err != nil { + logger.Warn(context.Background(), "failed to unmarshal provider metadata", slog.Error(err)) + return nil + } + opts, err := fantasy.UnmarshalProviderOptions(intermediate) + if err != nil { + logger.Warn(context.Background(), "failed to decode provider options", slog.Error(err)) + return nil + } + return opts +} + +// safeToolCallArgs ensures tool call args are valid JSON. Returns +// nil for empty or invalid input so the field is omitted. +func safeToolCallArgs(input string) json.RawMessage { + input = strings.TrimSpace(input) + if input == "" { + return nil + } + raw := json.RawMessage(input) + if !json.Valid(raw) { + return nil + } + return raw +} + +// fileReferencePartToText formats a file-reference SDK part as +// plain text for LLM consumption. LLMs don't understand +// file-reference natively, so we convert to a readable text +// representation. +func fileReferencePartToText(part codersdk.ChatMessagePart) string { + lineRange := fmt.Sprintf("%d", part.StartLine) + if part.StartLine != part.EndLine { + lineRange = fmt.Sprintf("%d-%d", part.StartLine, part.EndLine) + } + var sb strings.Builder + _, _ = fmt.Fprintf(&sb, "[file-reference] %s:%s", part.FileName, lineRange) + if content := strings.TrimSpace(part.Content); content != "" { + _, _ = fmt.Fprintf(&sb, "\n```%s\n%s\n```", part.FileName, content) + } + return sb.String() +} + +// toolResultPartToMessagePart converts an SDK tool-result part +// into a fantasy ToolResultPart for LLM dispatch. +func toolResultPartToMessagePart(logger slog.Logger, part codersdk.ChatMessagePart) fantasy.ToolResultPart { + toolCallID := sanitizeToolCallID(part.ToolCallID) + resultText := string(part.Result) + if resultText == "" || resultText == "null" { + resultText = "{}" + } + + opts := providerMetadataToOptions(logger, part.ProviderMetadata) + + if part.IsError { + message := strings.TrimSpace(resultText) + if extracted := extractErrorString(part.Result); extracted != "" { + message = extracted + } + return fantasy.ToolResultPart{ + ToolCallID: toolCallID, + ProviderExecuted: part.ProviderExecuted, + Output: fantasy.ToolResultOutputContentError{ + Error: xerrors.New(message), + }, + ProviderOptions: opts, + } + } + + return fantasy.ToolResultPart{ + ToolCallID: toolCallID, + ProviderExecuted: part.ProviderExecuted, + Output: fantasy.ToolResultOutputContentText{ + Text: resultText, + }, + ProviderOptions: opts, + } +} + +// partsToMessageParts converts SDK chat message parts into fantasy +// message parts for LLM dispatch. It handles file data injection +// from resolved files, file-reference to text conversion, and +// source part skipping. +func partsToMessageParts( + logger slog.Logger, + parts []codersdk.ChatMessagePart, + resolved map[uuid.UUID]FileData, +) []fantasy.MessagePart { + result := make([]fantasy.MessagePart, 0, len(parts)) + for _, part := range parts { + switch part.Type { + case codersdk.ChatMessagePartTypeText: + result = append(result, fantasy.TextPart{ + Text: part.Text, + ProviderOptions: providerMetadataToOptions(logger, part.ProviderMetadata), + }) + case codersdk.ChatMessagePartTypeReasoning: + result = append(result, fantasy.ReasoningPart{ + Text: part.Text, + ProviderOptions: providerMetadataToOptions(logger, part.ProviderMetadata), + }) + case codersdk.ChatMessagePartTypeToolCall: + result = append(result, fantasy.ToolCallPart{ + ToolCallID: sanitizeToolCallID(part.ToolCallID), + ToolName: part.ToolName, + Input: string(part.Args), + ProviderExecuted: part.ProviderExecuted, + ProviderOptions: providerMetadataToOptions(logger, part.ProviderMetadata), + }) + case codersdk.ChatMessagePartTypeToolResult: + result = append(result, toolResultPartToMessagePart(logger, part)) + case codersdk.ChatMessagePartTypeFile: + data := part.Data + mediaType := part.MediaType + if part.FileID.Valid { + if fd, ok := resolved[part.FileID.UUID]; ok { + data = fd.Data + if mediaType == "" { + mediaType = fd.MediaType + } + } + } + result = append(result, fantasy.FilePart{ + Data: data, + MediaType: mediaType, + ProviderOptions: providerMetadataToOptions(logger, part.ProviderMetadata), + }) + case codersdk.ChatMessagePartTypeFileReference: + // LLMs don't understand file-reference natively. + result = append(result, fantasy.TextPart{ + Text: fileReferencePartToText(part), + }) + case codersdk.ChatMessagePartTypeSource: + // Source parts are metadata-only, not sent to LLM. + continue + } + } + return result } diff --git a/coderd/chatd/chatprompt/chatprompt_test.go b/coderd/chatd/chatprompt/chatprompt_test.go index ab593e3926..5a4e41beff 100644 --- a/coderd/chatd/chatprompt/chatprompt_test.go +++ b/coderd/chatd/chatprompt/chatprompt_test.go @@ -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) diff --git a/coderd/chatd/quickgen.go b/coderd/chatd/quickgen.go index 321796208f..f5747404f3 100644 --- a/coderd/chatd/quickgen.go +++ b/coderd/chatd/quickgen.go @@ -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 } diff --git a/coderd/chatd/subagent.go b/coderd/chatd/subagent.go index cbc390c94e..8288e8ef3c 100644 --- a/coderd/chatd/subagent.go +++ b/coderd/chatd/subagent.go @@ -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 } diff --git a/coderd/chats.go b/coderd/chats.go index 959cc764aa..48aea9ca4b 100644 --- a/coderd/chats.go +++ b/coderd/chats.go @@ -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 { diff --git a/coderd/chats_test.go b/coderd/chats_test.go index 5b452f6e85..d765120e2f 100644 --- a/coderd/chats_test.go +++ b/coderd/chats_test.go @@ -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 { diff --git a/coderd/database/db2sdk/db2sdk.go b/coderd/database/db2sdk/db2sdk.go index 4b6139ca95..ee4cca1129 100644 --- a/coderd/database/db2sdk/db2sdk.go +++ b/coderd/database/db2sdk/db2sdk.go @@ -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 { diff --git a/coderd/database/db2sdk/db2sdk_test.go b/coderd/database/db2sdk/db2sdk_test.go index ef05cf5e6c..95cc8cf2da 100644 --- a/coderd/database/db2sdk/db2sdk_test.go +++ b/coderd/database/db2sdk/db2sdk_test.go @@ -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, diff --git a/codersdk/chats.go b/codersdk/chats.go index 9117714789..9881b748be 100644 --- a/codersdk/chats.go +++ b/codersdk/chats.go @@ -52,7 +52,7 @@ type ChatMessage struct { CreatedBy *uuid.UUID `json:"created_by,omitempty" format:"uuid"` ModelConfigID *uuid.UUID `json:"model_config_id,omitempty" format:"uuid"` CreatedAt time.Time `json:"created_at" format:"date-time"` - Role string `json:"role"` + Role ChatMessageRole `json:"role"` Content []ChatMessagePart `json:"content,omitempty"` Usage *ChatMessageUsage `json:"usage,omitempty"` } @@ -68,6 +68,17 @@ type ChatMessageUsage struct { ContextLimit *int64 `json:"context_limit,omitempty"` } +// ChatMessageRole represents the role of a chat message sender. +type ChatMessageRole string + +// ChatMessageRole enums. +const ( + ChatMessageRoleSystem ChatMessageRole = "system" + ChatMessageRoleUser ChatMessageRole = "user" + ChatMessageRoleAssistant ChatMessageRole = "assistant" + ChatMessageRoleTool ChatMessageRole = "tool" +) + // ChatMessagePartType represents a structured message part type. type ChatMessagePartType string @@ -82,24 +93,30 @@ const ( ) // ChatMessagePart is a structured chunk of a chat message. +// +// WARNING: This type is both an API wire type and a database +// persistence format. Its JSON layout is stored in the +// chat_messages.content column. Field additions, renames, type +// changes, and omitempty behavior all affect backward-compatible +// deserialization of stored rows. Treat changes to this struct +// with the same care as a database migration. type ChatMessagePart struct { - Type ChatMessagePartType `json:"type"` - Text string `json:"text,omitempty"` - Signature string `json:"signature,omitempty"` - ToolCallID string `json:"tool_call_id,omitempty"` - ToolName string `json:"tool_name,omitempty"` - Args json.RawMessage `json:"args,omitempty"` - ArgsDelta string `json:"args_delta,omitempty"` - Result json.RawMessage `json:"result,omitempty"` - ResultDelta string `json:"result_delta,omitempty"` - IsError bool `json:"is_error,omitempty"` - ProviderExecuted bool `json:"provider_executed,omitempty"` - SourceID string `json:"source_id,omitempty"` - URL string `json:"url,omitempty"` - Title string `json:"title,omitempty"` - MediaType string `json:"media_type,omitempty"` - Data []byte `json:"data,omitempty"` - FileID uuid.NullUUID `json:"file_id,omitempty" format:"uuid"` + Type ChatMessagePartType `json:"type"` + Text string `json:"text,omitempty"` + Signature string `json:"signature,omitempty"` + ToolCallID string `json:"tool_call_id,omitempty"` + ToolName string `json:"tool_name,omitempty"` + Args json.RawMessage `json:"args,omitempty"` + ArgsDelta string `json:"args_delta,omitempty"` + Result json.RawMessage `json:"result,omitempty"` + ResultDelta string `json:"result_delta,omitempty"` + IsError bool `json:"is_error,omitempty"` + SourceID string `json:"source_id,omitempty"` + URL string `json:"url,omitempty"` + Title string `json:"title,omitempty"` + MediaType string `json:"media_type,omitempty"` + Data []byte `json:"data,omitempty"` + FileID uuid.NullUUID `json:"file_id,omitempty" format:"uuid"` // The following fields are only set when Type is // ChatInputPartTypeFileReference. FileName string `json:"file_name,omitempty"` @@ -107,6 +124,87 @@ type ChatMessagePart struct { EndLine int `json:"end_line,omitempty"` // The code content from the diff that was commented on. Content string `json:"content,omitempty"` + // ProviderMetadata holds provider-specific response metadata + // (e.g. Anthropic cache control hints) as raw JSON. Internal + // only: stripped by db2sdk before API responses. + ProviderMetadata json.RawMessage `json:"provider_metadata,omitempty" typescript:"-"` + // ProviderExecuted indicates the tool call was executed by + // the provider (e.g. Anthropic computer use). + ProviderExecuted bool `json:"provider_executed,omitempty"` +} + +// StripInternal removes internal-only fields that must not be +// sent to API clients. Call before publishing via REST or SSE. +// +// Note: ArgsDelta and ResultDelta are intentionally preserved. +// They are streaming-only fields consumed by the frontend via +// SSE message_part events (see processStepStream in chatloop). +func (p *ChatMessagePart) StripInternal() { + p.ProviderMetadata = nil + if p.FileID.Valid { + p.Data = nil + } +} + +// ChatMessageText builds a text chat message part. +func ChatMessageText(text string) ChatMessagePart { + return ChatMessagePart{Type: ChatMessagePartTypeText, Text: text} +} + +// ChatMessageReasoning builds a reasoning chat message part. +func ChatMessageReasoning(text string) ChatMessagePart { + return ChatMessagePart{Type: ChatMessagePartTypeReasoning, Text: text} +} + +// ChatMessageToolCall builds a tool-call chat message part. +func ChatMessageToolCall(toolCallID, toolName string, args json.RawMessage) ChatMessagePart { + return ChatMessagePart{ + Type: ChatMessagePartTypeToolCall, + ToolCallID: toolCallID, + ToolName: toolName, + Args: args, + } +} + +// ChatMessageToolResult builds a tool-result chat message part. +func ChatMessageToolResult(toolCallID, toolName string, result json.RawMessage, isError bool) ChatMessagePart { + return ChatMessagePart{ + Type: ChatMessagePartTypeToolResult, + ToolCallID: toolCallID, + ToolName: toolName, + Result: result, + IsError: isError, + } +} + +// ChatMessageFile builds a file chat message part. +func ChatMessageFile(fileID uuid.UUID, mediaType string) ChatMessagePart { + return ChatMessagePart{ + Type: ChatMessagePartTypeFile, + FileID: uuid.NullUUID{UUID: fileID, Valid: true}, + MediaType: mediaType, + } +} + +// ChatMessageFileReference builds a file-reference chat message part. +func ChatMessageFileReference(fileName string, startLine, endLine int, content string) ChatMessagePart { + return ChatMessagePart{ + Type: ChatMessagePartTypeFileReference, + FileName: fileName, + StartLine: startLine, + EndLine: endLine, + Content: content, + } +} + +// ChatMessageSource builds a source chat message part. +func ChatMessageSource(sourceID, url, title string) ChatMessagePart { + return ChatMessagePart{ + Type: ChatMessagePartTypeSource, + SourceID: sourceID, + URL: url, + Title: title, + } } // ChatInputPartType represents an input part type for user chat input. @@ -568,7 +666,7 @@ type ChatQueuedMessage struct { // ChatStreamMessagePart is a streamed message part update. type ChatStreamMessagePart struct { - Role string `json:"role,omitempty"` + Role ChatMessageRole `json:"role,omitempty"` Part ChatMessagePart `json:"part"` } diff --git a/codersdk/chats_test.go b/codersdk/chats_test.go index 7cc32e9683..5e61aa81b5 100644 --- a/codersdk/chats_test.go +++ b/codersdk/chats_test.go @@ -4,6 +4,8 @@ import ( "encoding/json" "testing" + "github.com/google/uuid" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/coder/coder/v2/codersdk" @@ -51,3 +53,59 @@ func TestChatModelProviderOptions_UnmarshalJSON_ParsesPlainProviderPayloads(t *t *decoded.Anthropic.Effort, ) } + +func TestChatMessagePart_StripInternal(t *testing.T) { + t.Parallel() + + t.Run("StripsProviderMetadata", func(t *testing.T) { + t.Parallel() + part := codersdk.ChatMessagePart{ + Type: codersdk.ChatMessagePartTypeToolCall, + ToolCallID: "call-1", + ToolName: "some_tool", + Args: json.RawMessage(`{"key":"value"}`), + ProviderMetadata: json.RawMessage(`{"type":"ephemeral"}`), + } + part.StripInternal() + assert.Nil(t, part.ProviderMetadata) + // Public fields preserved. + assert.Equal(t, codersdk.ChatMessagePartTypeToolCall, part.Type) + assert.Equal(t, "call-1", part.ToolCallID) + assert.Equal(t, "some_tool", part.ToolName) + assert.JSONEq(t, `{"key":"value"}`, string(part.Args)) + }) + + t.Run("StripsFileDataWhenFileIDSet", func(t *testing.T) { + t.Parallel() + id := uuid.New() + part := codersdk.ChatMessagePart{ + Type: codersdk.ChatMessagePartTypeFile, + FileID: uuid.NullUUID{UUID: id, Valid: true}, + MediaType: "image/png", + Data: []byte("binary-payload"), + } + part.StripInternal() + assert.Nil(t, part.Data) + assert.Equal(t, id, part.FileID.UUID) + assert.Equal(t, "image/png", part.MediaType) + }) + + t.Run("PreservesDataWhenNoFileID", func(t *testing.T) { + t.Parallel() + part := codersdk.ChatMessagePart{ + Type: codersdk.ChatMessagePartTypeFile, + MediaType: "image/png", + Data: []byte("inline-data"), + } + part.StripInternal() + assert.Equal(t, []byte("inline-data"), part.Data) + }) + + t.Run("NoopOnCleanPart", func(t *testing.T) { + t.Parallel() + part := codersdk.ChatMessageText("hello") + part.StripInternal() + assert.Equal(t, "hello", part.Text) + assert.Equal(t, codersdk.ChatMessagePartTypeText, part.Type) + }) +} diff --git a/enterprise/coderd/chatd/chatd_test.go b/enterprise/coderd/chatd/chatd_test.go index 65cd410349..30dd161a01 100644 --- a/enterprise/coderd/chatd/chatd_test.go +++ b/enterprise/coderd/chatd/chatd_test.go @@ -10,7 +10,6 @@ import ( "testing" "time" - "charm.land/fantasy" "github.com/google/uuid" "github.com/stretchr/testify/require" "golang.org/x/xerrors" @@ -118,7 +117,7 @@ func TestSubscribeRelayReconnectsOnDrop(t *testing.T) { Type: codersdk.ChatStreamEventTypeMessagePart, MessagePart: &codersdk.ChatStreamMessagePart{ Role: "assistant", - Part: codersdk.ChatMessagePart{Type: codersdk.ChatMessagePartTypeText, Text: "first-relay"}, + Part: codersdk.ChatMessageText("first-relay"), }, } close(ch) @@ -128,7 +127,7 @@ func TestSubscribeRelayReconnectsOnDrop(t *testing.T) { Type: codersdk.ChatStreamEventTypeMessagePart, MessagePart: &codersdk.ChatStreamMessagePart{ Role: "assistant", - Part: codersdk.ChatMessagePart{Type: codersdk.ChatMessagePartTypeText, Text: "second-relay"}, + Part: codersdk.ChatMessageText("second-relay"), }, } // Don't close — keep alive so the subscriber stays connected. @@ -152,7 +151,7 @@ func TestSubscribeRelayReconnectsOnDrop(t *testing.T) { OwnerID: user.ID, Title: "relay-reconnect", ModelConfigID: model.ID, - InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}}, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, }) require.NoError(t, err) @@ -246,7 +245,7 @@ func TestSubscribeRelayAsyncDoesNotBlock(t *testing.T) { OwnerID: user.ID, Title: "relay-async-nonblock", ModelConfigID: model.ID, - InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}}, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, }) require.NoError(t, err) @@ -320,14 +319,14 @@ func TestSubscribeRelaySnapshotDelivered(t *testing.T) { Type: codersdk.ChatStreamEventTypeMessagePart, MessagePart: &codersdk.ChatStreamMessagePart{ Role: "assistant", - Part: codersdk.ChatMessagePart{Type: codersdk.ChatMessagePartTypeText, Text: "snap-one"}, + Part: codersdk.ChatMessageText("snap-one"), }, }, { Type: codersdk.ChatStreamEventTypeMessagePart, MessagePart: &codersdk.ChatStreamMessagePart{ Role: "assistant", - Part: codersdk.ChatMessagePart{Type: codersdk.ChatMessagePartTypeText, Text: "snap-two"}, + Part: codersdk.ChatMessageText("snap-two"), }, }, } @@ -337,7 +336,7 @@ func TestSubscribeRelaySnapshotDelivered(t *testing.T) { Type: codersdk.ChatStreamEventTypeMessagePart, MessagePart: &codersdk.ChatStreamMessagePart{ Role: "assistant", - Part: codersdk.ChatMessagePart{Type: codersdk.ChatMessagePartTypeText, Text: "live-part"}, + Part: codersdk.ChatMessageText("live-part"), }, } return snapshot, ch, func() {}, nil @@ -353,7 +352,7 @@ func TestSubscribeRelaySnapshotDelivered(t *testing.T) { OwnerID: user.ID, Title: "relay-snapshot", ModelConfigID: model.ID, - InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}}, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, }) require.NoError(t, err) @@ -440,7 +439,7 @@ func TestSubscribeRelayStaleDialDiscardedAfterInterrupt(t *testing.T) { Type: codersdk.ChatStreamEventTypeMessagePart, MessagePart: &codersdk.ChatStreamMessagePart{ Role: "assistant", - Part: codersdk.ChatMessagePart{Type: codersdk.ChatMessagePartTypeText, Text: "stale-part"}, + Part: codersdk.ChatMessageText("stale-part"), }, } close(ch) @@ -451,7 +450,7 @@ func TestSubscribeRelayStaleDialDiscardedAfterInterrupt(t *testing.T) { Type: codersdk.ChatStreamEventTypeMessagePart, MessagePart: &codersdk.ChatStreamMessagePart{ Role: "assistant", - Part: codersdk.ChatMessagePart{Type: codersdk.ChatMessagePartTypeText, Text: "new-worker-part"}, + Part: codersdk.ChatMessageText("new-worker-part"), }, } return nil, ch, func() {}, nil @@ -466,7 +465,7 @@ func TestSubscribeRelayStaleDialDiscardedAfterInterrupt(t *testing.T) { OwnerID: user.ID, Title: "stale-dial-test", ModelConfigID: model.ID, - InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}}, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, }) require.NoError(t, err) @@ -629,7 +628,7 @@ func TestSubscribeCancelDuringInFlightDial(t *testing.T) { OwnerID: user.ID, Title: "cancel-inflight-dial", ModelConfigID: model.ID, - InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}}, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, }) require.NoError(t, err) @@ -712,7 +711,7 @@ func TestSubscribeRelayRunningToRunningSwitch(t *testing.T) { Type: codersdk.ChatStreamEventTypeMessagePart, MessagePart: &codersdk.ChatStreamMessagePart{ Role: "assistant", - Part: codersdk.ChatMessagePart{Type: codersdk.ChatMessagePartTypeText, Text: "worker-b-part"}, + Part: codersdk.ChatMessageText("worker-b-part"), }, } return nil, ch, func() {}, nil @@ -727,7 +726,7 @@ func TestSubscribeRelayRunningToRunningSwitch(t *testing.T) { OwnerID: user.ID, Title: "running-to-running", ModelConfigID: model.ID, - InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}}, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, }) require.NoError(t, err) @@ -827,7 +826,7 @@ func TestSubscribeRelayFailedDialRetries(t *testing.T) { Type: codersdk.ChatStreamEventTypeMessagePart, MessagePart: &codersdk.ChatStreamMessagePart{ Role: "assistant", - Part: codersdk.ChatMessagePart{Type: codersdk.ChatMessagePartTypeText, Text: "retry-success"}, + Part: codersdk.ChatMessageText("retry-success"), }, } return nil, ch, func() {}, nil @@ -849,7 +848,7 @@ func TestSubscribeRelayFailedDialRetries(t *testing.T) { OwnerID: user.ID, Title: "failed-dial-retry", ModelConfigID: model.ID, - InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}}, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, }) require.NoError(t, err) @@ -940,7 +939,7 @@ func TestSubscribeRunningLocalWorkerClosesRelay(t *testing.T) { Type: codersdk.ChatStreamEventTypeMessagePart, MessagePart: &codersdk.ChatStreamMessagePart{ Role: "assistant", - Part: codersdk.ChatMessagePart{Type: codersdk.ChatMessagePartTypeText, Text: "remote-part"}, + Part: codersdk.ChatMessageText("remote-part"), }, } // Keep channel open so the relay stays active. @@ -959,7 +958,7 @@ func TestSubscribeRunningLocalWorkerClosesRelay(t *testing.T) { OwnerID: user.ID, Title: "local-worker-closes-relay", ModelConfigID: model.ID, - InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}}, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, }) require.NoError(t, err) @@ -1067,7 +1066,7 @@ func TestSubscribeRelayMultipleReconnects(t *testing.T) { OwnerID: user.ID, Title: "multiple-reconnects", ModelConfigID: model.ID, - InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}}, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, }) require.NoError(t, err) diff --git a/enterprise/coderd/chats_test.go b/enterprise/coderd/chats_test.go index 41ad871dc0..09b99a40db 100644 --- a/enterprise/coderd/chats_test.go +++ b/enterprise/coderd/chats_test.go @@ -153,14 +153,14 @@ func TestChatStreamRelay(t *testing.T) { firstChunkText := "relay-part-one" streamingChunks <- chattest.OpenAITextChunks(firstChunkText)[0] firstEvent := waitForStreamTextPart(ctx, t, firstEvents, firstChunkText) - require.Equal(t, "assistant", firstEvent.MessagePart.Role) + require.Equal(t, codersdk.ChatMessageRoleAssistant, firstEvent.MessagePart.Role) secondEvents, secondStream, err := relayClient.StreamChat(ctx, chat.ID, nil) require.NoError(t, err) defer secondStream.Close() secondSnapshotEvent := waitForStreamTextPart(ctx, t, secondEvents, firstChunkText) - require.Equal(t, "assistant", secondSnapshotEvent.MessagePart.Role) + require.Equal(t, codersdk.ChatMessageRoleAssistant, secondSnapshotEvent.MessagePart.Role) secondChunkText := "relay-part-two" streamingChunks <- chattest.OpenAITextChunks(secondChunkText)[0] @@ -344,7 +344,7 @@ func TestChatStreamRelay(t *testing.T) { firstChunkText := "tls-relay-part-one" streamingChunks <- chattest.OpenAITextChunks(firstChunkText)[0] firstEvent := waitForStreamTextPart(ctx, t, firstEvents, firstChunkText) - require.Equal(t, "assistant", firstEvent.MessagePart.Role) + require.Equal(t, codersdk.ChatMessageRoleAssistant, firstEvent.MessagePart.Role) // Subscribe from the non-worker replica. This triggers the // relay dial to the worker over TLS. With the bug, this @@ -357,7 +357,7 @@ func TestChatStreamRelay(t *testing.T) { // The relay should deliver the already-sent chunk as a // snapshot event. secondSnapshotEvent := waitForStreamTextPart(ctx, t, secondEvents, firstChunkText) - require.Equal(t, "assistant", secondSnapshotEvent.MessagePart.Role) + require.Equal(t, codersdk.ChatMessageRoleAssistant, secondSnapshotEvent.MessagePart.Role) // Send another chunk and verify it flows through the relay. secondChunkText := "tls-relay-part-two" @@ -512,7 +512,7 @@ func TestChatStreamRelay(t *testing.T) { firstChunkText := "cookie-relay-part-one" streamingChunks <- chattest.OpenAITextChunks(firstChunkText)[0] firstEvent := waitForStreamTextPart(ctx, t, firstEvents, firstChunkText) - require.Equal(t, "assistant", firstEvent.MessagePart.Role) + require.Equal(t, codersdk.ChatMessageRoleAssistant, firstEvent.MessagePart.Role) // Subscribe from the non-worker replica with cookie-only // auth. This triggers the relay dial. If the relay doesn't @@ -522,7 +522,7 @@ func TestChatStreamRelay(t *testing.T) { defer secondStream.Close() secondSnapshotEvent := waitForStreamTextPart(ctx, t, secondEvents, firstChunkText) - require.Equal(t, "assistant", secondSnapshotEvent.MessagePart.Role) + require.Equal(t, codersdk.ChatMessageRoleAssistant, secondSnapshotEvent.MessagePart.Role) secondChunkText := "cookie-relay-part-two" streamingChunks <- chattest.OpenAITextChunks(secondChunkText)[0] @@ -684,7 +684,7 @@ func TestChatStreamRelay(t *testing.T) { firstChunkText := "hostprefix-relay-part-one" streamingChunks <- chattest.OpenAITextChunks(firstChunkText)[0] firstEvent := waitForStreamTextPart(ctx, t, firstEvents, firstChunkText) - require.Equal(t, "assistant", firstEvent.MessagePart.Role) + require.Equal(t, codersdk.ChatMessageRoleAssistant, firstEvent.MessagePart.Role) // This subscribe triggers the relay. With the bug, the // worker replica's HTTPCookies.Middleware strips the bare @@ -695,7 +695,7 @@ func TestChatStreamRelay(t *testing.T) { defer secondStream.Close() secondSnapshotEvent := waitForStreamTextPart(ctx, t, secondEvents, firstChunkText) - require.Equal(t, "assistant", secondSnapshotEvent.MessagePart.Role) + require.Equal(t, codersdk.ChatMessageRoleAssistant, secondSnapshotEvent.MessagePart.Role) secondChunkText := "hostprefix-relay-part-two" streamingChunks <- chattest.OpenAITextChunks(secondChunkText)[0] @@ -854,7 +854,7 @@ func TestChatStreamRelay(t *testing.T) { // Verify every buffered part arrives on the relay subscriber. for _, text := range bufferedTexts { event := waitForStreamTextPart(ctx, t, relayEvents, text) - require.Equal(t, "assistant", event.MessagePart.Role) + require.Equal(t, codersdk.ChatMessageRoleAssistant, event.MessagePart.Role) } // Send one more chunk after the relay subscriber is connected diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index 0f60f47146..bcae027ddf 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -1158,7 +1158,7 @@ export interface ChatMessage { readonly created_by?: string; readonly model_config_id?: string; readonly created_at: string; - readonly role: string; + readonly role: ChatMessageRole; readonly content?: readonly ChatMessagePart[]; readonly usage?: ChatMessageUsage; } @@ -1166,6 +1166,13 @@ export interface ChatMessage { // From codersdk/chats.go /** * ChatMessagePart is a structured chunk of a chat message. + * + * WARNING: This type is both an API wire type and a database + * persistence format. Its JSON layout is stored in the + * chat_messages.content column. Field additions, renames, type + * changes, and omitempty behavior all affect backward-compatible + * deserialization of stored rows. Treat changes to this struct + * with the same care as a database migration. */ export interface ChatMessagePart { readonly type: ChatMessagePartType; @@ -1178,7 +1185,6 @@ export interface ChatMessagePart { readonly result?: Record; readonly result_delta?: string; readonly is_error?: boolean; - readonly provider_executed?: boolean; readonly source_id?: string; readonly url?: string; readonly title?: string; @@ -1196,6 +1202,11 @@ export interface ChatMessagePart { * The code content from the diff that was commented on. */ readonly content?: string; + /** + * ProviderExecuted indicates the tool call was executed by + * the provider (e.g. Anthropic computer use). + */ + readonly provider_executed?: boolean; } // From codersdk/chats.go @@ -1218,6 +1229,16 @@ export const ChatMessagePartTypes: ChatMessagePartType[] = [ "tool-result", ]; +// From codersdk/chats.go +export type ChatMessageRole = "assistant" | "system" | "tool" | "user"; + +export const ChatMessageRoles: ChatMessageRole[] = [ + "assistant", + "system", + "tool", + "user", +]; + // From codersdk/chats.go /** * ChatMessageUsage contains token usage information for a chat message. @@ -1599,7 +1620,7 @@ export const ChatStreamEventTypes: ChatStreamEventType[] = [ * ChatStreamMessagePart is a streamed message part update. */ export interface ChatStreamMessagePart { - readonly role?: string; + readonly role?: ChatMessageRole; readonly part: ChatMessagePart; } diff --git a/site/src/pages/AgentsPage/AgentDetail/ChatContext.test.tsx b/site/src/pages/AgentsPage/AgentDetail/ChatContext.test.tsx index da7aba67dd..9035c4e068 100644 --- a/site/src/pages/AgentsPage/AgentDetail/ChatContext.test.tsx +++ b/site/src/pages/AgentsPage/AgentDetail/ChatContext.test.tsx @@ -194,7 +194,7 @@ const makeChat = (chatID: string): TypesGen.Chat => ({ const makeMessage = ( chatID: string, id: number, - role: string, + role: TypesGen.ChatMessageRole, text: string, ): TypesGen.ChatMessage => ({ id,