mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat(coderd/chatd): unify chat storage on SDK parts and fix file-reference rendering (#22958)
File-reference parts in user messages were flattened to `TextContent` at write time because fantasy has no file-reference content type. The frontend never saw them as structured parts. This moves all write paths (user, assistant, tool) from fantasy envelope format to `codersdk.ChatMessagePart`. The streaming layer (`chatloop`) is untouched, the conversion happens at the serialization boundary in `persistStep`. Old rows are still readable. `ParseContent` uses a structural heuristic (`isFantasyEnvelopeFormat`) to distinguish legacy envelopes from SDK parts. We chose this over try/fallback because fantasy envelopes partially unmarshal into `ChatMessagePart` (the `type` field matches) while silently losing content. A guard test enforces that no SDK part can produce the envelope shape. This is forward-only: new rows are unreadable by old code. Chat is behind a feature flag so rollback risk is contained. Also adds a typed `ChatMessageRole` to replace raw strings and `fantasy.MessageRole*` casts at the persistence boundary. The type covers `ChatMessage.Role`, `ChatStreamMessagePart.Role`, the `PublishMessagePart` callback chain, and all DB write sites. `fantasy.MessageRole*` remains only where we build `fantasy.Message` structs for LLM dispatch. Separately, `ProviderMetadata` was leaking to SSE clients via `publishMessagePart`. `StripInternal` now runs on both the SSE and REST paths, covering this. Other cleanup: - Old `db2sdk.contentBlockToPart` silently dropped metadata on text/reasoning/tool-call content. New code preserves it. - `providerMetadataToOptions` now logs warnings instead of silently returning nil. - `db2sdk` shrinks from ~250 lines of parallel conversion to ~15 lines delegating to `chatprompt.ParseContent()`, removing the `fantasy` import entirely. Refs #22821
This commit is contained in:
+42
-48
@@ -179,10 +179,7 @@ type CreateOptions struct {
|
||||
Title string
|
||||
ModelConfigID uuid.UUID
|
||||
SystemPrompt string
|
||||
InitialUserContent []fantasy.Content
|
||||
// ContentFileIDs maps content block indices to their chat_files IDs
|
||||
// so the file_id can be preserved in the stored message JSON.
|
||||
ContentFileIDs map[int]uuid.UUID
|
||||
InitialUserContent []codersdk.ChatMessagePart
|
||||
}
|
||||
|
||||
// SendMessageBusyBehavior controls what happens when a chat is already active.
|
||||
@@ -200,12 +197,11 @@ const (
|
||||
|
||||
// SendMessageOptions controls user message insertion with busy-state behavior.
|
||||
type SendMessageOptions struct {
|
||||
ChatID uuid.UUID
|
||||
CreatedBy uuid.UUID
|
||||
Content []fantasy.Content
|
||||
ContentFileIDs map[int]uuid.UUID
|
||||
ModelConfigID *uuid.UUID
|
||||
BusyBehavior SendMessageBusyBehavior
|
||||
ChatID uuid.UUID
|
||||
CreatedBy uuid.UUID
|
||||
Content []codersdk.ChatMessagePart
|
||||
ModelConfigID *uuid.UUID
|
||||
BusyBehavior SendMessageBusyBehavior
|
||||
}
|
||||
|
||||
// SendMessageResult contains the outcome of user message processing.
|
||||
@@ -221,8 +217,7 @@ type EditMessageOptions struct {
|
||||
ChatID uuid.UUID
|
||||
CreatedBy uuid.UUID
|
||||
EditedMessageID int64
|
||||
Content []fantasy.Content
|
||||
ContentFileIDs map[int]uuid.UUID
|
||||
Content []codersdk.ChatMessagePart
|
||||
}
|
||||
|
||||
// EditMessageResult contains the updated user message and chat status.
|
||||
@@ -284,7 +279,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
|
||||
UUID: opts.ModelConfigID,
|
||||
Valid: true,
|
||||
},
|
||||
Role: "system",
|
||||
Role: string(codersdk.ChatMessageRoleSystem),
|
||||
Content: pqtype.NullRawMessage{
|
||||
RawMessage: systemContent,
|
||||
Valid: len(systemContent) > 0,
|
||||
@@ -304,7 +299,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
|
||||
}
|
||||
}
|
||||
|
||||
userContent, err := chatprompt.MarshalContent(opts.InitialUserContent, opts.ContentFileIDs)
|
||||
userContent, err := chatprompt.MarshalParts(opts.InitialUserContent)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("marshal initial user content: %w", err)
|
||||
}
|
||||
@@ -314,7 +309,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
|
||||
UUID: opts.ModelConfigID,
|
||||
Valid: true,
|
||||
},
|
||||
Role: "user",
|
||||
Role: string(codersdk.ChatMessageRoleUser),
|
||||
Content: userContent,
|
||||
CreatedBy: uuid.NullUUID{UUID: opts.OwnerID, Valid: opts.OwnerID != uuid.Nil},
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
@@ -372,7 +367,7 @@ func (p *Server) SendMessage(
|
||||
return SendMessageResult{}, xerrors.Errorf("invalid busy behavior %q", opts.BusyBehavior)
|
||||
}
|
||||
|
||||
content, err := chatprompt.MarshalContent(opts.Content, opts.ContentFileIDs)
|
||||
content, err := chatprompt.MarshalParts(opts.Content)
|
||||
if err != nil {
|
||||
return SendMessageResult{}, xerrors.Errorf("marshal message content: %w", err)
|
||||
}
|
||||
@@ -506,7 +501,7 @@ func (p *Server) EditMessage(
|
||||
return EditMessageResult{}, xerrors.New("content is required")
|
||||
}
|
||||
|
||||
content, err := chatprompt.MarshalContent(opts.Content, opts.ContentFileIDs)
|
||||
content, err := chatprompt.MarshalParts(opts.Content)
|
||||
if err != nil {
|
||||
return EditMessageResult{}, xerrors.Errorf("marshal message content: %w", err)
|
||||
}
|
||||
@@ -902,7 +897,7 @@ func insertUserMessageAndSetPending(
|
||||
message, err := insertChatMessageWithStore(ctx, store, database.InsertChatMessageParams{
|
||||
ChatID: lockedChat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
Role: "user",
|
||||
Role: string(codersdk.ChatMessageRoleUser),
|
||||
Content: content,
|
||||
CreatedBy: uuid.NullUUID{UUID: createdBy, Valid: createdBy != uuid.Nil},
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
@@ -1740,10 +1735,13 @@ func (p *Server) publishEditedMessage(chatID uuid.UUID, message database.ChatMes
|
||||
})
|
||||
}
|
||||
|
||||
func (p *Server) publishMessagePart(chatID uuid.UUID, role string, part codersdk.ChatMessagePart) {
|
||||
func (p *Server) publishMessagePart(chatID uuid.UUID, role codersdk.ChatMessageRole, part codersdk.ChatMessagePart) {
|
||||
if part.Type == "" {
|
||||
return
|
||||
}
|
||||
// Strip internal-only fields before client delivery.
|
||||
// Mirrors db2sdk.chatMessageParts stripping for REST.
|
||||
part.StripInternal()
|
||||
p.publishEvent(chatID, codersdk.ChatStreamEvent{
|
||||
Type: codersdk.ChatStreamEventTypeMessagePart,
|
||||
MessagePart: &codersdk.ChatStreamMessagePart{
|
||||
@@ -1954,7 +1952,7 @@ func (p *Server) processChat(ctx context.Context, chat database.Chat) {
|
||||
msg, insertErr := tx.InsertChatMessage(cleanupCtx, database.InsertChatMessageParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: latestChat.LastModelConfigID, Valid: true},
|
||||
Role: "user",
|
||||
Role: string(codersdk.ChatMessageRoleUser),
|
||||
Content: pqtype.NullRawMessage{
|
||||
RawMessage: nextQueued.Content,
|
||||
Valid: len(nextQueued.Content) > 0,
|
||||
@@ -2336,7 +2334,11 @@ func (p *Server) runChat(
|
||||
}
|
||||
|
||||
if len(assistantBlocks) > 0 {
|
||||
assistantContent, marshalErr := chatprompt.MarshalContent(assistantBlocks, nil)
|
||||
sdkParts := make([]codersdk.ChatMessagePart, 0, len(assistantBlocks))
|
||||
for _, block := range assistantBlocks {
|
||||
sdkParts = append(sdkParts, chatprompt.PartFromContent(block))
|
||||
}
|
||||
assistantContent, marshalErr := chatprompt.MarshalParts(sdkParts)
|
||||
if marshalErr != nil {
|
||||
return marshalErr
|
||||
}
|
||||
@@ -2346,7 +2348,7 @@ func (p *Server) runChat(
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: uuid.NullUUID{},
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleAssistant),
|
||||
Role: string(codersdk.ChatMessageRoleAssistant),
|
||||
Content: assistantContent,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
InputTokens: usageNullInt64(step.Usage.InputTokens, hasUsage),
|
||||
@@ -2371,7 +2373,8 @@ func (p *Server) runChat(
|
||||
}
|
||||
|
||||
for _, tr := range toolResults {
|
||||
resultContent, marshalErr := chatprompt.MarshalToolResultContent(tr)
|
||||
trPart := chatprompt.PartFromContent(tr)
|
||||
resultContent, marshalErr := chatprompt.MarshalParts([]codersdk.ChatMessagePart{trPart})
|
||||
if marshalErr != nil {
|
||||
return marshalErr
|
||||
}
|
||||
@@ -2380,7 +2383,7 @@ func (p *Server) runChat(
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: uuid.NullUUID{},
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleTool),
|
||||
Role: string(codersdk.ChatMessageRoleTool),
|
||||
Content: resultContent,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
InputTokens: sql.NullInt64{},
|
||||
@@ -2461,8 +2464,8 @@ func (p *Server) runChat(
|
||||
},
|
||||
ToolCallID: compactionToolCallID,
|
||||
ToolName: "chat_summarized",
|
||||
PublishMessagePart: func(role fantasy.MessageRole, part codersdk.ChatMessagePart) {
|
||||
p.publishMessagePart(chat.ID, string(role), part)
|
||||
PublishMessagePart: func(role codersdk.ChatMessageRole, part codersdk.ChatMessagePart) {
|
||||
p.publishMessagePart(chat.ID, role, part)
|
||||
},
|
||||
OnError: func(err error) {
|
||||
logger.Warn(ctx, "failed to compact chat context", slog.Error(err))
|
||||
@@ -2550,10 +2553,10 @@ func (p *Server) runChat(
|
||||
|
||||
PersistStep: persistStep,
|
||||
PublishMessagePart: func(
|
||||
role fantasy.MessageRole,
|
||||
role codersdk.ChatMessageRole,
|
||||
part codersdk.ChatMessagePart,
|
||||
) {
|
||||
p.publishMessagePart(chat.ID, string(role), part)
|
||||
p.publishMessagePart(chat.ID, role, part)
|
||||
},
|
||||
Compaction: compactionOptions,
|
||||
ReloadMessages: func(reloadCtx context.Context) ([]fantasy.Message, error) {
|
||||
@@ -2686,13 +2689,9 @@ func (p *Server) persistChatContextSummary(
|
||||
return xerrors.Errorf("encode summary tool args: %w", err)
|
||||
}
|
||||
|
||||
assistantContent, err := chatprompt.MarshalContent([]fantasy.Content{
|
||||
fantasy.ToolCallContent{
|
||||
ToolCallID: toolCallID,
|
||||
ToolName: "chat_summarized",
|
||||
Input: string(args),
|
||||
},
|
||||
}, nil)
|
||||
assistantContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageToolCall(toolCallID, "chat_summarized", args),
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("encode summary tool call: %w", err)
|
||||
}
|
||||
@@ -2708,14 +2707,9 @@ func (p *Server) persistChatContextSummary(
|
||||
if err != nil {
|
||||
return xerrors.Errorf("encode summary result payload: %w", err)
|
||||
}
|
||||
toolResult, err := chatprompt.MarshalToolResult(
|
||||
toolCallID,
|
||||
"chat_summarized",
|
||||
summaryResult,
|
||||
false,
|
||||
false,
|
||||
nil,
|
||||
)
|
||||
toolResult, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageToolResult(toolCallID, "chat_summarized", summaryResult, false),
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("encode summary tool result: %w", err)
|
||||
}
|
||||
@@ -2727,7 +2721,7 @@ func (p *Server) persistChatContextSummary(
|
||||
ChatID: chatID,
|
||||
CreatedBy: uuid.NullUUID{},
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleUser),
|
||||
Role: string(codersdk.ChatMessageRoleUser),
|
||||
Content: pqtype.NullRawMessage{
|
||||
RawMessage: systemContent,
|
||||
Valid: len(systemContent) > 0,
|
||||
@@ -2750,7 +2744,7 @@ func (p *Server) persistChatContextSummary(
|
||||
ChatID: chatID,
|
||||
CreatedBy: uuid.NullUUID{},
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleAssistant),
|
||||
Role: string(codersdk.ChatMessageRoleAssistant),
|
||||
Content: assistantContent,
|
||||
Visibility: database.ChatMessageVisibilityUser,
|
||||
Compressed: sql.NullBool{
|
||||
@@ -2774,7 +2768,7 @@ func (p *Server) persistChatContextSummary(
|
||||
ChatID: chatID,
|
||||
CreatedBy: uuid.NullUUID{},
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
Role: string(fantasy.MessageRoleTool),
|
||||
Role: string(codersdk.ChatMessageRoleTool),
|
||||
Content: toolResult,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
Compressed: sql.NullBool{
|
||||
@@ -3140,10 +3134,10 @@ func (p *Server) maybeSendPushNotification(
|
||||
|
||||
msg, err := p.db.GetLastChatMessageByRole(pushCtx, database.GetLastChatMessageByRoleParams{
|
||||
ChatID: chat.ID,
|
||||
Role: "assistant",
|
||||
Role: string(codersdk.ChatMessageRoleAssistant),
|
||||
})
|
||||
if err == nil {
|
||||
content, parseErr := chatprompt.ParseContent(msg.Role, msg.Content)
|
||||
content, parseErr := chatprompt.ParseContent(codersdk.ChatMessageRoleAssistant, msg.Content)
|
||||
if parseErr == nil {
|
||||
assistantText := strings.TrimSpace(contentBlocksToText(content))
|
||||
if assistantText != "" {
|
||||
|
||||
+25
-26
@@ -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)
|
||||
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -149,7 +149,7 @@ func TestRun_Compaction(t *testing.T) {
|
||||
SummaryPrompt: "summarize now",
|
||||
ToolCallID: "test-tool-call-id",
|
||||
ToolName: "chat_summarized",
|
||||
PublishMessagePart: func(role fantasy.MessageRole, part codersdk.ChatMessagePart) {
|
||||
PublishMessagePart: func(role codersdk.ChatMessageRole, part codersdk.ChatMessagePart) {
|
||||
switch part.Type {
|
||||
case codersdk.ChatMessagePartTypeToolCall:
|
||||
callOrder = append(callOrder, "publish_tool_call")
|
||||
@@ -218,7 +218,7 @@ func TestRun_Compaction(t *testing.T) {
|
||||
ThresholdPercent: 70,
|
||||
ToolCallID: "test-tool-call-id",
|
||||
ToolName: "chat_summarized",
|
||||
PublishMessagePart: func(_ fantasy.MessageRole, _ codersdk.ChatMessagePart) {
|
||||
PublishMessagePart: func(_ codersdk.ChatMessageRole, _ codersdk.ChatMessagePart) {
|
||||
publishCalled = true
|
||||
},
|
||||
Persist: func(_ context.Context, _ CompactionResult) error {
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,18 +1,23 @@
|
||||
package chatprompt_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"charm.land/fantasy"
|
||||
fantasyanthropic "charm.land/fantasy/providers/anthropic"
|
||||
"github.com/google/uuid"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog/v3/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/coderd/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/db2sdk"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
func TestConvertMessages_NormalizesAssistantToolCallInput(t *testing.T) {
|
||||
@@ -46,7 +51,6 @@ func TestConvertMessages_NormalizesAssistantToolCallInput(t *testing.T) {
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -153,22 +157,37 @@ func TestConvertMessagesWithFiles_ResolvesFileData(t *testing.T) {
|
||||
func TestConvertMessagesWithFiles_BackwardCompat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// A message with inline data and a file_id should use the
|
||||
// inline data even when the resolver returns nothing.
|
||||
// A legacy message with inline data and a file_id: ParseContent
|
||||
// extracts the file_id and clears inline data (resolved at LLM
|
||||
// dispatch time). When a resolver provides data, the file part
|
||||
// in the LLM prompt should contain the resolved data.
|
||||
fileID := uuid.New()
|
||||
inlineData := []byte("inline-image-data")
|
||||
resolvedData := []byte("resolved-image-data")
|
||||
|
||||
rawContent := mustJSON(t, []json.RawMessage{
|
||||
mustJSON(t, map[string]any{
|
||||
"type": "file",
|
||||
"data": map[string]any{
|
||||
"media_type": "image/png",
|
||||
"data": inlineData,
|
||||
"data": []byte("inline-image-data"),
|
||||
"file_id": fileID.String(),
|
||||
},
|
||||
}),
|
||||
})
|
||||
|
||||
resolver := func(_ context.Context, ids []uuid.UUID) (map[uuid.UUID]chatprompt.FileData, error) {
|
||||
result := make(map[uuid.UUID]chatprompt.FileData)
|
||||
for _, id := range ids {
|
||||
if id == fileID {
|
||||
result[id] = chatprompt.FileData{
|
||||
Data: resolvedData,
|
||||
MediaType: "image/png",
|
||||
}
|
||||
}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
prompt, err := chatprompt.ConvertMessagesWithFiles(
|
||||
context.Background(),
|
||||
[]database.ChatMessage{
|
||||
@@ -178,7 +197,7 @@ func TestConvertMessagesWithFiles_BackwardCompat(t *testing.T) {
|
||||
Content: pqtype.NullRawMessage{RawMessage: rawContent, Valid: true},
|
||||
},
|
||||
},
|
||||
nil, // No resolver.
|
||||
resolver,
|
||||
slogtest.Make(t, nil),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
@@ -187,7 +206,8 @@ func TestConvertMessagesWithFiles_BackwardCompat(t *testing.T) {
|
||||
|
||||
filePart, ok := fantasy.AsMessagePart[fantasy.FilePart](prompt[0].Content[0])
|
||||
require.True(t, ok, "expected FilePart")
|
||||
require.Equal(t, inlineData, filePart.Data)
|
||||
require.Equal(t, resolvedData, filePart.Data)
|
||||
require.Equal(t, "image/png", filePart.MediaType)
|
||||
}
|
||||
|
||||
func TestInjectFileID_StripsInlineData(t *testing.T) {
|
||||
@@ -587,6 +607,750 @@ func TestProviderExecutedResult_LegacyToolRow(t *testing.T) {
|
||||
require.Equal(t, []string{"toolu_exec"}, toolIDs)
|
||||
}
|
||||
|
||||
// TestSDKPartsNeverProduceFantasyEnvelopeShape guards the structural
|
||||
// invariant that isFantasyEnvelopeFormat relies on: no SDK part type
|
||||
// serializes with a top-level "data" field containing a JSON object
|
||||
// (starting with '{'). Fantasy envelopes always have
|
||||
// "data":{object}, while ChatMessagePart.Data is []byte which
|
||||
// serializes to a base64 string or is omitted. If this test fails,
|
||||
// the format discriminator can no longer distinguish legacy fantasy
|
||||
// content from SDK parts, and parseAssistantRole / parseUserRole
|
||||
// would silently lose data on legacy rows.
|
||||
func TestSDKPartsNeverProduceFantasyEnvelopeShape(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
parts := []codersdk.ChatMessagePart{
|
||||
{Type: codersdk.ChatMessagePartTypeText, Text: "hello"},
|
||||
{Type: codersdk.ChatMessagePartTypeFile, FileID: uuid.NullUUID{UUID: uuid.New(), Valid: true}, MediaType: "image/png"},
|
||||
{Type: codersdk.ChatMessagePartTypeFile, MediaType: "image/png", Data: []byte("fake-image-data")},
|
||||
{Type: codersdk.ChatMessagePartTypeFileReference, FileName: "main.go", StartLine: 1, EndLine: 10, Content: "func main() {}"},
|
||||
{Type: codersdk.ChatMessagePartTypeReasoning, Text: "thinking..."},
|
||||
{Type: codersdk.ChatMessagePartTypeToolCall, ToolCallID: "abc", ToolName: "read_file", Args: json.RawMessage(`{"path":"main.go"}`)},
|
||||
{Type: codersdk.ChatMessagePartTypeToolResult, ToolCallID: "abc", ToolName: "read_file", Result: json.RawMessage(`{"output":"code"}`)},
|
||||
{Type: codersdk.ChatMessagePartTypeSource, SourceID: "s1", URL: "https://example.com", Title: "Example"},
|
||||
}
|
||||
for _, part := range parts {
|
||||
raw, err := json.Marshal(part)
|
||||
require.NoError(t, err)
|
||||
var fields map[string]json.RawMessage
|
||||
require.NoError(t, json.Unmarshal(raw, &fields))
|
||||
if data, ok := fields["data"]; ok {
|
||||
trimmed := bytes.TrimSpace(data)
|
||||
require.NotEmpty(t, trimmed)
|
||||
assert.NotEqual(t, byte('{'), trimmed[0],
|
||||
"SDK part type %q serializes with data field starting with '{', "+
|
||||
"would be misidentified as fantasy envelope by isFantasyEnvelopeFormat",
|
||||
part.Type)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// nullRaw wraps raw JSON bytes in a NullRawMessage for test input.
|
||||
func nullRaw(data json.RawMessage) pqtype.NullRawMessage {
|
||||
return pqtype.NullRawMessage{RawMessage: data, Valid: true}
|
||||
}
|
||||
|
||||
func TestParseContent_BackwardCompat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fileID := uuid.New()
|
||||
|
||||
// Build legacy fantasy assistant content using MarshalContent.
|
||||
legacyAssistantReasoning, err := chatprompt.MarshalContent([]fantasy.Content{
|
||||
fantasy.ReasoningContent{
|
||||
Text: "let me think...",
|
||||
ProviderMetadata: fantasy.ProviderMetadata{
|
||||
"anthropic": &fantasyanthropic.ProviderCacheControlOptions{
|
||||
CacheControl: fantasyanthropic.CacheControl{Type: "ephemeral"},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
legacyAssistantSource, err := chatprompt.MarshalContent([]fantasy.Content{
|
||||
fantasy.SourceContent{
|
||||
ID: "src_001",
|
||||
URL: "https://example.com/doc",
|
||||
Title: "Example Doc",
|
||||
},
|
||||
}, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
legacyAssistantToolCall, err := chatprompt.MarshalContent([]fantasy.Content{
|
||||
fantasy.ToolCallContent{
|
||||
ToolCallID: "call_123",
|
||||
ToolName: "read_file",
|
||||
Input: `{"path":"main.go"}`,
|
||||
},
|
||||
}, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Build new SDK format using MarshalParts.
|
||||
sdkMetadata := json.RawMessage(`{"anthropic":{"type":"anthropic.cache_control_options","data":{"cache_control":{"type":"ephemeral"}}}}`)
|
||||
|
||||
newAssistantWithMeta, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{{
|
||||
Type: codersdk.ChatMessagePartTypeText,
|
||||
Text: "here is my answer",
|
||||
ProviderMetadata: sdkMetadata,
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
|
||||
newAssistantToolCall, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{{
|
||||
Type: codersdk.ChatMessagePartTypeToolCall,
|
||||
ToolCallID: "call_456",
|
||||
ToolName: "execute",
|
||||
Args: json.RawMessage(`{"cmd":"ls"}`),
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
|
||||
newToolResult, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{{
|
||||
Type: codersdk.ChatMessagePartTypeToolResult,
|
||||
ToolCallID: "call_456",
|
||||
ToolName: "execute",
|
||||
Result: json.RawMessage(`{"output":"file1.go"}`),
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
role codersdk.ChatMessageRole
|
||||
raw pqtype.NullRawMessage
|
||||
check func(t *testing.T, parts []codersdk.ChatMessagePart)
|
||||
}{
|
||||
{
|
||||
name: "system/plain_string",
|
||||
role: codersdk.ChatMessageRoleSystem,
|
||||
raw: nullRaw(mustJSON(t, "You are helpful.")),
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeText, parts[0].Type)
|
||||
assert.Equal(t, "You are helpful.", parts[0].Text)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "user/fantasy_text",
|
||||
role: codersdk.ChatMessageRoleUser,
|
||||
raw: nullRaw(mustJSON(t, []json.RawMessage{
|
||||
mustJSON(t, map[string]any{
|
||||
"type": "text",
|
||||
"data": map[string]any{"text": "hello from user"},
|
||||
}),
|
||||
})),
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeText, parts[0].Type)
|
||||
assert.Equal(t, "hello from user", parts[0].Text)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "assistant/fantasy_text",
|
||||
role: codersdk.ChatMessageRoleAssistant,
|
||||
raw: nullRaw(mustJSON(t, []json.RawMessage{
|
||||
mustJSON(t, map[string]any{
|
||||
"type": "text",
|
||||
"data": map[string]any{"text": "hello from assistant"},
|
||||
}),
|
||||
})),
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeText, parts[0].Type)
|
||||
assert.Equal(t, "hello from assistant", parts[0].Text)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "user/plain_string",
|
||||
role: codersdk.ChatMessageRoleUser,
|
||||
raw: nullRaw(mustJSON(t, "just a plain string")),
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeText, parts[0].Type)
|
||||
assert.Equal(t, "just a plain string", parts[0].Text)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "user/fantasy_file_with_file_id",
|
||||
role: codersdk.ChatMessageRoleUser,
|
||||
raw: nullRaw(mustJSON(t, []json.RawMessage{
|
||||
mustJSON(t, map[string]any{
|
||||
"type": "file",
|
||||
"data": map[string]any{
|
||||
"media_type": "image/png",
|
||||
"file_id": fileID.String(),
|
||||
},
|
||||
}),
|
||||
})),
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeFile, parts[0].Type)
|
||||
assert.Equal(t, "image/png", parts[0].MediaType)
|
||||
assert.True(t, parts[0].FileID.Valid)
|
||||
assert.Equal(t, fileID, parts[0].FileID.UUID)
|
||||
assert.Nil(t, parts[0].Data, "inline data cleared when file_id present")
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "assistant/fantasy_reasoning_with_metadata",
|
||||
role: codersdk.ChatMessageRoleAssistant,
|
||||
raw: legacyAssistantReasoning,
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeReasoning, parts[0].Type)
|
||||
assert.Equal(t, "let me think...", parts[0].Text)
|
||||
require.NotNil(t, parts[0].ProviderMetadata, "ProviderMetadata must be preserved")
|
||||
assert.Contains(t, string(parts[0].ProviderMetadata), "anthropic")
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "assistant/fantasy_source",
|
||||
role: codersdk.ChatMessageRoleAssistant,
|
||||
raw: legacyAssistantSource,
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeSource, parts[0].Type)
|
||||
assert.Equal(t, "src_001", parts[0].SourceID)
|
||||
assert.Equal(t, "https://example.com/doc", parts[0].URL)
|
||||
assert.Equal(t, "Example Doc", parts[0].Title)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "assistant/fantasy_tool_call",
|
||||
role: codersdk.ChatMessageRoleAssistant,
|
||||
raw: legacyAssistantToolCall,
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeToolCall, parts[0].Type)
|
||||
assert.Equal(t, "call_123", parts[0].ToolCallID)
|
||||
assert.Equal(t, "read_file", parts[0].ToolName)
|
||||
assert.JSONEq(t, `{"path":"main.go"}`, string(parts[0].Args))
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "tool/legacy_result_row",
|
||||
role: codersdk.ChatMessageRoleTool,
|
||||
raw: nullRaw(mustJSON(t, []map[string]any{{
|
||||
"tool_call_id": "call_123",
|
||||
"tool_name": "read_file",
|
||||
"result": json.RawMessage(`{"output":"package main"}`),
|
||||
}})),
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeToolResult, parts[0].Type)
|
||||
assert.Equal(t, "call_123", parts[0].ToolCallID)
|
||||
assert.Equal(t, "read_file", parts[0].ToolName)
|
||||
assert.JSONEq(t, `{"output":"package main"}`, string(parts[0].Result))
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "user/sdk_text",
|
||||
role: codersdk.ChatMessageRoleUser,
|
||||
raw: nullRaw(mustJSON(t, []codersdk.ChatMessagePart{
|
||||
{Type: codersdk.ChatMessagePartTypeText, Text: "hello sdk"},
|
||||
})),
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeText, parts[0].Type)
|
||||
assert.Equal(t, "hello sdk", parts[0].Text)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "user/sdk_file_reference",
|
||||
role: codersdk.ChatMessageRoleUser,
|
||||
raw: nullRaw(mustJSON(t, []codersdk.ChatMessagePart{
|
||||
{Type: codersdk.ChatMessagePartTypeFileReference, FileName: "main.go", StartLine: 1, EndLine: 10, Content: "func main() {}"},
|
||||
})),
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeFileReference, parts[0].Type)
|
||||
assert.Equal(t, "main.go", parts[0].FileName)
|
||||
assert.Equal(t, 1, parts[0].StartLine)
|
||||
assert.Equal(t, 10, parts[0].EndLine)
|
||||
assert.Equal(t, "func main() {}", parts[0].Content)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "user/sdk_file",
|
||||
role: codersdk.ChatMessageRoleUser,
|
||||
raw: nullRaw(mustJSON(t, []codersdk.ChatMessagePart{
|
||||
{Type: codersdk.ChatMessagePartTypeFile, FileID: uuid.NullUUID{UUID: fileID, Valid: true}, MediaType: "image/png"},
|
||||
})),
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeFile, parts[0].Type)
|
||||
assert.True(t, parts[0].FileID.Valid)
|
||||
assert.Equal(t, fileID, parts[0].FileID.UUID)
|
||||
assert.Equal(t, "image/png", parts[0].MediaType)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "assistant/sdk_text_with_metadata",
|
||||
role: codersdk.ChatMessageRoleAssistant,
|
||||
raw: newAssistantWithMeta,
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeText, parts[0].Type)
|
||||
assert.Equal(t, "here is my answer", parts[0].Text)
|
||||
assert.JSONEq(t, string(sdkMetadata), string(parts[0].ProviderMetadata))
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "assistant/sdk_tool_call",
|
||||
role: codersdk.ChatMessageRoleAssistant,
|
||||
raw: newAssistantToolCall,
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeToolCall, parts[0].Type)
|
||||
assert.Equal(t, "call_456", parts[0].ToolCallID)
|
||||
assert.Equal(t, "execute", parts[0].ToolName)
|
||||
assert.JSONEq(t, `{"cmd":"ls"}`, string(parts[0].Args))
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "tool/sdk_tool_result",
|
||||
role: codersdk.ChatMessageRoleTool,
|
||||
raw: newToolResult,
|
||||
check: func(t *testing.T, parts []codersdk.ChatMessagePart) {
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeToolResult, parts[0].Type)
|
||||
assert.Equal(t, "call_456", parts[0].ToolCallID)
|
||||
assert.Equal(t, "execute", parts[0].ToolName)
|
||||
assert.JSONEq(t, `{"output":"file1.go"}`, string(parts[0].Result))
|
||||
},
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
parts, err := chatprompt.ParseContent(tc.role, tc.raw)
|
||||
require.NoError(t, err)
|
||||
tc.check(t, parts)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestProviderMetadataRoundTrip verifies that Anthropic cache
|
||||
// control hints survive the full path: legacy fantasy DB row →
|
||||
// ParseContent → SDK part (ProviderMetadata) → partsToMessageParts
|
||||
// → fantasy.MessagePart (ProviderOptions).
|
||||
func TestProviderMetadataRoundTrip(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
legacyContent, err := chatprompt.MarshalContent([]fantasy.Content{
|
||||
fantasy.TextContent{
|
||||
Text: "cached response",
|
||||
ProviderMetadata: fantasy.ProviderMetadata{
|
||||
"anthropic": &fantasyanthropic.ProviderCacheControlOptions{
|
||||
CacheControl: fantasyanthropic.CacheControl{Type: "ephemeral"},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Step 1: ParseContent preserves metadata on the SDK part.
|
||||
parts, err := chatprompt.ParseContent(codersdk.ChatMessageRoleAssistant, legacyContent)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, parts, 1)
|
||||
require.NotNil(t, parts[0].ProviderMetadata,
|
||||
"ProviderMetadata must survive ParseContent")
|
||||
|
||||
// Step 2: ConvertMessagesWithFiles reconstructs typed
|
||||
// ProviderOptions on the fantasy part.
|
||||
prompt, err := chatprompt.ConvertMessagesWithFiles(
|
||||
context.Background(),
|
||||
[]database.ChatMessage{{
|
||||
Role: "assistant",
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
Content: legacyContent,
|
||||
}},
|
||||
nil,
|
||||
slogtest.Make(t, nil),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, prompt, 1)
|
||||
require.Len(t, prompt[0].Content, 1)
|
||||
|
||||
textPart, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[0].Content[0])
|
||||
require.True(t, ok, "expected TextPart")
|
||||
require.Equal(t, "cached response", textPart.Text)
|
||||
|
||||
cc := fantasyanthropic.GetCacheControl(textPart.ProviderOptions)
|
||||
require.NotNil(t, cc, "Anthropic cache control must survive round-trip")
|
||||
require.Equal(t, "ephemeral", cc.Type)
|
||||
}
|
||||
|
||||
// TestFileReferencePreservation verifies file-reference parts
|
||||
// survive the storage round-trip and convert to text for LLMs.
|
||||
func TestFileReferencePreservation(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
raw, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{{
|
||||
Type: codersdk.ChatMessagePartTypeFileReference,
|
||||
FileName: "main.go",
|
||||
StartLine: 10,
|
||||
EndLine: 20,
|
||||
Content: "func main() {}",
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Storage round-trip: all fields intact.
|
||||
parts, err := chatprompt.ParseContent(codersdk.ChatMessageRoleUser, raw)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeFileReference, parts[0].Type)
|
||||
assert.Equal(t, "main.go", parts[0].FileName)
|
||||
assert.Equal(t, 10, parts[0].StartLine)
|
||||
assert.Equal(t, 20, parts[0].EndLine)
|
||||
assert.Equal(t, "func main() {}", parts[0].Content)
|
||||
|
||||
// LLM dispatch: file-reference becomes a TextPart.
|
||||
prompt, err := chatprompt.ConvertMessagesWithFiles(
|
||||
context.Background(),
|
||||
[]database.ChatMessage{{
|
||||
Role: "user",
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
Content: raw,
|
||||
}},
|
||||
nil,
|
||||
slogtest.Make(t, nil),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, prompt, 1)
|
||||
require.Len(t, prompt[0].Content, 1)
|
||||
|
||||
textPart, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[0].Content[0])
|
||||
require.True(t, ok, "file-reference should become TextPart for LLM")
|
||||
assert.Contains(t, textPart.Text, "[file-reference]")
|
||||
assert.Contains(t, textPart.Text, "main.go")
|
||||
assert.Contains(t, textPart.Text, "10-20")
|
||||
assert.Contains(t, textPart.Text, "func main() {}")
|
||||
}
|
||||
|
||||
// TestAssistantWriteRoundTrip verifies the Stage 4 write path:
|
||||
// fantasy.Content (with ProviderMetadata) → PartFromContent →
|
||||
// MarshalParts → DB → ParseContent (SDK path) →
|
||||
// ConvertMessagesWithFiles → fantasy part with ProviderOptions.
|
||||
func TestAssistantWriteRoundTrip(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
original := fantasy.TextContent{
|
||||
Text: "response with cache hints",
|
||||
ProviderMetadata: fantasy.ProviderMetadata{
|
||||
"anthropic": &fantasyanthropic.ProviderCacheControlOptions{
|
||||
CacheControl: fantasyanthropic.CacheControl{Type: "ephemeral"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Simulate persistStep: PartFromContent → MarshalParts.
|
||||
sdkPart := chatprompt.PartFromContent(original)
|
||||
require.Equal(t, codersdk.ChatMessagePartTypeText, sdkPart.Type)
|
||||
require.NotNil(t, sdkPart.ProviderMetadata)
|
||||
|
||||
raw, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{sdkPart})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Read back via ParseContent (takes the new SDK path, not
|
||||
// the legacy fallback, because the stored format is flat).
|
||||
parts, err := chatprompt.ParseContent(codersdk.ChatMessageRoleAssistant, raw)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, parts, 1)
|
||||
assert.Equal(t, "response with cache hints", parts[0].Text)
|
||||
assert.JSONEq(t, string(sdkPart.ProviderMetadata), string(parts[0].ProviderMetadata))
|
||||
|
||||
// Full LLM dispatch: metadata reconstructed as typed options.
|
||||
prompt, err := chatprompt.ConvertMessagesWithFiles(
|
||||
context.Background(),
|
||||
[]database.ChatMessage{{
|
||||
Role: "assistant",
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
Content: raw,
|
||||
}},
|
||||
nil,
|
||||
slogtest.Make(t, nil),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, prompt, 1)
|
||||
require.Len(t, prompt[0].Content, 1)
|
||||
|
||||
textPart, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[0].Content[0])
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "response with cache hints", textPart.Text)
|
||||
|
||||
cc := fantasyanthropic.GetCacheControl(textPart.ProviderOptions)
|
||||
require.NotNil(t, cc, "cache control must survive new write → new read round-trip")
|
||||
require.Equal(t, "ephemeral", cc.Type)
|
||||
}
|
||||
|
||||
// TestMixedFormatConversation verifies ConvertMessagesWithFiles
|
||||
// handles a realistic post-deploy conversation where legacy and new
|
||||
// storage formats coexist.
|
||||
func TestMixedFormatConversation(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fileID := uuid.New()
|
||||
resolvedFileData := []byte("resolved-png-bytes")
|
||||
|
||||
resolver := func(_ context.Context, ids []uuid.UUID) (map[uuid.UUID]chatprompt.FileData, error) {
|
||||
out := make(map[uuid.UUID]chatprompt.FileData)
|
||||
for _, id := range ids {
|
||||
if id == fileID {
|
||||
out[id] = chatprompt.FileData{Data: resolvedFileData, MediaType: "image/png"}
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// 1. System (JSON string).
|
||||
systemRaw, err := json.Marshal("You are helpful.")
|
||||
require.NoError(t, err)
|
||||
|
||||
// 2. Old user (fantasy envelope: text + file with file_id).
|
||||
oldUserRaw := mustJSON(t, []json.RawMessage{
|
||||
mustJSON(t, map[string]any{
|
||||
"type": "text",
|
||||
"data": map[string]any{"text": "Look at this image."},
|
||||
}),
|
||||
mustJSON(t, map[string]any{
|
||||
"type": "file",
|
||||
"data": map[string]any{
|
||||
"media_type": "image/png",
|
||||
"file_id": fileID.String(),
|
||||
},
|
||||
}),
|
||||
})
|
||||
|
||||
// 3. Old assistant (fantasy envelope: tool-call).
|
||||
oldAssistantRaw, err := chatprompt.MarshalContent([]fantasy.Content{
|
||||
fantasy.ToolCallContent{
|
||||
ToolCallID: "call_1",
|
||||
ToolName: "analyze_image",
|
||||
Input: `{"detail":"high"}`,
|
||||
},
|
||||
}, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 4. Old tool (legacy result rows).
|
||||
oldToolRaw, err := chatprompt.MarshalToolResult(
|
||||
"call_1", "analyze_image",
|
||||
json.RawMessage(`{"description":"a cat"}`), false,
|
||||
false, nil,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 5. New user (SDK parts: text + file-reference).
|
||||
newUserRaw, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
|
||||
{Type: codersdk.ChatMessagePartTypeText, Text: "Check this diff."},
|
||||
{Type: codersdk.ChatMessagePartTypeFileReference, FileName: "main.go", StartLine: 5, EndLine: 15, Content: "func main() {}"},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// 6. New assistant (SDK parts: text with metadata).
|
||||
newAssistantMeta := json.RawMessage(`{"anthropic":{"type":"anthropic.cache_control_options","data":{"cache_control":{"type":"ephemeral"}}}}`)
|
||||
newAssistantRaw, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
|
||||
{Type: codersdk.ChatMessagePartTypeText, Text: "Here is my analysis.", ProviderMetadata: newAssistantMeta},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
messages := []database.ChatMessage{
|
||||
{Role: "system", Visibility: database.ChatMessageVisibilityModel, Content: pqtype.NullRawMessage{RawMessage: systemRaw, Valid: true}},
|
||||
{Role: "user", Visibility: database.ChatMessageVisibilityBoth, Content: pqtype.NullRawMessage{RawMessage: oldUserRaw, Valid: true}},
|
||||
{Role: "assistant", Visibility: database.ChatMessageVisibilityBoth, Content: oldAssistantRaw},
|
||||
{Role: "tool", Visibility: database.ChatMessageVisibilityBoth, Content: oldToolRaw},
|
||||
{Role: "user", Visibility: database.ChatMessageVisibilityBoth, Content: newUserRaw},
|
||||
{Role: "assistant", Visibility: database.ChatMessageVisibilityBoth, Content: newAssistantRaw},
|
||||
}
|
||||
|
||||
prompt, err := chatprompt.ConvertMessagesWithFiles(
|
||||
context.Background(), messages, resolver, slogtest.Make(t, nil),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, prompt, 6, "all 6 messages should produce prompt entries")
|
||||
|
||||
// 1. System.
|
||||
require.Equal(t, fantasy.MessageRoleSystem, prompt[0].Role)
|
||||
systemText, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[0].Content[0])
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "You are helpful.", systemText.Text)
|
||||
|
||||
// 2. Old user: text + file with resolved data.
|
||||
require.Equal(t, fantasy.MessageRoleUser, prompt[1].Role)
|
||||
require.Len(t, prompt[1].Content, 2)
|
||||
userText, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[1].Content[0])
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "Look at this image.", userText.Text)
|
||||
filePart, ok := fantasy.AsMessagePart[fantasy.FilePart](prompt[1].Content[1])
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, resolvedFileData, filePart.Data)
|
||||
assert.Equal(t, "image/png", filePart.MediaType)
|
||||
|
||||
// 3. Old assistant: tool-call with normalized input.
|
||||
require.Equal(t, fantasy.MessageRoleAssistant, prompt[2].Role)
|
||||
toolCalls := chatprompt.ExtractToolCalls(prompt[2].Content)
|
||||
require.Len(t, toolCalls, 1)
|
||||
assert.Equal(t, "call_1", toolCalls[0].ToolCallID)
|
||||
assert.Equal(t, "analyze_image", toolCalls[0].ToolName)
|
||||
assert.JSONEq(t, `{"detail":"high"}`, toolCalls[0].Input)
|
||||
|
||||
// 4. Old tool: result paired with call_1.
|
||||
require.Equal(t, fantasy.MessageRoleTool, prompt[3].Role)
|
||||
require.Len(t, prompt[3].Content, 1)
|
||||
toolResult, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[3].Content[0])
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "call_1", toolResult.ToolCallID)
|
||||
|
||||
// 5. New user: text + file-reference (converted to TextPart).
|
||||
require.Equal(t, fantasy.MessageRoleUser, prompt[4].Role)
|
||||
require.Len(t, prompt[4].Content, 2)
|
||||
newUserText, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[4].Content[0])
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "Check this diff.", newUserText.Text)
|
||||
refText, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[4].Content[1])
|
||||
require.True(t, ok)
|
||||
assert.Contains(t, refText.Text, "[file-reference]")
|
||||
assert.Contains(t, refText.Text, "main.go")
|
||||
|
||||
// 6. New assistant: text with ProviderMetadata → ProviderOptions.
|
||||
require.Equal(t, fantasy.MessageRoleAssistant, prompt[5].Role)
|
||||
require.Len(t, prompt[5].Content, 1)
|
||||
newAssistantText, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[5].Content[0])
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "Here is my analysis.", newAssistantText.Text)
|
||||
cc := fantasyanthropic.GetCacheControl(newAssistantText.ProviderOptions)
|
||||
require.NotNil(t, cc, "ProviderMetadata must survive on new-format assistant messages")
|
||||
assert.Equal(t, "ephemeral", cc.Type)
|
||||
}
|
||||
|
||||
// TestQueuedMessageRoundTrip verifies that a user message with
|
||||
// file-reference parts survives the queue → promote cycle. The
|
||||
// queued path stores MarshalParts output as raw JSON in
|
||||
// chat_queued_messages, db2sdk.ChatQueuedMessage parses it for
|
||||
// display while queued, then PromoteQueued copies the same raw
|
||||
// bytes into chat_messages where ParseContent reads them.
|
||||
func TestQueuedMessageRoundTrip(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Simulate the write path: user sends a message with text +
|
||||
// file-reference, which gets queued.
|
||||
parts := []codersdk.ChatMessagePart{
|
||||
{Type: codersdk.ChatMessagePartTypeText, Text: "Review this change."},
|
||||
{Type: codersdk.ChatMessagePartTypeFileReference, FileName: "api.go", StartLine: 42, EndLine: 58, Content: "func handleRequest() {}"},
|
||||
}
|
||||
raw, err := chatprompt.MarshalParts(parts)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Step 1: While queued, db2sdk.ChatQueuedMessage parses the
|
||||
// content for display. Verify it produces correct parts
|
||||
// (with internal fields stripped).
|
||||
queuedMsg := db2sdk.ChatQueuedMessage(database.ChatQueuedMessage{
|
||||
ID: 1,
|
||||
ChatID: uuid.New(),
|
||||
Content: raw.RawMessage,
|
||||
})
|
||||
require.Len(t, queuedMsg.Content, 2)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeText, queuedMsg.Content[0].Type)
|
||||
assert.Equal(t, "Review this change.", queuedMsg.Content[0].Text)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeFileReference, queuedMsg.Content[1].Type)
|
||||
assert.Equal(t, "api.go", queuedMsg.Content[1].FileName)
|
||||
assert.Equal(t, 42, queuedMsg.Content[1].StartLine)
|
||||
assert.Equal(t, 58, queuedMsg.Content[1].EndLine)
|
||||
assert.Equal(t, "func handleRequest() {}", queuedMsg.Content[1].Content)
|
||||
|
||||
// Step 2: PromoteQueued copies the raw bytes into
|
||||
// chat_messages. ParseContent must handle them identically.
|
||||
promoted, err := chatprompt.ParseContent(codersdk.ChatMessageRoleUser, pqtype.NullRawMessage{
|
||||
RawMessage: raw.RawMessage,
|
||||
Valid: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, promoted, 2)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeText, promoted[0].Type)
|
||||
assert.Equal(t, "Review this change.", promoted[0].Text)
|
||||
assert.Equal(t, codersdk.ChatMessagePartTypeFileReference, promoted[1].Type)
|
||||
assert.Equal(t, "api.go", promoted[1].FileName)
|
||||
assert.Equal(t, 42, promoted[1].StartLine)
|
||||
assert.Equal(t, 58, promoted[1].EndLine)
|
||||
assert.Equal(t, "func handleRequest() {}", promoted[1].Content)
|
||||
|
||||
// Step 3: The promoted message is used for LLM dispatch.
|
||||
// File-reference becomes a TextPart.
|
||||
prompt, err := chatprompt.ConvertMessagesWithFiles(
|
||||
context.Background(),
|
||||
[]database.ChatMessage{{
|
||||
Role: "user",
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
Content: pqtype.NullRawMessage{RawMessage: raw.RawMessage, Valid: true},
|
||||
}},
|
||||
nil,
|
||||
slogtest.Make(t, nil),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, prompt, 1)
|
||||
require.Len(t, prompt[0].Content, 2)
|
||||
|
||||
textPart, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[0].Content[0])
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "Review this change.", textPart.Text)
|
||||
|
||||
refPart, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[0].Content[1])
|
||||
require.True(t, ok)
|
||||
assert.Contains(t, refPart.Text, "[file-reference]")
|
||||
assert.Contains(t, refPart.Text, "api.go")
|
||||
}
|
||||
|
||||
func TestParseContent_ErrorPaths(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("null_content_returns_nil", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
parts, err := chatprompt.ParseContent(codersdk.ChatMessageRoleUser, pqtype.NullRawMessage{})
|
||||
require.NoError(t, err)
|
||||
assert.Nil(t, parts)
|
||||
})
|
||||
|
||||
t.Run("empty_content_returns_nil", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
parts, err := chatprompt.ParseContent(codersdk.ChatMessageRoleAssistant, pqtype.NullRawMessage{
|
||||
RawMessage: []byte{},
|
||||
Valid: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Nil(t, parts)
|
||||
})
|
||||
|
||||
t.Run("unknown_role", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, err := chatprompt.ParseContent(codersdk.ChatMessageRole("banana"), nullRaw(json.RawMessage(`"hello"`)))
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "unsupported chat message role")
|
||||
})
|
||||
|
||||
t.Run("system/malformed_json", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, err := chatprompt.ParseContent(codersdk.ChatMessageRoleSystem, nullRaw(json.RawMessage(`not json`)))
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "parse system content")
|
||||
})
|
||||
|
||||
t.Run("user/malformed_json", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, err := chatprompt.ParseContent(codersdk.ChatMessageRoleUser, nullRaw(json.RawMessage(`{not json`)))
|
||||
require.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("assistant/malformed_json", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, err := chatprompt.ParseContent(codersdk.ChatMessageRoleAssistant, nullRaw(json.RawMessage(`{not json`)))
|
||||
require.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("tool/malformed_json", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, err := chatprompt.ParseContent(codersdk.ChatMessageRoleTool, nullRaw(json.RawMessage(`{not json`)))
|
||||
require.Error(t, err)
|
||||
})
|
||||
}
|
||||
|
||||
func mustJSON(t *testing.T, v any) json.RawMessage {
|
||||
t.Helper()
|
||||
data, err := json.Marshal(v)
|
||||
|
||||
+20
-14
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user