mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user