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
|
||||
}
|
||||
|
||||
+25
-34
@@ -18,7 +18,6 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"github.com/go-chi/chi/v5"
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
@@ -222,7 +221,7 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
contentBlocks, contentFileIDs, titleSource, inputError := createChatInputFromRequest(ctx, api.Database, req)
|
||||
contentBlocks, titleSource, inputError := createChatInputFromRequest(ctx, api.Database, req)
|
||||
if inputError != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, *inputError)
|
||||
return
|
||||
@@ -257,7 +256,6 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
|
||||
ModelConfigID: modelConfigID,
|
||||
SystemPrompt: api.resolvedChatSystemPrompt(ctx),
|
||||
InitialUserContent: contentBlocks,
|
||||
ContentFileIDs: contentFileIDs,
|
||||
})
|
||||
if err != nil {
|
||||
if database.IsForeignKeyViolation(
|
||||
@@ -630,7 +628,7 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
contentBlocks, contentFileIDs, _, inputError := createChatInputFromParts(ctx, api.Database, req.Content, "content")
|
||||
contentBlocks, _, inputError := createChatInputFromParts(ctx, api.Database, req.Content, "content")
|
||||
if inputError != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: inputError.Message,
|
||||
@@ -642,12 +640,11 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
|
||||
sendResult, sendErr := api.chatDaemon.SendMessage(
|
||||
ctx,
|
||||
chatd.SendMessageOptions{
|
||||
ChatID: chatID,
|
||||
CreatedBy: apiKey.UserID,
|
||||
Content: contentBlocks,
|
||||
ContentFileIDs: contentFileIDs,
|
||||
ModelConfigID: req.ModelConfigID,
|
||||
BusyBehavior: chatd.SendMessageBusyBehaviorQueue,
|
||||
ChatID: chatID,
|
||||
CreatedBy: apiKey.UserID,
|
||||
Content: contentBlocks,
|
||||
ModelConfigID: req.ModelConfigID,
|
||||
BusyBehavior: chatd.SendMessageBusyBehaviorQueue,
|
||||
},
|
||||
)
|
||||
if sendErr != nil {
|
||||
@@ -707,7 +704,7 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
contentBlocks, contentFileIDs, _, inputError := createChatInputFromParts(ctx, api.Database, req.Content, "content")
|
||||
contentBlocks, _, inputError := createChatInputFromParts(ctx, api.Database, req.Content, "content")
|
||||
if inputError != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: inputError.Message,
|
||||
@@ -721,7 +718,6 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) {
|
||||
CreatedBy: apiKey.UserID,
|
||||
EditedMessageID: messageID,
|
||||
Content: contentBlocks,
|
||||
ContentFileIDs: contentFileIDs,
|
||||
})
|
||||
if editErr != nil {
|
||||
switch {
|
||||
@@ -1985,8 +1981,7 @@ func (api *API) chatFileByID(rw http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
func createChatInputFromRequest(ctx context.Context, db database.Store, req codersdk.CreateChatRequest) (
|
||||
[]fantasy.Content,
|
||||
map[int]uuid.UUID,
|
||||
[]codersdk.ChatMessagePart,
|
||||
string,
|
||||
*codersdk.Response,
|
||||
) {
|
||||
@@ -1998,32 +1993,31 @@ func createChatInputFromParts(
|
||||
db database.Store,
|
||||
parts []codersdk.ChatInputPart,
|
||||
fieldName string,
|
||||
) ([]fantasy.Content, map[int]uuid.UUID, string, *codersdk.Response) {
|
||||
) ([]codersdk.ChatMessagePart, string, *codersdk.Response) {
|
||||
if len(parts) == 0 {
|
||||
return nil, nil, "", &codersdk.Response{
|
||||
return nil, "", &codersdk.Response{
|
||||
Message: "Content is required.",
|
||||
Detail: "Content cannot be empty.",
|
||||
}
|
||||
}
|
||||
|
||||
content := make([]fantasy.Content, 0, len(parts))
|
||||
fileIDs := make(map[int]uuid.UUID)
|
||||
content := make([]codersdk.ChatMessagePart, 0, len(parts))
|
||||
textParts := make([]string, 0, len(parts))
|
||||
for i, part := range parts {
|
||||
switch strings.ToLower(strings.TrimSpace(string(part.Type))) {
|
||||
case string(codersdk.ChatInputPartTypeText):
|
||||
text := strings.TrimSpace(part.Text)
|
||||
if text == "" {
|
||||
return nil, nil, "", &codersdk.Response{
|
||||
return nil, "", &codersdk.Response{
|
||||
Message: "Invalid input part.",
|
||||
Detail: fmt.Sprintf("%s[%d].text cannot be empty.", fieldName, i),
|
||||
}
|
||||
}
|
||||
content = append(content, fantasy.TextContent{Text: text})
|
||||
content = append(content, codersdk.ChatMessageText(text))
|
||||
textParts = append(textParts, text)
|
||||
case string(codersdk.ChatInputPartTypeFile):
|
||||
if part.FileID == uuid.Nil {
|
||||
return nil, nil, "", &codersdk.Response{
|
||||
return nil, "", &codersdk.Response{
|
||||
Message: "Invalid input part.",
|
||||
Detail: fmt.Sprintf("%s[%d].file_id is required for file parts.", fieldName, i),
|
||||
}
|
||||
@@ -2034,27 +2028,26 @@ func createChatInputFromParts(
|
||||
chatFile, err := db.GetChatFileByID(ctx, part.FileID)
|
||||
if err != nil {
|
||||
if httpapi.Is404Error(err) {
|
||||
return nil, nil, "", &codersdk.Response{
|
||||
return nil, "", &codersdk.Response{
|
||||
Message: "Invalid input part.",
|
||||
Detail: fmt.Sprintf("%s[%d].file_id references a file that does not exist.", fieldName, i),
|
||||
}
|
||||
}
|
||||
return nil, nil, "", &codersdk.Response{
|
||||
return nil, "", &codersdk.Response{
|
||||
Message: "Internal error.",
|
||||
Detail: fmt.Sprintf("Failed to retrieve file for %s[%d].", fieldName, i),
|
||||
}
|
||||
}
|
||||
content = append(content, fantasy.FileContent{
|
||||
MediaType: chatFile.Mimetype,
|
||||
})
|
||||
fileIDs[len(content)-1] = part.FileID
|
||||
content = append(content, codersdk.ChatMessageFile(part.FileID, chatFile.Mimetype))
|
||||
case string(codersdk.ChatInputPartTypeFileReference):
|
||||
if part.FileName == "" {
|
||||
return nil, nil, "", &codersdk.Response{
|
||||
return nil, "", &codersdk.Response{
|
||||
Message: "Invalid input part.",
|
||||
Detail: fmt.Sprintf("%s[%d].file_name cannot be empty for file-reference.", fieldName, i),
|
||||
}
|
||||
}
|
||||
content = append(content, codersdk.ChatMessageFileReference(part.FileName, part.StartLine, part.EndLine, part.Content))
|
||||
// Build text representation for title generation.
|
||||
lineRange := fmt.Sprintf("%d", part.StartLine)
|
||||
if part.StartLine != part.EndLine {
|
||||
lineRange = fmt.Sprintf("%d-%d", part.StartLine, part.EndLine)
|
||||
@@ -2064,11 +2057,9 @@ func createChatInputFromParts(
|
||||
if strings.TrimSpace(part.Content) != "" {
|
||||
_, _ = fmt.Fprintf(&sb, "\n```%s\n%s\n```", part.FileName, strings.TrimSpace(part.Content))
|
||||
}
|
||||
text := sb.String()
|
||||
content = append(content, fantasy.TextContent{Text: text})
|
||||
textParts = append(textParts, text)
|
||||
textParts = append(textParts, sb.String())
|
||||
default:
|
||||
return nil, nil, "", &codersdk.Response{
|
||||
return nil, "", &codersdk.Response{
|
||||
Message: "Invalid input part.",
|
||||
Detail: fmt.Sprintf(
|
||||
"%s[%d].type %q is not supported.",
|
||||
@@ -2083,13 +2074,13 @@ func createChatInputFromParts(
|
||||
// Allow file-only messages. The titleSource may be empty
|
||||
// when only file parts are provided, callers handle this.
|
||||
if len(content) == 0 {
|
||||
return nil, nil, "", &codersdk.Response{
|
||||
return nil, "", &codersdk.Response{
|
||||
Message: "Content is required.",
|
||||
Detail: fmt.Sprintf("%s must include at least one text or file part.", fieldName),
|
||||
}
|
||||
}
|
||||
titleSource := strings.TrimSpace(strings.Join(textParts, " "))
|
||||
return content, fileIDs, titleSource, nil
|
||||
return content, titleSource, nil
|
||||
}
|
||||
|
||||
func chatTitleFromMessage(message string) string {
|
||||
|
||||
+92
-58
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user