mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat(site): display file attachments in chat UI (#24281)
Renders the durable file attachments introduced in #24280 in the chat interface. Without this, attachments were stored and served correctly but the UI showed raw file parts with no previews or download UX. Every attachment gets a download affordance, split into three rendering tiers: - **Images** — thumbnail with a hover/focus overlay containing a download link. `onFocusCapture`/`onBlurCapture` with `contains(relatedTarget)` keeps the overlay open while tabbing between the image and its download link. - **Text-like files** (`text/*`, `application/json`) — expandable preview button with loading + error-with-retry states and the same download overlay. Preview fetches throw a typed `FetchTextAttachmentError` with a `.status` field instead of a stringly-typed error. - **Everything else** — compact `FileCard` with extension badge, filename, and download link. User-side and assistant-side rendering now share `AttachmentBlocks.tsx` (`AttachmentPreviewFrame`, `TextAttachmentButton`, `ImageAttachmentButton`, `FileCard`, plus `getAttachmentHref`/`getAttachmentName`) instead of two near-duplicate implementations. The text-attachment overlay anchors to the preview surface so the download button stays pinned even when a loading/error status line widens the row below. `ComputerRenderer` detects when a screenshot was stored as a durable attachment (`attachment_file_id`) and suppresses the stale base64 rendering — the screenshot appears as a proper file part instead. `ToolLabel` shows the attached filename for `attach_file` tool calls. Storybook coverage in `ConversationTimeline.stories.tsx` was expanded to cover every tier (single/multiple images, inline + file-id text, JSON, download-only files, fetch-failure retry, mixed attachments + file references) with play-function assertions. <img width="811" height="150" alt="image" src="https://github.com/user-attachments/assets/27c71081-3502-4e80-92a7-d8adf1ff9323" /> ## Cleanup Per Mathias' post-merge suggestion on #24280, this PR also relocates `coderd/chatfiles` → `coderd/x/chatfiles` so the durable-attachment helpers live beside the rest of the `chatd` experimental surface. Closes CODAGT-91
This commit is contained in:
@@ -23,7 +23,7 @@ func buildAssistantPartsForPersist(
|
||||
) []codersdk.ChatMessagePart {
|
||||
parts := make([]codersdk.ChatMessagePart, 0, len(assistantBlocks)+len(toolResults))
|
||||
for _, block := range assistantBlocks {
|
||||
part := chatprompt.PartFromContent(block)
|
||||
part := chatprompt.PartFromContentWithLogger(ctx, logger, block)
|
||||
if part.ToolName != "" {
|
||||
if configID, ok := toolNameToConfigID[part.ToolName]; ok {
|
||||
part.MCPServerConfigID = uuid.NullUUID{UUID: configID, Valid: true}
|
||||
|
||||
@@ -6015,7 +6015,7 @@ func (p *Server) runChat(
|
||||
// FOR UPDATE lock is held only for the INSERT statements.
|
||||
// Marshaling is pure CPU work with no database dependency.
|
||||
assistantParts := buildAssistantPartsForPersist(
|
||||
ctx,
|
||||
persistCtx,
|
||||
p.logger,
|
||||
assistantBlocks,
|
||||
toolResults,
|
||||
@@ -6035,7 +6035,7 @@ func (p *Server) runChat(
|
||||
|
||||
toolResultContents := make([]pqtype.NullRawMessage, len(toolResults))
|
||||
for i, tr := range toolResults {
|
||||
trPart := chatprompt.PartFromContent(tr)
|
||||
trPart := chatprompt.PartFromContentWithLogger(ctx, logger, tr)
|
||||
if trPart.ToolName != "" {
|
||||
if configID, ok := toolNameToConfigID[trPart.ToolName]; ok {
|
||||
trPart.MCPServerConfigID = uuid.NullUUID{UUID: configID, Valid: true}
|
||||
@@ -6496,6 +6496,7 @@ func (p *Server) runChat(
|
||||
}
|
||||
p.publishMessagePart(chat.ID, role, part)
|
||||
},
|
||||
Logger: logger,
|
||||
Compaction: compactionOptions,
|
||||
ReloadMessages: func(reloadCtx context.Context) ([]fantasy.Message, error) {
|
||||
reloadedMsgs, err := p.db.GetChatMessagesForPromptByChatID(reloadCtx, chat.ID)
|
||||
|
||||
@@ -19,11 +19,13 @@ import (
|
||||
"charm.land/fantasy/schema"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog/v3"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtime"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatdebug"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chaterror"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatretry"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
@@ -146,6 +148,7 @@ type RunOptions struct {
|
||||
role codersdk.ChatMessageRole,
|
||||
part codersdk.ChatMessagePart,
|
||||
)
|
||||
Logger slog.Logger
|
||||
Compaction *CompactionOptions
|
||||
ReloadMessages func(context.Context) ([]fantasy.Message, error)
|
||||
DisableChainMode func()
|
||||
@@ -492,7 +495,8 @@ func Run(ctx context.Context, opts RunOptions) error {
|
||||
// Execute only built-in tools.
|
||||
toolResults = executeTools(ctx, opts.Tools, opts.ActiveTools, opts.ProviderTools, builtinCalls, opts.Metrics, provider, modelName, opts.BuiltinToolNames, func(tr fantasy.ToolResultContent, completedAt time.Time) {
|
||||
recordToolResultTimestamp(&result, tr.ToolCallID, completedAt)
|
||||
ssePart := chatprompt.PartFromContent(tr)
|
||||
publishToolAttachments(ctx, opts.Logger, tr, completedAt, publishMessagePart)
|
||||
ssePart := chatprompt.PartFromContentWithLogger(ctx, opts.Logger, tr)
|
||||
ssePart.CreatedAt = &completedAt
|
||||
publishMessagePart(codersdk.ChatMessageRoleTool, ssePart)
|
||||
})
|
||||
@@ -1545,6 +1549,33 @@ func recordToolResultTimestamp(result *stepResult, toolCallID string, ts time.Ti
|
||||
result.toolResultCreatedAt[toolCallID] = ts
|
||||
}
|
||||
|
||||
func publishToolAttachments(
|
||||
ctx context.Context,
|
||||
logger slog.Logger,
|
||||
tr fantasy.ToolResultContent,
|
||||
createdAt time.Time,
|
||||
publishMessagePart func(codersdk.ChatMessageRole, codersdk.ChatMessagePart),
|
||||
) {
|
||||
attachments, err := chattool.AttachmentsFromMetadata(tr.ClientMetadata)
|
||||
if err != nil {
|
||||
logger.Warn(ctx, "skipping malformed tool attachment metadata",
|
||||
slog.F("tool_name", tr.ToolName),
|
||||
slog.F("tool_call_id", tr.ToolCallID),
|
||||
slog.Error(err),
|
||||
)
|
||||
return
|
||||
}
|
||||
for _, attachment := range attachments {
|
||||
filePart := codersdk.ChatMessageFile(
|
||||
attachment.FileID,
|
||||
attachment.MediaType,
|
||||
attachment.Name,
|
||||
)
|
||||
filePart.CreatedAt = &createdAt
|
||||
publishMessagePart(codersdk.ChatMessageRoleAssistant, filePart)
|
||||
}
|
||||
}
|
||||
|
||||
func extractContextLimit(metadata fantasy.ProviderMetadata) sql.NullInt64 {
|
||||
if len(metadata) == 0 {
|
||||
return sql.NullInt64{}
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
|
||||
"cdr.dev/slog/v3"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
@@ -702,6 +703,29 @@ func MarshalToolResult(toolCallID, toolName string, result json.RawMessage, isEr
|
||||
// PartFromContent converts fantasy content into a SDK chat message
|
||||
// part, preserving ProviderMetadata and ProviderExecuted fields.
|
||||
func PartFromContent(block fantasy.Content) codersdk.ChatMessagePart {
|
||||
return sdkPartFromContent(block, nil)
|
||||
}
|
||||
|
||||
// PartFromContentWithLogger is for call sites that can surface malformed
|
||||
// attachment metadata immediately instead of dropping it silently.
|
||||
func PartFromContentWithLogger(
|
||||
ctx context.Context,
|
||||
logger slog.Logger,
|
||||
block fantasy.Content,
|
||||
) codersdk.ChatMessagePart {
|
||||
return sdkPartFromContent(block, func(content fantasy.ToolResultContent, err error) {
|
||||
logger.Warn(ctx, "skipping malformed tool attachment metadata",
|
||||
slog.F("tool_name", content.ToolName),
|
||||
slog.F("tool_call_id", content.ToolCallID),
|
||||
slog.Error(err),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
func sdkPartFromContent(
|
||||
block fantasy.Content,
|
||||
logMalformedAttachmentMetadata func(fantasy.ToolResultContent, error),
|
||||
) codersdk.ChatMessagePart {
|
||||
switch value := block.(type) {
|
||||
case fantasy.TextContent:
|
||||
return codersdk.ChatMessagePart{
|
||||
@@ -776,9 +800,9 @@ func PartFromContent(block fantasy.Content) codersdk.ChatMessagePart {
|
||||
ProviderMetadata: marshalProviderMetadata(value.ProviderMetadata),
|
||||
}
|
||||
case fantasy.ToolResultContent:
|
||||
return toolResultContentToPart(value)
|
||||
return toolResultContentToPart(value, logMalformedAttachmentMetadata)
|
||||
case *fantasy.ToolResultContent:
|
||||
return toolResultContentToPart(*value)
|
||||
return toolResultContentToPart(*value, logMalformedAttachmentMetadata)
|
||||
default:
|
||||
return codersdk.ChatMessagePart{}
|
||||
}
|
||||
@@ -794,7 +818,10 @@ func ToolResultToPart(toolCallID, toolName string, result json.RawMessage, isErr
|
||||
|
||||
// toolResultContentToPart converts a fantasy ToolResultContent into a
|
||||
// ChatMessagePart.
|
||||
func toolResultContentToPart(content fantasy.ToolResultContent) codersdk.ChatMessagePart {
|
||||
func toolResultContentToPart(
|
||||
content fantasy.ToolResultContent,
|
||||
logMalformedAttachmentMetadata func(fantasy.ToolResultContent, error),
|
||||
) codersdk.ChatMessagePart {
|
||||
var result json.RawMessage
|
||||
var isError bool
|
||||
var isMedia bool
|
||||
@@ -820,11 +847,24 @@ func toolResultContentToPart(content fantasy.ToolResultContent) codersdk.ChatMes
|
||||
}
|
||||
case fantasy.ToolResultOutputContentMedia:
|
||||
isMedia = true
|
||||
result, _ = json.Marshal(persistedMediaResult{
|
||||
persisted := persistedMediaResult{
|
||||
Data: output.Data,
|
||||
MimeType: output.MediaType,
|
||||
Text: output.Text,
|
||||
})
|
||||
}
|
||||
// Tool renderers only receive the persisted result JSON, while
|
||||
// ClientMetadata is consumed later to append sibling file parts.
|
||||
// Mirror attachment identity here so promoted media can be
|
||||
// recognized as the same durable attachment downstream.
|
||||
if attachment, ok := matchingAttachmentForMedia(
|
||||
content,
|
||||
output.MediaType,
|
||||
logMalformedAttachmentMetadata,
|
||||
); ok {
|
||||
persisted.AttachmentFileID = attachment.FileID.String()
|
||||
persisted.AttachmentName = attachment.Name
|
||||
}
|
||||
result, _ = json.Marshal(persisted)
|
||||
default:
|
||||
result = []byte(`{}`)
|
||||
}
|
||||
@@ -835,6 +875,26 @@ func toolResultContentToPart(content fantasy.ToolResultContent) codersdk.ChatMes
|
||||
return part
|
||||
}
|
||||
|
||||
func matchingAttachmentForMedia(
|
||||
content fantasy.ToolResultContent,
|
||||
mediaType string,
|
||||
logMalformedAttachmentMetadata func(fantasy.ToolResultContent, error),
|
||||
) (chattool.AttachmentMetadata, bool) {
|
||||
attachments, err := chattool.AttachmentsFromMetadata(content.ClientMetadata)
|
||||
if err != nil {
|
||||
if logMalformedAttachmentMetadata != nil {
|
||||
logMalformedAttachmentMetadata(content, err)
|
||||
}
|
||||
return chattool.AttachmentMetadata{}, false
|
||||
}
|
||||
for _, attachment := range attachments {
|
||||
if attachment.MediaType == mediaType {
|
||||
return attachment, true
|
||||
}
|
||||
}
|
||||
return chattool.AttachmentMetadata{}, false
|
||||
}
|
||||
|
||||
// Keep in sync with coderd/x/chatd/subagent.go.
|
||||
func isSubagentLifecycleToolName(name string) bool {
|
||||
switch name {
|
||||
@@ -1267,10 +1327,11 @@ func toolResultPartToMessagePart(logger slog.Logger, part codersdk.ChatMessagePa
|
||||
// IsError takes precedence and is handled above.
|
||||
// Detect media content flagged by toolResultContentToPart.
|
||||
// Screenshots from the computer use tool are stored as
|
||||
// {"data":"<base64>","mime_type":"image/png","text":"..."}.
|
||||
// Without this detection, the entire base64 payload is sent
|
||||
// as text tokens, which quickly exceeds the context limit
|
||||
// on follow-up messages.
|
||||
// {"data":"<base64>","mime_type":"image/png","text":"..."}
|
||||
// with optional attachment identity fields when the same image
|
||||
// was also promoted into a durable file part. Without this
|
||||
// detection, the entire base64 payload is sent as text tokens,
|
||||
// which quickly exceeds the context limit on follow-up messages.
|
||||
if part.IsMedia {
|
||||
var media persistedMediaResult
|
||||
unmarshalErr := json.Unmarshal(part.Result, &media)
|
||||
@@ -1319,12 +1380,17 @@ func toolResultPartToMessagePart(logger slog.Logger, part codersdk.ChatMessagePa
|
||||
// cannot drift.
|
||||
//
|
||||
// The "mime_type" key intentionally diverges from the fantasy
|
||||
// struct tag (json:"media_type"). Do not change it without
|
||||
// updating both paths.
|
||||
// struct tag (json:"media_type"). Optional attachment identity
|
||||
// fields are UI hints only. They let the frontend recognize when the
|
||||
// same media was also promoted into a durable file part, but the prompt
|
||||
// reconstruction path must continue to ignore them. Keep additions
|
||||
// backwards-compatible because existing rows may omit these fields.
|
||||
type persistedMediaResult struct {
|
||||
Data string `json:"data"`
|
||||
MimeType string `json:"mime_type"`
|
||||
Text string `json:"text"`
|
||||
Data string `json:"data"`
|
||||
MimeType string `json:"mime_type"`
|
||||
Text string `json:"text"`
|
||||
AttachmentFileID string `json:"attachment_file_id,omitempty"`
|
||||
AttachmentName string `json:"attachment_name,omitempty"`
|
||||
}
|
||||
|
||||
type missingFilePolicy uint8
|
||||
|
||||
@@ -21,6 +21,7 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
@@ -2384,6 +2385,103 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
|
||||
require.Equal(t, mimeType, mediaOutput.MediaType)
|
||||
})
|
||||
|
||||
t.Run("MediaResultCarriesPromotedAttachmentMetadata", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const callID = "call-screenshot-promoted"
|
||||
const toolName = "computer"
|
||||
const mimeType = "image/png"
|
||||
const attachmentName = "screenshot-2026-04-21T00-00-00Z.png"
|
||||
|
||||
attachmentID := uuid.MustParse("aaaaaaaa-bbbb-cccc-dddd-eeeeeeeeeeee")
|
||||
response := chattool.WithAttachments(
|
||||
fantasy.NewImageResponse([]byte(imageData), mimeType),
|
||||
chattool.AttachmentMetadata{
|
||||
FileID: attachmentID,
|
||||
MediaType: mimeType,
|
||||
Name: attachmentName,
|
||||
},
|
||||
)
|
||||
|
||||
sdkPart := chatprompt.PartFromContent(fantasy.ToolResultContent{
|
||||
ToolCallID: callID,
|
||||
ToolName: toolName,
|
||||
ClientMetadata: response.Metadata,
|
||||
Result: fantasy.ToolResultOutputContentMedia{
|
||||
Data: imageData,
|
||||
MediaType: mimeType,
|
||||
},
|
||||
})
|
||||
|
||||
var persisted struct {
|
||||
Data string `json:"data"`
|
||||
MimeType string `json:"mime_type"`
|
||||
Text string `json:"text"`
|
||||
AttachmentFileID string `json:"attachment_file_id"`
|
||||
AttachmentName string `json:"attachment_name"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal(sdkPart.Result, &persisted))
|
||||
require.Equal(t, imageData, persisted.Data)
|
||||
require.Equal(t, mimeType, persisted.MimeType)
|
||||
require.Equal(t, attachmentID.String(), persisted.AttachmentFileID)
|
||||
require.Equal(t, attachmentName, persisted.AttachmentName)
|
||||
|
||||
chat := insertPair(t, callID, toolName, []codersdk.ChatMessagePart{sdkPart})
|
||||
|
||||
prompt := loadPrompt(t, chat)
|
||||
require.Len(t, prompt, 2)
|
||||
|
||||
resultPart, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[1].Content[0])
|
||||
require.True(t, ok, "expected ToolResultPart")
|
||||
|
||||
mediaOutput, ok := fantasy.AsToolResultOutputType[fantasy.ToolResultOutputContentMedia](resultPart.Output)
|
||||
require.True(t, ok, "expected ToolResultOutputContentMedia, got %T", resultPart.Output)
|
||||
require.Equal(t, imageData, mediaOutput.Data)
|
||||
require.Equal(t, mimeType, mediaOutput.MediaType)
|
||||
})
|
||||
t.Run("MediaResultUsesMatchingAttachmentMetadata", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const callID = "call-screenshot-matching-attachment"
|
||||
const toolName = "computer"
|
||||
const mimeType = "image/png"
|
||||
const attachmentName = "screenshot-2026-04-21T00-00-01Z.png"
|
||||
|
||||
mismatchedAttachmentID := uuid.MustParse("11111111-2222-3333-4444-555555555555")
|
||||
matchingAttachmentID := uuid.MustParse("aaaaaaaa-bbbb-cccc-dddd-ffffffffffff")
|
||||
response := chattool.WithAttachments(
|
||||
fantasy.NewImageResponse([]byte(imageData), mimeType),
|
||||
chattool.AttachmentMetadata{
|
||||
FileID: mismatchedAttachmentID,
|
||||
MediaType: "application/pdf",
|
||||
Name: "report.pdf",
|
||||
},
|
||||
chattool.AttachmentMetadata{
|
||||
FileID: matchingAttachmentID,
|
||||
MediaType: mimeType,
|
||||
Name: attachmentName,
|
||||
},
|
||||
)
|
||||
|
||||
sdkPart := chatprompt.PartFromContent(fantasy.ToolResultContent{
|
||||
ToolCallID: callID,
|
||||
ToolName: toolName,
|
||||
ClientMetadata: response.Metadata,
|
||||
Result: fantasy.ToolResultOutputContentMedia{
|
||||
Data: imageData,
|
||||
MediaType: mimeType,
|
||||
},
|
||||
})
|
||||
|
||||
var persisted struct {
|
||||
AttachmentFileID string `json:"attachment_file_id"`
|
||||
AttachmentName string `json:"attachment_name"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal(sdkPart.Result, &persisted))
|
||||
require.Equal(t, matchingAttachmentID.String(), persisted.AttachmentFileID)
|
||||
require.Equal(t, attachmentName, persisted.AttachmentName)
|
||||
})
|
||||
|
||||
t.Run("MediaResultWithText", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -11,10 +11,10 @@ import (
|
||||
"github.com/google/uuid"
|
||||
|
||||
"cdr.dev/slog/v3"
|
||||
"github.com/coder/coder/v2/coderd/chatfiles"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbauthz"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
|
||||
"github.com/coder/coder/v2/coderd/x/chatfiles"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk"
|
||||
)
|
||||
|
||||
|
||||
@@ -6,9 +6,9 @@ import (
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/chatfiles"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
|
||||
"github.com/coder/coder/v2/coderd/x/chatfiles"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
|
||||
@@ -8,10 +8,10 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/chatfiles"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbmock"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
|
||||
"github.com/coder/coder/v2/coderd/x/chatfiles"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,250 @@
|
||||
package chatfiles
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"encoding/xml"
|
||||
"maps"
|
||||
"mime"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strings"
|
||||
"unicode"
|
||||
|
||||
"github.com/gabriel-vasile/mimetype"
|
||||
"golang.org/x/xerrors"
|
||||
)
|
||||
|
||||
const MaxStoredFileNameBytes = 255
|
||||
|
||||
var (
|
||||
// ErrStoredFileNameRequired indicates that a durable file name is empty
|
||||
// after normalization.
|
||||
ErrStoredFileNameRequired = xerrors.New("stored file name is required")
|
||||
|
||||
// ErrUnsupportedStoredFileType indicates that classified file bytes do not
|
||||
// map to an allowed durable file type.
|
||||
ErrUnsupportedStoredFileType = xerrors.New("unsupported attachment type")
|
||||
|
||||
utf8BOM = []byte{0xEF, 0xBB, 0xBF}
|
||||
|
||||
allowedStoredMediaTypes = map[string]struct{}{
|
||||
"image/png": {},
|
||||
"image/jpeg": {},
|
||||
"image/gif": {},
|
||||
"image/webp": {},
|
||||
"text/plain": {},
|
||||
"text/markdown": {},
|
||||
"text/csv": {},
|
||||
"application/json": {},
|
||||
"application/pdf": {},
|
||||
}
|
||||
|
||||
recordingArtifactMediaTypes = map[string]struct{}{
|
||||
"video/mp4": {},
|
||||
"image/jpeg": {},
|
||||
}
|
||||
)
|
||||
|
||||
// DetectMediaType detects the base media type of the given file contents.
|
||||
func DetectMediaType(data []byte) string {
|
||||
return BaseMediaType(mimetype.Detect(data).String())
|
||||
}
|
||||
|
||||
// BaseMediaType strips parameters from a media type.
|
||||
func BaseMediaType(mediaType string) string {
|
||||
if parsed, _, err := mime.ParseMediaType(mediaType); err == nil {
|
||||
return parsed
|
||||
}
|
||||
return mediaType
|
||||
}
|
||||
|
||||
// AllowedStoredMediaTypesString returns the supported durable chat file media
|
||||
// types as a comma-separated list.
|
||||
func AllowedStoredMediaTypesString() string {
|
||||
return strings.Join(slices.Sorted(maps.Keys(allowedStoredMediaTypes)), ", ")
|
||||
}
|
||||
|
||||
// IsAllowedStoredMediaType reports whether the media type is supported for
|
||||
// durable chat file storage.
|
||||
func IsAllowedStoredMediaType(mediaType string) bool {
|
||||
_, ok := allowedStoredMediaTypes[BaseMediaType(mediaType)]
|
||||
return ok
|
||||
}
|
||||
|
||||
// IsInlineRenderableStoredMediaType reports whether a stored chat file may be
|
||||
// served with Content-Disposition: inline. PDFs remain storable but
|
||||
// download-only because browser PDF viewers have a broader active-content
|
||||
// attack surface than the other media types we allow inline.
|
||||
func IsInlineRenderableStoredMediaType(mediaType string) bool {
|
||||
mediaType = BaseMediaType(mediaType)
|
||||
if !IsAllowedStoredMediaType(mediaType) {
|
||||
return false
|
||||
}
|
||||
return mediaType != "application/pdf"
|
||||
}
|
||||
|
||||
// NormalizeStoredFileName trims surrounding whitespace, strips control
|
||||
// characters, and truncates the name to the durable storage byte limit
|
||||
// without splitting UTF-8 runes.
|
||||
func NormalizeStoredFileName(name string) string {
|
||||
name = strings.Map(func(r rune) rune {
|
||||
if unicode.IsControl(r) {
|
||||
return -1
|
||||
}
|
||||
return r
|
||||
}, name)
|
||||
name = strings.TrimSpace(name)
|
||||
return truncateUTF8Bytes(name, MaxStoredFileNameBytes)
|
||||
}
|
||||
|
||||
// PrepareStoredFile normalizes the display name, rejects empty normalized
|
||||
// names, and classifies the file bytes using detectName when provided, so
|
||||
// callers can preserve subtype detection even when the user-facing filename is
|
||||
// overridden.
|
||||
func PrepareStoredFile(name, detectName string, data []byte) (storedName, mediaType string, err error) {
|
||||
storedName = NormalizeStoredFileName(name)
|
||||
if storedName == "" {
|
||||
return "", "", ErrStoredFileNameRequired
|
||||
}
|
||||
if strings.TrimSpace(detectName) == "" {
|
||||
detectName = storedName
|
||||
}
|
||||
mediaType = ClassifyStoredMediaType(detectName, data)
|
||||
if !IsAllowedStoredMediaType(mediaType) {
|
||||
return "", "", xerrors.Errorf("%w %q", ErrUnsupportedStoredFileType, mediaType)
|
||||
}
|
||||
return storedName, mediaType, nil
|
||||
}
|
||||
|
||||
// PrepareRecordingArtifact normalizes the recording artifact name, rejects
|
||||
// empty normalized names, and verifies that the bytes match the expected
|
||||
// recording media type.
|
||||
func PrepareRecordingArtifact(name, expectedMediaType string, data []byte) (storedName, mediaType string, err error) {
|
||||
expectedMediaType = BaseMediaType(expectedMediaType)
|
||||
if _, ok := recordingArtifactMediaTypes[expectedMediaType]; !ok {
|
||||
return "", "", xerrors.Errorf("unsupported recording artifact type %q", expectedMediaType)
|
||||
}
|
||||
|
||||
storedName = NormalizeStoredFileName(name)
|
||||
if storedName == "" {
|
||||
return "", "", ErrStoredFileNameRequired
|
||||
}
|
||||
mediaType = DetectMediaType(data)
|
||||
if mediaType != expectedMediaType {
|
||||
return "", "", xerrors.Errorf("recording artifact type mismatch: expected %q, detected %q", expectedMediaType, mediaType)
|
||||
}
|
||||
return storedName, mediaType, nil
|
||||
}
|
||||
|
||||
// IsCompatibleUploadMediaType reports whether an upload request that declared
|
||||
// declaredMediaType may be stored as storedMediaType after byte
|
||||
// classification. Exact matches are always compatible; the compatibility
|
||||
// table only covers explicit refinements like text/plain uploads that safely
|
||||
// store as richer text subtypes.
|
||||
func IsCompatibleUploadMediaType(declaredMediaType, storedMediaType string) bool {
|
||||
declaredMediaType = BaseMediaType(declaredMediaType)
|
||||
storedMediaType = BaseMediaType(storedMediaType)
|
||||
|
||||
if declaredMediaType == storedMediaType {
|
||||
return true
|
||||
}
|
||||
if declaredMediaType != "text/plain" {
|
||||
return false
|
||||
}
|
||||
|
||||
switch storedMediaType {
|
||||
case "text/markdown", "text/csv", "application/json":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// HasSVGRootElement reports whether the provided file bytes decode to an SVG
|
||||
// root element. This catches SVG content even when generic sniffers classify it
|
||||
// as text or XML.
|
||||
func HasSVGRootElement(data []byte) bool {
|
||||
data = bytes.TrimPrefix(data, utf8BOM)
|
||||
if len(data) == 0 {
|
||||
return false
|
||||
}
|
||||
|
||||
decoder := xml.NewDecoder(bytes.NewReader(data))
|
||||
for {
|
||||
token, err := decoder.Token()
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
|
||||
switch token := token.(type) {
|
||||
case xml.ProcInst, xml.Directive, xml.Comment:
|
||||
continue
|
||||
case xml.CharData:
|
||||
if len(bytes.TrimSpace(token)) == 0 {
|
||||
continue
|
||||
}
|
||||
return false
|
||||
case xml.StartElement:
|
||||
return strings.EqualFold(token.Name.Local, "svg")
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ClassifyStoredMediaType returns the media type that durable chat storage
|
||||
// would use for the given filename and bytes. Unsupported or blocked content is
|
||||
// returned as its detected media type so callers can report the specific type.
|
||||
func ClassifyStoredMediaType(name string, data []byte) string {
|
||||
if HasSVGRootElement(data) {
|
||||
return "image/svg+xml"
|
||||
}
|
||||
|
||||
mediaType := DetectMediaType(data)
|
||||
switch mediaType {
|
||||
case "image/png", "image/jpeg", "image/gif", "image/webp",
|
||||
"text/markdown", "text/csv", "application/json",
|
||||
"application/pdf", "application/xml", "text/xml":
|
||||
return mediaType
|
||||
case "text/plain":
|
||||
return refineTextMediaType(name, data)
|
||||
default:
|
||||
if strings.HasPrefix(mediaType, "text/") {
|
||||
return "text/plain"
|
||||
}
|
||||
return mediaType
|
||||
}
|
||||
}
|
||||
|
||||
func refineTextMediaType(name string, data []byte) string {
|
||||
switch strings.ToLower(filepath.Ext(name)) {
|
||||
case ".json":
|
||||
if json.Valid(data) {
|
||||
return "application/json"
|
||||
}
|
||||
case ".md", ".markdown":
|
||||
return "text/markdown"
|
||||
case ".csv":
|
||||
return "text/csv"
|
||||
}
|
||||
return "text/plain"
|
||||
}
|
||||
|
||||
func truncateUTF8Bytes(value string, maxBytes int) string {
|
||||
if maxBytes <= 0 || value == "" {
|
||||
return ""
|
||||
}
|
||||
if len(value) <= maxBytes {
|
||||
return value
|
||||
}
|
||||
|
||||
cut := 0
|
||||
for idx := range value {
|
||||
if idx > maxBytes {
|
||||
break
|
||||
}
|
||||
cut = idx
|
||||
}
|
||||
return value[:cut]
|
||||
}
|
||||
@@ -0,0 +1,345 @@
|
||||
package chatfiles_test
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/x/chatfiles"
|
||||
)
|
||||
|
||||
func TestDetectMediaType_WebP(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
data := append([]byte("RIFF"), []byte{0x24, 0x00, 0x00, 0x00}...)
|
||||
data = append(data, []byte("WEBPVP8 ")...)
|
||||
require.Equal(t, "image/webp", chatfiles.DetectMediaType(data))
|
||||
}
|
||||
|
||||
func TestClassifyStoredMediaType(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
fileName string
|
||||
data []byte
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "PlainText",
|
||||
fileName: "build.log",
|
||||
data: []byte("build succeeded\n"),
|
||||
want: "text/plain",
|
||||
},
|
||||
{
|
||||
name: "MarkdownFromExtension",
|
||||
fileName: "notes.md",
|
||||
data: []byte("# Release notes\n"),
|
||||
want: "text/markdown",
|
||||
},
|
||||
{
|
||||
name: "CSVFromDetector",
|
||||
fileName: "report.txt",
|
||||
data: []byte("name,count\nwidgets,3\n"),
|
||||
want: "text/csv",
|
||||
},
|
||||
{
|
||||
name: "JSONFromDetector",
|
||||
fileName: "payload.txt",
|
||||
data: []byte(`{"ok":true}`),
|
||||
want: "application/json",
|
||||
},
|
||||
{
|
||||
name: "UppercaseJSONExtension",
|
||||
fileName: "data.JSON",
|
||||
data: []byte(`{"ok":true}`),
|
||||
want: "application/json",
|
||||
},
|
||||
{
|
||||
name: "InvalidJSONExtensionFallsBackToPlainText",
|
||||
fileName: "broken.json",
|
||||
data: []byte("not json"),
|
||||
want: "text/plain",
|
||||
},
|
||||
{
|
||||
name: "UppercaseMDExtension",
|
||||
fileName: "NOTES.MD",
|
||||
data: []byte("# Notes\n"),
|
||||
want: "text/markdown",
|
||||
},
|
||||
{
|
||||
name: "PDF",
|
||||
fileName: "report.pdf",
|
||||
data: []byte("%PDF-1.7\n"),
|
||||
want: "application/pdf",
|
||||
},
|
||||
{
|
||||
name: "BinaryOctetStream",
|
||||
fileName: "data.bin",
|
||||
data: []byte{0x00, 0x01, 0x02, 0x03, 0x04, 0x05},
|
||||
want: "application/octet-stream",
|
||||
},
|
||||
{
|
||||
name: "HTMLFallsBackToTextPlain",
|
||||
fileName: "snippet.txt",
|
||||
data: []byte("<!DOCTYPE html><html><body>hello</body></html>"),
|
||||
want: "text/plain",
|
||||
},
|
||||
{
|
||||
name: "XMLStaysBlocked",
|
||||
fileName: "note.xml",
|
||||
data: []byte(`<?xml version="1.0"?><note><to>Tove</to></note>`),
|
||||
want: "text/xml",
|
||||
},
|
||||
{
|
||||
name: "SVGBlockedEvenWhenNamedText",
|
||||
fileName: "notes.txt",
|
||||
data: []byte(`<svg xmlns="http://www.w3.org/2000/svg"><text>Hello</text></svg>`),
|
||||
want: "image/svg+xml",
|
||||
},
|
||||
{
|
||||
name: "MarkdownMentioningSVGStaysMarkdown",
|
||||
fileName: "notes.md",
|
||||
data: []byte("# SVG Example\n<svg width=\"100\">...</svg>"),
|
||||
want: "text/markdown",
|
||||
},
|
||||
{
|
||||
name: "CSVMentioningSVGStaysCSV",
|
||||
fileName: "report.csv",
|
||||
data: []byte("name,icon\nlogo,<svg><rect/></svg>\n"),
|
||||
want: "text/csv",
|
||||
},
|
||||
{
|
||||
name: "TextMentioningSVGStaysPlainText",
|
||||
fileName: "main.go",
|
||||
data: []byte("package main\n// renders <svg> tags\n"),
|
||||
want: "text/plain",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
require.Equal(t, tt.want, chatfiles.ClassifyStoredMediaType(tt.fileName, tt.data))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrepareStoredFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("UsesDetectNameForSubtypeRefinement", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
name, mediaType, err := chatfiles.PrepareStoredFile(
|
||||
"payload.txt",
|
||||
"report.json",
|
||||
[]byte(`{"ok":true}`),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "payload.txt", name)
|
||||
require.Equal(t, "application/json", mediaType)
|
||||
})
|
||||
|
||||
t.Run("StripsControlCharactersAndTrimsExposedWhitespace", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
name, mediaType, err := chatfiles.PrepareStoredFile(
|
||||
"\x00 release\t notes.txt \x00",
|
||||
"release-notes.txt",
|
||||
[]byte("hello"),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "release notes.txt", name)
|
||||
require.Equal(t, "text/plain", mediaType)
|
||||
})
|
||||
|
||||
t.Run("RejectsEmptyNormalizedName", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
_, _, err := chatfiles.PrepareStoredFile(
|
||||
" \r\n\t ",
|
||||
"notes.txt",
|
||||
[]byte("hello"),
|
||||
)
|
||||
require.ErrorIs(t, err, chatfiles.ErrStoredFileNameRequired)
|
||||
})
|
||||
|
||||
t.Run("RejectsUnsupportedStoredFileType", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
_, _, err := chatfiles.PrepareStoredFile(
|
||||
"evil.svg",
|
||||
"evil.svg",
|
||||
[]byte(`<svg xmlns="http://www.w3.org/2000/svg"><rect/></svg>`),
|
||||
)
|
||||
require.ErrorIs(t, err, chatfiles.ErrUnsupportedStoredFileType)
|
||||
require.ErrorContains(t, err, "image/svg+xml")
|
||||
})
|
||||
|
||||
t.Run("TruncatesNamesAtRuneBoundaries", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
name, _, err := chatfiles.PrepareStoredFile(
|
||||
strings.Repeat("界", 100),
|
||||
"notes.txt",
|
||||
[]byte("hello"),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, strings.Repeat("界", 85), name)
|
||||
require.Equal(t, 255, len(name))
|
||||
})
|
||||
}
|
||||
|
||||
func TestPrepareRecordingArtifact(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("MP4", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
name, mediaType, err := chatfiles.PrepareRecordingArtifact(
|
||||
"recording.mp4",
|
||||
"video/mp4",
|
||||
[]byte{0x00, 0x00, 0x00, 0x18, 'f', 't', 'y', 'p', 'm', 'p', '4', '2', 0x00, 0x00, 0x00, 0x00, 'm', 'p', '4', '1', 'i', 's', 'o', 'm'},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "recording.mp4", name)
|
||||
require.Equal(t, "video/mp4", mediaType)
|
||||
})
|
||||
|
||||
t.Run("JPEG", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
name, mediaType, err := chatfiles.PrepareRecordingArtifact(
|
||||
"thumbnail.jpg",
|
||||
"image/jpeg",
|
||||
[]byte{0xFF, 0xD8, 0xFF, 0xE0, 0x00, 0x10, 'J', 'F', 'I', 'F', 0x00, 0x01, 0x01, 0x00, 0x00, 0x01, 0x00, 0x01, 0x00, 0x00},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "thumbnail.jpg", name)
|
||||
require.Equal(t, "image/jpeg", mediaType)
|
||||
})
|
||||
|
||||
t.Run("TypeMismatch", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
_, _, err := chatfiles.PrepareRecordingArtifact(
|
||||
"recording.mp4",
|
||||
"video/mp4",
|
||||
[]byte{0xFF, 0xD8, 0xFF, 0xE0, 0x00, 0x10, 'J', 'F', 'I', 'F', 0x00, 0x01, 0x01, 0x00, 0x00, 0x01, 0x00, 0x01, 0x00, 0x00},
|
||||
)
|
||||
require.ErrorContains(t, err, "recording artifact type mismatch")
|
||||
})
|
||||
|
||||
t.Run("RejectsEmptyNormalizedName", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
_, _, err := chatfiles.PrepareRecordingArtifact(
|
||||
" \r\n\t ",
|
||||
"video/mp4",
|
||||
[]byte{0x00, 0x00, 0x00, 0x18, 'f', 't', 'y', 'p', 'm', 'p', '4', '2', 0x00, 0x00, 0x00, 0x00, 'm', 'p', '4', '1', 'i', 's', 'o', 'm'},
|
||||
)
|
||||
require.ErrorIs(t, err, chatfiles.ErrStoredFileNameRequired)
|
||||
})
|
||||
|
||||
t.Run("UnsupportedExpectedType", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
_, _, err := chatfiles.PrepareRecordingArtifact(
|
||||
"recording.webm",
|
||||
"video/webm",
|
||||
[]byte("webm"),
|
||||
)
|
||||
require.ErrorContains(t, err, "unsupported recording artifact type")
|
||||
})
|
||||
}
|
||||
|
||||
func TestIsCompatibleUploadMediaType(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
declared string
|
||||
stored string
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "ExactMatch",
|
||||
declared: "text/plain",
|
||||
stored: "text/plain",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "TextPlainRefinesToMarkdown",
|
||||
declared: "text/plain",
|
||||
stored: "text/markdown",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "TextPlainRefinesToCSV",
|
||||
declared: "text/plain",
|
||||
stored: "text/csv",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "TextPlainRefinesToJSON",
|
||||
declared: "text/plain",
|
||||
stored: "application/json",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "TextPlainDoesNotRefineToPNG",
|
||||
declared: "text/plain",
|
||||
stored: "image/png",
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "JSONDoesNotRefineToPlainText",
|
||||
declared: "application/json",
|
||||
stored: "text/plain",
|
||||
want: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
require.Equal(t, tt.want, chatfiles.IsCompatibleUploadMediaType(tt.declared, tt.stored))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsAllowedStoredMediaType(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.True(t, chatfiles.IsAllowedStoredMediaType("text/plain; charset=utf-8"))
|
||||
require.True(t, chatfiles.IsAllowedStoredMediaType("text/markdown"))
|
||||
require.True(t, chatfiles.IsAllowedStoredMediaType("text/csv"))
|
||||
require.True(t, chatfiles.IsAllowedStoredMediaType("application/json"))
|
||||
require.True(t, chatfiles.IsAllowedStoredMediaType("application/pdf"))
|
||||
require.True(t, chatfiles.IsAllowedStoredMediaType("image/png"))
|
||||
require.False(t, chatfiles.IsAllowedStoredMediaType("image/svg+xml"))
|
||||
require.False(t, chatfiles.IsAllowedStoredMediaType("image/avif"))
|
||||
require.False(t, chatfiles.IsAllowedStoredMediaType("application/zip"))
|
||||
}
|
||||
|
||||
func TestIsInlineRenderableStoredMediaType(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.True(t, chatfiles.IsInlineRenderableStoredMediaType("text/plain; charset=utf-8"))
|
||||
require.True(t, chatfiles.IsInlineRenderableStoredMediaType("text/markdown"))
|
||||
require.True(t, chatfiles.IsInlineRenderableStoredMediaType("image/png"))
|
||||
require.False(t, chatfiles.IsInlineRenderableStoredMediaType("application/pdf"))
|
||||
require.False(t, chatfiles.IsInlineRenderableStoredMediaType("image/svg+xml"))
|
||||
}
|
||||
|
||||
func TestHasSVGRootElement(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.True(t, chatfiles.HasSVGRootElement([]byte(`<?xml version="1.0"?><svg xmlns="http://www.w3.org/2000/svg"></svg>`)))
|
||||
require.True(t, chatfiles.HasSVGRootElement([]byte("\xef\xbb\xbf<svg></svg>")))
|
||||
require.False(t, chatfiles.HasSVGRootElement([]byte("<html><body>not svg</body></html>")))
|
||||
require.False(t, chatfiles.HasSVGRootElement([]byte("# SVG Example\n<svg width=\"100\">...</svg>")))
|
||||
require.False(t, chatfiles.HasSVGRootElement([]byte("name,icon\nlogo,<svg><rect/></svg>\n")))
|
||||
}
|
||||
Reference in New Issue
Block a user