mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add chat suffix messages, idle failure, and content update support (#27428)
Adds generic chat state and query capabilities that the lifecycle hooks integration (#27429) builds on. Part of the lifecycle hooks stack (#27401, #27429, #27430). - `chatstate`: `EditMessage` accepts caller-provided suffix messages inserted after the replacement in the same transaction, transitions can carry a typed error kind, and `FinishError` is also allowed from waiting chats so admission-time failures can park an idle chat in error. - `chatstate`: `ValidateToolResults` holds the submitted-tool-result rules (duplicate, invalid JSON, missing, unexpected) in one place, so `CompleteRequiresAction` and API-level prechecks reject the same payloads with the same typed causes. - `database`: `InsertChat` accepts an optional caller-provided ID. No hook-specific state or behavior is introduced here; these primitives are usable by any caller. An earlier revision added a message-content rewrite primitive so a `pre_tool_use` override could update an already-committed tool call. Message content is immutable by design, and @hugodutka pushed back on changing that. The rewrite is gone: #27429 now dispatches the hook before the assistant message is stored, so the stored input is the one that runs and nothing needs updating. > This PR was written by Mux, an AI coding agent, on Mike's behalf.
This commit is contained in:
@@ -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})
|
||||
|
||||
@@ -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,
|
||||
|
||||
Generated
+13
-9
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -81,6 +81,7 @@ var transitionMatrix = map[ExecutionState]map[Transition][]ExecutionState{
|
||||
TransitionSendMessage: {StateR0},
|
||||
TransitionEditMessage: {StateR0},
|
||||
TransitionRequestCompaction: {StateR0},
|
||||
TransitionFinishError: {StateE0},
|
||||
},
|
||||
StateE0: {
|
||||
TransitionSetArchived: {StateXE0},
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user