diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index ecc58968c2f62..07278858db400 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -1286,6 +1286,7 @@ func (s *MethodTestSuite) TestChats() { dbm.EXPECT().InsertChatMessages(gomock.Any(), arg).Return(msgs, nil).AnyTimes() check.Args(arg).Asserts(chat, policy.ActionUpdate).Returns(msgs) })) + s.Run("InsertChatQueuedMessage", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { chat := testutil.Fake(s.T(), faker, database.Chat{}) arg := testutil.Fake(s.T(), faker, database.InsertChatQueuedMessageParams{ChatID: chat.ID}) diff --git a/coderd/database/dbgen/dbgen.go b/coderd/database/dbgen/dbgen.go index a85e500986cb5..1c0cf593d4457 100644 --- a/coderd/database/dbgen/dbgen.go +++ b/coderd/database/dbgen/dbgen.go @@ -93,6 +93,7 @@ func Chat(t testing.TB, db database.Store, seed database.Chat) database.Chat { } chat, err := db.InsertChat(genCtx, database.InsertChatParams{ + ID: uuid.NullUUID{UUID: seed.ID, Valid: seed.ID != uuid.Nil}, OrganizationID: takeFirst(seed.OrganizationID, uuid.New()), OwnerID: takeFirst(seed.OwnerID, uuid.New()), WorkspaceID: seed.WorkspaceID, diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index d55bbc5a5fadb..b86afa77bd23b 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -10270,6 +10270,7 @@ func (q *sqlQuerier) InsertAgentContextResourcesIntoChat(ctx context.Context, ar const insertChat = `-- name: InsertChat :one WITH inserted_chat AS ( INSERT INTO chats ( + id, organization_id, owner_id, workspace_id, @@ -10287,7 +10288,7 @@ INSERT INTO chats ( dynamic_tools, client_type ) VALUES ( - $1::uuid, + COALESCE($1::uuid, gen_random_uuid()), $2::uuid, $3::uuid, $4::uuid, @@ -10295,14 +10296,15 @@ INSERT INTO chats ( $6::uuid, $7::uuid, $8::uuid, - $9::text, - $10::chat_mode, - $11::chat_plan_mode, - $12::chat_status, - COALESCE($13::uuid[], '{}'::uuid[]), - COALESCE($14::jsonb, '{}'::jsonb), - $15::jsonb, - $16::chat_client_type + $9::uuid, + $10::text, + $11::chat_mode, + $12::chat_plan_mode, + $13::chat_status, + COALESCE($14::uuid[], '{}'::uuid[]), + COALESCE($15::jsonb, '{}'::jsonb), + $16::jsonb, + $17::chat_client_type ) RETURNING id, owner_id, workspace_id, title, status, worker_id, started_at, heartbeat_at, created_at, updated_at, parent_chat_id, root_chat_id, last_model_config_id, archived, last_error, mode, mcp_server_ids, labels, build_id, agent_id, pin_order, last_read_message_id, dynamic_tools, organization_id, plan_mode, client_type, last_turn_summary, user_acl, group_acl, snapshot_version, history_version, queue_version, generation_attempt, retry_state, retry_state_version, runner_id, requires_action_deadline_at, context_aggregate_hash, context_dirty_since, context_dirty_resources, context_error, last_reasoning_effort, compaction_requested_at, summary, summary_generated_at ), @@ -10365,6 +10367,7 @@ FROM chats_expanded ` type InsertChatParams struct { + ID uuid.NullUUID `db:"id" json:"id"` OrganizationID uuid.UUID `db:"organization_id" json:"organization_id"` OwnerID uuid.UUID `db:"owner_id" json:"owner_id"` WorkspaceID uuid.NullUUID `db:"workspace_id" json:"workspace_id"` @@ -10385,6 +10388,7 @@ type InsertChatParams struct { func (q *sqlQuerier) InsertChat(ctx context.Context, arg InsertChatParams) (Chat, error) { row := q.db.QueryRowContext(ctx, insertChat, + arg.ID, arg.OrganizationID, arg.OwnerID, arg.WorkspaceID, diff --git a/coderd/database/queries/chats.sql b/coderd/database/queries/chats.sql index fec4a006aecbc..5b6469b788b69 100644 --- a/coderd/database/queries/chats.sql +++ b/coderd/database/queries/chats.sql @@ -778,6 +778,7 @@ ORDER BY -- name: InsertChat :one WITH inserted_chat AS ( INSERT INTO chats ( + id, organization_id, owner_id, workspace_id, @@ -795,6 +796,7 @@ INSERT INTO chats ( dynamic_tools, client_type ) VALUES ( + COALESCE(sqlc.narg('id')::uuid, gen_random_uuid()), @organization_id::uuid, @owner_id::uuid, sqlc.narg('workspace_id')::uuid, diff --git a/coderd/x/chatd/ARCHITECTURE.md b/coderd/x/chatd/ARCHITECTURE.md index c144316bb8ccc..d0c140a614577 100644 --- a/coderd/x/chatd/ARCHITECTURE.md +++ b/coderd/x/chatd/ARCHITECTURE.md @@ -115,7 +115,7 @@ I don't recommend reading the rest of section thoroughly if this is your first t - `DeleteQueuedMessage(qid)` removes one queued message without changing the active history. - `PromoteQueuedMessage(qid)` makes a queued message the next message to process. It reorders the queue, interrupts active work, cancels pending dynamic-tool action, or promotes into history immediately as required by the input state. - `Interrupt(reason)` requests cancellation of an active generation or closes pending dynamic-tool action. It preserves queued backlog. -- `CompleteRequiresAction(results)` inserts submitted tool-result messages, clears `requires_action_deadline_at`, and lands in `running`. It preserves queued messages. +- `CompleteRequiresAction(results)` inserts submitted tool-result messages followed by any caller-provided suffix messages, clears `requires_action_deadline_at`, and lands in `running`. It preserves queued messages. - `RequestCompaction` records a manual compaction request on an idle chat by setting `compaction_requested_at` and landing in `running` without inserting any message. The chat worker picks the chat up like any other running chat and consumes the request. See [Manual compaction](#manual-compaction). ### Transitions used by the chat worker @@ -128,7 +128,7 @@ I don't recommend reading the rest of section thoroughly if this is your first t - `RecordGenerationAttempt` verifies the chat is still `running`, increments `generation_attempt`, and returns the updated chat snapshot. - `RecordRetryState(payload)` verifies the chat is still `running`, stores the retry payload sent to clients as `retry_state`, and returns the updated chat snapshot. - `FinishTurn` completes the current generation turn atomically. If the queue is empty, it lands in `waiting`. If the queue is non-empty, it removes the queue head, inserts it into history as a user turn, and lands in `running`. -- `FinishError(err)` ends a running chat in `error` and persists `last_error = err`, overwriting any prior stored error. +- `FinishError(err)` parks the chat in `error` and persists `last_error = err`, replacing any previously stored error. It is allowed when an unarchived chat is waiting or running. - `CancelRequiresAction(reason)` closes pending dynamic tool calls with synthetic cancellation tool results, satisfies the pending-action projection, clears `requires_action_deadline_at`, and lands in `running`. - `ReconcileInvalidState` reconciles a chat in an invalid state by setting it to a valid state. Defined in the [Invalid states](#invalid-states) section. @@ -149,6 +149,7 @@ stateDiagram-v2 W --> R0: SendMessage W --> R0: EditMessage W --> R0: RequestCompaction + W --> E0: FinishError W --> XW: SetArchived(true) E0 --> R0: SendMessage diff --git a/coderd/x/chatd/chatstate/toolresults.go b/coderd/x/chatd/chatstate/toolresults.go new file mode 100644 index 0000000000000..ab543abd7e61d --- /dev/null +++ b/coderd/x/chatd/chatstate/toolresults.go @@ -0,0 +1,29 @@ +package chatstate + +import "encoding/json" + +// ValidateToolResults returns the first validation error for results +// submitted against pending dynamic tool calls. +func ValidateToolResults(results []ToolResultInput, pending map[string]string) *ToolResultValidationError { + submitted := make(map[string]struct{}, len(results)) + for _, result := range results { + if _, dup := submitted[result.ToolCallID]; dup { + return &ToolResultValidationError{Cause: ErrToolResultDuplicate, ToolCallID: result.ToolCallID} + } + if !json.Valid(result.Output) { + return &ToolResultValidationError{Cause: ErrToolResultInvalidJSON, ToolCallID: result.ToolCallID} + } + submitted[result.ToolCallID] = struct{}{} + } + for toolCallID := range pending { + if _, ok := submitted[toolCallID]; !ok { + return &ToolResultValidationError{Cause: ErrToolResultMissing, ToolCallID: toolCallID} + } + } + for toolCallID := range submitted { + if _, ok := pending[toolCallID]; !ok { + return &ToolResultValidationError{Cause: ErrToolResultUnexpected, ToolCallID: toolCallID} + } + } + return nil +} diff --git a/coderd/x/chatd/chatstate/toolresults_test.go b/coderd/x/chatd/chatstate/toolresults_test.go new file mode 100644 index 0000000000000..ded7ecbed38b5 --- /dev/null +++ b/coderd/x/chatd/chatstate/toolresults_test.go @@ -0,0 +1,85 @@ +package chatstate_test + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/coderd/x/chatd/chatstate" +) + +func TestValidateToolResults(t *testing.T) { + t.Parallel() + + // Validation order is part of the contract because callers surface only + // the first violation. + pending := map[string]string{"call_a": "execute", "call_b": "read_file"} + resultA := chatstate.ToolResultInput{ToolCallID: "call_a", Output: json.RawMessage(`{"ok":true}`)} + resultB := chatstate.ToolResultInput{ToolCallID: "call_b", Output: json.RawMessage(`"done"`)} + badJSON := chatstate.ToolResultInput{ToolCallID: "call_b", Output: json.RawMessage(`{`)} + resultC := chatstate.ToolResultInput{ToolCallID: "call_c", Output: json.RawMessage(`{}`)} + + cases := []struct { + name string + results []chatstate.ToolResultInput + wantCause error + wantToolCallID string + }{ + { + name: "Complete", + results: []chatstate.ToolResultInput{resultA, resultB}, + }, + { + name: "Duplicate", + results: []chatstate.ToolResultInput{resultA, resultB, resultA}, + wantCause: chatstate.ErrToolResultDuplicate, + wantToolCallID: "call_a", + }, + { + name: "InvalidJSON", + results: []chatstate.ToolResultInput{resultA, badJSON}, + wantCause: chatstate.ErrToolResultInvalidJSON, + wantToolCallID: "call_b", + }, + { + name: "Missing", + results: []chatstate.ToolResultInput{resultA}, + wantCause: chatstate.ErrToolResultMissing, + wantToolCallID: "call_b", + }, + { + name: "Unexpected", + results: []chatstate.ToolResultInput{resultA, resultB, resultC}, + wantCause: chatstate.ErrToolResultUnexpected, + wantToolCallID: "call_c", + }, + { + name: "PerResultRulesOutrankSweeps", + results: []chatstate.ToolResultInput{resultC, badJSON}, + wantCause: chatstate.ErrToolResultInvalidJSON, + wantToolCallID: "call_b", + }, + { + name: "MissingOutranksUnexpected", + results: []chatstate.ToolResultInput{resultA, resultC}, + wantCause: chatstate.ErrToolResultMissing, + wantToolCallID: "call_b", + }, + } + + for _, test := range cases { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + invalid := chatstate.ValidateToolResults(test.results, pending) + if test.wantCause == nil { + require.Nil(t, invalid) + return + } + require.NotNil(t, invalid) + require.ErrorIs(t, invalid, test.wantCause) + require.Equal(t, test.wantToolCallID, invalid.ToolCallID) + }) + } +} diff --git a/coderd/x/chatd/chatstate/transition.go b/coderd/x/chatd/chatstate/transition.go index f7b6c3634c5c2..d6f4a03af904c 100644 --- a/coderd/x/chatd/chatstate/transition.go +++ b/coderd/x/chatd/chatstate/transition.go @@ -81,6 +81,7 @@ var transitionMatrix = map[ExecutionState]map[Transition][]ExecutionState{ TransitionSendMessage: {StateR0}, TransitionEditMessage: {StateR0}, TransitionRequestCompaction: {StateR0}, + TransitionFinishError: {StateE0}, }, StateE0: { TransitionSetArchived: {StateXE0}, diff --git a/coderd/x/chatd/chatstate/transitions.go b/coderd/x/chatd/chatstate/transitions.go index 5fbb4eeba8a57..30c3e31b4a5cb 100644 --- a/coderd/x/chatd/chatstate/transitions.go +++ b/coderd/x/chatd/chatstate/transitions.go @@ -60,6 +60,31 @@ func CreateChat( store database.Store, publisher Publisher, input CreateChatInput, +) (CreateChatResult, error) { + return insertChat(ctx, store, publisher, uuid.NullUUID{}, input) +} + +// CreateChatWithID creates a chat using a caller-minted ID so +// admission-time work can reference the chat before it exists. +func CreateChatWithID( + ctx context.Context, + store database.Store, + publisher Publisher, + chatID uuid.UUID, + input CreateChatInput, +) (CreateChatResult, error) { + if chatID == uuid.Nil { + return CreateChatResult{}, xerrors.New("chatstate: CreateChatWithID called with nil chat ID") + } + return insertChat(ctx, store, publisher, uuid.NullUUID{UUID: chatID, Valid: true}, input) +} + +func insertChat( + ctx context.Context, + store database.Store, + publisher Publisher, + chatID uuid.NullUUID, + input CreateChatInput, ) (CreateChatResult, error) { if store == nil { return CreateChatResult{}, xerrors.New("chatstate: CreateChat called with nil store") @@ -78,6 +103,7 @@ func CreateChat( defer buffer.Discard() err := store.InTx(func(store database.Store) error { chat, err := store.InsertChat(ctx, database.InsertChatParams{ + ID: chatID, OrganizationID: input.OrganizationID, OwnerID: input.OwnerID, WorkspaceID: input.WorkspaceID, @@ -355,48 +381,48 @@ func (tx *Tx) SendMessage(input SendMessageInput) (SendMessageResult, error) { // Idle / empty-queue error: insert directly into history, clear // last_error, leave queue alone. case StateW, StateE0: - return tx.sendMessageDirect(chat, input.Message) + return tx.sendMessageDirect(chat, input) // Error-with-queue: append to tail, promote previous head into // history, clear last_error. case StateE1: - return tx.sendMessageE1(chat, input.Message) + return tx.sendMessageE1(chat, input) // Running with no queue. case StateR0: if input.BusyBehavior == BusyBehaviorInterrupt { - return tx.sendMessageQueueAndSetStatus(chat, input.Message, database.ChatStatusInterrupting, chat.LastError, chat.RequiresActionDeadlineAt) + return tx.sendMessageQueueAndSetStatus(chat, input, database.ChatStatusInterrupting, chat.LastError, chat.RequiresActionDeadlineAt) } - return tx.sendMessageQueueAndSetStatus(chat, input.Message, chat.Status, chat.LastError, chat.RequiresActionDeadlineAt) + return tx.sendMessageQueueAndSetStatus(chat, input, chat.Status, chat.LastError, chat.RequiresActionDeadlineAt) // Running with queue. case StateR1: if input.BusyBehavior == BusyBehaviorInterrupt { - return tx.sendMessageQueueAndSetStatus(chat, input.Message, database.ChatStatusInterrupting, chat.LastError, chat.RequiresActionDeadlineAt) + return tx.sendMessageQueueAndSetStatus(chat, input, database.ChatStatusInterrupting, chat.LastError, chat.RequiresActionDeadlineAt) } - return tx.sendMessageQueueAndSetStatus(chat, input.Message, chat.Status, chat.LastError, chat.RequiresActionDeadlineAt) + return tx.sendMessageQueueAndSetStatus(chat, input, chat.Status, chat.LastError, chat.RequiresActionDeadlineAt) // Interrupting: queue regardless of busy behavior. case StateI0, StateI1: - return tx.sendMessageQueueAndSetStatus(chat, input.Message, chat.Status, chat.LastError, chat.RequiresActionDeadlineAt) + return tx.sendMessageQueueAndSetStatus(chat, input, chat.Status, chat.LastError, chat.RequiresActionDeadlineAt) // Requires-action: queue keeps A*; interrupt cancels pending // dynamic calls and resumes in running. case StateA0, StateA1: if input.BusyBehavior == BusyBehaviorInterrupt { - return tx.sendMessageInterruptRequiresAction(chat, input.Message) + return tx.sendMessageInterruptRequiresAction(chat, input) } - return tx.sendMessageQueueAndSetStatus(chat, input.Message, chat.Status, chat.LastError, chat.RequiresActionDeadlineAt) + return tx.sendMessageQueueAndSetStatus(chat, input, chat.Status, chat.LastError, chat.RequiresActionDeadlineAt) } return SendMessageResult{}, newTransitionError(TransitionSendMessage, from, "unhandled state in SendMessage") } -func (tx *Tx) sendMessageDirect(chat database.Chat, m Message) (SendMessageResult, error) { +func (tx *Tx) sendMessageDirect(chat database.Chat, input SendMessageInput) (SendMessageResult, error) { cancels, err := synthesizePendingToolCancellations(tx.ctx, tx.store, chat, "Tool execution interrupted by new user message", false) if err != nil { return SendMessageResult{}, err } - inserted, err := tx.insertMessages(append(cancels, m)) + inserted, err := tx.insertMessages(append(cancels, input.Message)) if err != nil { return SendMessageResult{}, xerrors.Errorf("insert direct user message: %w", err) } @@ -415,8 +441,8 @@ func (tx *Tx) sendMessageDirect(chat database.Chat, m Message) (SendMessageResul }, nil } -func (tx *Tx) sendMessageE1(chat database.Chat, m Message) (SendMessageResult, error) { - queued, err := tx.insertQueuedMessage(chat.OwnerID, m) +func (tx *Tx) sendMessageE1(chat database.Chat, input SendMessageInput) (SendMessageResult, error) { + queued, err := tx.insertQueuedMessage(chat.OwnerID, input.Message) if err != nil { return SendMessageResult{}, xerrors.Errorf("insert queued: %w", err) } @@ -457,12 +483,12 @@ func (tx *Tx) sendMessageE1(chat database.Chat, m Message) (SendMessageResult, e func (tx *Tx) sendMessageQueueAndSetStatus( chat database.Chat, - m Message, + input SendMessageInput, status database.ChatStatus, lastError pqtype.NullRawMessage, deadline sql.NullTime, ) (SendMessageResult, error) { - queued, err := tx.insertQueuedMessage(chat.OwnerID, m) + queued, err := tx.insertQueuedMessage(chat.OwnerID, input.Message) if err != nil { return SendMessageResult{}, xerrors.Errorf("insert queued: %w", err) } @@ -485,20 +511,32 @@ func (tx *Tx) sendMessageQueueAndSetStatus( }, nil } -func (tx *Tx) sendMessageInterruptRequiresAction(chat database.Chat, m Message) (SendMessageResult, error) { +func (tx *Tx) sendMessageInterruptRequiresAction(chat database.Chat, input SendMessageInput) (SendMessageResult, error) { cancels, err := synthesizePendingToolCancellations(tx.ctx, tx.store, chat, "Tool execution interrupted by user message", true) if err != nil { return SendMessageResult{}, err } - if _, err := tx.insertMessages(cancels); err != nil { + inserted, err := tx.insertMessages(cancels) + if err != nil { return SendMessageResult{}, xerrors.Errorf("insert requires-action cancellations: %w", err) } - return tx.sendMessageQueueAndSetStatus(chat, m, database.ChatStatusRunning, chat.LastError, sql.NullTime{}) + result, err := tx.sendMessageQueueAndSetStatus(chat, input, database.ChatStatusRunning, chat.LastError, sql.NullTime{}) + if err != nil { + return SendMessageResult{}, err + } + // Report the cancellation rows so API clients receive every + // user-visible message this send inserted. + result.InsertedMessages = inserted + return result, nil } // EditMessageInput configures [Tx.EditMessage]. type EditMessageInput struct { - MessageID int64 + MessageID int64 + // SuffixMessages are inserted after the replacement message in the + // same transaction, so a later edit's suffix truncation cleans them + // up together with the rest of the discarded turn. + SuffixMessages []Message CreatedBy uuid.UUID Content pqtype.NullRawMessage ModelConfigIDOverride uuid.NullUUID @@ -511,6 +549,7 @@ type EditMessageResult struct { DeletedMessageIDs []int64 DeletedQueuedMessageIDs []int64 CancellationMessages []database.ChatMessage + SuffixMessages []database.ChatMessage } // EditMessage replaces an earlier user message and discards the @@ -564,7 +603,6 @@ func (tx *Tx) EditMessage(input EditMessageInput) (EditMessageResult, error) { }); err != nil { return EditMessageResult{}, xerrors.Errorf("soft-delete suffix: %w", err) } - cancels, err := synthesizePendingToolCancellations(tx.ctx, tx.store, chat, "Tool execution interrupted by message edit", false) if err != nil { return EditMessageResult{}, err @@ -599,6 +637,10 @@ func (tx *Tx) EditMessage(input EditMessageInput) (EditMessageResult, error) { if len(insertedReplacement) == 1 { replacementRow = insertedReplacement[0] } + insertedSuffix, err := tx.insertMessages(input.SuffixMessages) + if err != nil { + return EditMessageResult{}, xerrors.Errorf("insert edit suffix messages: %w", err) + } deletedQueuedIDs, err := tx.clearQueue() if err != nil { @@ -620,6 +662,7 @@ func (tx *Tx) EditMessage(input EditMessageInput) (EditMessageResult, error) { DeletedMessageIDs: deletedIDs, DeletedQueuedMessageIDs: deletedQueuedIDs, CancellationMessages: cancellationMessages, + SuffixMessages: insertedSuffix, }, nil } @@ -877,9 +920,10 @@ type ToolResultInput struct { // CompleteRequiresActionInput configures [Tx.CompleteRequiresAction]. type CompleteRequiresActionInput struct { - CreatedBy uuid.UUID - ModelConfigID uuid.UUID - Results []ToolResultInput + CreatedBy uuid.UUID + ModelConfigID uuid.UUID + Results []ToolResultInput + SuffixMessages []Message } // CompleteRequiresActionResult is returned by [Tx.CompleteRequiresAction]. @@ -899,41 +943,11 @@ func (tx *Tx) CompleteRequiresAction(input CompleteRequiresActionInput) (Complet if err != nil { return CompleteRequiresActionResult{}, err } - submitted := make(map[string]ToolResultInput, len(input.Results)) - for _, r := range input.Results { - if _, dup := submitted[r.ToolCallID]; dup { - return CompleteRequiresActionResult{}, newTransitionErrorWithCause( - TransitionCompleteRequiresAction, from, - &ToolResultValidationError{Cause: ErrToolResultDuplicate, ToolCallID: r.ToolCallID}, - "duplicate tool_call_id submitted", - ) - } - if !json.Valid(r.Output) { - return CompleteRequiresActionResult{}, newTransitionErrorWithCause( - TransitionCompleteRequiresAction, from, - &ToolResultValidationError{Cause: ErrToolResultInvalidJSON, ToolCallID: r.ToolCallID}, - "tool result output is not valid JSON", - ) - } - submitted[r.ToolCallID] = r - } - for id := range pending { - if _, ok := submitted[id]; !ok { - return CompleteRequiresActionResult{}, newTransitionErrorWithCause( - TransitionCompleteRequiresAction, from, - &ToolResultValidationError{Cause: ErrToolResultMissing, ToolCallID: id}, - "submitted tool results do not match pending tool calls", - ) - } - } - for id := range submitted { - if _, ok := pending[id]; !ok { - return CompleteRequiresActionResult{}, newTransitionErrorWithCause( - TransitionCompleteRequiresAction, from, - &ToolResultValidationError{Cause: ErrToolResultUnexpected, ToolCallID: id}, - "submitted tool_call_id does not match a pending dynamic tool call", - ) - } + if invalid := ValidateToolResults(input.Results, pending); invalid != nil { + return CompleteRequiresActionResult{}, newTransitionErrorWithCause( + TransitionCompleteRequiresAction, from, invalid, + toolResultTransitionMessage(invalid.Cause), + ) } messages := make([]Message, 0, len(input.Results)) for _, r := range input.Results { @@ -957,7 +971,7 @@ func (tx *Tx) CompleteRequiresAction(input CompleteRequiresActionInput) (Complet ContentVersion: chatprompt.CurrentContentVersion, }) } - inserted, err := tx.insertMessages(messages) + inserted, err := tx.insertMessages(append(messages, input.SuffixMessages...)) if err != nil { return CompleteRequiresActionResult{}, xerrors.Errorf("insert tool results: %w", err) } @@ -976,6 +990,21 @@ func (tx *Tx) CompleteRequiresAction(input CompleteRequiresActionInput) (Complet }, nil } +func toolResultTransitionMessage(cause error) string { + switch { + case errors.Is(cause, ErrToolResultDuplicate): + return "duplicate tool_call_id submitted" + case errors.Is(cause, ErrToolResultInvalidJSON): + return "tool result output is not valid JSON" + case errors.Is(cause, ErrToolResultMissing): + return "submitted tool results do not match pending tool calls" + case errors.Is(cause, ErrToolResultUnexpected): + return "submitted tool_call_id does not match a pending dynamic tool call" + default: + return "submitted tool results are invalid" + } +} + // AcquireInput configures [Tx.Acquire]. type AcquireInput struct { WorkerID uuid.UUID @@ -1386,6 +1415,7 @@ type FinishErrorInput struct { type FinishErrorResult struct{} // FinishError parks the chat in error with the supplied last_error. +// It is allowed when an unarchived chat is waiting or running. func (tx *Tx) FinishError(input FinishErrorInput) (FinishErrorResult, error) { chat, _, err := tx.requireFromAllowed(TransitionFinishError) if err != nil { diff --git a/coderd/x/chatd/chatstate/transitions_matrix_test.go b/coderd/x/chatd/chatstate/transitions_matrix_test.go index 1f1597f498bf0..15258c53c26fd 100644 --- a/coderd/x/chatd/chatstate/transitions_matrix_test.go +++ b/coderd/x/chatd/chatstate/transitions_matrix_test.go @@ -868,6 +868,7 @@ func matrixCases() []transitionCaseSpec { // FinishError cases. finishErrorCase(chatstate.StateR0, chatstate.StateE0), finishErrorCase(chatstate.StateR1, chatstate.StateE1), + finishErrorCase(chatstate.StateW, chatstate.StateE0), // ReconcileInvalidState cases: Invalid with empty queue // lands in E0; Invalid with non-empty queue lands in E1. diff --git a/coderd/x/chatd/chatstate/transitions_test.go b/coderd/x/chatd/chatstate/transitions_test.go index c4df476d7ba99..6b5bfb1c3af43 100644 --- a/coderd/x/chatd/chatstate/transitions_test.go +++ b/coderd/x/chatd/chatstate/transitions_test.go @@ -413,6 +413,36 @@ func TestSendMessageQueueCapRejectsQueueAppend(t *testing.T) { "failed queue append must not bump queue_version") } +func TestSendMessageInterruptRequiresActionReturnsCancellations(t *testing.T) { + t.Parallel() + f := newTestFixture(t) + ctx := testutil.Context(t, testutil.WaitShort) + seeded := seedAOrA1(t, f, 0, "interrupt_cancels") + require.Equal(t, chatstate.StateA0, f.classify(ctx, t, seeded.chatID)) + + m := chatstate.NewChatMachine(f.DB, f.Pub, seeded.chatID) + var send chatstate.SendMessageResult + require.NoError(t, m.Update(ctx, func(tx *chatstate.Tx, store database.Store) error { + var err error + send, err = tx.SendMessage(chatstate.SendMessageInput{ + Message: userTextMessage("interrupt", f.User.ID, f.Model.ID), + BusyBehavior: chatstate.BusyBehaviorInterrupt, + }) + return err + })) + + require.NotNil(t, send.QueuedMessage, "interrupt from A0 queues the user message") + require.Len(t, send.InsertedMessages, 1, + "the synthetic tool cancellation must be reported to callers") + cancel := send.InsertedMessages[0] + require.Equal(t, database.ChatMessageRoleTool, cancel.Role) + require.Equal(t, database.ChatMessageVisibilityBoth, cancel.Visibility) + parts, err := chatprompt.ParseContent(cancel) + require.NoError(t, err) + require.Len(t, parts, 1) + require.Equal(t, seeded.pendingToolCallID, parts[0].ToolCallID) +} + // TestEditMessageNonUserReturnsSentinel asserts that editing a // non-user message returns chatstate.ErrEditedMessageNotUser via // the TransitionError cause chain, and still matches the generic @@ -461,6 +491,82 @@ func TestEditMessageNonUserReturnsSentinel(t *testing.T) { "ErrEditedMessageNotUser still matches the generic transition sentinel") } +func TestEditMessageInsertsSuffixMessages(t *testing.T) { + t.Parallel() + f := newTestFixture(t) + ctx := testutil.Context(t, testutil.WaitShort) + created := createTestChat(t, f) + m := chatstate.NewChatMachine(f.DB, f.Pub, created.Chat.ID) + + target := userTextMessage("original prompt", f.User.ID, f.Model.ID) + + var targetID int64 + require.NoError(t, m.Update(ctx, func(tx *chatstate.Tx, store database.Store) error { + step, err := tx.CommitStep(chatstate.CommitStepInput{ + Messages: []chatstate.Message{target}, + }) + if err != nil { + return err + } + require.Len(t, step.InsertedMessages, 1) + targetID = step.InsertedMessages[0].ID + return nil + })) + + firstSuffix := userTextMessage("first notice", f.User.ID, f.Model.ID) + firstSuffix.Role = database.ChatMessageRoleSystem + secondSuffix := userTextMessage("second notice", f.User.ID, f.Model.ID) + secondSuffix.Role = database.ChatMessageRoleSystem + rawContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{ + codersdk.ChatMessageText("edited prompt"), + }) + require.NoError(t, err) + + var result chatstate.EditMessageResult + require.NoError(t, m.Update(ctx, func(tx *chatstate.Tx, store database.Store) error { + result, err = tx.EditMessage(chatstate.EditMessageInput{ + MessageID: targetID, + SuffixMessages: []chatstate.Message{firstSuffix, secondSuffix}, + CreatedBy: f.User.ID, + Content: rawContent, + }) + return err + })) + + require.Contains(t, result.DeletedMessageIDs, targetID) + require.Len(t, result.SuffixMessages, 2) + assertChatMessageText(t, result.SuffixMessages[0], "first notice") + assertChatMessageText(t, result.SuffixMessages[1], "second notice") + require.Greater(t, result.SuffixMessages[0].ID, result.ReplacementMessage.ID, + "the suffix messages must follow the replacement") + require.Greater(t, result.SuffixMessages[1].ID, result.SuffixMessages[0].ID, + "suffix message IDs must ascend with the input array") + _, err = f.DB.GetChatMessageByID(ctx, result.ReplacementMessage.ID) + require.NoError(t, err, "the replacement message must stay active") + + // Reloading is what the generation path does, so the input order has to + // survive the round trip and not just the returned slice. + history, err := f.DB.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ + ChatID: created.Chat.ID, + }) + require.NoError(t, err) + require.GreaterOrEqual(t, len(history), 3) + tail := history[len(history)-3:] + require.Equal(t, []int64{ + result.ReplacementMessage.ID, + result.SuffixMessages[0].ID, + result.SuffixMessages[1].ID, + }, []int64{tail[0].ID, tail[1].ID, tail[2].ID}) + assertChatMessageText(t, tail[0], "edited prompt") + assertChatMessageText(t, tail[1], "first notice") + assertChatMessageText(t, tail[2], "second notice") + + for _, message := range history { + require.NotEqual(t, targetID, message.ID, + "the edited message must not stay in active history") + } +} + // TestTransitionAbandon_RejectsUnowned verifies that calling Abandon // on a chat the runner does not own returns ErrTransitionNotAllowed // wrapped in a TransitionError that records the loaded from-state,