mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: track chat file associations with chat_file_links on chats (#23537)
Needed by #23833 Adds a `chat_file_links` association table to track which files are associated with each chat. - `AppendChatFileIDs` query links a file to a chat with deduplication - `GetChatFileMetadataByIDs` query returns lightweight file metadata by IDs - Tool-created files (e.g. `propose_plan`) are linked to the chat after insert - User-uploaded files are linked to the chat when the referencing message is sent - Single-chat GET endpoint hydrates `files: ChatFileMetadata[]` on the response > 🤖 Created by Coder Agents and massaged into shape by a human.
This commit is contained in:
+139
-19
@@ -413,7 +413,7 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
contentBlocks, titleSource, inputError := createChatInputFromRequest(ctx, api.Database, req)
|
||||
contentBlocks, titleSource, fileIDs, inputError := createChatInputFromRequest(ctx, api.Database, req)
|
||||
if inputError != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, *inputError)
|
||||
return
|
||||
@@ -524,7 +524,32 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
httpapi.Write(ctx, rw, http.StatusCreated, db2sdk.Chat(chat, nil))
|
||||
// Link any user-uploaded files referenced in the initial
|
||||
// message to this newly created chat (best-effort; cap
|
||||
// enforced in SQL).
|
||||
unlinked, capExceeded := api.linkFilesToChat(ctx, chat.ID, fileIDs)
|
||||
|
||||
// Re-read the chat so the response reflects the authoritative
|
||||
// database state (file links are deduped in the join table).
|
||||
chat, err = api.Database.GetChatByID(ctx, chat.ID)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Failed to read back chat after creation.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
chatFiles := api.fetchChatFileMetadata(ctx, chat.ID)
|
||||
response := db2sdk.Chat(chat, nil, chatFiles)
|
||||
if len(unlinked) > 0 {
|
||||
if capExceeded {
|
||||
response.Warnings = append(response.Warnings, fileLinkCapWarning(len(unlinked)))
|
||||
} else {
|
||||
response.Warnings = append(response.Warnings, fileLinkErrorWarning(len(unlinked)))
|
||||
}
|
||||
}
|
||||
httpapi.Write(ctx, rw, http.StatusCreated, response)
|
||||
}
|
||||
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
@@ -1301,7 +1326,11 @@ func (api *API) getChat(rw http.ResponseWriter, r *http.Request) {
|
||||
slog.Error(err),
|
||||
)
|
||||
}
|
||||
httpapi.Write(ctx, rw, http.StatusOK, db2sdk.Chat(chat, diffStatus))
|
||||
|
||||
// Hydrate file metadata for all files linked to this chat.
|
||||
chatFiles := api.fetchChatFileMetadata(ctx, chat.ID)
|
||||
|
||||
httpapi.Write(ctx, rw, http.StatusOK, db2sdk.Chat(chat, diffStatus, chatFiles))
|
||||
}
|
||||
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
@@ -1791,7 +1820,7 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
contentBlocks, _, inputError := createChatInputFromParts(ctx, api.Database, req.Content, "content")
|
||||
contentBlocks, _, fileIDs, inputError := createChatInputFromParts(ctx, api.Database, req.Content, "content")
|
||||
if inputError != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: inputError.Message,
|
||||
@@ -1873,6 +1902,9 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
// Link any user-uploaded files referenced in this message
|
||||
// to the chat (best-effort; cap enforced in SQL).
|
||||
unlinked, capExceeded := api.linkFilesToChat(ctx, chatID, fileIDs)
|
||||
response := codersdk.CreateChatMessageResponse{Queued: sendResult.Queued}
|
||||
if sendResult.Queued {
|
||||
if sendResult.QueuedMessage != nil {
|
||||
@@ -1882,6 +1914,13 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
|
||||
message := convertChatMessage(sendResult.Message)
|
||||
response.Message = &message
|
||||
}
|
||||
if len(unlinked) > 0 {
|
||||
if capExceeded {
|
||||
response.Warnings = append(response.Warnings, fileLinkCapWarning(len(unlinked)))
|
||||
} else {
|
||||
response.Warnings = append(response.Warnings, fileLinkErrorWarning(len(unlinked)))
|
||||
}
|
||||
}
|
||||
|
||||
httpapi.Write(ctx, rw, http.StatusOK, response)
|
||||
}
|
||||
@@ -1915,7 +1954,7 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
contentBlocks, _, inputError := createChatInputFromParts(ctx, api.Database, req.Content, "content")
|
||||
contentBlocks, _, fileIDs, inputError := createChatInputFromParts(ctx, api.Database, req.Content, "content")
|
||||
if inputError != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: inputError.Message,
|
||||
@@ -1954,8 +1993,20 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
message := convertChatMessage(editResult.Message)
|
||||
httpapi.Write(ctx, rw, http.StatusOK, message)
|
||||
// Link any user-uploaded files referenced in the edited
|
||||
// message to the chat (best-effort; cap enforced in SQL).
|
||||
unlinked, capExceeded := api.linkFilesToChat(ctx, chat.ID, fileIDs)
|
||||
response := codersdk.EditChatMessageResponse{
|
||||
Message: convertChatMessage(editResult.Message),
|
||||
}
|
||||
if len(unlinked) > 0 {
|
||||
if capExceeded {
|
||||
response.Warnings = append(response.Warnings, fileLinkCapWarning(len(unlinked)))
|
||||
} else {
|
||||
response.Warnings = append(response.Warnings, fileLinkErrorWarning(len(unlinked)))
|
||||
}
|
||||
}
|
||||
httpapi.Write(ctx, rw, http.StatusOK, response)
|
||||
}
|
||||
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
@@ -2232,7 +2283,7 @@ func (api *API) interruptChat(rw http.ResponseWriter, r *http.Request) {
|
||||
chat = updatedChat
|
||||
}
|
||||
|
||||
httpapi.Write(ctx, rw, http.StatusOK, db2sdk.Chat(chat, nil))
|
||||
httpapi.Write(ctx, rw, http.StatusOK, db2sdk.Chat(chat, nil, nil))
|
||||
}
|
||||
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
@@ -2276,7 +2327,7 @@ func (api *API) regenerateChatTitle(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
httpapi.Write(ctx, rw, http.StatusOK, db2sdk.Chat(updatedChat, nil))
|
||||
httpapi.Write(ctx, rw, http.StatusOK, db2sdk.Chat(updatedChat, nil, nil))
|
||||
}
|
||||
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
@@ -3688,6 +3739,7 @@ func (api *API) chatFileByID(rw http.ResponseWriter, r *http.Request) {
|
||||
func createChatInputFromRequest(ctx context.Context, db database.Store, req codersdk.CreateChatRequest) (
|
||||
[]codersdk.ChatMessagePart,
|
||||
string,
|
||||
[]uuid.UUID,
|
||||
*codersdk.Response,
|
||||
) {
|
||||
return createChatInputFromParts(ctx, db, req.Content, "content")
|
||||
@@ -3698,14 +3750,15 @@ func createChatInputFromParts(
|
||||
db database.Store,
|
||||
parts []codersdk.ChatInputPart,
|
||||
fieldName string,
|
||||
) ([]codersdk.ChatMessagePart, string, *codersdk.Response) {
|
||||
) ([]codersdk.ChatMessagePart, string, []uuid.UUID, *codersdk.Response) {
|
||||
if len(parts) == 0 {
|
||||
return nil, "", &codersdk.Response{
|
||||
return nil, "", nil, &codersdk.Response{
|
||||
Message: "Content is required.",
|
||||
Detail: "Content cannot be empty.",
|
||||
}
|
||||
}
|
||||
|
||||
var fileIDs []uuid.UUID
|
||||
content := make([]codersdk.ChatMessagePart, 0, len(parts))
|
||||
textParts := make([]string, 0, len(parts))
|
||||
for i, part := range parts {
|
||||
@@ -3713,7 +3766,7 @@ func createChatInputFromParts(
|
||||
case string(codersdk.ChatInputPartTypeText):
|
||||
text := strings.TrimSpace(part.Text)
|
||||
if text == "" {
|
||||
return nil, "", &codersdk.Response{
|
||||
return nil, "", nil, &codersdk.Response{
|
||||
Message: "Invalid input part.",
|
||||
Detail: fmt.Sprintf("%s[%d].text cannot be empty.", fieldName, i),
|
||||
}
|
||||
@@ -3722,7 +3775,7 @@ func createChatInputFromParts(
|
||||
textParts = append(textParts, text)
|
||||
case string(codersdk.ChatInputPartTypeFile):
|
||||
if part.FileID == uuid.Nil {
|
||||
return nil, "", &codersdk.Response{
|
||||
return nil, "", nil, &codersdk.Response{
|
||||
Message: "Invalid input part.",
|
||||
Detail: fmt.Sprintf("%s[%d].file_id is required for file parts.", fieldName, i),
|
||||
}
|
||||
@@ -3733,20 +3786,23 @@ func createChatInputFromParts(
|
||||
chatFile, err := db.GetChatFileByID(ctx, part.FileID)
|
||||
if err != nil {
|
||||
if httpapi.Is404Error(err) {
|
||||
return nil, "", &codersdk.Response{
|
||||
return nil, "", 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, "", &codersdk.Response{
|
||||
return nil, "", nil, &codersdk.Response{
|
||||
Message: "Internal error.",
|
||||
Detail: fmt.Sprintf("Failed to retrieve file for %s[%d].", fieldName, i),
|
||||
}
|
||||
}
|
||||
content = append(content, codersdk.ChatMessageFile(part.FileID, chatFile.Mimetype))
|
||||
fileIDs = append(fileIDs, part.FileID)
|
||||
// file-reference parts carry inline code snippets, not uploaded
|
||||
// files. They have no FileID and are excluded from file tracking.
|
||||
case string(codersdk.ChatInputPartTypeFileReference):
|
||||
if part.FileName == "" {
|
||||
return nil, "", &codersdk.Response{
|
||||
return nil, "", nil, &codersdk.Response{
|
||||
Message: "Invalid input part.",
|
||||
Detail: fmt.Sprintf("%s[%d].file_name cannot be empty for file-reference.", fieldName, i),
|
||||
}
|
||||
@@ -3764,7 +3820,7 @@ func createChatInputFromParts(
|
||||
}
|
||||
textParts = append(textParts, sb.String())
|
||||
default:
|
||||
return nil, "", &codersdk.Response{
|
||||
return nil, "", nil, &codersdk.Response{
|
||||
Message: "Invalid input part.",
|
||||
Detail: fmt.Sprintf(
|
||||
"%s[%d].type %q is not supported.",
|
||||
@@ -3779,13 +3835,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, "", &codersdk.Response{
|
||||
return nil, "", 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, titleSource, nil
|
||||
return content, titleSource, fileIDs, nil
|
||||
}
|
||||
|
||||
func chatTitleFromMessage(message string) string {
|
||||
@@ -3820,6 +3876,70 @@ func truncateRunes(value string, maxLen int) string {
|
||||
return string(runes[:maxLen])
|
||||
}
|
||||
|
||||
// linkFilesToChat inserts file-link rows into the chat_file_links
|
||||
// join table. Cap enforcement and dedup are handled atomically in
|
||||
// SQL. On success returns (nil, false). On failure returns the full
|
||||
// input fileIDs slice — linking is all-or-nothing because the
|
||||
// SQL operates on the batch atomically. capExceeded indicates
|
||||
// whether the failure was due to the cap being exceeded (true)
|
||||
// or a database error (false).
|
||||
// Failures are logged but never block the caller.
|
||||
func (api *API) linkFilesToChat(ctx context.Context, chatID uuid.UUID, fileIDs []uuid.UUID) (unlinked []uuid.UUID, capExceeded bool) {
|
||||
if len(fileIDs) == 0 {
|
||||
return nil, false
|
||||
}
|
||||
rejected, err := api.Database.LinkChatFiles(ctx, database.LinkChatFilesParams{
|
||||
ChatID: chatID,
|
||||
MaxFileLinks: int32(codersdk.MaxChatFileIDs),
|
||||
FileIds: fileIDs,
|
||||
})
|
||||
if err != nil {
|
||||
api.Logger.Error(ctx, "failed to link files to chat",
|
||||
slog.F("chat_id", chatID),
|
||||
slog.F("file_ids", fileIDs),
|
||||
slog.Error(err),
|
||||
)
|
||||
return fileIDs, false
|
||||
}
|
||||
if rejected > 0 {
|
||||
api.Logger.Warn(ctx, "file cap reached, files not linked",
|
||||
slog.F("chat_id", chatID),
|
||||
slog.F("file_ids", fileIDs),
|
||||
slog.F("max_file_links", codersdk.MaxChatFileIDs),
|
||||
)
|
||||
return fileIDs, true
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// fileLinkCapWarning builds a user-facing warning when a batch
|
||||
// of file IDs was atomically rejected because the resulting
|
||||
// array would exceed the per-chat file cap.
|
||||
func fileLinkCapWarning(count int) string {
|
||||
return fmt.Sprintf("file linking skipped: batch of %d file(s) would exceed limit of %d", count, codersdk.MaxChatFileIDs)
|
||||
}
|
||||
|
||||
// fileLinkErrorWarning builds a user-facing warning when a
|
||||
// database error prevented linking files to a chat.
|
||||
func fileLinkErrorWarning(count int) string {
|
||||
return fmt.Sprintf("%d file(s) could not be linked due to a server error", count)
|
||||
}
|
||||
|
||||
// fetchChatFileMetadata returns metadata for all files linked to
|
||||
// the given chat. Errors are logged and result in a nil return
|
||||
// (callers treat file metadata as best-effort).
|
||||
func (api *API) fetchChatFileMetadata(ctx context.Context, chatID uuid.UUID) []database.GetChatFileMetadataByChatIDRow {
|
||||
rows, err := api.Database.GetChatFileMetadataByChatID(ctx, chatID)
|
||||
if err != nil {
|
||||
api.Logger.Error(ctx, "failed to fetch chat file metadata",
|
||||
slog.F("chat_id", chatID),
|
||||
slog.Error(err),
|
||||
)
|
||||
return nil
|
||||
}
|
||||
return rows
|
||||
}
|
||||
|
||||
func convertChatCostModelBreakdown(model database.GetChatCostPerModelRow) codersdk.ChatCostModelBreakdown {
|
||||
displayName := strings.TrimSpace(model.DisplayName)
|
||||
if displayName == "" {
|
||||
|
||||
Reference in New Issue
Block a user