fix: keep chat attachments while a linking chat exists

Fixes https://linear.app/codercom/issue/CODAGT-616/keep-chat-attachments-while-chats-remain-unarchived

Chat attachments could disappear even though the chat was still available. This happened when a message was saved without recording which attachments it used, or when cleanup deleted attachments before an archived chat itself was removed.

Creating a chat, sending or queuing a message, and editing a message now record both the message and which attachments it uses as one operation. If the chat is already at the 50-attachment limit, the chat change fails without being partially saved.

Concurrent attachment writes serialize the 50-file cap per chat. Cleanup locks candidates and checks again for new links before deleting. If a file becomes unavailable after input validation, create, send, and edit return a clear client error and roll back the chat change.

An attachment stays available while any chat that uses it still exists. After an archived chat reaches the end of its retention period and is deleted, an old attachment that no remaining chat uses can be cleaned up. The retention guide and unavailable-attachment UI text document this lifecycle. This change cannot restore attachments that were already deleted.

The database migration adds two indexes so attachment cleanup stays fast as attachments accumulate.

> This PR was authored by Mux (AI) on Mike's behalf.
This commit is contained in:
Michael Suchacz
2026-08-11 13:53:15 +02:00
committed by GitHub
parent 72ad835330
commit 57f38b5c24
33 changed files with 1130 additions and 368 deletions
+2
View File
@@ -47,6 +47,8 @@ There is other data that is held in the database and is associated with a chat,
We call it **metadata**. The core state machine concerns itself with **execution state**. As a general guideline, a piece of data is execution state if the core state machine needs it to decide what the next state transition may be, or if it's directly modified by a state transition. For example, a queued message is part of the execution state because it impacts what the next action of the agent loop can be. If the agent loop finishes processing a user message and would otherwise stop, but there's a queued message, the agent loop will start processing the queued message instead. On the other hand, a chat's title does not impact the agent loop at all - it's just a label that helps the user identify the chat.
File links are metadata, but one invariant is enforced at transition time: if a transition persists message content that references uploaded files (chat create, message send, queued send, or message edit), it records the file links in the same transaction. If linking would exceed the per-chat attachment cap, the whole transition is rejected. File retention skips files that are still linked to existing chats, so a persisted message must never reference a file without a link.
If the distinction isn't completely clear to you at this point, don't worry. It should become clearer as you learn more about the core state machine.
## Execution states
+9
View File
@@ -1378,6 +1378,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
},
ClientType: opts.ClientType,
InitialMessages: initialMessages,
FileIDs: chatprompt.FileIDs(contentParts),
})
if err != nil {
return database.Chat{}, err
@@ -1557,6 +1558,11 @@ func (p *Server) SendMessage(
// previous queue head into history; report those inserts so
// clients can update their caches.
result.InsertedMessages = sendResult.InsertedMessages
// File-link errors must roll back the message.
if err := chatstate.LinkFiles(ctx, store, opts.ChatID, chatprompt.FileIDs(contentParts)); err != nil {
return err
}
// Capture the post-transition chat inside the same
// transaction so the returned chat and the watch event
// reflect the snapshot bump and status change produced by
@@ -1866,6 +1872,9 @@ func (p *Server) EditMessage(
inserted = append(inserted, editResult.SuffixMessages...)
result.InsertedMessages = inserted
result.DeletedMessageIDs = editResult.DeletedMessageIDs
if err := chatstate.LinkFiles(ctx, store, opts.ChatID, chatprompt.FileIDs(contentParts)); err != nil {
return err
}
// Capture the post-edit chat inside the same transaction so
// the returned chat and the debug-cleanup cutoff use the
// snapshot bump and updated_at stamped by the transition.
+208
View File
@@ -1646,6 +1646,214 @@ func TestSendMessageQueueBehaviorQueuesWhenBusy(t *testing.T) {
require.Len(t, messages, 1)
}
func TestMessageFileLinking(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
replica := newTestServer(t, db, ps, uuid.New())
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(t, db)
insertFile := func(name string) uuid.UUID {
t.Helper()
row, err := db.InsertChatFile(ctx, database.InsertChatFileParams{
OwnerID: user.ID,
OrganizationID: org.ID,
Name: name,
Mimetype: "image/png",
Data: []byte("png-bytes"),
})
require.NoError(t, err)
return row.ID
}
linkedFileIDs := func(chatID uuid.UUID) []uuid.UUID {
t.Helper()
rows, err := db.GetChatFileMetadataByChatID(ctx, chatID)
require.NoError(t, err)
ids := make([]uuid.UUID, 0, len(rows))
for _, row := range rows {
ids = append(ids, row.ID)
}
return ids
}
fileCreate := insertFile("create.png")
fileSend := insertFile("send.png")
fileQueued := insertFile("queued.png")
fileEdit := insertFile("edit.png")
chat, err := replica.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
OwnerID: user.ID,
Title: "file-linking",
ModelConfigID: model.ID,
InitialUserContent: []codersdk.ChatMessagePart{
codersdk.ChatMessageText("with attachment"),
codersdk.ChatMessageFile(fileCreate, "image/png", "create.png"),
},
})
require.NoError(t, err)
require.Contains(t, linkedFileIDs(chat.ID), fileCreate)
chat, err = db.UpdateChatStatus(ctx, database.UpdateChatStatusParams{
ID: chat.ID,
Status: database.ChatStatusWaiting,
})
require.NoError(t, err)
sendResult, err := replica.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: chat.ID,
Content: []codersdk.ChatMessagePart{
codersdk.ChatMessageText("another attachment"),
codersdk.ChatMessageFile(fileSend, "image/png", "send.png"),
},
})
require.NoError(t, err)
require.False(t, sendResult.Queued)
require.Contains(t, linkedFileIDs(chat.ID), fileSend)
chat, err = db.UpdateChatStatus(ctx, database.UpdateChatStatusParams{
ID: chat.ID,
Status: database.ChatStatusWaiting,
})
require.NoError(t, err)
messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
ChatID: chat.ID,
AfterID: 0,
})
require.NoError(t, err)
var userMessageID int64
for _, msg := range messages {
if msg.Role == database.ChatMessageRoleUser {
userMessageID = msg.ID
break
}
}
require.NotZero(t, userMessageID)
_, err = replica.EditMessage(ctx, chatd.EditMessageOptions{
ChatID: chat.ID,
CreatedBy: user.ID,
EditedMessageID: userMessageID,
Content: []codersdk.ChatMessagePart{
codersdk.ChatMessageText("edited attachment"),
codersdk.ChatMessageFile(fileEdit, "image/png", "edit.png"),
},
})
require.NoError(t, err)
editedLinks := linkedFileIDs(chat.ID)
require.Contains(t, editedLinks, fileEdit)
require.Contains(t, editedLinks, fileCreate)
// Queued files must be linked before promotion to prevent purge.
_, err = db.UpdateChatStatus(ctx, database.UpdateChatStatusParams{
ID: chat.ID,
Status: database.ChatStatusRunning,
WorkerID: uuid.NullUUID{UUID: uuid.New(), Valid: true},
StartedAt: sql.NullTime{Time: time.Now(), Valid: true},
HeartbeatAt: sql.NullTime{Time: time.Now(), Valid: true},
})
require.NoError(t, err)
queuedResult, err := replica.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: chat.ID,
Content: []codersdk.ChatMessagePart{
codersdk.ChatMessageText("queued attachment"),
codersdk.ChatMessageFile(fileQueued, "image/png", "queued.png"),
},
BusyBehavior: chatd.SendMessageBusyBehaviorQueue,
})
require.NoError(t, err)
require.True(t, queuedResult.Queued)
require.Contains(t, linkedFileIDs(chat.ID), fileQueued)
}
func TestMessageFileLinkingCapRollsBack(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
replica := newTestServer(t, db, ps, uuid.New())
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(t, db)
chat, err := replica.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
OwnerID: user.ID,
Title: "cap-rollback",
ModelConfigID: model.ID,
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
})
require.NoError(t, err)
capFileIDs := make([]uuid.UUID, 0, codersdk.MaxChatFileIDs)
for i := range codersdk.MaxChatFileIDs {
row, err := db.InsertChatFile(ctx, database.InsertChatFileParams{
OwnerID: user.ID,
OrganizationID: org.ID,
Name: fmt.Sprintf("cap-%d.png", i),
Mimetype: "image/png",
Data: []byte("png-bytes"),
})
require.NoError(t, err)
capFileIDs = append(capFileIDs, row.ID)
}
rejected, err := db.LinkChatFiles(ctx, database.LinkChatFilesParams{
ChatID: chat.ID,
MaxFileLinks: int32(codersdk.MaxChatFileIDs),
FileIds: capFileIDs,
})
require.NoError(t, err)
require.Zero(t, rejected)
extra, err := db.InsertChatFile(ctx, database.InsertChatFileParams{
OwnerID: user.ID,
OrganizationID: org.ID,
Name: "extra.png",
Mimetype: "image/png",
Data: []byte("png-bytes"),
})
require.NoError(t, err)
chat, err = db.UpdateChatStatus(ctx, database.UpdateChatStatusParams{
ID: chat.ID,
Status: database.ChatStatusWaiting,
})
require.NoError(t, err)
messagesBefore, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
ChatID: chat.ID,
AfterID: 0,
})
require.NoError(t, err)
_, err = replica.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: chat.ID,
Content: []codersdk.ChatMessagePart{
codersdk.ChatMessageText("one too many"),
codersdk.ChatMessageFile(extra.ID, "image/png", "extra.png"),
},
})
require.ErrorIs(t, err, chatstate.ErrChatFileCapExceeded)
messagesAfter, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
ChatID: chat.ID,
AfterID: 0,
})
require.NoError(t, err)
require.Len(t, messagesAfter, len(messagesBefore), "rejected send must not persist a message")
files, err := db.GetChatFileMetadataByChatID(ctx, chat.ID)
require.NoError(t, err)
require.Len(t, files, codersdk.MaxChatFileIDs)
sendResult, err := replica.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: chat.ID,
Content: []codersdk.ChatMessagePart{
codersdk.ChatMessageText("re-reference"),
codersdk.ChatMessageFile(capFileIDs[0], "image/png", "cap-0.png"),
},
})
require.NoError(t, err)
require.False(t, sendResult.Queued)
}
func TestPlanTurnPromptContract(t *testing.T) {
t.Parallel()
+12
View File
@@ -94,6 +94,18 @@ func ExtractFileID(raw json.RawMessage) (uuid.UUID, error) {
return uuid.Parse(envelope.Data.FileID)
}
// FileIDs returns the valid file IDs referenced by file parts.
func FileIDs(parts []codersdk.ChatMessagePart) []uuid.UUID {
var ids []uuid.UUID
for _, part := range parts {
if part.Type != codersdk.ChatMessagePartTypeFile || !part.FileID.Valid {
continue
}
ids = append(ids, part.FileID.UUID)
}
return ids
}
// ConvertMessagesWithFiles converts persisted chat messages into LLM
// prompt messages, resolving user file references via the provided
// resolver. Missing-data placeholders are emitted only for replayed
+6
View File
@@ -49,6 +49,12 @@ var (
// wraps this sentinel.
ErrMessageQueueFull = xerrors.New("chat message queue is full")
// ErrChatFileCapExceeded reports a [LinkFiles] cap rejection.
ErrChatFileCapExceeded = xerrors.New("chat attachment cap exceeded")
// ErrChatFileUnavailable reports a missing file passed to [LinkFiles].
ErrChatFileUnavailable = xerrors.New("chat attachment unavailable")
// ErrToolResultDuplicate is returned by [Tx.CompleteRequiresAction]
// when the same tool_call_id appears more than once in the
// submitted results.
+37
View File
@@ -0,0 +1,37 @@
package chatstate
import (
"context"
"errors"
"github.com/google/uuid"
"golang.org/x/xerrors"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/codersdk"
)
// LinkFiles links files, returning [ErrChatFileCapExceeded] for cap rejections
// and [ErrChatFileUnavailable] for missing files. Use the caller's transaction
// so failures roll back related writes; existing links use no additional slots.
func LinkFiles(ctx context.Context, store database.Store, chatID uuid.UUID, fileIDs []uuid.UUID) error {
if len(fileIDs) == 0 {
return nil
}
rejected, err := store.LinkChatFiles(ctx, database.LinkChatFilesParams{
ChatID: chatID,
MaxFileLinks: int32(codersdk.MaxChatFileIDs),
FileIds: fileIDs,
})
if err != nil {
wrapped := xerrors.Errorf("link chat files: %w", err)
if database.IsForeignKeyViolation(err, database.ForeignKeyChatFileLinksFileID) {
return errors.Join(ErrChatFileUnavailable, wrapped)
}
return wrapped
}
if rejected > 0 {
return ErrChatFileCapExceeded
}
return nil
}
+38
View File
@@ -0,0 +1,38 @@
package chatstate_test
import (
"context"
"testing"
"github.com/google/uuid"
"github.com/lib/pq"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbmock"
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
"github.com/coder/coder/v2/codersdk"
)
func TestLinkFilesUnavailable(t *testing.T) {
t.Parallel()
ctrl := gomock.NewController(t)
store := dbmock.NewMockStore(ctrl)
chatID := uuid.New()
fileID := uuid.New()
foreignKeyErr := &pq.Error{
Code: pq.ErrorCode("23503"),
Constraint: string(database.ForeignKeyChatFileLinksFileID),
}
store.EXPECT().LinkChatFiles(gomock.Any(), database.LinkChatFilesParams{
ChatID: chatID,
MaxFileLinks: int32(codersdk.MaxChatFileIDs),
FileIds: []uuid.UUID{fileID},
}).Return(int32(0), foreignKeyErr)
err := chatstate.LinkFiles(context.Background(), store, chatID, []uuid.UUID{fileID})
require.ErrorIs(t, err, chatstate.ErrChatFileUnavailable)
require.ErrorIs(t, err, foreignKeyErr)
}
+5
View File
@@ -35,6 +35,8 @@ type CreateChatInput struct {
DynamicTools pqtype.NullRawMessage
ClientType database.ChatClientType
InitialMessages []Message
// FileIDs are linked atomically with the initial messages.
FileIDs []uuid.UUID
}
// CreateChatResult is the value returned by [CreateChat]. It carries
@@ -131,6 +133,9 @@ func insertChat(
if err != nil {
return xerrors.Errorf("insert initial messages: %w", err)
}
if err := LinkFiles(ctx, store, chat.ID, input.FileIDs); err != nil {
return err
}
refreshed, err := store.GetChatByID(ctx, chat.ID)
if err != nil {
return xerrors.Errorf("reload chat after initial messages: %w", err)
+14 -1
View File
@@ -477,7 +477,15 @@ func TestEditMessageUserPromptSubmitHook(t *testing.T) {
t.Cleanup(consumer.Close)
server := newHookTestServer(t, db, ps, consumer)
upload := codersdk.ChatMessageFile(uuid.New(), "image/png", "edited.png")
chatFile, err := db.InsertChatFile(ctx, database.InsertChatFileParams{
OwnerID: user.ID,
OrganizationID: org.ID,
Name: "edited.png",
Mimetype: "image/png",
Data: []byte("png-bytes"),
})
require.NoError(t, err)
upload := codersdk.ChatMessageFile(chatFile.ID, chatFile.Mimetype, chatFile.Name)
reference := codersdk.ChatMessageFileReference("main.go", 1, 3, "package main")
result, err := server.EditMessage(ctx, chatd.EditMessageOptions{
ChatID: chat.ID,
@@ -490,6 +498,11 @@ func TestEditMessageUserPromptSubmitHook(t *testing.T) {
},
})
require.NoError(t, err)
linkedFiles, err := db.GetChatFileMetadataByChatID(ctx, chat.ID)
require.NoError(t, err)
require.Len(t, linkedFiles, 1)
require.Equal(t, chatFile.ID, linkedFiles[0].ID)
parts, err := chatprompt.ParseContent(result.Message)
require.NoError(t, err)
require.Equal(t, []codersdk.ChatMessagePart{
+7 -10
View File
@@ -2,11 +2,13 @@ package chatd
import (
"context"
"errors"
"github.com/google/uuid"
"golang.org/x/xerrors"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
"github.com/coder/coder/v2/coderd/x/chatfiles"
"github.com/coder/coder/v2/codersdk"
@@ -91,16 +93,11 @@ func storeLinkedChatFileTx(
return chattool.AttachmentMetadata{}, xerrors.Errorf("insert chat file: %w", err)
}
rejected, err := tx.LinkChatFiles(ctx, database.LinkChatFilesParams{
ChatID: chatID,
MaxFileLinks: int32(codersdk.MaxChatFileIDs),
FileIds: []uuid.UUID{row.ID},
})
if err != nil {
return chattool.AttachmentMetadata{}, xerrors.Errorf("link chat file: %w", err)
}
if rejected > 0 {
return chattool.AttachmentMetadata{}, xerrors.Errorf("chat already has the maximum of %d linked files", codersdk.MaxChatFileIDs)
if err := chatstate.LinkFiles(ctx, tx, chatID, []uuid.UUID{row.ID}); err != nil {
if errors.Is(err, chatstate.ErrChatFileCapExceeded) {
return chattool.AttachmentMetadata{}, xerrors.Errorf("chat already has the maximum of %d linked files", codersdk.MaxChatFileIDs)
}
return chattool.AttachmentMetadata{}, err
}
return chattool.AttachmentMetadata{
@@ -10,9 +10,12 @@ import (
"go.uber.org/mock/gomock"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbgen"
"github.com/coder/coder/v2/coderd/database/dbmock"
"github.com/coder/coder/v2/coderd/database/dbtestutil"
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/testutil"
)
func TestStoreChatAttachment_Success(t *testing.T) {
@@ -34,7 +37,7 @@ func TestStoreChatAttachment_Success(t *testing.T) {
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
}
expectStoreChatAttachmentTx(t, db, tx)
expectStoreChatAttachmentInTx(t, db, tx)
tx.EXPECT().GetWorkspaceByID(gomock.Any(), workspaceID).Return(database.Workspace{ID: workspaceID, OrganizationID: orgID}, nil)
tx.EXPECT().InsertChatFile(gomock.Any(), gomock.AssignableToTypeOf(database.InsertChatFileParams{})).DoAndReturn(
func(_ context.Context, arg database.InsertChatFileParams) (database.InsertChatFileRow, error) {
@@ -80,7 +83,7 @@ func TestStoreChatAttachment_UsesDetectNameForClassification(t *testing.T) {
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
}
expectStoreChatAttachmentTx(t, db, tx)
expectStoreChatAttachmentInTx(t, db, tx)
tx.EXPECT().GetWorkspaceByID(gomock.Any(), workspaceID).Return(database.Workspace{ID: workspaceID, OrganizationID: orgID}, nil)
tx.EXPECT().InsertChatFile(gomock.Any(), gomock.AssignableToTypeOf(database.InsertChatFileParams{})).DoAndReturn(
func(_ context.Context, arg database.InsertChatFileParams) (database.InsertChatFileRow, error) {
@@ -121,7 +124,7 @@ func TestStoreChatAttachment_AllowsUnsupportedPromptInputType(t *testing.T) {
}
data := []byte(`<svg xmlns="http://www.w3.org/2000/svg"><rect/></svg>`)
expectStoreChatAttachmentTx(t, db, tx)
expectStoreChatAttachmentInTx(t, db, tx)
tx.EXPECT().GetWorkspaceByID(gomock.Any(), workspaceID).Return(database.Workspace{ID: workspaceID, OrganizationID: orgID}, nil)
tx.EXPECT().InsertChatFile(gomock.Any(), gomock.AssignableToTypeOf(database.InsertChatFileParams{})).DoAndReturn(
func(_ context.Context, arg database.InsertChatFileParams) (database.InsertChatFileRow, error) {
@@ -175,7 +178,7 @@ func TestStoreChatAttachment_WorkspaceLookupError(t *testing.T) {
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
}
expectStoreChatAttachmentTx(t, db, tx)
expectStoreChatAttachmentInTx(t, db, tx)
tx.EXPECT().GetWorkspaceByID(gomock.Any(), workspaceID).Return(database.Workspace{}, context.DeadlineExceeded)
attachment, err := server.storeChatAttachment(context.Background(), chatSnapshot, "build.log", "build.log", []byte("build output"))
@@ -199,7 +202,7 @@ func TestStoreChatAttachment_InsertError(t *testing.T) {
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
}
expectStoreChatAttachmentTx(t, db, tx)
expectStoreChatAttachmentInTx(t, db, tx)
tx.EXPECT().GetWorkspaceByID(gomock.Any(), workspaceID).Return(database.Workspace{ID: workspaceID, OrganizationID: uuid.New()}, nil)
tx.EXPECT().InsertChatFile(gomock.Any(), gomock.Any()).Return(database.InsertChatFileRow{}, context.DeadlineExceeded)
@@ -228,7 +231,7 @@ func TestStoreChatAttachment_StrictCapError(t *testing.T) {
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
}
expectStoreChatAttachmentTx(t, db, tx)
expectStoreChatAttachmentInTx(t, db, tx)
tx.EXPECT().GetWorkspaceByID(gomock.Any(), workspaceID).Return(database.Workspace{ID: workspaceID, OrganizationID: orgID}, nil)
tx.EXPECT().InsertChatFile(gomock.Any(), gomock.AssignableToTypeOf(database.InsertChatFileParams{})).Return(database.InsertChatFileRow{ID: fileID}, nil)
tx.EXPECT().LinkChatFiles(gomock.Any(), database.LinkChatFilesParams{
@@ -261,7 +264,7 @@ func TestStoreChatAttachment_LinkError(t *testing.T) {
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
}
expectStoreChatAttachmentTx(t, db, tx)
expectStoreChatAttachmentInTx(t, db, tx)
tx.EXPECT().GetWorkspaceByID(gomock.Any(), workspaceID).Return(database.Workspace{ID: workspaceID, OrganizationID: orgID}, nil)
tx.EXPECT().InsertChatFile(gomock.Any(), gomock.Any()).Return(database.InsertChatFileRow{ID: fileID}, nil)
tx.EXPECT().LinkChatFiles(gomock.Any(), database.LinkChatFilesParams{
@@ -276,7 +279,126 @@ func TestStoreChatAttachment_LinkError(t *testing.T) {
require.Equal(t, chattool.AttachmentMetadata{}, attachment)
}
func expectStoreChatAttachmentTx(t *testing.T, db, tx *dbmock.MockStore) {
func TestStoreChatAttachment_SerializesCapCheck(t *testing.T) {
t.Parallel()
ctx := chatdTestContext(t)
db, _, rawDB := dbtestutil.NewDBWithSQLDB(t)
user, _, model := seedInternalChatDeps(t, db)
workspace, _, _ := seedWorkspaceBinding(t, db, user.ID)
chat := dbgen.Chat(t, db, database.Chat{
OrganizationID: workspace.OrganizationID,
OwnerID: user.ID,
WorkspaceID: uuid.NullUUID{UUID: workspace.ID, Valid: true},
LastModelConfigID: model.ID,
})
for i := range codersdk.MaxChatFileIDs - 1 {
insertLinkedChatFile(
ctx,
t,
db,
chat.ID,
user.ID,
workspace.OrganizationID,
fmt.Sprintf("existing-%02d.txt", i),
"text/plain",
[]byte("existing"),
)
}
lockKey := int64(uuid.New().ID())
_, err := rawDB.ExecContext(ctx, fmt.Sprintf(`
CREATE FUNCTION test_block_chat_file_link() RETURNS trigger AS $$
BEGIN
PERFORM pg_advisory_xact_lock(%d);
RETURN NEW;
END;
$$ LANGUAGE plpgsql;
CREATE TRIGGER test_block_chat_file_link
BEFORE INSERT ON chat_file_links
FOR EACH ROW EXECUTE FUNCTION test_block_chat_file_link();
`, lockKey))
require.NoError(t, err)
barrierConn, err := rawDB.Conn(ctx)
require.NoError(t, err)
barrierReleased := false
t.Cleanup(func() {
if !barrierReleased {
_, _ = barrierConn.ExecContext(context.Background(), "SELECT pg_advisory_unlock($1)", lockKey)
}
_ = barrierConn.Close()
})
_, err = barrierConn.ExecContext(ctx, "SELECT pg_advisory_lock($1)", lockKey)
require.NoError(t, err)
server := &Server{db: db}
attachmentResults := make(chan error, 2)
for i := range 2 {
go func() {
_, err := server.storeChatAttachment(
ctx,
chat,
fmt.Sprintf("concurrent-%d.txt", i),
"attachment.txt",
[]byte("attachment"),
)
attachmentResults <- err
}()
}
var linkWaits, chatLockWaits int
testutil.Eventually(ctx, t, func(ctx context.Context) bool {
err := rawDB.QueryRowContext(ctx, `
SELECT
COUNT(*) FILTER (
WHERE query LIKE '%-- name: LinkChatFilesAfterLock%'
AND wait_event = 'advisory'
),
COUNT(*) FILTER (
WHERE query LIKE '%-- name: LockChatByID%'
AND wait_event_type = 'Lock'
)
FROM pg_stat_activity
WHERE datname = current_database()
AND pid <> pg_backend_pid()
`).Scan(&linkWaits, &chatLockWaits)
return err == nil && linkWaits >= 1 && linkWaits+chatLockWaits == 2
}, testutil.IntervalFast, "wait for both attachment transactions")
require.NoError(t, ctx.Err(), "waiting for attachment transactions")
_, err = barrierConn.ExecContext(ctx, "SELECT pg_advisory_unlock($1)", lockKey)
barrierReleased = true
require.NoError(t, err)
var successes, capRejections int
for range 2 {
select {
case err := <-attachmentResults:
if err == nil {
successes++
continue
}
require.ErrorContains(t, err, fmt.Sprintf("chat already has the maximum of %d linked files", codersdk.MaxChatFileIDs))
capRejections++
case <-ctx.Done():
require.Failf(t, "attachment store did not finish", "context ended: %v", ctx.Err())
}
}
require.Equal(t, 1, successes)
require.Equal(t, 1, capRejections)
files, err := db.GetChatFileMetadataByChatID(ctx, chat.ID)
require.NoError(t, err)
require.Len(t, files, codersdk.MaxChatFileIDs)
var fileCount int
require.NoError(t, rawDB.QueryRowContext(ctx, "SELECT COUNT(*) FROM chat_files").Scan(&fileCount))
require.Equal(t, codersdk.MaxChatFileIDs, fileCount)
}
func expectStoreChatAttachmentInTx(t *testing.T, db, tx *dbmock.MockStore) {
t.Helper()
db.EXPECT().InTx(gomock.Any(), gomock.AssignableToTypeOf(&database.TxOptions{})).DoAndReturn(