mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: wire chat lifecycle hooks into chatd (#27429)
Wires chat lifecycle hooks into chatd, gated by the `agent-lifecycle-hooks` experiment. Part of the lifecycle hooks stack (#27401, #27428, #27430). See `docs/admin/setup/chat-lifecycle-hooks.md` for the consumer-facing contract. ## Summary When a hook URL is configured, chatd dispatches `session_start`, `user_prompt_submit`, `pre_tool_use`, `post_tool_use`, `pre_compact`, `post_compact`, and `stop` events to the consumer and applies its responses. ## Design - **Stateless**: Coder stores no hook dispatch or decision state. Delivery is at least once; consumers deduplicate on stable payload identifiers (chat ID, event type, tool-use ID) and answer duplicates with the same decision. - **Admission-time prompt effects**: `user_prompt_submit` dispatches exactly once per submission (create, send, queue, edit, subagent spawn) and folds its effects into the stored prompt as typed message parts: original-or-overridden user parts, then model-only `hook-context`, then a user-visible `hook-notice`. Hook context is stripped from every client-facing conversion; hook notices are excluded from model prompts. The server rejects hook parts in client-submitted content. - **Tool gating**: `pre_tool_use` allow can override tool input; deny becomes a synthetic denied tool result, with any returned model context persisted as a model-only transcript row so it never reaches clients. The denial text identifies an external policy (the deployment's lifecycle hook) as the source and marks the decision as persistent, so the model explains the denial instead of retrying it or misreporting it as an infrastructure failure. - **Fail closed**: a dispatch failure rejects the triggering request or moves the chat to the error state in the same transaction as the affected step, so a runnable state is never published with unapproved content. - **Admission before persistence**: `pre_tool_use` is dispatched for the calls the model produced, before the assistant message is stored. See "Staged tool admission" below. - **Fresh dispatch per tool call**: every non-provider-executed tool call is decided by its own `pre_tool_use` dispatch; Coder never reuses an earlier decision on the consumer's behalf. Retries re-dispatch the same logical event. ## Structure All hook dispatch flows through one seam: entry points build a `chathooks.Chat` (chat identity) and a `chathooks.Message` (event details) and call `Trigger.Trigger`, the only component that talks to the dispatcher. The integration lives in the `coderd/x/chatd/chathooks` subpackage, split by responsibility: - `trigger.go`: the trigger seam; builds the wire envelope per event, normalizes deny into a typed error, and holds the package's single enabled-check. - `effects.go`: pure conversion of hook results into transcript rows and prompt parts. - `errors.go`: failure classification (dispatch error messages, denial mapping, tool-result dispatch-failure scanning). - `tooluse.go`: the tool-call gate (`pre_tool_use` preflight, `post_tool_use` payloads, applying admitted input to the step). Server-bound glue stays in `coderd/x/chatd/hook_server.go`: the chat-parking dispatch error handlers, the step-commit row insertion wrappers, and the dynamic post-tool-use state loader, which depends on chatd validation types. This PR adopts the `codersdk/x/agenthooks` and `coderd/x/agenthooks/dispatch` import paths introduced at the tip of #27401; intermediate commits still reference the pre-move paths and are not individually buildable. ## Staged tool admission `pre_tool_use` originally ran at tool execution time, which is after the assistant message carrying the tool call was already committed. An `input_override` therefore had to rewrite stored message content in place. @hugodutka pointed out that chatd treats message content as immutable, and that the rewrite was a shortcut rather than a requirement. It was also a correctness problem in its own right: the rewrite only updated the database, so the transcript could show one input while a different one had executed. The hook now runs before the step is persisted: ```text provider stream ends (tool calls complete, in memory) -> pre_tool_use dispatch per call -> ONE transaction: assistant row with admitted inputs, synthetic denials, hook rows -> execute ``` The step is inserted once, carrying the input the tool runs with. `UpdateChatMessageContentByID` and `Tx.UpdateMessageContent` are deleted from #27428, so message content stays immutable. Two consequences, both intentional: - **Clients converge rather than wait.** Tool-call parts still stream live, so a rewritten call briefly shows the model's proposed input before the committed message replaces it. The chat store already clears stream state when an assistant message arrives, so the stored input wins with no frontend change and no added latency before tool cards appear. - **A call already in history was already admitted.** Execution consumes the stored input instead of dispatching a second decision, which keeps one dispatch and one set of hook effects per call. A consumer policy change between admission and execution applies to later calls, not to calls already admitted. The per-chat debug endpoint still records the provider's original tool input. Its purpose is to report provider behavior, and it requires an explicit per-chat debug flag; the invariant here covers the transcript. ## Configuration Adds `chat-hook-url`, `chat-hook-secret`, `chat-hook-timeout`, and `chat-hook-enabled` deployment options with startup validation. The flags are hidden from `coder server --help` while the feature is experimental; the setup guide documents them. ## Tool input validation Built-in tool arguments reach a consumer as raw JSON with key spelling preserved, but the tools decode those bytes with Go, which matches struct fields case-insensitively and keeps the last match. A policy reading `path` could therefore authorize one value while the tool executed another, and a lone case variant such as `{"PATH":"/secret"}` was invisible to a policy checking for `path`. Coder now rejects a built-in tool call whose input repeats a key or spells a schema property with different capitalization, before the `pre_tool_use` dispatch, so a consumer is never asked to authorize bytes whose meaning depends on the reader. Rejected calls produce an error result the model can retry; unambiguous calls in the same batch still run. A consumer-authored `input_override` is rechecked after the dispatch and fails the turn closed, because the model cannot correct it. Dynamic and MCP inputs are excluded because the client and the workspace agent execute those calls rather than coderd. Two paths needed more than a schema check. Execution resolves a deprecated tool name to its canonical tool, so validation resolves aliases first. The `edit_files` decoder also reads `search` and `replace`, which its schema does not advertise, so those aliases are now matched exactly and their case variants ignored. A hook denial now returns a structured 403 carrying `kind: "hook_denied"`, mirroring the dispatch-failure response that already carries its own kind. Without it a client cannot tell a policy decision apart from a generic failure, and the chat UI titled a denial "Request failed". Adding a kind needs no migration: `ChatErrorKind` is persisted only inside the JSONB `chats.last_error` column, whose decoder accepts unknown kinds. The hook docs also correct the tool-input convergence window. A batch dispatches sequentially before the assistant row commits, so the original input stays visible for a span that scales with the number of tool calls in the step rather than a single hook timeout. > This PR was written by Mux, an AI coding agent, on Mike's behalf.
This commit is contained in:
Generated
+52
-7
@@ -17370,6 +17370,18 @@ const docTemplate = `{
|
||||
},
|
||||
"debug_logging_enabled": {
|
||||
"type": "boolean"
|
||||
},
|
||||
"hook_enabled": {
|
||||
"type": "boolean"
|
||||
},
|
||||
"hook_secret": {
|
||||
"type": "string"
|
||||
},
|
||||
"hook_timeout": {
|
||||
"type": "integer"
|
||||
},
|
||||
"hook_url": {
|
||||
"$ref": "#/definitions/serpent.URL"
|
||||
}
|
||||
}
|
||||
},
|
||||
@@ -17638,7 +17650,9 @@ const docTemplate = `{
|
||||
"usage_limit",
|
||||
"missing_key",
|
||||
"provider_disabled",
|
||||
"content_filter"
|
||||
"content_filter",
|
||||
"hook_dispatch_failed",
|
||||
"hook_denied"
|
||||
],
|
||||
"x-enum-varnames": [
|
||||
"ChatErrorKindGeneric",
|
||||
@@ -17651,7 +17665,9 @@ const docTemplate = `{
|
||||
"ChatErrorKindUsageLimit",
|
||||
"ChatErrorKindMissingKey",
|
||||
"ChatErrorKindProviderDisabled",
|
||||
"ChatErrorKindContentFilter"
|
||||
"ChatErrorKindContentFilter",
|
||||
"ChatErrorKindHookDispatchFailed",
|
||||
"ChatErrorKindHookDenied"
|
||||
]
|
||||
},
|
||||
"codersdk.ChatFileMetadata": {
|
||||
@@ -17995,7 +18011,9 @@ const docTemplate = `{
|
||||
"file",
|
||||
"file-reference",
|
||||
"context-file",
|
||||
"skill"
|
||||
"skill",
|
||||
"hook-context",
|
||||
"hook-notice"
|
||||
],
|
||||
"x-enum-varnames": [
|
||||
"ChatMessagePartTypeText",
|
||||
@@ -18006,7 +18024,9 @@ const docTemplate = `{
|
||||
"ChatMessagePartTypeFile",
|
||||
"ChatMessagePartTypeFileReference",
|
||||
"ChatMessagePartTypeContextFile",
|
||||
"ChatMessagePartTypeSkill"
|
||||
"ChatMessagePartTypeSkill",
|
||||
"ChatMessagePartTypeHookContext",
|
||||
"ChatMessagePartTypeHookNotice"
|
||||
]
|
||||
},
|
||||
"codersdk.ChatMessageRole": {
|
||||
@@ -18764,6 +18784,13 @@ const docTemplate = `{
|
||||
"message": {
|
||||
"$ref": "#/definitions/codersdk.ChatMessage"
|
||||
},
|
||||
"messages": {
|
||||
"description": "Messages contains all user-visible messages inserted by the send, in\ninsertion order. A queued send on an errored chat may promote the\nprevious queue head, so clients must upsert the full batch.",
|
||||
"type": "array",
|
||||
"items": {
|
||||
"$ref": "#/definitions/codersdk.ChatMessage"
|
||||
}
|
||||
},
|
||||
"queued": {
|
||||
"type": "boolean"
|
||||
},
|
||||
@@ -20200,9 +20227,23 @@ const docTemplate = `{
|
||||
"codersdk.EditChatMessageResponse": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"deleted_message_ids": {
|
||||
"description": "DeletedMessageIDs holds the IDs of previously visible messages the\nedit removed, including stale hook notices from the edited turn.\nClients should drop them from local caches.",
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "integer"
|
||||
}
|
||||
},
|
||||
"message": {
|
||||
"$ref": "#/definitions/codersdk.ChatMessage"
|
||||
},
|
||||
"messages": {
|
||||
"description": "Messages holds every user-visible message inserted by the edit, in\ninsertion order. Hook-generated suffix messages may follow Message,\nso clients must upsert the full batch.",
|
||||
"type": "array",
|
||||
"items": {
|
||||
"$ref": "#/definitions/codersdk.ChatMessage"
|
||||
}
|
||||
},
|
||||
"warnings": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
@@ -20276,11 +20317,13 @@ const docTemplate = `{
|
||||
"ai-gateway-seat-exclusion",
|
||||
"ai-gateway-cost-control",
|
||||
"chat-advisor",
|
||||
"chat-virtual-desktop"
|
||||
"chat-virtual-desktop",
|
||||
"agent-lifecycle-hooks"
|
||||
],
|
||||
"x-enum-comments": {
|
||||
"ExperimentAIGatewayCostControl": "Enables AI Gateway cost control functionality.",
|
||||
"ExperimentAIGatewaySeatExclusion": "Excludes AI Gateway (AI Bridge) usage from AI Governance seat consumption.",
|
||||
"ExperimentAgentLifecycleHooks": "Enables chat lifecycle hook webhooks for agent chats.",
|
||||
"ExperimentAutoFillParameters": "This should not be taken out of experiments until we have redesigned the feature.",
|
||||
"ExperimentChatAdvisor": "Enables the advisor tool for root agent chats.",
|
||||
"ExperimentChatVirtualDesktop": "Enables virtual desktop and computer use provider for agents.",
|
||||
@@ -20308,7 +20351,8 @@ const docTemplate = `{
|
||||
"Excludes AI Gateway (AI Bridge) usage from AI Governance seat consumption.",
|
||||
"Enables AI Gateway cost control functionality.",
|
||||
"Enables the advisor tool for root agent chats.",
|
||||
"Enables virtual desktop and computer use provider for agents."
|
||||
"Enables virtual desktop and computer use provider for agents.",
|
||||
"Enables chat lifecycle hook webhooks for agent chats."
|
||||
],
|
||||
"x-enum-varnames": [
|
||||
"ExperimentExample",
|
||||
@@ -20324,7 +20368,8 @@ const docTemplate = `{
|
||||
"ExperimentAIGatewaySeatExclusion",
|
||||
"ExperimentAIGatewayCostControl",
|
||||
"ExperimentChatAdvisor",
|
||||
"ExperimentChatVirtualDesktop"
|
||||
"ExperimentChatVirtualDesktop",
|
||||
"ExperimentAgentLifecycleHooks"
|
||||
]
|
||||
},
|
||||
"codersdk.ExternalAPIKeyScopes": {
|
||||
|
||||
Generated
+52
-7
@@ -15607,6 +15607,18 @@
|
||||
},
|
||||
"debug_logging_enabled": {
|
||||
"type": "boolean"
|
||||
},
|
||||
"hook_enabled": {
|
||||
"type": "boolean"
|
||||
},
|
||||
"hook_secret": {
|
||||
"type": "string"
|
||||
},
|
||||
"hook_timeout": {
|
||||
"type": "integer"
|
||||
},
|
||||
"hook_url": {
|
||||
"$ref": "#/definitions/serpent.URL"
|
||||
}
|
||||
}
|
||||
},
|
||||
@@ -15864,7 +15876,9 @@
|
||||
"usage_limit",
|
||||
"missing_key",
|
||||
"provider_disabled",
|
||||
"content_filter"
|
||||
"content_filter",
|
||||
"hook_dispatch_failed",
|
||||
"hook_denied"
|
||||
],
|
||||
"x-enum-varnames": [
|
||||
"ChatErrorKindGeneric",
|
||||
@@ -15877,7 +15891,9 @@
|
||||
"ChatErrorKindUsageLimit",
|
||||
"ChatErrorKindMissingKey",
|
||||
"ChatErrorKindProviderDisabled",
|
||||
"ChatErrorKindContentFilter"
|
||||
"ChatErrorKindContentFilter",
|
||||
"ChatErrorKindHookDispatchFailed",
|
||||
"ChatErrorKindHookDenied"
|
||||
]
|
||||
},
|
||||
"codersdk.ChatFileMetadata": {
|
||||
@@ -16215,7 +16231,9 @@
|
||||
"file",
|
||||
"file-reference",
|
||||
"context-file",
|
||||
"skill"
|
||||
"skill",
|
||||
"hook-context",
|
||||
"hook-notice"
|
||||
],
|
||||
"x-enum-varnames": [
|
||||
"ChatMessagePartTypeText",
|
||||
@@ -16226,7 +16244,9 @@
|
||||
"ChatMessagePartTypeFile",
|
||||
"ChatMessagePartTypeFileReference",
|
||||
"ChatMessagePartTypeContextFile",
|
||||
"ChatMessagePartTypeSkill"
|
||||
"ChatMessagePartTypeSkill",
|
||||
"ChatMessagePartTypeHookContext",
|
||||
"ChatMessagePartTypeHookNotice"
|
||||
]
|
||||
},
|
||||
"codersdk.ChatMessageRole": {
|
||||
@@ -16952,6 +16972,13 @@
|
||||
"message": {
|
||||
"$ref": "#/definitions/codersdk.ChatMessage"
|
||||
},
|
||||
"messages": {
|
||||
"description": "Messages contains all user-visible messages inserted by the send, in\ninsertion order. A queued send on an errored chat may promote the\nprevious queue head, so clients must upsert the full batch.",
|
||||
"type": "array",
|
||||
"items": {
|
||||
"$ref": "#/definitions/codersdk.ChatMessage"
|
||||
}
|
||||
},
|
||||
"queued": {
|
||||
"type": "boolean"
|
||||
},
|
||||
@@ -18338,9 +18365,23 @@
|
||||
"codersdk.EditChatMessageResponse": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"deleted_message_ids": {
|
||||
"description": "DeletedMessageIDs holds the IDs of previously visible messages the\nedit removed, including stale hook notices from the edited turn.\nClients should drop them from local caches.",
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "integer"
|
||||
}
|
||||
},
|
||||
"message": {
|
||||
"$ref": "#/definitions/codersdk.ChatMessage"
|
||||
},
|
||||
"messages": {
|
||||
"description": "Messages holds every user-visible message inserted by the edit, in\ninsertion order. Hook-generated suffix messages may follow Message,\nso clients must upsert the full batch.",
|
||||
"type": "array",
|
||||
"items": {
|
||||
"$ref": "#/definitions/codersdk.ChatMessage"
|
||||
}
|
||||
},
|
||||
"warnings": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
@@ -18410,11 +18451,13 @@
|
||||
"ai-gateway-seat-exclusion",
|
||||
"ai-gateway-cost-control",
|
||||
"chat-advisor",
|
||||
"chat-virtual-desktop"
|
||||
"chat-virtual-desktop",
|
||||
"agent-lifecycle-hooks"
|
||||
],
|
||||
"x-enum-comments": {
|
||||
"ExperimentAIGatewayCostControl": "Enables AI Gateway cost control functionality.",
|
||||
"ExperimentAIGatewaySeatExclusion": "Excludes AI Gateway (AI Bridge) usage from AI Governance seat consumption.",
|
||||
"ExperimentAgentLifecycleHooks": "Enables chat lifecycle hook webhooks for agent chats.",
|
||||
"ExperimentAutoFillParameters": "This should not be taken out of experiments until we have redesigned the feature.",
|
||||
"ExperimentChatAdvisor": "Enables the advisor tool for root agent chats.",
|
||||
"ExperimentChatVirtualDesktop": "Enables virtual desktop and computer use provider for agents.",
|
||||
@@ -18442,7 +18485,8 @@
|
||||
"Excludes AI Gateway (AI Bridge) usage from AI Governance seat consumption.",
|
||||
"Enables AI Gateway cost control functionality.",
|
||||
"Enables the advisor tool for root agent chats.",
|
||||
"Enables virtual desktop and computer use provider for agents."
|
||||
"Enables virtual desktop and computer use provider for agents.",
|
||||
"Enables chat lifecycle hook webhooks for agent chats."
|
||||
],
|
||||
"x-enum-varnames": [
|
||||
"ExperimentExample",
|
||||
@@ -18458,7 +18502,8 @@
|
||||
"ExperimentAIGatewaySeatExclusion",
|
||||
"ExperimentAIGatewayCostControl",
|
||||
"ExperimentChatAdvisor",
|
||||
"ExperimentChatVirtualDesktop"
|
||||
"ExperimentChatVirtualDesktop",
|
||||
"ExperimentAgentLifecycleHooks"
|
||||
]
|
||||
},
|
||||
"codersdk.ExternalAPIKeyScopes": {
|
||||
|
||||
@@ -98,6 +98,7 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/workspacestats"
|
||||
"github.com/coder/coder/v2/coderd/wsbuilder"
|
||||
"github.com/coder/coder/v2/coderd/wsbuildorchestrator"
|
||||
"github.com/coder/coder/v2/coderd/x/agenthooks/dispatch"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/mcpclient"
|
||||
@@ -879,6 +880,27 @@ func New(options *Options) *API {
|
||||
// the chat daemon stays nil and chat HTTP handlers return a
|
||||
// service-unavailable error with a clear remediation message.
|
||||
if options.DeploymentValues.AI.BridgeConfig.Enabled.Value() {
|
||||
var hookDispatcher *dispatch.Dispatcher
|
||||
chatConfig := options.DeploymentValues.AI.Chat
|
||||
hooksConfigured := chatConfig.HookURL.String() != "" && chatConfig.HookEnabled.Value()
|
||||
hooksExperimentEnabled := experiments.Enabled(codersdk.ExperimentAgentLifecycleHooks)
|
||||
if hooksConfigured && !hooksExperimentEnabled {
|
||||
options.Logger.Warn(ctx, "chat lifecycle hooks are configured but inactive; enable the agent-lifecycle-hooks experiment to activate them",
|
||||
slog.F("experiment", codersdk.ExperimentAgentLifecycleHooks),
|
||||
)
|
||||
}
|
||||
if hooksConfigured && hooksExperimentEnabled {
|
||||
hookDispatcher = dispatch.New(
|
||||
options.Logger,
|
||||
nil,
|
||||
chatConfig.HookURL.String(),
|
||||
chatConfig.HookSecret.Value(),
|
||||
chatConfig.HookTimeout.Value(),
|
||||
api.DeploymentID,
|
||||
buildinfo.Version(),
|
||||
options.PrometheusRegistry,
|
||||
)
|
||||
}
|
||||
api.chatDaemon = chatd.New(options.Pubsub, chatd.Config{
|
||||
Logger: options.Logger.Named("chatd"),
|
||||
Database: options.Database,
|
||||
@@ -898,6 +920,7 @@ func New(options *Options) *API {
|
||||
StartWorkspace: api.chatStartWorkspace,
|
||||
StopWorkspace: api.chatStopWorkspace,
|
||||
WebpushDispatcher: options.WebPushDispatcher,
|
||||
HookDispatcher: hookDispatcher,
|
||||
UsageTracker: options.WorkspaceUsageTracker,
|
||||
PrometheusRegistry: options.PrometheusRegistry,
|
||||
OIDCTokenSource: oidcMCPSrc,
|
||||
|
||||
@@ -1641,11 +1641,17 @@ func chatMessageParts(m database.ChatMessage) ([]codersdk.ChatMessagePart, error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Strip internal-only fields before API responses.
|
||||
// Strip internal-only fields before API responses. Hook context
|
||||
// parts are model-only and must never reach clients.
|
||||
filtered := parts[:0]
|
||||
for i := range parts {
|
||||
if parts[i].Type == codersdk.ChatMessagePartTypeHookContext {
|
||||
continue
|
||||
}
|
||||
parts[i].StripInternal()
|
||||
filtered = append(filtered, parts[i])
|
||||
}
|
||||
return parts, nil
|
||||
return filtered, nil
|
||||
}
|
||||
|
||||
func nullUUIDPtr(v uuid.NullUUID) *uuid.UUID {
|
||||
|
||||
+135
-27
@@ -47,9 +47,11 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/util/xjson"
|
||||
"github.com/coder/coder/v2/coderd/workspaceapps"
|
||||
"github.com/coder/coder/v2/coderd/wsbuilder"
|
||||
"github.com/coder/coder/v2/coderd/x/agenthooks/dispatch"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/agentselect"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chaterror"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chathooks"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
|
||||
@@ -110,6 +112,39 @@ func writeChatUsageLimitExceeded(
|
||||
})
|
||||
}
|
||||
|
||||
// Avoid returning raw dispatch errors, which may expose deployment internals.
|
||||
func writeChatHookDispatchFailed(ctx context.Context, rw http.ResponseWriter, hookErr *dispatch.Error) {
|
||||
httpapi.Write(ctx, rw, http.StatusBadGateway, codersdk.ChatHookDispatchFailedResponse{
|
||||
Response: codersdk.Response{
|
||||
Message: "Chat lifecycle hook dispatch failed.",
|
||||
Detail: fmt.Sprintf("Lifecycle hook dispatch %s failed (%s).", hookErr.DispatchID, hookErr.Class),
|
||||
},
|
||||
Kind: codersdk.ChatErrorKindHookDispatchFailed,
|
||||
})
|
||||
}
|
||||
|
||||
// writeChatHookErr writes the response for lifecycle hook denials and
|
||||
// dispatch failures, reporting whether it handled the error. The fallback
|
||||
// message is used when the hook denies without a user message.
|
||||
func writeChatHookErr(ctx context.Context, rw http.ResponseWriter, err error, deniedFallback string) bool {
|
||||
if denied, ok := errors.AsType[*chathooks.UserPromptDeniedError](err); ok {
|
||||
message := denied.UserMessage
|
||||
if message == "" {
|
||||
message = deniedFallback
|
||||
}
|
||||
httpapi.Write(ctx, rw, http.StatusForbidden, codersdk.ChatHookDeniedResponse{
|
||||
Response: codersdk.Response{Message: message},
|
||||
Kind: codersdk.ChatErrorKindHookDenied,
|
||||
})
|
||||
return true
|
||||
}
|
||||
if hookErr, ok := errors.AsType[*dispatch.Error](err); ok {
|
||||
writeChatHookDispatchFailed(ctx, rw, hookErr)
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func maybeWriteLimitErr(ctx context.Context, rw http.ResponseWriter, err error) bool {
|
||||
var limitErr *chatd.UsageLimitExceededError
|
||||
if errors.As(err, &limitErr) {
|
||||
@@ -1227,8 +1262,7 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
// Cap the raw request body to prevent excessive memory use
|
||||
// from large dynamic tool schemas.
|
||||
// Limit memory used to decode dynamic tool schemas.
|
||||
r.Body = http.MaxBytesReader(rw, r.Body, int64(2*maxSystemPromptLenBytes))
|
||||
|
||||
var req codersdk.CreateChatRequest
|
||||
@@ -1424,23 +1458,27 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
chat, err := api.chatDaemon.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: req.OrganizationID,
|
||||
OwnerID: apiKey.UserID,
|
||||
WorkspaceID: workspaceSelection.WorkspaceID,
|
||||
Title: title,
|
||||
ModelConfigID: modelConfigID,
|
||||
ReasoningEffort: reasoningEffort,
|
||||
PlanMode: planModeToNullChatPlanMode(req.PlanMode),
|
||||
ClientType: clientType,
|
||||
SystemPrompt: req.SystemPrompt,
|
||||
InitialUserContent: contentBlocks,
|
||||
MCPServerIDs: mcpServerIDs,
|
||||
Labels: labels,
|
||||
DynamicTools: dynamicToolsJSON,
|
||||
OrganizationID: req.OrganizationID,
|
||||
OwnerID: apiKey.UserID,
|
||||
WorkspaceID: workspaceSelection.WorkspaceID,
|
||||
Title: title,
|
||||
TitleDerivedFromContent: true,
|
||||
ModelConfigID: modelConfigID,
|
||||
ReasoningEffort: reasoningEffort,
|
||||
PlanMode: planModeToNullChatPlanMode(req.PlanMode),
|
||||
ClientType: clientType,
|
||||
SystemPrompt: req.SystemPrompt,
|
||||
InitialUserContent: contentBlocks,
|
||||
MCPServerIDs: mcpServerIDs,
|
||||
Labels: labels,
|
||||
DynamicTools: dynamicToolsJSON,
|
||||
// IMPORTANT: users can only create root chats at the time of writing.
|
||||
ParentChatID: uuid.NullUUID{},
|
||||
})
|
||||
if err != nil {
|
||||
if writeChatHookErr(ctx, rw, err, "Chat creation denied by lifecycle hook.") {
|
||||
return
|
||||
}
|
||||
if maybeWriteLimitErr(ctx, rw, err) {
|
||||
return
|
||||
}
|
||||
@@ -1483,10 +1521,22 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
// Link any user-uploaded files referenced in the initial
|
||||
// message to this newly created chat (best-effort; cap
|
||||
// enforced in SQL).
|
||||
unlinked, capExceeded := api.linkFilesToChat(ctx, chat.ID, fileIDs)
|
||||
linkFileIDs := fileIDs
|
||||
if len(fileIDs) > 0 {
|
||||
initialUser, err := api.Database.GetLastChatMessageByRole(ctx, database.GetLastChatMessageByRoleParams{
|
||||
ChatID: chat.ID,
|
||||
Role: database.ChatMessageRoleUser,
|
||||
})
|
||||
if err != nil {
|
||||
api.Logger.Warn(ctx, "load initial message for file linking",
|
||||
slog.F("chat_id", chat.ID),
|
||||
slog.Error(err),
|
||||
)
|
||||
} else {
|
||||
linkFileIDs = api.linkedFileIDsFromContent(ctx, initialUser, fileIDs)
|
||||
}
|
||||
}
|
||||
unlinked, capExceeded := api.linkFilesToChat(ctx, chat.ID, linkFileIDs)
|
||||
|
||||
// Re-read the chat so the response reflects the authoritative
|
||||
// database state (file links are deduped in the join table).
|
||||
@@ -3422,6 +3472,9 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
|
||||
},
|
||||
)
|
||||
if sendErr != nil {
|
||||
if writeChatHookErr(ctx, rw, sendErr, "Chat message denied by lifecycle hook.") {
|
||||
return
|
||||
}
|
||||
if maybeWriteLimitErr(ctx, rw, sendErr) {
|
||||
return
|
||||
}
|
||||
@@ -3476,9 +3529,19 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
// Link any user-uploaded files referenced in this message
|
||||
// to the chat (best-effort; cap enforced in SQL).
|
||||
unlinked, capExceeded := api.linkFilesToChat(ctx, chatID, fileIDs)
|
||||
linkFileIDs := fileIDs
|
||||
if sendResult.Queued {
|
||||
if sendResult.QueuedMessage != nil {
|
||||
linkFileIDs = api.linkedFileIDsFromContent(ctx, database.ChatMessage{
|
||||
Role: database.ChatMessageRoleUser,
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
Content: pqtype.NullRawMessage{RawMessage: sendResult.QueuedMessage.Content, Valid: true},
|
||||
}, fileIDs)
|
||||
}
|
||||
} else {
|
||||
linkFileIDs = api.linkedFileIDsFromContent(ctx, sendResult.Message, fileIDs)
|
||||
}
|
||||
unlinked, capExceeded := api.linkFilesToChat(ctx, chatID, linkFileIDs)
|
||||
response := codersdk.CreateChatMessageResponse{Queued: sendResult.Queued}
|
||||
if sendResult.Queued {
|
||||
if sendResult.QueuedMessage != nil {
|
||||
@@ -3488,6 +3551,14 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
|
||||
message := convertChatMessage(sendResult.Message)
|
||||
response.Message = &message
|
||||
}
|
||||
// Return the full user-visible inserted batch. A queued send on an errored
|
||||
// chat can promote the previous queue head, which clients must cache.
|
||||
for _, inserted := range sendResult.InsertedMessages {
|
||||
if inserted.Visibility == database.ChatMessageVisibilityModel {
|
||||
continue
|
||||
}
|
||||
response.Messages = append(response.Messages, convertChatMessage(inserted))
|
||||
}
|
||||
if len(unlinked) > 0 {
|
||||
if capExceeded {
|
||||
response.Warnings = append(response.Warnings, fileLinkCapWarning(len(unlinked)))
|
||||
@@ -3591,6 +3662,9 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) {
|
||||
ReasoningEffort: editReasoningEffort,
|
||||
})
|
||||
if editErr != nil {
|
||||
if writeChatHookErr(ctx, rw, editErr, "Chat message denied by lifecycle hook.") {
|
||||
return
|
||||
}
|
||||
if maybeWriteLimitErr(ctx, rw, editErr) {
|
||||
return
|
||||
}
|
||||
@@ -3635,12 +3709,19 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
// Link any user-uploaded files referenced in the edited
|
||||
// message to the chat (best-effort; cap enforced in SQL).
|
||||
unlinked, capExceeded := api.linkFilesToChat(ctx, chat.ID, fileIDs)
|
||||
response := codersdk.EditChatMessageResponse{
|
||||
Message: convertChatMessage(editResult.Message),
|
||||
unlinked, capExceeded := api.linkFilesToChat(ctx, chat.ID, api.linkedFileIDsFromContent(ctx, editResult.Message, fileIDs))
|
||||
response := codersdk.EditChatMessageResponse{Message: convertChatMessage(editResult.Message)}
|
||||
// Synthetic cancellations precede the replacement with lower IDs;
|
||||
// clients that seed their transcript cache from this response need
|
||||
// all user-visible inserted rows, or a stream reconnect with
|
||||
// after_id set to the replacement would skip the earlier ones.
|
||||
for _, inserted := range editResult.InsertedMessages {
|
||||
if inserted.Visibility == database.ChatMessageVisibilityModel {
|
||||
continue
|
||||
}
|
||||
response.Messages = append(response.Messages, convertChatMessage(inserted))
|
||||
}
|
||||
response.DeletedMessageIDs = editResult.DeletedMessageIDs
|
||||
if len(unlinked) > 0 {
|
||||
if capExceeded {
|
||||
response.Warnings = append(response.Warnings, fileLinkCapWarning(len(unlinked)))
|
||||
@@ -6859,6 +6940,29 @@ func createChatInputFromParts(
|
||||
return content, pasteData, fileIDs, nil
|
||||
}
|
||||
|
||||
// A prompt override may remove file parts, so derive links from persisted
|
||||
// content. Fall back to request IDs if parsing fails.
|
||||
func (api *API) linkedFileIDsFromContent(ctx context.Context, msg database.ChatMessage, requestFileIDs []uuid.UUID) []uuid.UUID {
|
||||
if len(requestFileIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
parts, err := chatprompt.ParseContent(msg)
|
||||
if err != nil {
|
||||
api.Logger.Warn(ctx, "parse persisted message for file linking",
|
||||
slog.F("message_id", msg.ID),
|
||||
slog.Error(err),
|
||||
)
|
||||
return requestFileIDs
|
||||
}
|
||||
var ids []uuid.UUID
|
||||
for _, part := range parts {
|
||||
if part.Type == codersdk.ChatMessagePartTypeFile && part.FileID.Valid {
|
||||
ids = append(ids, part.FileID.UUID)
|
||||
}
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
// linkFilesToChat inserts file-link rows into the chat_file_links
|
||||
// join table. Cap enforcement and dedup are handled atomically in
|
||||
// SQL. On success returns (nil, false). On failure returns the full
|
||||
@@ -8196,6 +8300,10 @@ func (api *API) postChatToolResults(rw http.ResponseWriter, r *http.Request) {
|
||||
DynamicTools: dynamicTools,
|
||||
})
|
||||
if err != nil {
|
||||
if hookErr, ok := errors.AsType[*dispatch.Error](err); ok {
|
||||
writeChatHookDispatchFailed(ctx, rw, hookErr)
|
||||
return
|
||||
}
|
||||
var validationErr *chatd.ToolResultValidationError
|
||||
var conflictErr *chatd.ToolResultStatusConflictError
|
||||
switch {
|
||||
|
||||
@@ -0,0 +1,598 @@
|
||||
package coderd_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/coderdtest"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbauthz"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
"github.com/coder/serpent"
|
||||
)
|
||||
|
||||
func TestPostChatsInitialPromptHookErrors(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
statusCode int
|
||||
response string
|
||||
wantStatus int
|
||||
wantMessage string
|
||||
wantKind codersdk.ChatErrorKind
|
||||
}{
|
||||
{
|
||||
name: "deny",
|
||||
statusCode: http.StatusOK,
|
||||
response: `{"permission":{"decision":"deny"},"user_message":"blocked by policy"}`,
|
||||
wantStatus: http.StatusForbidden,
|
||||
wantMessage: "blocked by policy",
|
||||
wantKind: codersdk.ChatErrorKindHookDenied,
|
||||
},
|
||||
{
|
||||
name: "dispatch failure",
|
||||
statusCode: http.StatusInternalServerError,
|
||||
wantStatus: http.StatusBadGateway,
|
||||
wantKind: codersdk.ChatErrorKindHookDispatchFailed,
|
||||
},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
requests := make(chan agenthooks.Request, 2)
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
requests <- request
|
||||
w.WriteHeader(test.statusCode)
|
||||
if test.response != "" {
|
||||
_, err := w.Write([]byte(test.response))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
client, db := newChatClientWithDatabase(t, func(opts *coderdtest.Options) {
|
||||
opts.ChatWorkerDisabled = true
|
||||
require.NoError(t, opts.DeploymentValues.AI.Chat.HookURL.Set(consumer.URL))
|
||||
opts.DeploymentValues.AI.Chat.HookSecret = serpent.String("test-hook-secret-32-bytes-minimum!!")
|
||||
opts.DeploymentValues.AI.Chat.HookTimeout = serpent.Duration(time.Second)
|
||||
opts.DeploymentValues.AI.Chat.HookEnabled = serpent.Bool(true)
|
||||
})
|
||||
user := coderdtest.CreateFirstUser(t, client.Client)
|
||||
model := createAdditionalChatModelConfig(t, client, "openai", "gpt-4.1")
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
res, err := client.Request(ctx, http.MethodPost, "/api/experimental/chats", codersdk.CreateChatRequest{
|
||||
OrganizationID: user.OrganizationID,
|
||||
ModelConfigID: &model.ID,
|
||||
Content: []codersdk.ChatInputPart{{
|
||||
Type: codersdk.ChatInputPartTypeText,
|
||||
Text: "blocked prompt",
|
||||
}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
defer res.Body.Close()
|
||||
require.Equal(t, test.wantStatus, res.StatusCode)
|
||||
// Both outcomes share this wire shape, differing only in kind.
|
||||
var response struct {
|
||||
codersdk.Response
|
||||
Kind codersdk.ChatErrorKind `json:"kind"`
|
||||
}
|
||||
require.NoError(t, json.NewDecoder(res.Body).Decode(&response))
|
||||
require.Equal(t, test.wantKind, response.Kind)
|
||||
if test.wantMessage != "" {
|
||||
require.Equal(t, test.wantMessage, response.Message)
|
||||
}
|
||||
request := testutil.RequireReceive(ctx, t, requests)
|
||||
require.Equal(t, agenthooks.EventUserPromptSubmit, request.Type)
|
||||
require.NotEqual(t, uuid.Nil, request.Meta.ChatID)
|
||||
_, err = db.GetChatByID(dbauthz.AsSystemRestricted(ctx), request.Meta.ChatID)
|
||||
require.ErrorIs(t, err, sql.ErrNoRows)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatLifecycleHooksExperimentDisabled(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var hookRequests atomic.Int32
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
hookRequests.Add(1)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
client, _ := newChatClientWithDatabase(t, func(opts *coderdtest.Options) {
|
||||
opts.ChatWorkerDisabled = true
|
||||
opts.DeploymentValues.Experiments = serpent.StringArray{
|
||||
string(codersdk.ExperimentChatAdvisor),
|
||||
string(codersdk.ExperimentChatVirtualDesktop),
|
||||
}
|
||||
require.NoError(t, opts.DeploymentValues.AI.Chat.HookURL.Set(consumer.URL))
|
||||
opts.DeploymentValues.AI.Chat.HookSecret = serpent.String("test-hook-secret-32-bytes-minimum!!")
|
||||
opts.DeploymentValues.AI.Chat.HookTimeout = serpent.Duration(time.Second)
|
||||
opts.DeploymentValues.AI.Chat.HookEnabled = serpent.Bool(true)
|
||||
})
|
||||
user := coderdtest.CreateFirstUser(t, client.Client)
|
||||
model := createAdditionalChatModelConfig(t, client, "openai", "gpt-4.1")
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
_, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
|
||||
OrganizationID: user.OrganizationID,
|
||||
ModelConfigID: &model.ID,
|
||||
Content: []codersdk.ChatInputPart{{
|
||||
Type: codersdk.ChatInputPartTypeText,
|
||||
Text: "prompt with hooks disabled",
|
||||
}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Zero(t, hookRequests.Load())
|
||||
}
|
||||
|
||||
func TestChatPromptHookContextHiddenFromAPI(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const secret = "test-hook-secret-32-bytes-minimum!!"
|
||||
consumer := newHookConsumer(t, secret, agenthooks.Hooks{
|
||||
UserPromptSubmit: func(context.Context, agenthooks.Meta, agenthooks.UserPromptSubmitData) (agenthooks.Response, error) {
|
||||
return agenthooks.Response{
|
||||
ModelContext: "prompt context",
|
||||
UserMessage: "prompt notice",
|
||||
}, nil
|
||||
},
|
||||
})
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
client, _ := newChatClientWithDatabase(t, func(opts *coderdtest.Options) {
|
||||
opts.ChatWorkerDisabled = true
|
||||
require.NoError(t, opts.DeploymentValues.AI.Chat.HookURL.Set(consumer.URL))
|
||||
opts.DeploymentValues.AI.Chat.HookSecret = serpent.String(secret)
|
||||
opts.DeploymentValues.AI.Chat.HookTimeout = serpent.Duration(time.Second)
|
||||
opts.DeploymentValues.AI.Chat.HookEnabled = serpent.Bool(true)
|
||||
})
|
||||
user := coderdtest.CreateFirstUser(t, client.Client)
|
||||
model := createAdditionalChatModelConfig(t, client, "openai", "gpt-4.1")
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
|
||||
OrganizationID: user.OrganizationID,
|
||||
ModelConfigID: &model.ID,
|
||||
Content: []codersdk.ChatInputPart{{
|
||||
Type: codersdk.ChatInputPartTypeText,
|
||||
Text: "initial prompt",
|
||||
}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
messages, err := client.GetChatMessages(ctx, chat.ID, nil)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, messages.Messages, 1)
|
||||
require.Equal(t, []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("initial prompt"),
|
||||
{Type: codersdk.ChatMessagePartTypeHookNotice, Text: "prompt notice"},
|
||||
}, messages.Messages[0].Content)
|
||||
}
|
||||
|
||||
func TestChatLifecycleHooksWorkedExample(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const (
|
||||
secret = "test-hook-secret-32-bytes-minimum!!"
|
||||
deniedToolCallID = "call_denied"
|
||||
allowedToolCallID = "call_allowed"
|
||||
)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
var modelCalls atomic.Int32
|
||||
secondModelRequest := make(chan []byte, 1)
|
||||
thirdModelRequest := make(chan []byte, 1)
|
||||
modelURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("Lifecycle hooks")
|
||||
}
|
||||
switch modelCalls.Add(1) {
|
||||
case 1:
|
||||
chunk := chattest.OpenAIToolCallChunk("read_secret", `{"path":"/tmp/secret"}`)
|
||||
chunk.Choices[0].ToolCalls[0].ID = deniedToolCallID
|
||||
return chattest.OpenAIStreamingResponse(chunk)
|
||||
case 2:
|
||||
secondModelRequest <- bytes.Clone(req.RawBody)
|
||||
chunk := chattest.OpenAIToolCallChunk("search_docs", `{"query":"customer secret"}`)
|
||||
chunk.Choices[0].ToolCalls[0].ID = allowedToolCallID
|
||||
return chattest.OpenAIStreamingResponse(chunk)
|
||||
case 3:
|
||||
thirdModelRequest <- bytes.Clone(req.RawBody)
|
||||
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("done")...)
|
||||
default:
|
||||
return chattest.OpenAIErrorResponse(http.StatusInternalServerError, "unexpected_call", "unexpected model call")
|
||||
}
|
||||
})
|
||||
|
||||
hookEvents := make(chan agenthooks.EventType, 16)
|
||||
recordHook := func(event agenthooks.EventType) {
|
||||
hookEvents <- event
|
||||
}
|
||||
consumer := newHookConsumer(t, secret, agenthooks.Hooks{
|
||||
SessionStart: func(context.Context, agenthooks.Meta, agenthooks.SessionStartData) (agenthooks.Response, error) {
|
||||
recordHook(agenthooks.EventSessionStart)
|
||||
return agenthooks.Response{}, nil
|
||||
},
|
||||
UserPromptSubmit: func(context.Context, agenthooks.Meta, agenthooks.UserPromptSubmitData) (agenthooks.Response, error) {
|
||||
recordHook(agenthooks.EventUserPromptSubmit)
|
||||
return agenthooks.Response{}, nil
|
||||
},
|
||||
PreToolUse: func(_ context.Context, _ agenthooks.Meta, tool agenthooks.PreToolUseData) (agenthooks.Response, error) {
|
||||
recordHook(agenthooks.EventPreToolUse)
|
||||
switch tool.ToolUseID {
|
||||
case deniedToolCallID:
|
||||
return agenthooks.Response{Permission: &agenthooks.Permission{
|
||||
Decision: agenthooks.PermissionDeny,
|
||||
Reason: "secret reads are blocked",
|
||||
}}, nil
|
||||
case allowedToolCallID:
|
||||
return agenthooks.Response{Permission: &agenthooks.Permission{
|
||||
Decision: agenthooks.PermissionAllow,
|
||||
InputOverride: json.RawMessage(`{"query":"public documentation"}`),
|
||||
}}, nil
|
||||
default:
|
||||
return agenthooks.Response{}, nil
|
||||
}
|
||||
},
|
||||
PostToolUse: func(context.Context, agenthooks.Meta, agenthooks.PostToolUseData) (agenthooks.Response, error) {
|
||||
recordHook(agenthooks.EventPostToolUse)
|
||||
return agenthooks.Response{
|
||||
ModelContext: "The approved search result is safe to use.",
|
||||
UserMessage: "Search result approved by policy.",
|
||||
}, nil
|
||||
},
|
||||
Stop: func(context.Context, agenthooks.Meta, agenthooks.StopData) (agenthooks.Response, error) {
|
||||
recordHook(agenthooks.EventStop)
|
||||
return agenthooks.Response{}, nil
|
||||
},
|
||||
})
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
client, db := newChatClientWithDatabase(t, func(opts *coderdtest.Options) {
|
||||
require.NoError(t, opts.DeploymentValues.AI.Chat.HookURL.Set(consumer.URL))
|
||||
opts.DeploymentValues.AI.Chat.HookSecret = serpent.String(secret)
|
||||
opts.DeploymentValues.AI.Chat.HookTimeout = serpent.Duration(time.Second)
|
||||
opts.DeploymentValues.AI.Chat.HookEnabled = serpent.Bool(true)
|
||||
})
|
||||
user := coderdtest.CreateFirstUser(t, client.Client)
|
||||
model := createChatModelConfigWithBaseURL(t, client, modelURL)
|
||||
|
||||
chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
|
||||
OrganizationID: user.OrganizationID,
|
||||
ModelConfigID: &model.ID,
|
||||
Content: []codersdk.ChatInputPart{{
|
||||
Type: codersdk.ChatInputPartTypeText,
|
||||
Text: "Find the deployment documentation.",
|
||||
}},
|
||||
UnsafeDynamicTools: []codersdk.DynamicTool{
|
||||
{
|
||||
Name: "read_secret",
|
||||
Description: "Read a secret file.",
|
||||
InputSchema: json.RawMessage(`{"type":"object"}`),
|
||||
},
|
||||
{
|
||||
Name: "search_docs",
|
||||
Description: "Search public documentation.",
|
||||
InputSchema: json.RawMessage(`{"type":"object"}`),
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
var stored database.Chat
|
||||
testutil.Eventually(ctx, t, func(ctx context.Context) bool {
|
||||
stored, err = db.GetChatByID(dbauthz.AsSystemRestricted(ctx), chat.ID)
|
||||
return err == nil && stored.Status == database.ChatStatusRequiresAction
|
||||
}, testutil.IntervalFast)
|
||||
require.Equal(t, int32(2), modelCalls.Load())
|
||||
require.Contains(t, string(testutil.RequireReceive(ctx, t, secondModelRequest)), "Reason: secret reads are blocked.")
|
||||
|
||||
messages, err := client.GetChatMessages(ctx, chat.ID, nil)
|
||||
require.NoError(t, err)
|
||||
var allowedCall *codersdk.ChatMessagePart
|
||||
for _, message := range messages.Messages {
|
||||
for i := range message.Content {
|
||||
part := &message.Content[i]
|
||||
if part.Type == codersdk.ChatMessagePartTypeToolCall && part.ToolCallID == allowedToolCallID {
|
||||
allowedCall = part
|
||||
}
|
||||
}
|
||||
}
|
||||
require.NotNil(t, allowedCall)
|
||||
require.JSONEq(t, `{"query":"public documentation"}`, string(allowedCall.Args))
|
||||
|
||||
err = client.SubmitToolResults(ctx, chat.ID, codersdk.SubmitToolResultsRequest{
|
||||
Results: []codersdk.ToolResult{{
|
||||
ToolCallID: allowedToolCallID,
|
||||
Output: json.RawMessage(`{"matches":["agent hooks"]}`),
|
||||
}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
testutil.Eventually(ctx, t, func(ctx context.Context) bool {
|
||||
stored, err = db.GetChatByID(dbauthz.AsSystemRestricted(ctx), chat.ID)
|
||||
return err == nil && stored.Status == database.ChatStatusWaiting
|
||||
}, testutil.IntervalFast)
|
||||
require.Contains(t, string(testutil.RequireReceive(ctx, t, thirdModelRequest)), "The approved search result is safe to use.")
|
||||
require.Equal(t, int32(3), modelCalls.Load())
|
||||
|
||||
messages, err = client.GetChatMessages(ctx, chat.ID, nil)
|
||||
require.NoError(t, err)
|
||||
var foundPostToolNotice bool
|
||||
for _, message := range messages.Messages {
|
||||
if message.Role != codersdk.ChatMessageRoleSystem {
|
||||
continue
|
||||
}
|
||||
for _, part := range message.Content {
|
||||
if part.Type == codersdk.ChatMessagePartTypeText && part.Text == "Search result approved by policy." {
|
||||
foundPostToolNotice = true
|
||||
}
|
||||
}
|
||||
}
|
||||
require.True(t, foundPostToolNotice)
|
||||
|
||||
var seenEvents []agenthooks.EventType
|
||||
for {
|
||||
event := testutil.RequireReceive(ctx, t, hookEvents)
|
||||
seenEvents = append(seenEvents, event)
|
||||
if event == agenthooks.EventStop {
|
||||
break
|
||||
}
|
||||
}
|
||||
require.Contains(t, seenEvents, agenthooks.EventUserPromptSubmit)
|
||||
require.Contains(t, seenEvents, agenthooks.EventSessionStart)
|
||||
var preToolUseEvents int
|
||||
for _, event := range seenEvents {
|
||||
if event == agenthooks.EventPreToolUse {
|
||||
preToolUseEvents++
|
||||
}
|
||||
}
|
||||
require.GreaterOrEqual(t, preToolUseEvents, 2)
|
||||
require.Contains(t, seenEvents, agenthooks.EventPostToolUse)
|
||||
}
|
||||
|
||||
func TestChatHooksFileLinksAfterPromptOverride(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const secret = "test-hook-secret-32-bytes-minimum!!"
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
modelURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("done")...)
|
||||
})
|
||||
consumer := newHookConsumer(t, secret, agenthooks.Hooks{
|
||||
UserPromptSubmit: func(_ context.Context, _ agenthooks.Meta, data agenthooks.UserPromptSubmitData) (agenthooks.Response, error) {
|
||||
if strings.Contains(data.Prompt, "REDACTME") {
|
||||
return agenthooks.Response{Permission: &agenthooks.Permission{
|
||||
Decision: agenthooks.PermissionAllow,
|
||||
InputOverride: json.RawMessage(`{"prompt":"redacted"}`),
|
||||
}}, nil
|
||||
}
|
||||
return agenthooks.Response{}, nil
|
||||
},
|
||||
})
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
client, api := newChatClientWithAPI(t, func(opts *coderdtest.Options) {
|
||||
require.NoError(t, opts.DeploymentValues.AI.Chat.HookURL.Set(consumer.URL))
|
||||
opts.DeploymentValues.AI.Chat.HookSecret = serpent.String(secret)
|
||||
opts.DeploymentValues.AI.Chat.HookTimeout = serpent.Duration(time.Second)
|
||||
opts.DeploymentValues.AI.Chat.HookEnabled = serpent.Bool(true)
|
||||
})
|
||||
user := coderdtest.CreateFirstUser(t, client.Client)
|
||||
model := createChatModelConfigWithBaseURL(t, client, modelURL)
|
||||
|
||||
uploadFile := func(name string) uuid.UUID {
|
||||
pngData := append([]byte{0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A}, make([]byte, 16)...)
|
||||
resp, err := client.UploadChatFile(ctx, user.OrganizationID, "image/png", name, bytes.NewReader(pngData))
|
||||
require.NoError(t, err)
|
||||
return resp.ID
|
||||
}
|
||||
|
||||
redactedFile := uploadFile("redacted.png")
|
||||
chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
|
||||
OrganizationID: user.OrganizationID,
|
||||
ModelConfigID: &model.ID,
|
||||
Content: []codersdk.ChatInputPart{
|
||||
{Type: codersdk.ChatInputPartTypeText, Text: "REDACTME create"},
|
||||
{Type: codersdk.ChatInputPartTypeFile, FileID: redactedFile},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
created, err := client.GetChat(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, created.Files, "overridden create must not link dropped attachments")
|
||||
|
||||
coderdtest.WaitForChatSettled(ctx, t, api, chat.ID)
|
||||
|
||||
keptFile := uploadFile("kept.png")
|
||||
sendResp, err := client.CreateChatMessage(ctx, chat.ID, codersdk.CreateChatMessageRequest{
|
||||
Content: []codersdk.ChatInputPart{
|
||||
{Type: codersdk.ChatInputPartTypeText, Text: "keep this"},
|
||||
{Type: codersdk.ChatInputPartTypeFile, FileID: keptFile},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.False(t, sendResp.Queued)
|
||||
afterSend, err := client.GetChat(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, afterSend.Files, 1)
|
||||
require.Equal(t, keptFile, afterSend.Files[0].ID)
|
||||
|
||||
coderdtest.WaitForChatSettled(ctx, t, api, chat.ID)
|
||||
|
||||
droppedFile := uploadFile("dropped.png")
|
||||
_, err = client.CreateChatMessage(ctx, chat.ID, codersdk.CreateChatMessageRequest{
|
||||
Content: []codersdk.ChatInputPart{
|
||||
{Type: codersdk.ChatInputPartTypeText, Text: "REDACTME send"},
|
||||
{Type: codersdk.ChatInputPartTypeFile, FileID: droppedFile},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
afterOverride, err := client.GetChat(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, afterOverride.Files, 1, "overridden send must not link dropped attachments")
|
||||
require.Equal(t, keptFile, afterOverride.Files[0].ID)
|
||||
}
|
||||
|
||||
func TestChatHookNoticeMessagesInResponses(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const secret = "test-hook-secret-32-bytes-minimum!!"
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
modelURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("done")...)
|
||||
})
|
||||
|
||||
consumer := newHookConsumer(t, secret, agenthooks.Hooks{
|
||||
SessionStart: func(context.Context, agenthooks.Meta, agenthooks.SessionStartData) (agenthooks.Response, error) {
|
||||
return agenthooks.Response{UserMessage: "session notice"}, nil
|
||||
},
|
||||
UserPromptSubmit: func(_ context.Context, _ agenthooks.Meta, data agenthooks.UserPromptSubmitData) (agenthooks.Response, error) {
|
||||
response := agenthooks.Response{UserMessage: "prompt notice"}
|
||||
if data.Prompt == "edited prompt" {
|
||||
response.ModelContext = "prompt context"
|
||||
}
|
||||
return response, nil
|
||||
},
|
||||
})
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
client, db := newChatClientWithDatabase(t, func(opts *coderdtest.Options) {
|
||||
require.NoError(t, opts.DeploymentValues.AI.Chat.HookURL.Set(consumer.URL))
|
||||
opts.DeploymentValues.AI.Chat.HookSecret = serpent.String(secret)
|
||||
opts.DeploymentValues.AI.Chat.HookTimeout = serpent.Duration(time.Second)
|
||||
opts.DeploymentValues.AI.Chat.HookEnabled = serpent.Bool(true)
|
||||
})
|
||||
user := coderdtest.CreateFirstUser(t, client.Client)
|
||||
model := createChatModelConfigWithBaseURL(t, client, modelURL)
|
||||
|
||||
chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
|
||||
OrganizationID: user.OrganizationID,
|
||||
ModelConfigID: &model.ID,
|
||||
Content: []codersdk.ChatInputPart{{
|
||||
Type: codersdk.ChatInputPartTypeText,
|
||||
Text: "initial prompt",
|
||||
}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
waitForWaiting := func() {
|
||||
testutil.Eventually(ctx, t, func(ctx context.Context) bool {
|
||||
stored, err := db.GetChatByID(dbauthz.AsSystemRestricted(ctx), chat.ID)
|
||||
return err == nil && stored.Status == database.ChatStatusWaiting
|
||||
}, testutil.IntervalFast)
|
||||
}
|
||||
waitForWaiting()
|
||||
|
||||
assertPromptContent := func(message codersdk.ChatMessage, prompt string) {
|
||||
t.Helper()
|
||||
require.Equal(t, codersdk.ChatMessageRoleUser, message.Role)
|
||||
require.Equal(t, []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText(prompt),
|
||||
{Type: codersdk.ChatMessagePartTypeHookNotice, Text: "prompt notice"},
|
||||
}, message.Content)
|
||||
}
|
||||
|
||||
initialMessages, err := client.GetChatMessages(ctx, chat.ID, nil)
|
||||
require.NoError(t, err)
|
||||
var initialPrompt *codersdk.ChatMessage
|
||||
for i := range initialMessages.Messages {
|
||||
message := &initialMessages.Messages[i]
|
||||
if message.Role == codersdk.ChatMessageRoleUser && len(message.Content) > 0 && message.Content[0].Text == "initial prompt" {
|
||||
initialPrompt = message
|
||||
break
|
||||
}
|
||||
}
|
||||
require.NotNil(t, initialPrompt)
|
||||
assertPromptContent(*initialPrompt, "initial prompt")
|
||||
|
||||
sent, err := client.CreateChatMessage(ctx, chat.ID, codersdk.CreateChatMessageRequest{
|
||||
Content: []codersdk.ChatInputPart{{
|
||||
Type: codersdk.ChatInputPartTypeText,
|
||||
Text: "second prompt",
|
||||
}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.False(t, sent.Queued, "idle chat must insert directly")
|
||||
require.NotNil(t, sent.Message)
|
||||
require.NotEmpty(t, sent.Messages, "send response must carry the inserted batch")
|
||||
last := sent.Messages[len(sent.Messages)-1]
|
||||
require.Equal(t, sent.Message.ID, last.ID, "user message must be last in the batch")
|
||||
assertPromptContent(last, "second prompt")
|
||||
assertPromptContent(*sent.Message, "second prompt")
|
||||
|
||||
waitForWaiting()
|
||||
|
||||
edited, err := client.EditChatMessage(ctx, chat.ID, sent.Message.ID, codersdk.EditChatMessageRequest{
|
||||
Content: []codersdk.ChatInputPart{{
|
||||
Type: codersdk.ChatInputPartTypeText,
|
||||
Text: "edited prompt",
|
||||
}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotZero(t, edited.Message.ID, "successful edits must return the replacement message")
|
||||
require.NotEmpty(t, edited.Messages, "edit response must carry the inserted batch")
|
||||
var editedBatchMessage *codersdk.ChatMessage
|
||||
for i := range edited.Messages {
|
||||
if edited.Messages[i].ID == edited.Message.ID {
|
||||
editedBatchMessage = &edited.Messages[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
require.NotNil(t, editedBatchMessage)
|
||||
assertPromptContent(*editedBatchMessage, "edited prompt")
|
||||
assertPromptContent(edited.Message, "edited prompt")
|
||||
|
||||
allMessages, err := client.GetChatMessages(ctx, chat.ID, nil)
|
||||
require.NoError(t, err)
|
||||
var sessionNoticeFound bool
|
||||
for _, message := range allMessages.Messages {
|
||||
if message.Role != codersdk.ChatMessageRoleSystem {
|
||||
continue
|
||||
}
|
||||
for _, part := range message.Content {
|
||||
if part.Type == codersdk.ChatMessagePartTypeText && part.Text == "session notice" {
|
||||
sessionNoticeFound = true
|
||||
}
|
||||
}
|
||||
}
|
||||
require.True(t, sessionNoticeFound)
|
||||
}
|
||||
|
||||
// newHookConsumer serves hooks with its own URL as the configured audience,
|
||||
// which is the value Coder signs when it dispatches there. The listener is
|
||||
// allocated first because httptest.NewServer builds its handler before the
|
||||
// server has a URL.
|
||||
func newHookConsumer(t *testing.T, secret string, hooks agenthooks.Hooks) *httptest.Server {
|
||||
t.Helper()
|
||||
|
||||
server := httptest.NewUnstartedServer(nil)
|
||||
server.Config.Handler = agenthooks.NewHTTPHandler([]byte(secret), "http://"+server.Listener.Addr().String(), hooks)
|
||||
server.Start()
|
||||
return server
|
||||
}
|
||||
@@ -75,6 +75,7 @@ func newChatTestOptions(
|
||||
values.Experiments = serpent.StringArray{
|
||||
string(codersdk.ExperimentChatAdvisor),
|
||||
string(codersdk.ExperimentChatVirtualDesktop),
|
||||
string(codersdk.ExperimentAgentLifecycleHooks),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7189,7 +7190,6 @@ func TestSendMessageWithModelOverrideUpdatesLastModelConfigID(t *testing.T) {
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.False(t, resp.Queued)
|
||||
require.NotNil(t, resp.Message)
|
||||
require.NotNil(t, resp.Message.ModelConfigID)
|
||||
require.Equal(t, modelConfigB.ID, *resp.Message.ModelConfigID)
|
||||
|
||||
@@ -7470,7 +7470,6 @@ func TestSubsequentSendWithoutOverrideUsesPersistedModel(t *testing.T) {
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.False(t, resp.Queued)
|
||||
require.NotNil(t, resp.Message)
|
||||
require.NotNil(t, resp.Message.ModelConfigID)
|
||||
require.Equal(t, modelConfigB.ID, *resp.Message.ModelConfigID)
|
||||
|
||||
@@ -8033,7 +8032,6 @@ func TestChatMessageWithFiles(t *testing.T) {
|
||||
if resp.Queued {
|
||||
require.NotNil(t, resp.QueuedMessage)
|
||||
} else {
|
||||
require.NotNil(t, resp.Message)
|
||||
require.Equal(t, codersdk.ChatMessageRoleUser, resp.Message.Role)
|
||||
}
|
||||
})
|
||||
@@ -8081,7 +8079,6 @@ func TestChatMessageWithFiles(t *testing.T) {
|
||||
if resp.Queued {
|
||||
require.NotNil(t, resp.QueuedMessage)
|
||||
} else {
|
||||
require.NotNil(t, resp.Message)
|
||||
require.Equal(t, codersdk.ChatMessageRoleUser, resp.Message.Role)
|
||||
}
|
||||
|
||||
|
||||
@@ -111,7 +111,7 @@ I don't recommend reading the rest of section thoroughly if this is your first t
|
||||
- `Create(initialMessages)` creates a new chat, initializes `snapshot_version` to 1, inserts its initial history, and lands in `running`. The inserted initial history sets `history_version` to 1. Since the queue has not changed, `queue_version` remains 0. This transition is a special case: since the chat does not exist at the time it's run, the chat row cannot be locked before the transition is applied.
|
||||
- `SetArchived(archived)` sets or clears the archived marker for one chat.
|
||||
- `SendMessage(m, busy_behavior)` inserts a user message directly when the chat is idle, or queues it when the chat is busy. `busy_behavior` must be either `queue` or `interrupt`. With `busy_behavior=interrupt`, it also requests interruption or cancels a pending dynamic-tool action as needed.
|
||||
- `EditMessage(k, replacement)` clears queued messages, cancels or obsoletes active work, marks the truncated active-history suffix as deleted, inserts the replacement turn, and lands in `running`.
|
||||
- `EditMessage(k, replacement)` clears queued messages, cancels or obsoletes active work, marks the truncated active-history suffix as deleted, inserts the replacement turn followed by any caller-provided suffix messages, and lands in `running`.
|
||||
- `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.
|
||||
@@ -903,10 +903,20 @@ Users can also request a compaction on demand via `POST /api/experimental/chats/
|
||||
1. The endpoint applies the `RequestCompaction` transition: only allowed from `W`, sets `chats.compaction_requested_at = now()`, lands in `R0` without inserting any message, and publishes a status-change pubsub event to wake workers. A timestamp is used instead of a boolean for debuggability. AI Gateway attribution needs no per-request key: generation preparation resolves the owner's synthetic API key like any other turn.
|
||||
2. The generation goroutine's decision logic checks `compaction_requested_at` after the unresolved local/dynamic tool guards but before the history-completeness check (an idle chat's history is otherwise complete, which would end the turn). If the marker is set and at least one uncompressed assistant message exists after the latest compaction boundary, it selects a forced compaction; if there is nothing to compact, the marker is ignored and the turn finishes normally, clearing it.
|
||||
3. A forced compaction bypasses the automatic threshold gates (usage below threshold, unknown context window, and the threshold=100 disable) and stamps `source: "manual"` instead of `source: "automatic"` into the `chat_summarized` tool call arguments, tool result JSON, and streamed parts so clients can render manual compactions distinctly.
|
||||
4. The compaction `CommitStep` consumes the request by clearing `compaction_requested_at` in the same transaction that commits the summary triplet. The next decision pass finds the history complete and finishes the turn, so the chat returns to `waiting` with no assistant follow-up.
|
||||
4. The compaction `CommitStep` consumes the request by clearing `compaction_requested_at` in the same transaction that commits the summary triplet. The next decision pass finds the history complete and finishes the turn, so the chat returns to `waiting` with no assistant follow-up. A `post_compact` hook effect is the one exception: because the decision reads user-visible history, an effect that commits a user-visible message leaves the history incomplete and the turn continues with an assistant response. A model-only effect such as `model_context` reaches the model without resuming generation.
|
||||
|
||||
The `compaction_requested_at` marker is one-shot: transitions that keep an active turn alive (`Acquire`, `Abandon`, `SetArchived`, queueing a message on a busy chat) carry it forward, while every other transition that rewrites the execution state (`FinishTurn`, `FinishError`, `Interrupt`, `EditMessage`, `PromoteQueuedMessage`, `CancelRequiresAction`, `ReconcileInvalidState`, and so on) clears it by construction, so a stale request can never replay on a later turn.
|
||||
|
||||
# Lifecycle hooks
|
||||
|
||||
When the `agent-lifecycle-hooks` experiment is enabled and a hook URL is configured, chatd sends events to an external consumer at key points in a conversation: session start, prompt submission, tool use, compaction, and turn completion.
|
||||
|
||||
The consumer can observe activity, add model-only or user-visible context, replace supported prompt or tool input, and deny prompts or tool calls. Prompt submission is evaluated once when the submission is accepted, including queued messages and subagent prompts. Returned context becomes part of the conversation for its intended audience, except that context returned before a compaction guides the compaction summary instead.
|
||||
|
||||
Lifecycle hooks fail closed. If the consumer cannot be reached or returns an invalid response, Coder stops the triggering operation rather than continuing without the consumer's decision. Affected chats can enter an error state until the consumer recovers or hooks are disabled.
|
||||
|
||||
Coder stores no hook-specific dispatch or decision state. Delivery is best-effort and can duplicate, and a failed dispatch is never redelivered, so the consumer owns durable policy state, audit records, and deduplication based on stable event identifiers.
|
||||
|
||||
# Stream loop
|
||||
|
||||
The stream loop powers the `GET /api/experimental/chats/{chat}/stream` endpoint. It is scoped to one chat and one client WebSocket. It's responsible for delivering a stream of chat updates to the client, including:
|
||||
|
||||
+277
-52
@@ -39,10 +39,12 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/util/xjson"
|
||||
"github.com/coder/coder/v2/coderd/webpush"
|
||||
"github.com/coder/coder/v2/coderd/workspacestats"
|
||||
"github.com/coder/coder/v2/coderd/x/agenthooks/dispatch"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/agentselect"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatadvisor"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatdebug"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chaterror"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chathooks"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatloop"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatopenai"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
@@ -54,6 +56,7 @@ import (
|
||||
skillspkg "github.com/coder/coder/v2/coderd/x/skills"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
@@ -176,6 +179,7 @@ type Server struct {
|
||||
stopWorkspaceFn chattool.StopWorkspaceFn
|
||||
pubsub pubsub.Pubsub
|
||||
webpushDispatcher webpush.Dispatcher
|
||||
hooks *chathooks.Trigger
|
||||
providerAPIKeys chatprovider.ProviderAPIKeys
|
||||
allowBYOK bool
|
||||
oidcTokenSource mcpclient.UserOIDCTokenSource
|
||||
@@ -1164,24 +1168,25 @@ func (e *UsageLimitExceededError) Error() string {
|
||||
|
||||
// CreateOptions controls chat creation in the shared chat mutation path.
|
||||
type CreateOptions struct {
|
||||
OrganizationID uuid.UUID
|
||||
OwnerID uuid.UUID
|
||||
WorkspaceID uuid.NullUUID
|
||||
BuildID uuid.NullUUID
|
||||
AgentID uuid.NullUUID
|
||||
ParentChatID uuid.NullUUID
|
||||
RootChatID uuid.NullUUID
|
||||
Title string
|
||||
ModelConfigID uuid.UUID
|
||||
ReasoningEffort *string
|
||||
ChatMode database.NullChatMode
|
||||
PlanMode database.NullChatPlanMode
|
||||
ClientType database.ChatClientType
|
||||
SystemPrompt string
|
||||
InitialUserContent []codersdk.ChatMessagePart
|
||||
MCPServerIDs []uuid.UUID
|
||||
Labels database.StringMap
|
||||
DynamicTools json.RawMessage
|
||||
OrganizationID uuid.UUID
|
||||
OwnerID uuid.UUID
|
||||
WorkspaceID uuid.NullUUID
|
||||
BuildID uuid.NullUUID
|
||||
AgentID uuid.NullUUID
|
||||
ParentChatID uuid.NullUUID
|
||||
RootChatID uuid.NullUUID
|
||||
Title string
|
||||
TitleDerivedFromContent bool
|
||||
ModelConfigID uuid.UUID
|
||||
ReasoningEffort *string
|
||||
ChatMode database.NullChatMode
|
||||
PlanMode database.NullChatPlanMode
|
||||
ClientType database.ChatClientType
|
||||
SystemPrompt string
|
||||
InitialUserContent []codersdk.ChatMessagePart
|
||||
MCPServerIDs []uuid.UUID
|
||||
Labels database.StringMap
|
||||
DynamicTools json.RawMessage
|
||||
}
|
||||
|
||||
// SendMessageBusyBehavior controls what happens when a chat is already active.
|
||||
@@ -1214,7 +1219,11 @@ type SendMessageResult struct {
|
||||
Queued bool
|
||||
QueuedMessage *database.ChatQueuedMessage
|
||||
Message database.ChatMessage
|
||||
Chat database.Chat
|
||||
// InsertedMessages holds every message the send inserted, in
|
||||
// insertion order. A queued send on an errored chat can still
|
||||
// insert messages by promoting the previous queue head.
|
||||
InsertedMessages []database.ChatMessage
|
||||
Chat database.Chat
|
||||
}
|
||||
|
||||
// EditMessageOptions controls user message edits via soft-delete and re-insert.
|
||||
@@ -1233,7 +1242,14 @@ type EditMessageOptions struct {
|
||||
// EditMessageResult contains the replacement user message and chat status.
|
||||
type EditMessageResult struct {
|
||||
Message database.ChatMessage
|
||||
Chat database.Chat
|
||||
// InsertedMessages holds every message the edit inserted, in
|
||||
// insertion order: synthetic tool cancellations, the replacement
|
||||
// user message, then hook suffix messages.
|
||||
InsertedMessages []database.ChatMessage
|
||||
// DeletedMessageIDs holds every previously visible message the
|
||||
// edit soft-deleted.
|
||||
DeletedMessageIDs []int64
|
||||
Chat database.Chat
|
||||
}
|
||||
|
||||
// PromoteQueuedOptions controls queued-message promotion.
|
||||
@@ -1301,6 +1317,38 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
|
||||
return database.Chat{}, xerrors.Errorf("marshal labels: %w", err)
|
||||
}
|
||||
|
||||
chatID := uuid.New()
|
||||
contentParts := opts.InitialUserContent
|
||||
if p.hooks.Enabled() {
|
||||
// Validate model admission before dispatch, matching the insert path.
|
||||
if err := validateCreateModelConfigID(ctx, p.db, opts.ModelConfigID); err != nil {
|
||||
return database.Chat{}, err
|
||||
}
|
||||
turnID := uuid.New()
|
||||
promptMessage, err := chathooks.UserPromptMessage(contentParts)
|
||||
if err != nil {
|
||||
return database.Chat{}, err
|
||||
}
|
||||
promptResult, err := p.hooks.Trigger(ctx, chathooks.Chat{
|
||||
ID: chatID,
|
||||
OwnerID: opts.OwnerID,
|
||||
WorkspaceID: opts.WorkspaceID,
|
||||
TurnID: &turnID,
|
||||
}, promptMessage, agenthooks.EventUserPromptSubmit)
|
||||
if err != nil {
|
||||
return database.Chat{}, chathooks.UserPromptDenial(err)
|
||||
}
|
||||
composed, overridden, err := chathooks.ComposeUserPromptContent(contentParts, promptResult)
|
||||
if err != nil {
|
||||
return database.Chat{}, err
|
||||
}
|
||||
contentParts = composed
|
||||
// Avoid deriving titles from the prompt that policy replaced.
|
||||
if overridden && opts.TitleDerivedFromContent {
|
||||
opts.Title = chatprompt.FallbackTitle(chatprompt.TitleText(contentParts, nil))
|
||||
}
|
||||
}
|
||||
|
||||
userPrompt := SanitizePromptText(opts.SystemPrompt)
|
||||
workspaceAwareness := workspaceDetachedAwareness
|
||||
if opts.WorkspaceID.Valid {
|
||||
@@ -1312,7 +1360,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
|
||||
if err != nil {
|
||||
return database.Chat{}, xerrors.Errorf("marshal workspace awareness: %w", err)
|
||||
}
|
||||
userContent, err := chatprompt.MarshalParts(opts.InitialUserContent)
|
||||
userContent, err := chatprompt.MarshalParts(contentParts)
|
||||
if err != nil {
|
||||
return database.Chat{}, xerrors.Errorf("marshal initial user content: %w", err)
|
||||
}
|
||||
@@ -1339,7 +1387,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
|
||||
initialMessages = append(initialMessages, systemMessage(workspaceAwarenessContent, opts.ModelConfigID))
|
||||
initialMessages = append(initialMessages, userMessage(userContent, opts.ModelConfigID, opts.OwnerID, opts.ReasoningEffort))
|
||||
|
||||
result, err := chatstate.CreateChat(ctx, p.db, p.pubsub, chatstate.CreateChatInput{
|
||||
result, err := chatstate.CreateChatWithID(ctx, p.db, p.pubsub, chatID, chatstate.CreateChatInput{
|
||||
OrganizationID: opts.OrganizationID,
|
||||
OwnerID: opts.OwnerID,
|
||||
WorkspaceID: opts.WorkspaceID,
|
||||
@@ -1409,7 +1457,47 @@ func (p *Server) SendMessage(
|
||||
return SendMessageResult{}, xerrors.Errorf("invalid busy behavior %q", opts.BusyBehavior)
|
||||
}
|
||||
|
||||
content, err := chatprompt.MarshalParts(opts.Content)
|
||||
contentParts := opts.Content
|
||||
if p.hooks.Enabled() {
|
||||
turnID := uuid.New()
|
||||
chat, err := p.db.GetChatByID(ctx, opts.ChatID)
|
||||
if err != nil {
|
||||
return SendMessageResult{}, xerrors.Errorf("load chat for user_prompt_submit: %w", err)
|
||||
}
|
||||
// Repeat these admission checks under the transaction lock.
|
||||
if chat.Archived {
|
||||
return SendMessageResult{}, ErrChatArchived
|
||||
}
|
||||
if err := p.checkUsageLimit(ctx, p.db, chat.OwnerID, uuid.NullUUID{UUID: chat.OrganizationID, Valid: true}); err != nil {
|
||||
return SendMessageResult{}, err
|
||||
}
|
||||
if _, err := resolveSendMessageModelConfigID(ctx, p.db, chat, opts.ModelConfigID); err != nil {
|
||||
return SendMessageResult{}, err
|
||||
}
|
||||
// Check queue capacity before dispatch; the transaction
|
||||
// rechecks it under lock.
|
||||
queuedCount, err := p.db.CountChatQueuedMessages(ctx, opts.ChatID)
|
||||
if err != nil {
|
||||
return SendMessageResult{}, xerrors.Errorf("count queued messages: %w", err)
|
||||
}
|
||||
if queuedCount >= chatstate.MaxQueueSize {
|
||||
return SendMessageResult{}, &chatstate.MessageQueueFullError{Max: chatstate.MaxQueueSize}
|
||||
}
|
||||
promptMessage, err := chathooks.UserPromptMessage(contentParts)
|
||||
if err != nil {
|
||||
return SendMessageResult{}, err
|
||||
}
|
||||
promptResult, err := p.hooks.Trigger(ctx, chathooks.ChatFor(chat, &turnID), promptMessage, agenthooks.EventUserPromptSubmit)
|
||||
if err != nil {
|
||||
return SendMessageResult{}, p.handleUserPromptDispatchError(ctx, opts.ChatID, chathooks.UserPromptDenial(err))
|
||||
}
|
||||
contentParts, _, err = chathooks.ComposeUserPromptContent(contentParts, promptResult)
|
||||
if err != nil {
|
||||
return SendMessageResult{}, err
|
||||
}
|
||||
}
|
||||
|
||||
content, err := chatprompt.MarshalParts(contentParts)
|
||||
if err != nil {
|
||||
return SendMessageResult{}, xerrors.Errorf("marshal message content: %w", err)
|
||||
}
|
||||
@@ -1480,8 +1568,9 @@ func (p *Server) SendMessage(
|
||||
|
||||
// Queue capacity is enforced inside tx.SendMessage; this
|
||||
// wrapper only propagates the typed error.
|
||||
message := userMessage(content, modelConfigID, messageCreatedBy, opts.ReasoningEffort)
|
||||
sendResult, err := tx.SendMessage(chatstate.SendMessageInput{
|
||||
Message: userMessage(content, modelConfigID, messageCreatedBy, opts.ReasoningEffort),
|
||||
Message: message,
|
||||
BusyBehavior: busyBehaviorToChatState(busyBehavior),
|
||||
})
|
||||
if err != nil {
|
||||
@@ -1497,6 +1586,10 @@ func (p *Server) SendMessage(
|
||||
// last in the inserted slice.
|
||||
result.Message = sendResult.InsertedMessages[len(sendResult.InsertedMessages)-1]
|
||||
}
|
||||
// A queued send on an errored chat can also promote the
|
||||
// previous queue head into history; report those inserts so
|
||||
// clients can update their caches.
|
||||
result.InsertedMessages = sendResult.InsertedMessages
|
||||
// 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
|
||||
@@ -1571,6 +1664,20 @@ func requireEnabledChatModelConfig(
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateCreateModelConfigID(ctx context.Context, store database.Store, modelConfigID uuid.UUID) error {
|
||||
if modelConfigID == uuid.Nil {
|
||||
return xerrors.Errorf("%w: %s", ErrInvalidModelConfigID, modelConfigID)
|
||||
}
|
||||
chatdCtx := chatdModelConfigLookupContext(ctx)
|
||||
if _, err := store.GetChatModelConfigByID(chatdCtx, modelConfigID); err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return xerrors.Errorf("%w: %s", ErrInvalidModelConfigID, modelConfigID)
|
||||
}
|
||||
return xerrors.Errorf("get requested model config %s: %w", modelConfigID, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func resolveFallbackModelConfigID(
|
||||
ctx context.Context,
|
||||
store database.Store,
|
||||
@@ -1614,6 +1721,37 @@ func resolveFallbackModelConfigID(
|
||||
return defaultConfig.ID, nil
|
||||
}
|
||||
|
||||
func validateModelConfigOverride(
|
||||
ctx context.Context,
|
||||
store database.Store,
|
||||
requested uuid.UUID,
|
||||
) (uuid.NullUUID, error) {
|
||||
if requested == uuid.Nil {
|
||||
return uuid.NullUUID{}, nil
|
||||
}
|
||||
if err := requireEnabledChatModelConfig(ctx, store, requested); err != nil {
|
||||
return uuid.NullUUID{}, err
|
||||
}
|
||||
return uuid.NullUUID{UUID: requested, Valid: true}, nil
|
||||
}
|
||||
|
||||
func validateEditTarget(ctx context.Context, store database.Store, chatID uuid.UUID, messageID int64) error {
|
||||
target, err := store.GetChatMessageByID(ctx, messageID)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return ErrEditedMessageNotFound
|
||||
}
|
||||
return xerrors.Errorf("get edited message: %w", err)
|
||||
}
|
||||
if target.ChatID != chatID || target.Deleted {
|
||||
return ErrEditedMessageNotFound
|
||||
}
|
||||
if target.Role != database.ChatMessageRoleUser {
|
||||
return ErrEditedMessageNotUser
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// EditMessage replaces an earlier user message and discards the
|
||||
// active-history suffix through chatstate.EditMessage. Model-config
|
||||
// override validation and usage-limit admission run in the same
|
||||
@@ -1632,7 +1770,46 @@ func (p *Server) EditMessage(
|
||||
return EditMessageResult{}, xerrors.New("content is required")
|
||||
}
|
||||
|
||||
content, err := chatprompt.MarshalParts(opts.Content)
|
||||
contentParts := opts.Content
|
||||
var sessionStartHookResult *chathooks.Result
|
||||
if p.hooks.Enabled() {
|
||||
turnID := uuid.New()
|
||||
chat, err := p.db.GetChatByID(ctx, opts.ChatID)
|
||||
if err != nil {
|
||||
return EditMessageResult{}, xerrors.Errorf("load chat for edit hooks: %w", err)
|
||||
}
|
||||
// Repeat these admission checks under the transaction lock.
|
||||
if chat.Archived {
|
||||
return EditMessageResult{}, ErrChatArchived
|
||||
}
|
||||
if err := p.checkUsageLimit(ctx, p.db, chat.OwnerID, uuid.NullUUID{UUID: chat.OrganizationID, Valid: true}); err != nil {
|
||||
return EditMessageResult{}, err
|
||||
}
|
||||
if err := validateEditTarget(ctx, p.db, opts.ChatID, opts.EditedMessageID); err != nil {
|
||||
return EditMessageResult{}, err
|
||||
}
|
||||
if _, err := validateModelConfigOverride(ctx, p.db, opts.ModelConfigID); err != nil {
|
||||
return EditMessageResult{}, err
|
||||
}
|
||||
sessionStartHookResult, err = p.hooks.Trigger(ctx, chathooks.ChatFor(chat, &turnID), chathooks.Message{Source: chathooks.SessionStartSourceClear}, agenthooks.EventSessionStart)
|
||||
if err != nil {
|
||||
return EditMessageResult{}, p.handleAPIDispatchError(ctx, opts.ChatID, agenthooks.EventSessionStart, err)
|
||||
}
|
||||
promptMessage, err := chathooks.UserPromptMessage(contentParts)
|
||||
if err != nil {
|
||||
return EditMessageResult{}, err
|
||||
}
|
||||
promptResult, err := p.hooks.Trigger(ctx, chathooks.ChatFor(chat, &turnID), promptMessage, agenthooks.EventUserPromptSubmit)
|
||||
if err != nil {
|
||||
return EditMessageResult{}, p.handleUserPromptDispatchError(ctx, opts.ChatID, chathooks.UserPromptDenial(err))
|
||||
}
|
||||
contentParts, _, err = chathooks.ComposeUserPromptContent(contentParts, promptResult)
|
||||
if err != nil {
|
||||
return EditMessageResult{}, err
|
||||
}
|
||||
}
|
||||
|
||||
content, err := chatprompt.MarshalParts(contentParts)
|
||||
if err != nil {
|
||||
return EditMessageResult{}, xerrors.Errorf("marshal message content: %w", err)
|
||||
}
|
||||
@@ -1667,18 +1844,19 @@ func (p *Server) EditMessage(
|
||||
if target.ChatID != opts.ChatID {
|
||||
return ErrEditedMessageNotFound
|
||||
}
|
||||
if target.Deleted {
|
||||
return ErrEditedMessageNotFound
|
||||
}
|
||||
if target.Role != database.ChatMessageRoleUser {
|
||||
return ErrEditedMessageNotUser
|
||||
}
|
||||
editedMsg = target
|
||||
|
||||
// Validate the optional model-config override up front so
|
||||
// the user sees ErrInvalidModelConfigID instead of a
|
||||
// foreign-key error from the message-insert path.
|
||||
var modelOverride uuid.NullUUID
|
||||
if opts.ModelConfigID != uuid.Nil {
|
||||
if err := requireEnabledChatModelConfig(ctx, store, opts.ModelConfigID); err != nil {
|
||||
return err
|
||||
}
|
||||
modelOverride = uuid.NullUUID{UUID: opts.ModelConfigID, Valid: true}
|
||||
} else {
|
||||
modelOverride, err := validateModelConfigOverride(ctx, store, opts.ModelConfigID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !modelOverride.Valid {
|
||||
// Without an explicit override the transition preserves
|
||||
// the edited message's original model, which may have been
|
||||
// disabled since; resolve it like a normal message send.
|
||||
@@ -1695,6 +1873,19 @@ func (p *Server) EditMessage(
|
||||
}
|
||||
}
|
||||
|
||||
modelConfigID := target.ModelConfigID.UUID
|
||||
if modelOverride.Valid {
|
||||
modelConfigID = modelOverride.UUID
|
||||
}
|
||||
// The prompt response already rides in the replacement content;
|
||||
// only the session_start(clear) response needs transcript rows.
|
||||
// They insert after the replacement so a later edit's suffix
|
||||
// truncation cleans them up.
|
||||
suffixMessages, err := chathooks.EventMessages(sessionStartHookResult, modelConfigID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var reasoningEffortOverride database.NullChatReasoningEffort
|
||||
if opts.ReasoningEffort != nil && *opts.ReasoningEffort != "" {
|
||||
reasoningEffortOverride = database.NullChatReasoningEffort{ChatReasoningEffort: database.ChatReasoningEffort(*opts.ReasoningEffort), Valid: true}
|
||||
@@ -1702,6 +1893,7 @@ func (p *Server) EditMessage(
|
||||
|
||||
editResult, err := tx.EditMessage(chatstate.EditMessageInput{
|
||||
MessageID: opts.EditedMessageID,
|
||||
SuffixMessages: suffixMessages,
|
||||
CreatedBy: opts.CreatedBy,
|
||||
Content: content,
|
||||
ModelConfigIDOverride: modelOverride,
|
||||
@@ -1714,6 +1906,12 @@ func (p *Server) EditMessage(
|
||||
return err
|
||||
}
|
||||
result.Message = editResult.ReplacementMessage
|
||||
inserted := make([]database.ChatMessage, 0, len(editResult.CancellationMessages)+len(editResult.SuffixMessages)+1)
|
||||
inserted = append(inserted, editResult.CancellationMessages...)
|
||||
inserted = append(inserted, editResult.ReplacementMessage)
|
||||
inserted = append(inserted, editResult.SuffixMessages...)
|
||||
result.InsertedMessages = inserted
|
||||
result.DeletedMessageIDs = editResult.DeletedMessageIDs
|
||||
// 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.
|
||||
@@ -1958,21 +2156,38 @@ func (e *ToolResultStatusConflictError) Error() string {
|
||||
)
|
||||
}
|
||||
|
||||
// SubmitToolResults validates and persists client-provided tool
|
||||
// results, returning the chat to running through the chatstate state
|
||||
// machine. Validation runs inside the same transaction as the
|
||||
// transition so the assistant message and pending tool calls cannot
|
||||
// drift between reads.
|
||||
// SubmitToolResults dispatches hooks before completing the
|
||||
// requires_action transition.
|
||||
func (p *Server) SubmitToolResults(
|
||||
ctx context.Context,
|
||||
opts SubmitToolResultsOptions,
|
||||
) error {
|
||||
machine := p.newChatMachine(opts.ChatID)
|
||||
var hookSuffix []chatstate.Message
|
||||
if p.hooks.Enabled() {
|
||||
state, err := loadDynamicPostToolUseState(ctx, machine, opts)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, result := range opts.Results {
|
||||
response, err := p.hooks.Trigger(ctx, chathooks.ChatFor(state.chat, nil), chathooks.DynamicPostToolUseMessage(result, state.toolNames[result.ToolCallID]), agenthooks.EventPostToolUse)
|
||||
if err != nil {
|
||||
// Leave pending calls intact so the client can resubmit after recovery.
|
||||
return chathooks.GenerationDispatchError(agenthooks.EventPostToolUse, err)
|
||||
}
|
||||
responseMessages, err := chathooks.EventMessages(response, state.modelConfigID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
hookSuffix = append(hookSuffix, responseMessages...)
|
||||
}
|
||||
}
|
||||
|
||||
var (
|
||||
statusConflict *ToolResultStatusConflictError
|
||||
refreshChat database.Chat
|
||||
refreshedOK bool
|
||||
)
|
||||
machine := p.newChatMachine(opts.ChatID)
|
||||
updateErr := machine.Update(ctx, func(tx *chatstate.Tx, store database.Store) error {
|
||||
locked, err := store.GetChatByID(ctx, opts.ChatID)
|
||||
if err != nil {
|
||||
@@ -1983,11 +2198,11 @@ func (p *Server) SubmitToolResults(
|
||||
}
|
||||
|
||||
toolResults := make([]chatstate.ToolResultInput, 0, len(opts.Results))
|
||||
for _, r := range opts.Results {
|
||||
for _, result := range opts.Results {
|
||||
toolResults = append(toolResults, chatstate.ToolResultInput{
|
||||
ToolCallID: r.ToolCallID,
|
||||
Output: r.Output,
|
||||
IsError: r.IsError,
|
||||
ToolCallID: result.ToolCallID,
|
||||
Output: result.Output,
|
||||
IsError: result.IsError,
|
||||
})
|
||||
}
|
||||
modelConfigID := opts.ModelConfigID
|
||||
@@ -1995,9 +2210,10 @@ func (p *Server) SubmitToolResults(
|
||||
modelConfigID = locked.LastModelConfigID
|
||||
}
|
||||
if _, err := tx.CompleteRequiresAction(chatstate.CompleteRequiresActionInput{
|
||||
CreatedBy: opts.UserID,
|
||||
ModelConfigID: modelConfigID,
|
||||
Results: toolResults,
|
||||
CreatedBy: opts.UserID,
|
||||
ModelConfigID: modelConfigID,
|
||||
Results: toolResults,
|
||||
SuffixMessages: hookSuffix,
|
||||
}); err != nil {
|
||||
if !errors.Is(err, chatstate.ErrInvalidState) &&
|
||||
locked.Status != database.ChatStatusRequiresAction &&
|
||||
@@ -2009,9 +2225,6 @@ func (p *Server) SubmitToolResults(
|
||||
}
|
||||
return xerrors.Errorf("complete requires action: %w", err)
|
||||
}
|
||||
// Capture the chat inside the transaction so the watch event
|
||||
// uses the snapshot bump and status change produced by the
|
||||
// transition itself.
|
||||
refreshed, err := store.GetChatByID(ctx, opts.ChatID)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("reload chat after tool results: %w", err)
|
||||
@@ -2115,7 +2328,9 @@ func (p *Server) InterruptChat(
|
||||
// must be idle (waiting); the worker then generates and commits the
|
||||
// compaction summary through the normal generation loop, bypassing
|
||||
// the usage threshold, and the chat returns to waiting with no
|
||||
// assistant follow-up.
|
||||
// assistant follow-up unless a post_compact hook commits a
|
||||
// user-visible message, which leaves the history incomplete and
|
||||
// resumes generation.
|
||||
//
|
||||
// Returns the post-transition chat and an error so callers can map
|
||||
// state conflicts deliberately: archived chats return ErrChatArchived,
|
||||
@@ -2875,6 +3090,7 @@ type Config struct {
|
||||
AllowBYOKSet bool
|
||||
AlwaysEnableDebugLogs bool
|
||||
WebpushDispatcher webpush.Dispatcher
|
||||
HookDispatcher *dispatch.Dispatcher
|
||||
UsageTracker *workspacestats.UsageTracker
|
||||
Clock quartz.Clock
|
||||
AIBridgeTransportFactory *atomic.Pointer[aibridge.TransportFactory]
|
||||
@@ -2942,6 +3158,14 @@ func New(ps pubsub.Pubsub, cfg Config) *Server {
|
||||
if cfg.AllowBYOKSet {
|
||||
allowBYOK = cfg.AllowBYOK
|
||||
}
|
||||
|
||||
// Require the experiment even for injected dispatchers to
|
||||
// preserve explicit opt-in.
|
||||
hookDispatcher := cfg.HookDispatcher
|
||||
if hookDispatcher != nil && !cfg.Experiments.Enabled(codersdk.ExperimentAgentLifecycleHooks) {
|
||||
cfg.Logger.Warn(ctx, "ignoring chat lifecycle hook dispatcher; the agent-lifecycle-hooks experiment is not enabled")
|
||||
hookDispatcher = nil
|
||||
}
|
||||
p := &Server{
|
||||
cancel: cancel,
|
||||
db: cfg.Database,
|
||||
@@ -2956,6 +3180,7 @@ func New(ps pubsub.Pubsub, cfg Config) *Server {
|
||||
stopWorkspaceFn: cfg.StopWorkspace,
|
||||
pubsub: ps,
|
||||
webpushDispatcher: cfg.WebpushDispatcher,
|
||||
hooks: chathooks.NewTrigger(hookDispatcher),
|
||||
providerAPIKeys: cfg.ProviderAPIKeys,
|
||||
allowBYOK: allowBYOK,
|
||||
oidcTokenSource: cfg.OIDCTokenSource,
|
||||
|
||||
@@ -5523,7 +5523,7 @@ func TestActiveServer_ManualCompaction(t *testing.T) {
|
||||
|
||||
const compactionSummary = "manual compaction summary"
|
||||
|
||||
t.Run("compacts below threshold and returns to waiting", func(t *testing.T) {
|
||||
t.Run("compacts below threshold and returns to waiting without hooks", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
@@ -5571,7 +5571,7 @@ func TestActiveServer_ManualCompaction(t *testing.T) {
|
||||
"compaction commit must consume the request marker")
|
||||
require.Equal(t, int32(1), compactionRequests.Load(), "one forced compaction call")
|
||||
require.Equal(t, int32(1), streamCount.Load(),
|
||||
"manual compaction must not trigger an assistant follow-up")
|
||||
"manual compaction alone must not trigger an assistant follow-up")
|
||||
|
||||
messages := chatMessages(ctx, t, db, chat.ID)
|
||||
promptMessages, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
package chathooks
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
// EventMessages converts a turn-time hook result into ordinary
|
||||
// transcript rows: model context becomes a user-role, model-visible row
|
||||
// and the user message becomes a system-role, user-visible notice row.
|
||||
func EventMessages(result *Result, modelConfigID uuid.UUID) ([]chatstate.Message, error) {
|
||||
messages := make([]chatstate.Message, 0, 2)
|
||||
if strings.TrimSpace(result.GetModelContext()) != "" {
|
||||
content, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{codersdk.ChatMessageText(result.ModelContext)})
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("marshal hook model context: %w", err)
|
||||
}
|
||||
messages = append(messages, chatstate.Message{
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: content,
|
||||
Visibility: database.ChatMessageVisibilityModel,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: modelConfigID != uuid.Nil},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
})
|
||||
}
|
||||
if result.GetUserMessage() != "" {
|
||||
content, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{codersdk.ChatMessageText(result.UserMessage)})
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("marshal hook user message: %w", err)
|
||||
}
|
||||
messages = append(messages, chatstate.Message{
|
||||
Role: database.ChatMessageRoleSystem,
|
||||
Content: content,
|
||||
Visibility: database.ChatMessageVisibilityUser,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: modelConfigID != uuid.Nil},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
})
|
||||
}
|
||||
return messages, nil
|
||||
}
|
||||
|
||||
func EventMessagesForResults(
|
||||
results []*Result,
|
||||
modelConfigID uuid.UUID,
|
||||
) ([]chatstate.Message, error) {
|
||||
var messages []chatstate.Message
|
||||
for _, result := range results {
|
||||
resultMessages, err := EventMessages(result, modelConfigID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
messages = append(messages, resultMessages...)
|
||||
}
|
||||
return messages, nil
|
||||
}
|
||||
|
||||
// deniedToolResult synthesizes the denial as a tool result so the model
|
||||
// can replan within the same turn. The result is client-visible, so it
|
||||
// must never carry the consumer's model_context; that travels as a
|
||||
// model-only transcript row instead. The text must distinguish a policy
|
||||
// denial from a genuine tool failure, or the model retries the call and
|
||||
// misreports the denial as an infrastructure error.
|
||||
func deniedToolResult(toolCall fantasy.ToolCallContent, reason string) fantasy.ToolResultContent {
|
||||
message := "This tool usage was blocked by an external policy" +
|
||||
" (the deployment's lifecycle hook); the tool call was not executed."
|
||||
if reason = strings.TrimSpace(reason); reason != "" {
|
||||
message += " Reason: " + reason + "."
|
||||
}
|
||||
message += " This is an administrative policy decision, not a tool or" +
|
||||
" workspace failure; retrying the same call will be denied again." +
|
||||
" Explain the policy block to the user and adjust your approach."
|
||||
return fantasy.ToolResultContent{
|
||||
ToolCallID: toolCall.ToolCallID,
|
||||
ToolName: toolCall.ToolName,
|
||||
Result: fantasy.ToolResultOutputContentError{
|
||||
Error: xerrors.New(message),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// RestoreToolCallOrder reorders known tool results to match the assistant's
|
||||
// call order while preserving slots for unrelated entries.
|
||||
func RestoreToolCallOrder(content []fantasy.Content, calls []fantasy.ToolCallContent) {
|
||||
position := make(map[string]int, len(calls))
|
||||
for index, call := range calls {
|
||||
position[call.ToolCallID] = index
|
||||
}
|
||||
slots := make([]int, 0, len(content))
|
||||
results := make([]fantasy.ToolResultContent, 0, len(content))
|
||||
for index, entry := range content {
|
||||
result, ok := entry.(fantasy.ToolResultContent)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if _, known := position[result.ToolCallID]; !known {
|
||||
continue
|
||||
}
|
||||
slots = append(slots, index)
|
||||
results = append(results, result)
|
||||
}
|
||||
slices.SortStableFunc(results, func(a, b fantasy.ToolResultContent) int {
|
||||
return position[a.ToolCallID] - position[b.ToolCallID]
|
||||
})
|
||||
for index, slot := range slots {
|
||||
content[slot] = results[index]
|
||||
}
|
||||
}
|
||||
|
||||
func UserPromptOverride(result *Result) (string, bool, error) {
|
||||
if result == nil || len(result.InputOverride) == 0 {
|
||||
return "", false, nil
|
||||
}
|
||||
var override struct {
|
||||
Prompt *string `json:"prompt"`
|
||||
}
|
||||
decoder := json.NewDecoder(bytes.NewReader(result.InputOverride))
|
||||
decoder.DisallowUnknownFields()
|
||||
if err := decoder.Decode(&override); err != nil {
|
||||
return "", false, xerrors.Errorf("decode user prompt input override: %w", err)
|
||||
}
|
||||
if override.Prompt == nil {
|
||||
return "", false, xerrors.New("decode user prompt input override: prompt is required")
|
||||
}
|
||||
if err := decoder.Decode(&struct{}{}); err != io.EOF {
|
||||
return "", false, xerrors.New("decode user prompt input override: trailing JSON value")
|
||||
}
|
||||
return *override.Prompt, true, nil
|
||||
}
|
||||
|
||||
func UserPromptParts(result *Result) []codersdk.ChatMessagePart {
|
||||
parts := make([]codersdk.ChatMessagePart, 0, 2)
|
||||
if result.GetModelContext() != "" {
|
||||
parts = append(parts, codersdk.ChatMessagePart{
|
||||
Type: codersdk.ChatMessagePartTypeHookContext,
|
||||
Text: result.ModelContext,
|
||||
})
|
||||
}
|
||||
if result.GetUserMessage() != "" {
|
||||
parts = append(parts, codersdk.ChatMessagePart{
|
||||
Type: codersdk.ChatMessagePartTypeHookNotice,
|
||||
Text: result.UserMessage,
|
||||
})
|
||||
}
|
||||
return parts
|
||||
}
|
||||
|
||||
// ComposeUserPromptContent applies a user_prompt_submit result to the
|
||||
// submitted parts. The merge order is fixed: override-or-original user
|
||||
// parts first, then hook-context, then hook-notice. The composite
|
||||
// content then flows through the ordinary send, queue, and edit paths.
|
||||
func ComposeUserPromptContent(parts []codersdk.ChatMessagePart, result *Result) ([]codersdk.ChatMessagePart, bool, error) {
|
||||
override, overridden, err := UserPromptOverride(result)
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
userParts := parts
|
||||
if overridden {
|
||||
userParts = []codersdk.ChatMessagePart{codersdk.ChatMessageText(override)}
|
||||
}
|
||||
hookParts := UserPromptParts(result)
|
||||
if len(hookParts) == 0 {
|
||||
return userParts, overridden, nil
|
||||
}
|
||||
combined := make([]codersdk.ChatMessagePart, 0, len(userParts)+len(hookParts))
|
||||
combined = append(combined, userParts...)
|
||||
combined = append(combined, hookParts...)
|
||||
return combined, overridden, nil
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
package chathooks
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"charm.land/fantasy"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/x/agenthooks/dispatch"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chaterror"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
)
|
||||
|
||||
// deniedError is trigger's normalized form of a permission deny.
|
||||
// Callers translate it per event: user_prompt_submit sites map it to
|
||||
// UserPromptDeniedError, pre_tool_use sites fold it into a synthetic
|
||||
// tool result.
|
||||
type deniedError struct {
|
||||
Event agenthooks.EventType
|
||||
Reason string
|
||||
ModelContext string
|
||||
UserMessage string
|
||||
}
|
||||
|
||||
func (e *deniedError) Error() string {
|
||||
if e.Reason == "" {
|
||||
return fmt.Sprintf("%s denied by lifecycle hook", e.Event)
|
||||
}
|
||||
return fmt.Sprintf("%s denied by lifecycle hook: %s", e.Event, e.Reason)
|
||||
}
|
||||
|
||||
// UserPromptDeniedError reports that a lifecycle hook rejected a prompt.
|
||||
type UserPromptDeniedError struct {
|
||||
UserMessage string
|
||||
}
|
||||
|
||||
// Error includes UserMessage so callers that only surface the error
|
||||
// string, such as subagent tool responses, still expose the user-facing
|
||||
// denial message. The HTTP handlers unwrap the typed error instead.
|
||||
func (e *UserPromptDeniedError) Error() string {
|
||||
if e.UserMessage == "" {
|
||||
return "user prompt denied by lifecycle hook"
|
||||
}
|
||||
return "user prompt denied by lifecycle hook: " + e.UserMessage
|
||||
}
|
||||
|
||||
func UserPromptDenial(err error) error {
|
||||
if denied, ok := errors.AsType[*deniedError](err); ok {
|
||||
return &UserPromptDeniedError{UserMessage: denied.UserMessage}
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func DispatchErrorMessage(eventType agenthooks.EventType, dispatchErr error) (string, bool) {
|
||||
structured, ok := errors.AsType[*dispatch.Error](dispatchErr)
|
||||
if !ok {
|
||||
return "", false
|
||||
}
|
||||
return fmt.Sprintf(
|
||||
"hook dispatch failed: %s: %s (dispatch %s)",
|
||||
eventType,
|
||||
structured.Class,
|
||||
structured.DispatchID,
|
||||
), true
|
||||
}
|
||||
|
||||
func GenerationDispatchError(eventType agenthooks.EventType, dispatchErr error) error {
|
||||
message, ok := DispatchErrorMessage(eventType, dispatchErr)
|
||||
if !ok {
|
||||
message = dispatchErr.Error()
|
||||
}
|
||||
return chaterror.WithClassification(dispatchErr, chaterror.ClassifiedError{
|
||||
Message: message,
|
||||
Kind: codersdk.ChatErrorKindHookDispatchFailed,
|
||||
})
|
||||
}
|
||||
|
||||
// DispatchFailureFromResults returns the first tool result error
|
||||
// whose chain contains a hook dispatch failure. Tools that dispatch
|
||||
// lifecycle hooks inside Run (subagent spawn admission) must fail
|
||||
// closed, but the tool loop persists Run errors as ordinary tool
|
||||
// results the model can ignore, so the turn has to be failed even
|
||||
// though the step commits.
|
||||
func DispatchFailureFromResults(content []fantasy.Content) error {
|
||||
for _, block := range content {
|
||||
toolResult, ok := asToolResultContent(block)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if resultErr := dispatchFailureFromResult(toolResult); resultErr != nil {
|
||||
return resultErr
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func dispatchFailureFromResult(toolResult fantasy.ToolResultContent) error {
|
||||
var resultErr error
|
||||
switch output := toolResult.Result.(type) {
|
||||
case fantasy.ToolResultOutputContentError:
|
||||
resultErr = output.Error
|
||||
case *fantasy.ToolResultOutputContentError:
|
||||
if output != nil {
|
||||
resultErr = output.Error
|
||||
}
|
||||
}
|
||||
if resultErr == nil {
|
||||
return nil
|
||||
}
|
||||
if _, ok := errors.AsType[*dispatch.Error](resultErr); !ok {
|
||||
return nil
|
||||
}
|
||||
return resultErr
|
||||
}
|
||||
|
||||
func asToolResultContent(block fantasy.Content) (fantasy.ToolResultContent, bool) {
|
||||
if tr, ok := fantasy.AsContentType[fantasy.ToolResultContent](block); ok {
|
||||
return tr, true
|
||||
}
|
||||
if tr, ok := fantasy.AsContentType[*fantasy.ToolResultContent](block); ok && tr != nil {
|
||||
return *tr, true
|
||||
}
|
||||
return fantasy.ToolResultContent{}, false
|
||||
}
|
||||
@@ -0,0 +1,275 @@
|
||||
package chathooks
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"github.com/google/uuid"
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog/v3/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/x/agenthooks/dispatch"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
func TestSessionStartDispatchSources(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const secret = "test-hook-secret-32-bytes-minimum!!"
|
||||
type received struct {
|
||||
request agenthooks.Request
|
||||
claims agenthooks.Claims
|
||||
data agenthooks.SessionStartData
|
||||
}
|
||||
receivedCh := make(chan received, 2)
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
claims, err := agenthooks.Verify(r.Header.Get("Authorization"), []byte(secret))
|
||||
require.NoError(t, err)
|
||||
var data agenthooks.SessionStartData
|
||||
require.NoError(t, json.Unmarshal(request.Data, &data))
|
||||
receivedCh <- received{request: request, claims: claims, data: data}
|
||||
_, err = w.Write([]byte(`{}`))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
db, _ := dbtestutil.NewDB(t)
|
||||
dispatcher := dispatch.New(
|
||||
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
|
||||
consumer.Client(),
|
||||
consumer.URL,
|
||||
secret,
|
||||
time.Second,
|
||||
"test-deployment",
|
||||
"test-version",
|
||||
prometheus.NewRegistry(),
|
||||
)
|
||||
trigger := NewTrigger(dispatcher)
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{})
|
||||
chat := dbgen.Chat(t, db, database.Chat{OwnerID: user.ID, OrganizationID: org.ID, LastModelConfigID: model.ID})
|
||||
turnID := uuid.New()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
_, err := trigger.Trigger(ctx, ChatFor(chat, &turnID), Message{Source: SessionStartSource(nil)}, agenthooks.EventSessionStart)
|
||||
require.NoError(t, err)
|
||||
_, err = trigger.Trigger(ctx, ChatFor(chat, &turnID), Message{Source: SessionStartSource([]database.ChatMessage{{Role: database.ChatMessageRoleAssistant}})}, agenthooks.EventSessionStart)
|
||||
require.NoError(t, err)
|
||||
|
||||
startup := <-receivedCh
|
||||
resume := <-receivedCh
|
||||
require.Equal(t, agenthooks.EventSessionStart, startup.request.Type)
|
||||
require.Equal(t, SessionStartSourceStartup, startup.data.Source)
|
||||
require.Equal(t, startup.request.Meta.DispatchID, startup.claims.JTI)
|
||||
require.Equal(t, agenthooks.EventSessionStart, resume.request.Type)
|
||||
require.Equal(t, SessionStartSourceResume, resume.data.Source)
|
||||
require.Equal(t, resume.request.Meta.DispatchID, resume.claims.JTI)
|
||||
require.NotEqual(t, startup.claims.JTI, resume.claims.JTI)
|
||||
}
|
||||
|
||||
func TestRejectDuplicateToolUseIDs(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.NoError(t, RejectDuplicateToolUseIDs([]fantasy.ToolCallContent{
|
||||
{ToolCallID: "first", ToolName: "read_file", Input: `{}`},
|
||||
{ToolCallID: "second", ToolName: "execute", Input: `{}`},
|
||||
}))
|
||||
require.ErrorContains(t, RejectDuplicateToolUseIDs([]fantasy.ToolCallContent{
|
||||
{ToolCallID: "duplicate", ToolName: "read_file", Input: `{}`},
|
||||
{ToolCallID: "duplicate", ToolName: "execute", Input: `{}`},
|
||||
}), "duplicate tool use ID")
|
||||
}
|
||||
|
||||
func newTestTrigger(t *testing.T, handler http.Handler) *Trigger {
|
||||
t.Helper()
|
||||
consumer := httptest.NewServer(handler)
|
||||
t.Cleanup(consumer.Close)
|
||||
return NewTrigger(dispatch.New(
|
||||
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
|
||||
consumer.Client(),
|
||||
consumer.URL,
|
||||
"test-hook-secret-32-bytes-minimum!!",
|
||||
time.Second,
|
||||
"test-deployment",
|
||||
"test-version",
|
||||
prometheus.NewRegistry(),
|
||||
))
|
||||
}
|
||||
|
||||
func TestHookTriggerDisabled(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for name, trigger := range map[string]*Trigger{
|
||||
"NilTrigger": nil,
|
||||
"NilDispatcher": NewTrigger(nil),
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
require.False(t, trigger.Enabled())
|
||||
result, err := trigger.Trigger(t.Context(), Chat{ID: uuid.New()}, Message{}, agenthooks.EventStop)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, result.GetModelContext())
|
||||
require.Empty(t, result.GetUserMessage())
|
||||
require.Empty(t, result.InputOverride)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestHookTriggerDeny(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
trigger := newTestTrigger(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, err := w.Write([]byte(`{
|
||||
"permission": {"decision": "deny", "reason": "policy"},
|
||||
"model_context": "try another tool",
|
||||
"user_message": "blocked by policy"
|
||||
}`))
|
||||
assert.NoError(t, err)
|
||||
}))
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
result, err := trigger.Trigger(ctx, Chat{ID: uuid.New(), OwnerID: uuid.New()}, Message{
|
||||
ToolUseID: "call_1",
|
||||
ToolName: "execute",
|
||||
ToolInput: json.RawMessage(`{}`),
|
||||
}, agenthooks.EventPreToolUse)
|
||||
require.Nil(t, result)
|
||||
var denied *deniedError
|
||||
require.ErrorAs(t, err, &denied)
|
||||
require.Equal(t, agenthooks.EventPreToolUse, denied.Event)
|
||||
require.Equal(t, "policy", denied.Reason)
|
||||
require.Equal(t, "try another tool", denied.ModelContext)
|
||||
require.Equal(t, "blocked by policy", denied.UserMessage)
|
||||
}
|
||||
|
||||
func TestHookTriggerEventPayloads(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
requests := make(chan agenthooks.Request, 1)
|
||||
trigger := newTestTrigger(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
assert.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
requests <- request
|
||||
_, err := w.Write([]byte(`{}`))
|
||||
assert.NoError(t, err)
|
||||
}))
|
||||
chat := Chat{
|
||||
ID: uuid.New(),
|
||||
OwnerID: uuid.New(),
|
||||
WorkspaceID: uuid.NullUUID{UUID: uuid.New(), Valid: true},
|
||||
}
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
dispatchEvent := func(t *testing.T, msg Message, event agenthooks.EventType) agenthooks.Request {
|
||||
t.Helper()
|
||||
_, err := trigger.Trigger(ctx, chat, msg, event)
|
||||
require.NoError(t, err)
|
||||
request := <-requests
|
||||
require.Equal(t, event, request.Type)
|
||||
require.Equal(t, chat.ID, request.Meta.ChatID)
|
||||
require.Equal(t, chat.OwnerID, request.Meta.OwnerID)
|
||||
require.NotNil(t, request.Meta.WorkspaceID)
|
||||
require.Equal(t, chat.WorkspaceID.UUID, *request.Meta.WorkspaceID)
|
||||
return request
|
||||
}
|
||||
|
||||
sessionStart := dispatchEvent(t, Message{Source: SessionStartSourceClear}, agenthooks.EventSessionStart)
|
||||
var sessionStartData agenthooks.SessionStartData
|
||||
require.NoError(t, json.Unmarshal(sessionStart.Data, &sessionStartData))
|
||||
require.Equal(t, SessionStartSourceClear, sessionStartData.Source)
|
||||
|
||||
prompt := dispatchEvent(t, Message{Prompt: "hello", Parts: json.RawMessage(`[{"type":"text","text":"hello"}]`)}, agenthooks.EventUserPromptSubmit)
|
||||
var promptData agenthooks.UserPromptSubmitData
|
||||
require.NoError(t, json.Unmarshal(prompt.Data, &promptData))
|
||||
require.Equal(t, "hello", promptData.Prompt)
|
||||
require.JSONEq(t, `[{"type":"text","text":"hello"}]`, string(promptData.Parts))
|
||||
|
||||
preToolUse := dispatchEvent(t, Message{ToolUseID: "call_1", ToolName: "execute", ToolInput: json.RawMessage(`{"cmd":"ls"}`)}, agenthooks.EventPreToolUse)
|
||||
var preToolUseData agenthooks.PreToolUseData
|
||||
require.NoError(t, json.Unmarshal(preToolUse.Data, &preToolUseData))
|
||||
require.Equal(t, "call_1", preToolUseData.ToolUseID)
|
||||
require.Equal(t, "execute", preToolUseData.ToolName)
|
||||
require.JSONEq(t, `{"cmd":"ls"}`, string(preToolUseData.ToolInput))
|
||||
|
||||
postToolUse := dispatchEvent(t, Message{ToolUseID: "call_1", ToolName: "execute", ToolResponse: json.RawMessage(`{"ok":true}`), ToolError: "boom"}, agenthooks.EventPostToolUse)
|
||||
var postToolUseData agenthooks.PostToolUseData
|
||||
require.NoError(t, json.Unmarshal(postToolUse.Data, &postToolUseData))
|
||||
require.Equal(t, "call_1", postToolUseData.ToolUseID)
|
||||
require.Equal(t, "execute", postToolUseData.ToolName)
|
||||
require.JSONEq(t, `{"ok":true}`, string(postToolUseData.ToolResponse))
|
||||
require.Equal(t, "boom", postToolUseData.ToolError)
|
||||
|
||||
for _, event := range []agenthooks.EventType{agenthooks.EventPreCompact, agenthooks.EventPostCompact, agenthooks.EventStop} {
|
||||
dispatchEvent(t, Message{}, event)
|
||||
}
|
||||
|
||||
_, err := trigger.Trigger(ctx, chat, Message{}, agenthooks.EventType("bogus"))
|
||||
require.ErrorContains(t, err, "unsupported hook event")
|
||||
}
|
||||
|
||||
func TestRestoreToolCallOrder(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
calls := []fantasy.ToolCallContent{
|
||||
{ToolCallID: "call_a", ToolName: "write_file"},
|
||||
{ToolCallID: "call_b", ToolName: "read_file"},
|
||||
{ToolCallID: "call_c", ToolName: "execute"},
|
||||
}
|
||||
content := []fantasy.Content{
|
||||
fantasy.ToolResultContent{ToolCallID: "call_c", ToolName: "execute"},
|
||||
fantasy.ToolResultContent{ToolCallID: "call_b", ToolName: "read_file"},
|
||||
fantasy.ToolResultContent{ToolCallID: "call_a", ToolName: "write_file"},
|
||||
}
|
||||
RestoreToolCallOrder(content, calls)
|
||||
gotIDs := make([]string, 0, len(content))
|
||||
for _, entry := range content {
|
||||
result, ok := entry.(fantasy.ToolResultContent)
|
||||
require.True(t, ok)
|
||||
gotIDs = append(gotIDs, result.ToolCallID)
|
||||
}
|
||||
require.Equal(t, []string{"call_a", "call_b", "call_c"}, gotIDs)
|
||||
|
||||
mixed := []fantasy.Content{
|
||||
fantasy.ToolResultContent{ToolCallID: "call_b", ToolName: "read_file"},
|
||||
fantasy.TextContent{Text: "note"},
|
||||
fantasy.ToolResultContent{ToolCallID: "unknown", ToolName: "other"},
|
||||
fantasy.ToolResultContent{ToolCallID: "call_a", ToolName: "write_file"},
|
||||
}
|
||||
RestoreToolCallOrder(mixed, calls)
|
||||
first, ok := mixed[0].(fantasy.ToolResultContent)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "call_a", first.ToolCallID)
|
||||
_, ok = mixed[1].(fantasy.TextContent)
|
||||
require.True(t, ok)
|
||||
unknown, ok := mixed[2].(fantasy.ToolResultContent)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "unknown", unknown.ToolCallID)
|
||||
last, ok := mixed[3].(fantasy.ToolResultContent)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "call_b", last.ToolCallID)
|
||||
}
|
||||
|
||||
func TestEventMessagesSkipsBlankModelContext(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
modelConfigID := uuid.New()
|
||||
messages, err := EventMessages(&Result{ModelContext: " \n\t "}, modelConfigID)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, messages)
|
||||
|
||||
messages, err = EventMessages(&Result{ModelContext: "real context"}, modelConfigID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, messages, 1)
|
||||
require.Equal(t, database.ChatMessageVisibilityModel, messages[0].Visibility)
|
||||
}
|
||||
@@ -0,0 +1,205 @@
|
||||
package chathooks
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
)
|
||||
|
||||
// RejectDuplicateToolUseIDs fails closed because hook consumers key decisions
|
||||
// by tool-use ID; a duplicated ID in one step makes decisions unattributable.
|
||||
// Callers must check the complete set of pending calls before removing any,
|
||||
// because a filtered-out duplicate still shares its ID with a synthetic result.
|
||||
func RejectDuplicateToolUseIDs(toolCalls []fantasy.ToolCallContent) error {
|
||||
seen := make(map[string]struct{}, len(toolCalls))
|
||||
for _, toolCall := range toolCalls {
|
||||
if toolCall.ProviderExecuted {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[toolCall.ToolCallID]; ok {
|
||||
return xerrors.Errorf("duplicate tool use ID %q in one step; lifecycle hook decisions cannot be attributed unambiguously", toolCall.ToolCallID)
|
||||
}
|
||||
seen[toolCall.ToolCallID] = struct{}{}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// PreToolUseExecutionResult preserves hook results in tool-call order for
|
||||
// transcript injection.
|
||||
type PreToolUseExecutionResult struct {
|
||||
Allowed []fantasy.ToolCallContent
|
||||
Denied []fantasy.ToolResultContent
|
||||
Results []*Result
|
||||
Overrides map[string]json.RawMessage
|
||||
}
|
||||
|
||||
func (t *Trigger) PreflightPendingToolCalls(
|
||||
ctx context.Context,
|
||||
chat Chat,
|
||||
toolCalls []fantasy.ToolCallContent,
|
||||
) (PreToolUseExecutionResult, error) {
|
||||
if !t.Enabled() {
|
||||
return PreToolUseExecutionResult{Allowed: toolCalls}, nil
|
||||
}
|
||||
result := PreToolUseExecutionResult{
|
||||
Allowed: make([]fantasy.ToolCallContent, 0, len(toolCalls)),
|
||||
}
|
||||
if err := RejectDuplicateToolUseIDs(toolCalls); err != nil {
|
||||
return PreToolUseExecutionResult{}, err
|
||||
}
|
||||
|
||||
for _, toolCall := range toolCalls {
|
||||
callResult, err := t.Trigger(ctx, chat, Message{
|
||||
ToolUseID: toolCall.ToolCallID,
|
||||
ToolName: toolCall.ToolName,
|
||||
ToolInput: json.RawMessage(toolCall.Input),
|
||||
}, agenthooks.EventPreToolUse)
|
||||
if err != nil {
|
||||
denied, ok := errors.AsType[*deniedError](err)
|
||||
if !ok {
|
||||
return PreToolUseExecutionResult{}, err
|
||||
}
|
||||
// The synthetic tool result is client-visible, so the
|
||||
// denial's model context becomes a model-only transcript
|
||||
// row instead of riding in the result.
|
||||
result.Results = append(result.Results, &Result{
|
||||
ModelContext: denied.ModelContext,
|
||||
UserMessage: denied.UserMessage,
|
||||
})
|
||||
result.Denied = append(result.Denied, deniedToolResult(toolCall, denied.Reason))
|
||||
continue
|
||||
}
|
||||
result.Results = append(result.Results, callResult)
|
||||
if len(callResult.InputOverride) > 0 {
|
||||
toolCall.Input = string(callResult.InputOverride)
|
||||
if result.Overrides == nil {
|
||||
result.Overrides = make(map[string]json.RawMessage)
|
||||
}
|
||||
result.Overrides[toolCall.ToolCallID] = callResult.InputOverride
|
||||
}
|
||||
result.Allowed = append(result.Allowed, toolCall)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func postToolUseMessage(toolResult fantasy.ToolResultContent) (Message, error) {
|
||||
msg := Message{
|
||||
ToolUseID: toolResult.ToolCallID,
|
||||
ToolName: toolResult.ToolName,
|
||||
}
|
||||
switch output := toolResult.Result.(type) {
|
||||
case fantasy.ToolResultOutputContentError:
|
||||
if output.Error != nil {
|
||||
msg.ToolError = output.Error.Error()
|
||||
}
|
||||
case *fantasy.ToolResultOutputContentError:
|
||||
if output != nil && output.Error != nil {
|
||||
msg.ToolError = output.Error.Error()
|
||||
}
|
||||
default:
|
||||
encoded, err := json.Marshal(toolResult.Result)
|
||||
if err != nil {
|
||||
return Message{}, xerrors.Errorf("marshal post_tool_use response: %w", err)
|
||||
}
|
||||
msg.ToolResponse = encoded
|
||||
}
|
||||
return msg, nil
|
||||
}
|
||||
|
||||
func (t *Trigger) PostToolUseResults(
|
||||
ctx context.Context,
|
||||
chat Chat,
|
||||
content []fantasy.Content,
|
||||
) ([]*Result, error) {
|
||||
if !t.Enabled() {
|
||||
return nil, nil
|
||||
}
|
||||
results := make([]*Result, 0, len(content))
|
||||
var firstErr error
|
||||
for _, block := range content {
|
||||
toolResult, ok := asToolResultContent(block)
|
||||
if !ok || toolResult.ProviderExecuted {
|
||||
continue
|
||||
}
|
||||
// A hook dispatch failure means admission was refused, so the tool
|
||||
// never ran and there is no use to post-process.
|
||||
if dispatchFailureFromResult(toolResult) != nil {
|
||||
continue
|
||||
}
|
||||
msg, err := postToolUseMessage(toolResult)
|
||||
if err != nil {
|
||||
if firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
continue
|
||||
}
|
||||
result, err := t.Trigger(ctx, chat, msg, agenthooks.EventPostToolUse)
|
||||
if err != nil {
|
||||
if firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
continue
|
||||
}
|
||||
results = append(results, result)
|
||||
}
|
||||
return results, firstErr
|
||||
}
|
||||
|
||||
// ApplyAdmittedToolCalls applies admitted inputs before persistence and
|
||||
// appends denial results.
|
||||
func ApplyAdmittedToolCalls(content []fantasy.Content, preflight PreToolUseExecutionResult) []fantasy.Content {
|
||||
if len(preflight.Overrides) == 0 && len(preflight.Denied) == 0 {
|
||||
return content
|
||||
}
|
||||
rewritten := make([]fantasy.Content, 0, len(content)+len(preflight.Denied))
|
||||
for _, block := range content {
|
||||
toolCall, ok := fantasy.AsContentType[fantasy.ToolCallContent](block)
|
||||
if !ok {
|
||||
rewritten = append(rewritten, block)
|
||||
continue
|
||||
}
|
||||
if input, found := preflight.Overrides[toolCall.ToolCallID]; found {
|
||||
toolCall.Input = string(input)
|
||||
}
|
||||
rewritten = append(rewritten, toolCall)
|
||||
}
|
||||
for _, denied := range preflight.Denied {
|
||||
rewritten = append(rewritten, denied)
|
||||
}
|
||||
return rewritten
|
||||
}
|
||||
|
||||
// PendingToolCalls returns the calls a step leaves for Coder to run. Hooks
|
||||
// never see provider-executed calls because the provider runs them itself.
|
||||
func PendingToolCalls(content []fantasy.Content) []fantasy.ToolCallContent {
|
||||
toolCalls := make([]fantasy.ToolCallContent, 0, len(content))
|
||||
for _, block := range content {
|
||||
toolCall, ok := fantasy.AsContentType[fantasy.ToolCallContent](block)
|
||||
if !ok || toolCall.ProviderExecuted {
|
||||
continue
|
||||
}
|
||||
toolCalls = append(toolCalls, toolCall)
|
||||
}
|
||||
return toolCalls
|
||||
}
|
||||
|
||||
func DynamicPostToolUseMessage(result codersdk.ToolResult, toolName string) Message {
|
||||
msg := Message{
|
||||
ToolUseID: result.ToolCallID,
|
||||
ToolName: toolName,
|
||||
}
|
||||
if result.IsError {
|
||||
if err := json.Unmarshal(result.Output, &msg.ToolError); err != nil {
|
||||
msg.ToolError = string(result.Output)
|
||||
}
|
||||
} else {
|
||||
msg.ToolResponse = append(json.RawMessage(nil), result.Output...)
|
||||
}
|
||||
return msg
|
||||
}
|
||||
@@ -0,0 +1,202 @@
|
||||
// Package chathooks integrates chat lifecycle hooks into chatd: it
|
||||
// builds event envelopes, dispatches them, and converts consumer
|
||||
// responses into transcript effects and permission decisions.
|
||||
package chathooks
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/x/agenthooks/dispatch"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
)
|
||||
|
||||
const (
|
||||
SessionStartSourceStartup = "startup"
|
||||
SessionStartSourceResume = "resume"
|
||||
SessionStartSourceClear = "clear"
|
||||
)
|
||||
|
||||
func SessionStartSource(messages []database.ChatMessage) string {
|
||||
for _, message := range messages {
|
||||
if message.Role == database.ChatMessageRoleAssistant {
|
||||
return SessionStartSourceResume
|
||||
}
|
||||
}
|
||||
return SessionStartSourceStartup
|
||||
}
|
||||
|
||||
// Trigger is the only component that talks to the hook dispatcher.
|
||||
// Every lifecycle event flows through trigger, which builds the wire
|
||||
// envelope, dispatches, and normalizes the outcome.
|
||||
type Trigger struct {
|
||||
dispatcher *dispatch.Dispatcher
|
||||
}
|
||||
|
||||
func NewTrigger(dispatcher *dispatch.Dispatcher) *Trigger {
|
||||
return &Trigger{dispatcher: dispatcher}
|
||||
}
|
||||
|
||||
func (t *Trigger) Enabled() bool {
|
||||
return t != nil && t.dispatcher.Enabled()
|
||||
}
|
||||
|
||||
// Chat identifies the chat and turn an event belongs to. Admission
|
||||
// events for chats that do not exist yet (create, subagent spawn) fill
|
||||
// the fields directly instead of loading a row.
|
||||
type Chat struct {
|
||||
ID uuid.UUID
|
||||
OwnerID uuid.UUID
|
||||
WorkspaceID uuid.NullUUID
|
||||
ParentChatID uuid.NullUUID
|
||||
RootChatID uuid.NullUUID
|
||||
TurnID *uuid.UUID
|
||||
}
|
||||
|
||||
func ChatFor(chat database.Chat, turnID *uuid.UUID) Chat {
|
||||
return Chat{
|
||||
ID: chat.ID,
|
||||
OwnerID: chat.OwnerID,
|
||||
WorkspaceID: chat.WorkspaceID,
|
||||
ParentChatID: chat.ParentChatID,
|
||||
RootChatID: chat.RootChatID,
|
||||
TurnID: turnID,
|
||||
}
|
||||
}
|
||||
|
||||
func (c Chat) ref() agenthooks.ChatRef {
|
||||
ref := agenthooks.ChatRef{
|
||||
ChatID: c.ID,
|
||||
OwnerID: c.OwnerID,
|
||||
TurnID: c.TurnID,
|
||||
}
|
||||
if c.WorkspaceID.Valid {
|
||||
ref.WorkspaceID = &c.WorkspaceID.UUID
|
||||
}
|
||||
if c.ParentChatID.Valid {
|
||||
ref.ParentChatID = &c.ParentChatID.UUID
|
||||
}
|
||||
if c.RootChatID.Valid {
|
||||
ref.RootChatID = &c.RootChatID.UUID
|
||||
}
|
||||
return ref
|
||||
}
|
||||
|
||||
type Message struct {
|
||||
Source string
|
||||
Prompt string
|
||||
Parts json.RawMessage
|
||||
ToolUseID string
|
||||
ToolName string
|
||||
ToolInput json.RawMessage
|
||||
ToolResponse json.RawMessage
|
||||
ToolError string
|
||||
}
|
||||
|
||||
func UserPromptMessage(parts []codersdk.ChatMessagePart) (Message, error) {
|
||||
encoded, err := chatprompt.MarshalParts(parts)
|
||||
if err != nil {
|
||||
return Message{}, xerrors.Errorf("marshal prompt parts for hook: %w", err)
|
||||
}
|
||||
return Message{
|
||||
Prompt: textFromParts(parts),
|
||||
Parts: encoded.RawMessage,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Result is a consumer response normalized for callers: a non-empty
|
||||
// InputOverride means the permission decision was allow with a
|
||||
// replacement input (the wire contract rejects allow without one).
|
||||
// Denials surface as *deniedError instead.
|
||||
type Result struct {
|
||||
InputOverride json.RawMessage
|
||||
ModelContext string
|
||||
UserMessage string
|
||||
}
|
||||
|
||||
var emptyResult = &Result{}
|
||||
|
||||
func (r *Result) GetModelContext() string {
|
||||
if r == nil {
|
||||
return ""
|
||||
}
|
||||
return r.ModelContext
|
||||
}
|
||||
|
||||
func (r *Result) GetUserMessage() string {
|
||||
if r == nil {
|
||||
return ""
|
||||
}
|
||||
return r.UserMessage
|
||||
}
|
||||
|
||||
func (t *Trigger) Trigger(
|
||||
ctx context.Context,
|
||||
chat Chat,
|
||||
msg Message,
|
||||
event agenthooks.EventType,
|
||||
) (*Result, error) {
|
||||
if !t.Enabled() {
|
||||
return emptyResult, nil
|
||||
}
|
||||
var data any
|
||||
switch event {
|
||||
case agenthooks.EventSessionStart:
|
||||
data = agenthooks.SessionStartData{Source: msg.Source}
|
||||
case agenthooks.EventUserPromptSubmit:
|
||||
data = agenthooks.UserPromptSubmitData{Prompt: msg.Prompt, Parts: msg.Parts}
|
||||
case agenthooks.EventPreToolUse:
|
||||
data = agenthooks.PreToolUseData{ToolUseID: msg.ToolUseID, ToolName: msg.ToolName, ToolInput: msg.ToolInput}
|
||||
case agenthooks.EventPostToolUse:
|
||||
data = agenthooks.PostToolUseData{ToolUseID: msg.ToolUseID, ToolName: msg.ToolName, ToolResponse: msg.ToolResponse, ToolError: msg.ToolError}
|
||||
case agenthooks.EventPreCompact:
|
||||
data = agenthooks.PreCompactData{}
|
||||
case agenthooks.EventPostCompact:
|
||||
data = agenthooks.PostCompactData{}
|
||||
case agenthooks.EventStop:
|
||||
data = agenthooks.StopData{}
|
||||
default:
|
||||
return nil, xerrors.Errorf("unsupported hook event %q", event)
|
||||
}
|
||||
response, _, err := t.dispatcher.Dispatch(ctx, dispatch.Event{
|
||||
Type: event,
|
||||
ChatRef: chat.ref(),
|
||||
Data: data,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if response.Permission != nil && response.Permission.Decision == agenthooks.PermissionDeny {
|
||||
return nil, &deniedError{
|
||||
Event: event,
|
||||
Reason: response.Permission.Reason,
|
||||
ModelContext: response.ModelContext,
|
||||
UserMessage: response.UserMessage,
|
||||
}
|
||||
}
|
||||
result := &Result{
|
||||
ModelContext: response.ModelContext,
|
||||
UserMessage: response.UserMessage,
|
||||
}
|
||||
if response.Permission != nil && response.Permission.Decision == agenthooks.PermissionAllow {
|
||||
result.InputOverride = response.Permission.InputOverride
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func textFromParts(parts []codersdk.ChatMessagePart) string {
|
||||
var builder strings.Builder
|
||||
for _, part := range parts {
|
||||
if part.Type == codersdk.ChatMessagePartTypeText {
|
||||
_, _ = builder.WriteString(part.Text)
|
||||
}
|
||||
}
|
||||
return builder.String()
|
||||
}
|
||||
@@ -274,6 +274,7 @@ type GenerateCompactionOptions struct {
|
||||
ContextLimit int64
|
||||
ContextLimitFallback int64
|
||||
SummaryPrompt string
|
||||
SummaryHint string
|
||||
SystemSummaryPrefix string
|
||||
StepUsage fantasy.Usage
|
||||
StepMetadata fantasy.ProviderMetadata
|
||||
|
||||
@@ -87,6 +87,7 @@ type CompactionOptions struct {
|
||||
ThresholdPercent int32
|
||||
ContextLimit int64
|
||||
SummaryPrompt string
|
||||
SummaryHint string
|
||||
SystemSummaryPrefix string
|
||||
Persist func(context.Context, CompactionResult) error
|
||||
DebugSvc *chatdebug.Service
|
||||
@@ -213,6 +214,7 @@ func normalizedCompactionGenerateConfig(opts GenerateCompactionOptions) (Compact
|
||||
ThresholdPercent: opts.ThresholdPercent,
|
||||
ContextLimit: opts.ContextLimit,
|
||||
SummaryPrompt: opts.SummaryPrompt,
|
||||
SummaryHint: opts.SummaryHint,
|
||||
SystemSummaryPrefix: opts.SystemSummaryPrefix,
|
||||
DebugSvc: opts.DebugSvc,
|
||||
ChatID: opts.ChatID,
|
||||
@@ -416,11 +418,13 @@ func generateCompactionSummary(
|
||||
) (summary string, err error) {
|
||||
summaryPrompt := make([]fantasy.Message, 0, len(messages)+1)
|
||||
summaryPrompt = append(summaryPrompt, messages...)
|
||||
summaryParts := []fantasy.MessagePart{fantasy.TextPart{Text: options.SummaryPrompt}}
|
||||
if strings.TrimSpace(options.SummaryHint) != "" {
|
||||
summaryParts = append(summaryParts, fantasy.TextPart{Text: options.SummaryHint})
|
||||
}
|
||||
summaryPrompt = append(summaryPrompt, fantasy.Message{
|
||||
Role: fantasy.MessageRoleUser,
|
||||
Content: []fantasy.MessagePart{
|
||||
fantasy.TextPart{Text: options.SummaryPrompt},
|
||||
},
|
||||
Role: fantasy.MessageRoleUser,
|
||||
Content: summaryParts,
|
||||
})
|
||||
toolChoice := fantasy.ToolChoiceNone
|
||||
|
||||
|
||||
@@ -923,9 +923,20 @@ func hasErrorField(raw json.RawMessage) bool {
|
||||
return ok
|
||||
}
|
||||
|
||||
// injectMissingToolResults keeps tool results adjacent to the
|
||||
// assistant message that issued the calls. Hook effects, such as
|
||||
// pre_tool_use model context, can persist rows between an assistant
|
||||
// tool call and its result rows, so matching result rows found before
|
||||
// the next assistant message are hoisted back next to the call;
|
||||
// otherwise the interleaved row would make the real result look
|
||||
// orphaned. Unanswered local calls get synthetic interrupted results.
|
||||
func injectMissingToolResults(prompt []fantasy.Message) []fantasy.Message {
|
||||
result := make([]fantasy.Message, 0, len(prompt))
|
||||
hoisted := make(map[int]bool)
|
||||
for i := 0; i < len(prompt); i++ {
|
||||
if hoisted[i] {
|
||||
continue
|
||||
}
|
||||
msg := prompt[i]
|
||||
result = append(result, msg)
|
||||
|
||||
@@ -936,28 +947,38 @@ func injectMissingToolResults(prompt []fantasy.Message) []fantasy.Message {
|
||||
if len(toolCalls) == 0 {
|
||||
continue
|
||||
}
|
||||
callIDs := make(map[string]struct{}, len(toolCalls))
|
||||
for _, tc := range toolCalls {
|
||||
callIDs[tc.ToolCallID] = struct{}{}
|
||||
}
|
||||
|
||||
// Collect the tool call IDs that have results in the
|
||||
// following tool message(s).
|
||||
// Hoist tool rows answering this assistant's calls, in
|
||||
// persisted order, from anywhere before the next assistant
|
||||
// message. Interleaved non-tool rows keep their relative
|
||||
// order after the results.
|
||||
answered := make(map[string]struct{})
|
||||
j := i + 1
|
||||
for ; j < len(prompt); j++ {
|
||||
if prompt[j].Role != fantasy.MessageRoleTool {
|
||||
for j := i + 1; j < len(prompt); j++ {
|
||||
if prompt[j].Role == fantasy.MessageRoleAssistant {
|
||||
break
|
||||
}
|
||||
if prompt[j].Role != fantasy.MessageRoleTool {
|
||||
continue
|
||||
}
|
||||
answersThisCall := false
|
||||
for _, part := range prompt[j].Content {
|
||||
tr, ok := safeAsToolResultPart(part)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
answered[tr.ToolCallID] = struct{}{}
|
||||
if _, ok := callIDs[tr.ToolCallID]; ok {
|
||||
answersThisCall = true
|
||||
answered[tr.ToolCallID] = struct{}{}
|
||||
}
|
||||
}
|
||||
if answersThisCall {
|
||||
result = append(result, prompt[j])
|
||||
hoisted[j] = true
|
||||
}
|
||||
}
|
||||
if i+1 < j {
|
||||
// Preserve persisted tool result ordering and inject any
|
||||
// synthetic results after the existing contiguous tool messages.
|
||||
result = append(result, prompt[i+1:j]...)
|
||||
i = j - 1
|
||||
}
|
||||
|
||||
// Build synthetic results for any unanswered tool calls.
|
||||
@@ -1642,6 +1663,16 @@ func partsToMessageParts(
|
||||
_, _ = sb.WriteString(part.ContextFileContent)
|
||||
_, _ = sb.WriteString("\n</workspace-context>")
|
||||
result = append(result, fantasy.TextPart{Text: sb.String()})
|
||||
case codersdk.ChatMessagePartTypeHookContext:
|
||||
// Lifecycle hook model context rides inside the user
|
||||
// message and is sent to the model as plain text.
|
||||
if strings.TrimSpace(part.Text) == "" {
|
||||
continue
|
||||
}
|
||||
result = append(result, fantasy.TextPart{Text: part.Text})
|
||||
case codersdk.ChatMessagePartTypeHookNotice:
|
||||
// Client-only hook notice, never sent to the model.
|
||||
continue
|
||||
case codersdk.ChatMessagePartTypeSource:
|
||||
// Source parts are metadata-only, not sent to LLM.
|
||||
continue
|
||||
|
||||
@@ -927,6 +927,44 @@ func TestInjectMissingToolUses_DropsProviderExecutedOrphans(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestInjectMissingToolResults_HookContextBetweenCallAndResult(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
assistantContent := mustMarshalContent(t, []fantasy.Content{
|
||||
fantasy.ToolCallContent{
|
||||
ToolCallID: "toolu_gated",
|
||||
ToolName: "execute",
|
||||
Input: `{"command":"ls"}`,
|
||||
},
|
||||
})
|
||||
hookContext := mustMarshalContent(t, []fantasy.Content{
|
||||
fantasy.TextContent{Text: "hook approval context"},
|
||||
})
|
||||
result := mustMarshalToolResult(t,
|
||||
"toolu_gated", "execute",
|
||||
json.RawMessage(`{"output":"ok"}`),
|
||||
false, false, false,
|
||||
)
|
||||
|
||||
prompt := convertMessagesWithoutFiles(t, []database.ChatMessage{
|
||||
{Role: database.ChatMessageRoleAssistant, Visibility: database.ChatMessageVisibilityBoth, Content: assistantContent},
|
||||
{Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityModel, Content: hookContext},
|
||||
{Role: database.ChatMessageRoleTool, Visibility: database.ChatMessageVisibilityBoth, Content: result},
|
||||
})
|
||||
|
||||
require.Len(t, prompt, 3)
|
||||
require.Equal(t, fantasy.MessageRoleAssistant, prompt[0].Role)
|
||||
require.Equal(t, fantasy.MessageRoleTool, prompt[1].Role)
|
||||
require.Equal(t, fantasy.MessageRoleUser, prompt[2].Role)
|
||||
require.Equal(t, []string{"toolu_gated"}, extractToolResultIDs(t, prompt[1]))
|
||||
for _, part := range prompt[1].Content {
|
||||
tr, ok := asToolResultPartForTest(part)
|
||||
require.True(t, ok)
|
||||
_, isError := tr.Output.(fantasy.ToolResultOutputContentError)
|
||||
require.False(t, isError, "expected no synthetic interrupted result")
|
||||
}
|
||||
}
|
||||
|
||||
// TestInjectMissingToolUses_DropsOnlyProviderExecutedMessage verifies
|
||||
// that a tool message containing only a provider-executed result is
|
||||
// entirely dropped.
|
||||
|
||||
@@ -9,11 +9,13 @@ import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
)
|
||||
|
||||
// OpenAIHandler handles OpenAI API requests and returns a response.
|
||||
@@ -87,6 +89,44 @@ type OpenAIMessage struct {
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
// UnmarshalJSON accepts both string content and the structured
|
||||
// content-part array the SDK emits for multi-part messages,
|
||||
// concatenating the text items with newlines.
|
||||
func (m *OpenAIMessage) UnmarshalJSON(data []byte) error {
|
||||
var raw struct {
|
||||
Role string `json:"role"`
|
||||
Content json.RawMessage `json:"content"`
|
||||
}
|
||||
if err := json.Unmarshal(data, &raw); err != nil {
|
||||
return err
|
||||
}
|
||||
m.Role = raw.Role
|
||||
if len(raw.Content) == 0 {
|
||||
m.Content = ""
|
||||
return nil
|
||||
}
|
||||
var text string
|
||||
if err := json.Unmarshal(raw.Content, &text); err == nil {
|
||||
m.Content = text
|
||||
return nil
|
||||
}
|
||||
var parts []struct {
|
||||
Type string `json:"type"`
|
||||
Text string `json:"text"`
|
||||
}
|
||||
if err := json.Unmarshal(raw.Content, &parts); err != nil {
|
||||
return xerrors.Errorf("decode message content: %w", err)
|
||||
}
|
||||
var texts []string
|
||||
for _, part := range parts {
|
||||
if part.Type == "text" && part.Text != "" {
|
||||
texts = append(texts, part.Text)
|
||||
}
|
||||
}
|
||||
m.Content = strings.Join(texts, "\n")
|
||||
return nil
|
||||
}
|
||||
|
||||
// OpenAIToolFunction represents the function definition inside a tool.
|
||||
type OpenAIToolFunction struct {
|
||||
Name string `json:"name"`
|
||||
|
||||
@@ -42,24 +42,38 @@ type editFileEdit struct {
|
||||
func (e *editFileEdit) UnmarshalJSON(data []byte) error {
|
||||
var raw struct {
|
||||
OldText string `json:"old_text"`
|
||||
Search string `json:"search"`
|
||||
NewText string `json:"new_text"`
|
||||
Replace string `json:"replace"`
|
||||
ReplaceAll bool `json:"replace_all"`
|
||||
}
|
||||
if err := json.Unmarshal(data, &raw); err != nil {
|
||||
return err
|
||||
}
|
||||
e.OldText = raw.OldText
|
||||
if e.OldText == "" {
|
||||
e.OldText = raw.Search
|
||||
}
|
||||
e.NewText = raw.NewText
|
||||
if e.NewText == "" {
|
||||
e.NewText = raw.Replace
|
||||
}
|
||||
e.ReplaceAll = raw.ReplaceAll
|
||||
return nil
|
||||
if e.OldText != "" && e.NewText != "" {
|
||||
return nil
|
||||
}
|
||||
// The aliases are absent from the advertised schema, so tool-input
|
||||
// validation cannot reject case variants of them the way it does for
|
||||
// declared properties. Matching them exactly keeps the keys this
|
||||
// decoder reads identical to the ones a policy sees.
|
||||
var exact map[string]json.RawMessage
|
||||
if err := json.Unmarshal(data, &exact); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := decodeAlias(exact, "search", &e.OldText); err != nil {
|
||||
return err
|
||||
}
|
||||
return decodeAlias(exact, "replace", &e.NewText)
|
||||
}
|
||||
|
||||
func decodeAlias(raw map[string]json.RawMessage, key string, target *string) error {
|
||||
value, present := raw[key]
|
||||
if *target != "" || !present {
|
||||
return nil
|
||||
}
|
||||
return json.Unmarshal(value, target)
|
||||
}
|
||||
|
||||
func (a EditFilesArgs) toSDKFiles() []workspacesdk.FileEdits {
|
||||
|
||||
@@ -582,6 +582,37 @@ func TestEditFiles_DeprecatedSearchReplaceFieldsStillWork(t *testing.T) {
|
||||
assert.False(t, resp.IsError)
|
||||
}
|
||||
|
||||
func TestEditFiles_DeprecatedFieldsAreCaseSensitive(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
mockConn := agentconnmock.NewMockAgentConn(ctrl)
|
||||
targetPath := "/home/coder/main.go"
|
||||
mockConn.EXPECT().
|
||||
EditFiles(gomock.Any(), workspacesdk.FileEditRequest{
|
||||
Files: []workspacesdk.FileEdits{{
|
||||
Path: targetPath,
|
||||
Edits: []workspacesdk.FileEdit{{}},
|
||||
}},
|
||||
IncludeDiff: true,
|
||||
}).
|
||||
Return(workspacesdk.FileEditResponse{}, nil)
|
||||
|
||||
tool := chattool.EditFiles(chattool.EditFilesOptions{
|
||||
GetWorkspaceConn: func(context.Context) (workspacesdk.AgentConn, error) {
|
||||
return mockConn, nil
|
||||
},
|
||||
})
|
||||
|
||||
resp, err := tool.Run(context.Background(), fantasy.ToolCall{
|
||||
ID: "call-1",
|
||||
Name: "edit_files",
|
||||
Input: `{"files":[{"path":"` + targetPath + `","edits":[{"SEARCH":"old","REPLACE":"replacement"}]}]}`,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.False(t, resp.IsError)
|
||||
}
|
||||
|
||||
func TestEditFiles_NewFieldNamesTakePrecedenceOverOld(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -0,0 +1,331 @@
|
||||
package chatd_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk/agentconnmock"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
func TestCompactionHooksHintAndPostCommitResponses(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var postSawCommitted atomic.Bool
|
||||
fixture := startCompactionHookChat(t,
|
||||
func(t *testing.T, db database.Store, request agenthooks.Request) (int, string) {
|
||||
switch request.Type {
|
||||
case agenthooks.EventPreCompact:
|
||||
return http.StatusOK, `{"model_context":"preserve deployment constraints","user_message":"compaction starting"}`
|
||||
case agenthooks.EventPostCompact:
|
||||
postSawCommitted.Store(hasCompactionRows(t, db, request.Meta.ChatID))
|
||||
return http.StatusOK, `{"model_context":"post compact context","user_message":"compaction complete"}`
|
||||
default:
|
||||
return http.StatusOK, `{}`
|
||||
}
|
||||
},
|
||||
func(t *testing.T, body string) {
|
||||
require.Contains(t, body, "preserve deployment constraints")
|
||||
},
|
||||
)
|
||||
|
||||
waitCtx := testutil.Context(t, testutil.WaitLong)
|
||||
testutil.Eventually(waitCtx, t, func(context.Context) bool {
|
||||
updated, err := fixture.db.GetChatByID(waitCtx, fixture.chat.ID)
|
||||
return err == nil && updated.Status == database.ChatStatusWaiting && !updated.Archived
|
||||
}, testutil.IntervalFast)
|
||||
// post_compact runs before its effects commit with the compaction step.
|
||||
require.False(t, postSawCommitted.Load())
|
||||
require.Equal(t, int32(1), fixture.compactionCalls.Load())
|
||||
require.Equal(t, int32(2), fixture.streamCalls.Load(),
|
||||
"automatic compaction continues the turn, and hook effects must not suppress that")
|
||||
|
||||
userMessages := chatMessages(fixture.ctx, t, fixture.db, fixture.chat.ID)
|
||||
promptMessages, err := fixture.db.GetChatMessagesForPromptByChatID(fixture.ctx, fixture.chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.True(t, hasMessageText(t, userMessages, "compaction starting", database.ChatMessageVisibilityUser))
|
||||
require.True(t, hasMessageText(t, userMessages, "compaction complete", database.ChatMessageVisibilityUser))
|
||||
require.True(t, hasMessageText(t, promptMessages, "post compact context", database.ChatMessageVisibilityModel))
|
||||
require.False(t, hasMessageText(t, promptMessages, "preserve deployment constraints", database.ChatMessageVisibilityModel))
|
||||
}
|
||||
|
||||
func TestPreCompactHookFailureAbortsCompaction(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fixture := startCompactionHookChat(t,
|
||||
func(_ *testing.T, _ database.Store, request agenthooks.Request) (int, string) {
|
||||
if request.Type == agenthooks.EventPreCompact {
|
||||
return http.StatusInternalServerError, ""
|
||||
}
|
||||
return http.StatusOK, `{}`
|
||||
},
|
||||
func(t *testing.T, _ string) {
|
||||
require.FailNow(t, "compaction model called after pre_compact failure")
|
||||
},
|
||||
)
|
||||
waitCtx := testutil.Context(t, testutil.WaitLong)
|
||||
failed := waitForChatStatus(waitCtx, t, fixture.db, fixture.chat.ID, database.ChatStatusError)
|
||||
require.Equal(t, int32(0), fixture.compactionCalls.Load())
|
||||
require.False(t, hasCompactionRows(t, fixture.db, fixture.chat.ID))
|
||||
require.Equal(t, int32(1), fixture.preCompactCalls.Load())
|
||||
require.Zero(t, fixture.postCompactCalls.Load())
|
||||
lastError := chatLastErrorMessage(failed.LastError)
|
||||
require.Contains(t, lastError, "hook dispatch failed: pre_compact: http_error")
|
||||
}
|
||||
|
||||
func TestPostCompactHookFailureKeepsCompaction(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var postSawCommitted atomic.Bool
|
||||
fixture := startCompactionHookChat(t,
|
||||
func(t *testing.T, db database.Store, request agenthooks.Request) (int, string) {
|
||||
if request.Type == agenthooks.EventPostCompact {
|
||||
postSawCommitted.Store(hasCompactionRows(t, db, request.Meta.ChatID))
|
||||
return http.StatusInternalServerError, ""
|
||||
}
|
||||
return http.StatusOK, `{}`
|
||||
},
|
||||
func(*testing.T, string) {},
|
||||
)
|
||||
waitCtx := testutil.Context(t, testutil.WaitLong)
|
||||
failed := waitForChatStatus(waitCtx, t, fixture.db, fixture.chat.ID, database.ChatStatusError)
|
||||
require.False(t, postSawCommitted.Load())
|
||||
require.Equal(t, int32(1), fixture.compactionCalls.Load())
|
||||
require.True(t, hasCompactionRows(t, fixture.db, fixture.chat.ID))
|
||||
require.Equal(t, int32(1), fixture.preCompactCalls.Load())
|
||||
require.Equal(t, int32(1), fixture.postCompactCalls.Load())
|
||||
lastError := chatLastErrorMessage(failed.LastError)
|
||||
require.Contains(t, lastError, "hook dispatch failed: post_compact: http_error")
|
||||
}
|
||||
|
||||
func TestManualCompactionPostCompactEffects(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
postCompact string
|
||||
wantFollowUp bool
|
||||
wantVisibleMsg string
|
||||
}{
|
||||
{
|
||||
name: "user message resumes generation",
|
||||
postCompact: `{"user_message":"compaction complete"}`,
|
||||
wantFollowUp: true,
|
||||
wantVisibleMsg: "compaction complete",
|
||||
},
|
||||
{
|
||||
name: "model context alone finishes the turn",
|
||||
postCompact: `{"model_context":"post compact context"}`,
|
||||
wantFollowUp: false,
|
||||
},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
var streamCalls atomic.Int32
|
||||
var compactionCalls atomic.Int32
|
||||
anthropicURL := chattest.NewAnthropic(t, func(req *chattest.AnthropicRequest) chattest.AnthropicResponse {
|
||||
body := anthropicRequestBody(t, *req)
|
||||
if !req.Stream {
|
||||
if strings.Contains(body, "You are performing a context compaction") {
|
||||
compactionCalls.Add(1)
|
||||
return anthropicCompactionResponse("manual hook compaction summary")
|
||||
}
|
||||
return chattest.AnthropicNonStreamingResponse("title")
|
||||
}
|
||||
// Low usage keeps automatic compaction out of the way, so
|
||||
// only the manual request can trigger one.
|
||||
streamCalls.Add(1)
|
||||
return chattest.AnthropicStreamingResponse(chattest.AnthropicTextChunksWithCacheUsage(chattest.AnthropicUsage{
|
||||
InputTokens: 10,
|
||||
OutputTokens: 5,
|
||||
}, "assistant answer")...)
|
||||
})
|
||||
user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL)
|
||||
model = updateChatModelCompressionThreshold(t, db, model, 100, 70)
|
||||
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
body := `{}`
|
||||
if request.Type == agenthooks.EventPostCompact {
|
||||
body = test.postCompact
|
||||
}
|
||||
_, err := w.Write([]byte(body))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
|
||||
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
})
|
||||
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "hello from the user")
|
||||
chat = waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
|
||||
require.Equal(t, int32(1), streamCalls.Load())
|
||||
|
||||
_, err := server.CompactChat(ctx, chat)
|
||||
require.NoError(t, err)
|
||||
chat = waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
|
||||
require.False(t, chat.LastError.Valid)
|
||||
require.Equal(t, int32(1), compactionCalls.Load())
|
||||
|
||||
wantStreams := int32(1)
|
||||
if test.wantFollowUp {
|
||||
wantStreams = 2
|
||||
}
|
||||
require.Equal(t, wantStreams, streamCalls.Load())
|
||||
if test.wantVisibleMsg != "" {
|
||||
messages := chatMessages(ctx, t, db, chat.ID)
|
||||
require.True(t, hasMessageText(t, messages, test.wantVisibleMsg, database.ChatMessageVisibilityUser))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
type compactionHookFixture struct {
|
||||
ctx context.Context
|
||||
db database.Store
|
||||
chat database.Chat
|
||||
compactionCalls *atomic.Int32
|
||||
streamCalls *atomic.Int32
|
||||
preCompactCalls *atomic.Int32
|
||||
postCompactCalls *atomic.Int32
|
||||
}
|
||||
|
||||
func startCompactionHookChat(
|
||||
t *testing.T,
|
||||
hookResponse func(*testing.T, database.Store, agenthooks.Request) (int, string),
|
||||
inspectCompaction func(*testing.T, string),
|
||||
) compactionHookFixture {
|
||||
t.Helper()
|
||||
|
||||
const (
|
||||
contextLimit = int64(100)
|
||||
thresholdPercent = int32(70)
|
||||
)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
var compactionCalls atomic.Int32
|
||||
var preCompactCalls atomic.Int32
|
||||
var postCompactCalls atomic.Int32
|
||||
var streamCalls atomic.Int32
|
||||
anthropicURL := chattest.NewAnthropic(t, func(req *chattest.AnthropicRequest) chattest.AnthropicResponse {
|
||||
body := anthropicRequestBody(t, *req)
|
||||
if !req.Stream {
|
||||
if strings.Contains(body, "You are performing a context compaction") {
|
||||
compactionCalls.Add(1)
|
||||
inspectCompaction(t, body)
|
||||
return anthropicCompactionResponse("hook compaction summary")
|
||||
}
|
||||
return chattest.AnthropicNonStreamingResponse("title")
|
||||
}
|
||||
if streamCalls.Add(1) == 1 {
|
||||
return highUsageReadFileResponse("/tmp/hook.txt")
|
||||
}
|
||||
return chattest.AnthropicStreamingResponse(chattest.AnthropicTextChunksWithCacheUsage(chattest.AnthropicUsage{
|
||||
InputTokens: 20,
|
||||
OutputTokens: 5,
|
||||
}, "continued after compaction")...)
|
||||
})
|
||||
user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL)
|
||||
model = updateChatModelCompressionThreshold(t, db, model, contextLimit, thresholdPercent)
|
||||
ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID)
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
switch request.Type {
|
||||
case agenthooks.EventPreCompact:
|
||||
preCompactCalls.Add(1)
|
||||
case agenthooks.EventPostCompact:
|
||||
postCompactCalls.Add(1)
|
||||
}
|
||||
status, body := hookResponse(t, db, request)
|
||||
w.WriteHeader(status)
|
||||
if body != "" {
|
||||
_, err := w.Write([]byte(body))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
mockConn := agentconnmock.NewMockAgentConn(ctrl)
|
||||
setupToolExecutionAgentConn(t, mockConn)
|
||||
mockConn.EXPECT().ReadFileLines(gomock.Any(), "/tmp/hook.txt", int64(1), int64(0), gomock.Any()).
|
||||
Return(workspacesdk.ReadFileLinesResponse{
|
||||
Success: true, FileSize: 12, TotalLines: 1, LinesRead: 1, Content: "1\tpackage main",
|
||||
}, nil).
|
||||
Times(1)
|
||||
|
||||
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
|
||||
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
|
||||
require.Equal(t, dbAgent.ID, agentID)
|
||||
return mockConn, func() {}, nil
|
||||
}
|
||||
})
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true},
|
||||
AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true},
|
||||
Title: "compaction-hooks",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("trigger compaction hooks"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
return compactionHookFixture{
|
||||
ctx: ctx,
|
||||
db: db,
|
||||
chat: chat,
|
||||
compactionCalls: &compactionCalls,
|
||||
streamCalls: &streamCalls,
|
||||
preCompactCalls: &preCompactCalls,
|
||||
postCompactCalls: &postCompactCalls,
|
||||
}
|
||||
}
|
||||
|
||||
func hasCompactionRows(t *testing.T, db database.Store, chatID uuid.UUID) bool {
|
||||
t.Helper()
|
||||
userMessages := chatMessages(t.Context(), t, db, chatID)
|
||||
promptMessages, err := db.GetChatMessagesForPromptByChatID(t.Context(), chatID)
|
||||
require.NoError(t, err)
|
||||
compressed := compressedChatSummarizedMessages(t, append(promptMessages, userMessages...))
|
||||
return len(compressed.summaries) > 0 && len(compressed.calls) > 0 && len(compressed.results) > 0
|
||||
}
|
||||
|
||||
func hasMessageText(t *testing.T, messages []database.ChatMessage, text string, visibility database.ChatMessageVisibility) bool {
|
||||
t.Helper()
|
||||
for _, message := range messages {
|
||||
parts, err := chatprompt.ParseContent(message)
|
||||
require.NoError(t, err)
|
||||
if message.Visibility == visibility && len(parts) == 1 && parts[0].Text == text {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,236 @@
|
||||
package chatd_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/x/agenthooks/dispatch"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chathooks"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
func TestCreateChatUserPromptSubmitHook(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("passthrough", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
server, requests := newCreateHookTestServer(t, db, ps, http.StatusOK, `{}`)
|
||||
|
||||
chat, err := server.CreateChat(ctx, createHookOptions(t, db, user.ID, org.ID, model.ID, "passthrough"))
|
||||
require.NoError(t, err)
|
||||
request := testutil.RequireReceive(ctx, t, requests)
|
||||
require.Equal(t, agenthooks.EventUserPromptSubmit, request.Type)
|
||||
require.Equal(t, chat.ID, request.Meta.ChatID)
|
||||
require.Equal(t, user.ID, request.Meta.OwnerID)
|
||||
require.NotNil(t, request.Meta.TurnID)
|
||||
data := decodeHookData[agenthooks.UserPromptSubmitData](t, request)
|
||||
require.Equal(t, "passthrough", data.Prompt)
|
||||
var hookParts []codersdk.ChatMessagePart
|
||||
require.NoError(t, json.Unmarshal(data.Parts, &hookParts))
|
||||
require.Equal(t, []codersdk.ChatMessagePart{codersdk.ChatMessageText("passthrough")}, hookParts)
|
||||
|
||||
messages := chatMessages(ctx, t, db, chat.ID)
|
||||
initialUser := messages[len(messages)-1]
|
||||
require.Equal(t, database.ChatMessageRoleUser, initialUser.Role)
|
||||
require.Equal(t, database.ChatMessageVisibilityBoth, initialUser.Visibility)
|
||||
require.Equal(t, "passthrough", hookMessageText(t, initialUser))
|
||||
})
|
||||
|
||||
t.Run("override", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
server, requests := newCreateHookTestServer(t, db, ps, http.StatusOK, `{"permission":{"decision":"allow","input_override":{"prompt":"redacted"}}}`)
|
||||
|
||||
opts := createHookOptions(t, db, user.ID, org.ID, model.ID, "secret")
|
||||
opts.Title = chatprompt.FallbackTitle(chatprompt.TitleText(opts.InitialUserContent, nil))
|
||||
opts.TitleDerivedFromContent = true
|
||||
chat, err := server.CreateChat(ctx, opts)
|
||||
require.NoError(t, err)
|
||||
request := testutil.RequireReceive(ctx, t, requests)
|
||||
require.NotNil(t, request.Meta.TurnID)
|
||||
messages := chatMessages(ctx, t, db, chat.ID)
|
||||
initialUser := messages[len(messages)-1]
|
||||
require.Equal(t, "redacted", hookMessageText(t, initialUser))
|
||||
require.Equal(t, "redacted", chat.Title, "prompt-derived title must be recomputed from the override")
|
||||
})
|
||||
|
||||
t.Run("override keeps explicit title", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
server, _ := newCreateHookTestServer(t, db, ps, http.StatusOK, `{"permission":{"decision":"allow","input_override":{"prompt":"redacted"}}}`)
|
||||
|
||||
chat, err := server.CreateChat(ctx, createHookOptions(t, db, user.ID, org.ID, model.ID, "secret"))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "create hook test", chat.Title)
|
||||
})
|
||||
|
||||
t.Run("invalid model config rejected before dispatch", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
server, requests := newCreateHookTestServer(t, db, ps, http.StatusOK, `{}`)
|
||||
|
||||
opts := createHookOptions(t, db, user.ID, org.ID, model.ID, "prompt")
|
||||
opts.ModelConfigID = uuid.New()
|
||||
_, err := server.CreateChat(ctx, opts)
|
||||
require.ErrorIs(t, err, chatd.ErrInvalidModelConfigID)
|
||||
select {
|
||||
case request := <-requests:
|
||||
t.Fatalf("unexpected hook dispatch %s for rejected create", request.Type)
|
||||
default:
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("override recomputes paste-derived title", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
server, _ := newCreateHookTestServer(t, db, ps, http.StatusOK, `{"permission":{"decision":"allow","input_override":{"prompt":"redacted"}}}`)
|
||||
|
||||
opts := createHookOptions(t, db, user.ID, org.ID, model.ID, " ")
|
||||
opts.Title = chatprompt.FallbackTitle("secret paste content")
|
||||
opts.TitleDerivedFromContent = true
|
||||
chat, err := server.CreateChat(ctx, opts)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "redacted", chat.Title,
|
||||
"paste-derived title must be recomputed from the override")
|
||||
})
|
||||
|
||||
t.Run("response messages", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
server, _ := newCreateHookTestServer(t, db, ps, http.StatusOK, `{"model_context":"model only","user_message":"user only"}`)
|
||||
|
||||
chat, err := server.CreateChat(ctx, createHookOptions(t, db, user.ID, org.ID, model.ID, "prompt"))
|
||||
require.NoError(t, err)
|
||||
promptMessages, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, promptMessages)
|
||||
initialUser := promptMessages[len(promptMessages)-1]
|
||||
require.Equal(t, database.ChatMessageRoleUser, initialUser.Role)
|
||||
require.Equal(t, database.ChatMessageVisibilityBoth, initialUser.Visibility)
|
||||
parts, err := chatprompt.ParseContent(initialUser)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("prompt"),
|
||||
{Type: codersdk.ChatMessagePartTypeHookContext, Text: "model only"},
|
||||
{Type: codersdk.ChatMessagePartTypeHookNotice, Text: "user only"},
|
||||
}, parts)
|
||||
})
|
||||
|
||||
t.Run("deny", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
server, requests := newCreateHookTestServer(t, db, ps, http.StatusOK, `{"permission":{"decision":"deny"},"user_message":"blocked"}`)
|
||||
|
||||
_, err := server.CreateChat(ctx, createHookOptions(t, db, user.ID, org.ID, model.ID, "prompt"))
|
||||
var denied *chathooks.UserPromptDeniedError
|
||||
require.ErrorAs(t, err, &denied)
|
||||
require.Equal(t, "blocked", denied.UserMessage)
|
||||
request := testutil.RequireReceive(ctx, t, requests)
|
||||
requireCreateHookChatMissing(ctx, t, db, request.Meta.ChatID)
|
||||
})
|
||||
|
||||
t.Run("dispatch failure", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
server, requests := newCreateHookTestServer(t, db, ps, http.StatusInternalServerError, "")
|
||||
|
||||
_, err := server.CreateChat(ctx, createHookOptions(t, db, user.ID, org.ID, model.ID, "prompt"))
|
||||
var dispatchErr *dispatch.Error
|
||||
require.ErrorAs(t, err, &dispatchErr)
|
||||
require.Equal(t, dispatch.ResultHTTPError, dispatchErr.Class)
|
||||
request := testutil.RequireReceive(ctx, t, requests)
|
||||
requireCreateHookChatMissing(ctx, t, db, request.Meta.ChatID)
|
||||
})
|
||||
}
|
||||
|
||||
func TestCreateChatHooksDisabledUnchanged(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
server := newTestServer(t, db, ps, uuid.New())
|
||||
|
||||
chat, err := server.CreateChat(ctx, createHookOptions(t, db, user.ID, org.ID, model.ID, "unchanged"))
|
||||
require.NoError(t, err)
|
||||
messages := chatMessages(ctx, t, db, chat.ID)
|
||||
initialUser := messages[len(messages)-1]
|
||||
require.Equal(t, "unchanged", hookMessageText(t, initialUser))
|
||||
}
|
||||
|
||||
func newCreateHookTestServer(
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
ps dbpubsub.Pubsub,
|
||||
statusCode int,
|
||||
response string,
|
||||
) (*chatd.Server, <-chan agenthooks.Request) {
|
||||
t.Helper()
|
||||
requests := make(chan agenthooks.Request, 2)
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
requests <- request
|
||||
w.WriteHeader(statusCode)
|
||||
if response != "" {
|
||||
_, err := w.Write([]byte(response))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
return newHookTestServer(t, db, ps, consumer), requests
|
||||
}
|
||||
|
||||
func createHookOptions(
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
userID uuid.UUID,
|
||||
organizationID uuid.UUID,
|
||||
modelConfigID uuid.UUID,
|
||||
prompt string,
|
||||
) chatd.CreateOptions {
|
||||
t.Helper()
|
||||
return chatd.CreateOptions{
|
||||
OrganizationID: organizationID,
|
||||
OwnerID: userID,
|
||||
Title: "create hook test",
|
||||
ModelConfigID: modelConfigID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText(prompt)},
|
||||
}
|
||||
}
|
||||
|
||||
func requireCreateHookChatMissing(ctx context.Context, t *testing.T, db database.Store, chatID uuid.UUID) {
|
||||
t.Helper()
|
||||
_, err := db.GetChatByID(ctx, chatID)
|
||||
require.ErrorIs(t, err, sql.ErrNoRows)
|
||||
}
|
||||
+426
-61
@@ -16,12 +16,14 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatdebug"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chaterror"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chathooks"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatloop"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatretry"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/messagepartbuffer"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
)
|
||||
|
||||
// generationPrepareInput contains the committed state used to prepare one
|
||||
@@ -129,11 +131,10 @@ var errCompactionStillOverLimit = chaterror.WithClassification(
|
||||
)
|
||||
|
||||
type generationDecision struct {
|
||||
kind generationActionKind
|
||||
localToolCalls []fantasy.ToolCallContent
|
||||
pendingDynamicToolCalls []pendingDynamicToolCall
|
||||
finishReason generationFinishReason
|
||||
promotedMessageID int64
|
||||
kind generationActionKind
|
||||
localToolCalls []fantasy.ToolCallContent
|
||||
finishReason generationFinishReason
|
||||
promotedMessageID int64
|
||||
// forced marks a compact action triggered by a manual
|
||||
// compaction request rather than the usage threshold.
|
||||
forced bool
|
||||
@@ -199,12 +200,11 @@ func decideGenerationAction(input generationDecisionInput) (generationDecision,
|
||||
Input: dynamicCall.Args,
|
||||
})
|
||||
}
|
||||
dynamicCalls = nil
|
||||
}
|
||||
return generationDecision{kind: generationActionExecuteLocalTools, localToolCalls: localCalls, pendingDynamicToolCalls: dynamicCalls}, nil
|
||||
return generationDecision{kind: generationActionExecuteLocalTools, localToolCalls: localCalls}, nil
|
||||
}
|
||||
if len(dynamicCalls) > 0 {
|
||||
return generationDecision{kind: generationActionEnterRequiresAction, pendingDynamicToolCalls: dynamicCalls}, nil
|
||||
return generationDecision{kind: generationActionEnterRequiresAction}, nil
|
||||
}
|
||||
|
||||
// A manual compaction request wins over every non-tool decision:
|
||||
@@ -317,6 +317,15 @@ func unresolvedToolCallsFromHistory(
|
||||
return localCalls, dynamicCalls, nil
|
||||
}
|
||||
|
||||
// exclusiveBatchRejected reports whether the exclusive-tool policy will
|
||||
// reject the whole batch, which mirrors the condition chatloop applies
|
||||
// when it decides that nothing in the batch may execute. Callers must ask
|
||||
// before filtering the batch, because dropping calls from it can leave the
|
||||
// exclusive call alone and admissible.
|
||||
func exclusiveBatchRejected(toolCalls []fantasy.ToolCallContent, exclusiveToolNames map[string]bool) bool {
|
||||
return len(toolCalls) > 1 && hasExclusiveToolCall(toolCalls, exclusiveToolNames)
|
||||
}
|
||||
|
||||
func hasExclusiveToolCall(toolCalls []fantasy.ToolCallContent, exclusiveToolNames map[string]bool) bool {
|
||||
if len(exclusiveToolNames) == 0 {
|
||||
return false
|
||||
@@ -329,13 +338,105 @@ func hasExclusiveToolCall(toolCalls []fantasy.ToolCallContent, exclusiveToolName
|
||||
return false
|
||||
}
|
||||
|
||||
type sessionStartResult struct {
|
||||
Chat database.Chat
|
||||
}
|
||||
|
||||
func applySessionStartResponse(
|
||||
ctx context.Context,
|
||||
machine *chatstate.ChatMachine,
|
||||
input chatWorkerTaskStartInput,
|
||||
chat database.Chat,
|
||||
result *chathooks.Result,
|
||||
) (sessionStartResult, error) {
|
||||
if result.GetModelContext() == "" && result.GetUserMessage() == "" {
|
||||
return sessionStartResult{Chat: chat}, nil
|
||||
}
|
||||
|
||||
eventMessages, err := chathooks.EventMessages(result, chat.LastModelConfigID)
|
||||
if err != nil {
|
||||
return sessionStartResult{}, err
|
||||
}
|
||||
|
||||
var applied sessionStartResult
|
||||
err = machine.Update(ctx, func(tx *chatstate.Tx, store database.Store) error {
|
||||
if _, err := loadChatForGeneration(ctx, store, input, generationAttemptNotRequired); err != nil {
|
||||
return xerrors.Errorf("load chat for session_start response: %w", err)
|
||||
}
|
||||
if len(eventMessages) > 0 {
|
||||
if _, err := tx.CommitStep(chatstate.CommitStepInput{Messages: eventMessages}); err != nil {
|
||||
return xerrors.Errorf("insert session_start response messages: %w", err)
|
||||
}
|
||||
}
|
||||
applied.Chat, err = store.GetChatByID(ctx, input.ChatID)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("reload chat after session_start response: %w", err)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return sessionStartResult{}, normalizeTaskTransitionError(err, "apply session_start response")
|
||||
}
|
||||
return applied, nil
|
||||
}
|
||||
|
||||
func (s *taskStarter) startGenerationSession(
|
||||
ctx context.Context,
|
||||
machine *chatstate.ChatMachine,
|
||||
input chatWorkerTaskStartInput,
|
||||
chat database.Chat,
|
||||
messages []database.ChatMessage,
|
||||
) (result sessionStartResult, dispatched bool, err error) {
|
||||
dispatched, complete, err := input.SessionStart.claim(ctx)
|
||||
if err != nil {
|
||||
return sessionStartResult{}, false, errors.Join(errTaskExpectedExit, xerrors.Errorf("claim session_start: %w", err))
|
||||
}
|
||||
if !dispatched {
|
||||
return sessionStartResult{Chat: chat}, false, nil
|
||||
}
|
||||
|
||||
completed := false
|
||||
// Re-arm the claim until its response is applied so a replacement task
|
||||
// can replay session_start effects.
|
||||
defer func() { complete(completed) }()
|
||||
response, err := s.server.hooks.Trigger(ctx, chathooks.ChatFor(chat, input.hookTurnID()), chathooks.Message{Source: chathooks.SessionStartSource(messages)}, agenthooks.EventSessionStart)
|
||||
if err != nil {
|
||||
return sessionStartResult{}, true, chathooks.GenerationDispatchError(agenthooks.EventSessionStart, err)
|
||||
}
|
||||
result, err = applySessionStartResponse(ctx, machine, input, chat, response)
|
||||
if err != nil {
|
||||
return sessionStartResult{}, true, err
|
||||
}
|
||||
completed = true
|
||||
return result, true, nil
|
||||
}
|
||||
|
||||
func (s *taskStarter) StartGeneration(ctx context.Context, input chatWorkerTaskStartInput) error {
|
||||
if input.StopNudges == nil {
|
||||
input.StopNudges = &stopNudgeTracker{}
|
||||
}
|
||||
if input.TurnID == uuid.Nil {
|
||||
input.TurnID = uuid.New()
|
||||
}
|
||||
machine := chatstate.NewChatMachine(s.opts.Store, s.opts.Pubsub, input.ChatID)
|
||||
for {
|
||||
chat, messages, err := loadGenerationState(ctx, machine, input)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("load generation state: %w", err)
|
||||
}
|
||||
if s.server.hooks.Enabled() {
|
||||
result, dispatched, err := s.startGenerationSession(ctx, machine, input, chat, messages)
|
||||
if err != nil {
|
||||
if errors.Is(err, errTaskExpectedExit) {
|
||||
return err
|
||||
}
|
||||
return s.finishGenerationError(ctx, machine, input, err, generationAttemptNotRequired)
|
||||
}
|
||||
if dispatched {
|
||||
input.HistoryVersion = result.Chat.HistoryVersion
|
||||
continue
|
||||
}
|
||||
}
|
||||
prepareInput := generationPrepareInput{
|
||||
Chat: chat,
|
||||
Messages: messages,
|
||||
@@ -350,20 +451,25 @@ func (s *taskStarter) StartGeneration(ctx context.Context, input chatWorkerTaskS
|
||||
return s.finishGenerationError(ctx, machine, input, err, generationAttemptNotRequired)
|
||||
}
|
||||
cleanup := prepared.Cleanup
|
||||
decision, err := retryGenerationPhase(ctx, s, "decide", func() (generationDecision, error) {
|
||||
return decideGenerationAction(generationDecisionInput{
|
||||
chat: prepared.Chat,
|
||||
messages: prepared.Messages,
|
||||
dynamicToolNames: prepared.DynamicToolNames,
|
||||
exclusiveToolNames: prepared.ExclusiveToolNames,
|
||||
stopAfterTools: prepared.StopAfterTools,
|
||||
maxSteps: prepared.MaxSteps,
|
||||
compactionEnabled: prepared.Compaction != nil,
|
||||
compactionNeeded: prepared.Compaction != nil && prepared.Compaction.Required,
|
||||
compactionThresholdPercent: generationCompactionThreshold(prepared.Compaction),
|
||||
compactionContextLimit: generationCompactionContextLimit(prepared.Compaction),
|
||||
var decision generationDecision
|
||||
if input.StopNudges.consume(stopNudgeKey(prepared.Messages)) {
|
||||
decision = generationDecision{kind: generationActionGenerateAssistant}
|
||||
} else {
|
||||
decision, err = retryGenerationPhase(ctx, s, "decide", func() (generationDecision, error) {
|
||||
return decideGenerationAction(generationDecisionInput{
|
||||
chat: prepared.Chat,
|
||||
messages: prepared.Messages,
|
||||
dynamicToolNames: prepared.DynamicToolNames,
|
||||
exclusiveToolNames: prepared.ExclusiveToolNames,
|
||||
stopAfterTools: prepared.StopAfterTools,
|
||||
maxSteps: prepared.MaxSteps,
|
||||
compactionEnabled: prepared.Compaction != nil,
|
||||
compactionNeeded: prepared.Compaction != nil && prepared.Compaction.Required,
|
||||
compactionThresholdPercent: generationCompactionThreshold(prepared.Compaction),
|
||||
compactionContextLimit: generationCompactionContextLimit(prepared.Compaction),
|
||||
})
|
||||
})
|
||||
})
|
||||
}
|
||||
if err != nil {
|
||||
cleanup()
|
||||
if errors.Is(err, errTaskExpectedExit) || errors.Is(err, errTaskRetryable) {
|
||||
@@ -636,6 +742,11 @@ func (s *taskStarter) generateAssistant(
|
||||
if len(outcome.Step.Content) == 0 {
|
||||
return s.finishGenerationTurn(ctx, machine, input, generationDecision{kind: generationActionFinishTurn, finishReason: generationFinishReasonComplete}, requireGenerationAttempt(attempt.number))
|
||||
}
|
||||
preflight, err := s.admitStepToolCalls(ctx, input, prepared, outcome.Step.Content)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
outcome.Step.Content = chathooks.ApplyAdmittedToolCalls(outcome.Step.Content, preflight)
|
||||
messages, err := buildCommitStepMessages(buildCommitStepMessagesInput{
|
||||
modelConfigID: prepared.ModelConfigID,
|
||||
modelCallConfig: prepared.ModelConfig,
|
||||
@@ -647,7 +758,41 @@ func (s *taskStarter) generateAssistant(
|
||||
if err != nil {
|
||||
return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number))
|
||||
}
|
||||
return s.commitGenerationStep(ctx, machine, input, attempt.number, generationActionGenerateAssistant, messages)
|
||||
messages, err = appendHookResultMessages(messages, preflight.Results, prepared.ModelConfigID)
|
||||
if err != nil {
|
||||
return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number))
|
||||
}
|
||||
return s.commitGenerationStep(ctx, machine, input, attempt.number, generationActionGenerateAssistant, messages, generationCommitHooks{})
|
||||
}
|
||||
|
||||
func (s *taskStarter) admitStepToolCalls(
|
||||
ctx context.Context,
|
||||
input chatWorkerTaskStartInput,
|
||||
prepared generationPrepared,
|
||||
content []fantasy.Content,
|
||||
) (chathooks.PreToolUseExecutionResult, error) {
|
||||
if !s.server.hooks.Enabled() {
|
||||
return chathooks.PreToolUseExecutionResult{}, nil
|
||||
}
|
||||
toolCalls := chathooks.PendingToolCalls(content)
|
||||
if len(toolCalls) == 0 || exclusiveBatchRejected(toolCalls, prepared.ExclusiveToolNames) {
|
||||
return chathooks.PreToolUseExecutionResult{}, nil
|
||||
}
|
||||
// Check the full batch first: a call removed below still occupies its ID
|
||||
// in the step, so filtering before this would hide the collision.
|
||||
if err := chathooks.RejectDuplicateToolUseIDs(toolCalls); err != nil {
|
||||
return chathooks.PreToolUseExecutionResult{}, chathooks.GenerationDispatchError(agenthooks.EventPreToolUse, err)
|
||||
}
|
||||
unambiguous, ambiguous := partitionAmbiguousToolCalls(prepared, toolCalls)
|
||||
preflight, err := s.server.hooks.PreflightPendingToolCalls(ctx, chathooks.ChatFor(prepared.Chat, input.hookTurnID()), unambiguous)
|
||||
if err != nil {
|
||||
return chathooks.PreToolUseExecutionResult{}, chathooks.GenerationDispatchError(agenthooks.EventPreToolUse, err)
|
||||
}
|
||||
if err := validateOverriddenToolInputs(prepared, preflight); err != nil {
|
||||
return chathooks.PreToolUseExecutionResult{}, chathooks.GenerationDispatchError(agenthooks.EventPreToolUse, err)
|
||||
}
|
||||
preflight.Denied = append(preflight.Denied, ambiguous...)
|
||||
return preflight, nil
|
||||
}
|
||||
|
||||
func (s *taskStarter) executeLocalTools(
|
||||
@@ -657,6 +802,11 @@ func (s *taskStarter) executeLocalTools(
|
||||
prepared generationPrepared,
|
||||
decision generationDecision,
|
||||
) error {
|
||||
allowed := decision.localToolCalls
|
||||
var denied []fantasy.ToolResultContent
|
||||
if !exclusiveBatchRejected(decision.localToolCalls, prepared.ExclusiveToolNames) {
|
||||
allowed, denied = partitionAmbiguousToolCalls(prepared, decision.localToolCalls)
|
||||
}
|
||||
attempt, err := s.beginGenerationAttempt(ctx, machine, input)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("beginGenerationAttempt: %w", err)
|
||||
@@ -668,25 +818,41 @@ func (s *taskStarter) executeLocalTools(
|
||||
provider = prepared.Model.Provider()
|
||||
modelName = prepared.Model.Model()
|
||||
}
|
||||
outcome, err := chatloop.ExecuteLocalTools(ctx, chatloop.ExecuteLocalToolsOptions{
|
||||
Tools: prepared.Tools,
|
||||
ActiveTools: prepared.ActiveTools,
|
||||
ProviderTools: prepared.ProviderTools,
|
||||
ToolCalls: decision.localToolCalls,
|
||||
ExclusiveToolNames: prepared.ExclusiveToolNames,
|
||||
BuiltinToolNames: prepared.BuiltinToolNames,
|
||||
ModelProvider: provider,
|
||||
ModelName: modelName,
|
||||
ContextLimit: prepared.ContextLimitFallback,
|
||||
ToolNameAliases: subagentToolNameAliases,
|
||||
PublishMessagePart: attempt.publish,
|
||||
Logger: s.opts.Logger,
|
||||
Metrics: s.server.metrics,
|
||||
Clock: s.opts.Clock,
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("execute local tools: %w", err)
|
||||
var outcome chatloop.ToolExecutionOutcome
|
||||
var spawnDispatchErr error
|
||||
if len(allowed) > 0 {
|
||||
outcome, err = chatloop.ExecuteLocalTools(ctx, chatloop.ExecuteLocalToolsOptions{
|
||||
Tools: prepared.Tools,
|
||||
ActiveTools: prepared.ActiveTools,
|
||||
ProviderTools: prepared.ProviderTools,
|
||||
ToolCalls: allowed,
|
||||
ExclusiveToolNames: prepared.ExclusiveToolNames,
|
||||
BuiltinToolNames: prepared.BuiltinToolNames,
|
||||
ModelProvider: provider,
|
||||
ModelName: modelName,
|
||||
ContextLimit: prepared.ContextLimitFallback,
|
||||
ToolNameAliases: subagentToolNameAliases,
|
||||
PublishMessagePart: attempt.publish,
|
||||
Logger: s.opts.Logger,
|
||||
Metrics: s.server.metrics,
|
||||
Clock: s.opts.Clock,
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("execute local tools: %w", err)
|
||||
}
|
||||
// Subagent spawn admission dispatches user_prompt_submit inside
|
||||
// the tool run; its failure surfaces as a tool result error. The
|
||||
// step still commits so a sibling tool that already ran keeps its
|
||||
// result and is not re-executed, and the turn fails afterwards.
|
||||
if hookErr := chathooks.DispatchFailureFromResults(outcome.Step.Content); hookErr != nil {
|
||||
spawnDispatchErr = chathooks.GenerationDispatchError(agenthooks.EventUserPromptSubmit, hookErr)
|
||||
}
|
||||
}
|
||||
postResults, postDispatchErr := s.server.hooks.PostToolUseResults(ctx, chathooks.ChatFor(prepared.Chat, input.hookTurnID()), outcome.Step.Content)
|
||||
for _, result := range denied {
|
||||
outcome.Step.Content = append(outcome.Step.Content, result)
|
||||
}
|
||||
chathooks.RestoreToolCallOrder(outcome.Step.Content, decision.localToolCalls)
|
||||
messages, err := buildCommitStepMessages(buildCommitStepMessagesInput{
|
||||
modelConfigID: prepared.ModelConfigID,
|
||||
modelCallConfig: prepared.ModelConfig,
|
||||
@@ -698,7 +864,29 @@ func (s *taskStarter) executeLocalTools(
|
||||
if err != nil {
|
||||
return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number))
|
||||
}
|
||||
return s.commitGenerationStep(ctx, machine, input, attempt.number, generationActionExecuteLocalTools, messages)
|
||||
messages, err = appendHookResultMessages(messages, postResults, prepared.ModelConfigID)
|
||||
if err != nil {
|
||||
return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number))
|
||||
}
|
||||
var postCommitErr error
|
||||
switch {
|
||||
case spawnDispatchErr != nil:
|
||||
// Spawn admission is the causal root: post_tool_use only ran
|
||||
// because the batch reached execution at all.
|
||||
postCommitErr = spawnDispatchErr
|
||||
if postDispatchErr != nil {
|
||||
s.opts.Logger.Warn(ctx, "post_tool_use hook dispatch failed alongside spawn admission",
|
||||
slog.F("chat_id", input.ChatID),
|
||||
slog.F("worker_id", input.WorkerID),
|
||||
slog.Error(postDispatchErr),
|
||||
)
|
||||
}
|
||||
case postDispatchErr != nil:
|
||||
postCommitErr = chathooks.GenerationDispatchError(agenthooks.EventPostToolUse, postDispatchErr)
|
||||
}
|
||||
return s.commitGenerationStep(ctx, machine, input, attempt.number, generationActionExecuteLocalTools, messages, generationCommitHooks{
|
||||
PostCommitError: postCommitErr,
|
||||
})
|
||||
}
|
||||
|
||||
// compactionSourceForDecision maps a compact decision to the
|
||||
@@ -751,6 +939,11 @@ func (s *taskStarter) generateCompaction(
|
||||
overrideModel.modelConfig,
|
||||
)
|
||||
}
|
||||
preResult, err := s.server.hooks.Trigger(ctx, chathooks.ChatFor(prepared.Chat, input.hookTurnID()), chathooks.Message{}, agenthooks.EventPreCompact)
|
||||
if err != nil {
|
||||
return chathooks.GenerationDispatchError(agenthooks.EventPreCompact, err)
|
||||
}
|
||||
compactionOpts.SummaryHint = preResult.GetModelContext()
|
||||
compactionOpts.PublishMessagePart = attempt.publish
|
||||
compactionOpts.Source = source
|
||||
compactionOpts.Force = source == chatloop.CompactionSourceManual
|
||||
@@ -779,14 +972,37 @@ func (s *taskStarter) generateCompaction(
|
||||
s.server.metrics.RecordCompaction(metricProvider, metricModel, false, err)
|
||||
return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number))
|
||||
}
|
||||
err = s.commitGenerationStep(ctx, machine, input, attempt.number, generationActionCompact, stepMessagesForCommit{
|
||||
// The summary hint already consumed the pre_compact model context.
|
||||
persistedPreResult := &chathooks.Result{UserMessage: preResult.GetUserMessage()}
|
||||
commitMessages, err := applyHookResultMessages(stepMessagesForCommit{
|
||||
Messages: messages.Messages,
|
||||
VisibleIndexes: visibleMessageIndexes(messages.Messages),
|
||||
ConsumeCompactionRequest: true,
|
||||
}, []*chathooks.Result{persistedPreResult}, prepared.ModelConfigID)
|
||||
if err != nil {
|
||||
s.server.metrics.RecordCompaction(metricProvider, metricModel, false, err)
|
||||
return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number))
|
||||
}
|
||||
// Hook effects and fail-closed errors must commit atomically with
|
||||
// compaction; a separate commit races the runner and can be dropped
|
||||
// on crash.
|
||||
postResult, postDispatchErr := s.server.hooks.Trigger(ctx, chathooks.ChatFor(prepared.Chat, input.hookTurnID()), chathooks.Message{}, agenthooks.EventPostCompact)
|
||||
var postCommitErr error
|
||||
if postDispatchErr != nil {
|
||||
postCommitErr = chathooks.GenerationDispatchError(agenthooks.EventPostCompact, postDispatchErr)
|
||||
} else {
|
||||
commitMessages, err = appendHookResultMessages(commitMessages, []*chathooks.Result{postResult}, prepared.ModelConfigID)
|
||||
if err != nil {
|
||||
s.server.metrics.RecordCompaction(metricProvider, metricModel, false, err)
|
||||
return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number))
|
||||
}
|
||||
}
|
||||
err = s.commitGenerationStep(ctx, machine, input, attempt.number, generationActionCompact, commitMessages, generationCommitHooks{
|
||||
PostCommitError: postCommitErr,
|
||||
})
|
||||
s.server.metrics.RecordCompaction(metricProvider, metricModel, err == nil, err)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("commit generation step: %w", err)
|
||||
return xerrors.Errorf("commit compaction step: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -871,6 +1087,10 @@ func (s *taskStarter) beginGenerationAttempt(
|
||||
}, nil
|
||||
}
|
||||
|
||||
type generationCommitHooks struct {
|
||||
PostCommitError error
|
||||
}
|
||||
|
||||
func (s *taskStarter) commitGenerationStep(
|
||||
ctx context.Context,
|
||||
machine *chatstate.ChatMachine,
|
||||
@@ -878,10 +1098,31 @@ func (s *taskStarter) commitGenerationStep(
|
||||
attempt int64,
|
||||
kind generationActionKind,
|
||||
messages stepMessagesForCommit,
|
||||
commitHooks generationCommitHooks,
|
||||
) error {
|
||||
if len(messages.Messages) == 0 {
|
||||
if commitHooks.PostCommitError != nil {
|
||||
return s.finishGenerationError(ctx, machine, input, commitHooks.PostCommitError, requireGenerationAttempt(attempt))
|
||||
}
|
||||
return s.finishGenerationTurn(ctx, machine, input, generationDecision{kind: generationActionFinishTurn, finishReason: generationFinishReasonComplete}, requireGenerationAttempt(attempt))
|
||||
}
|
||||
failClosed := commitHooks.PostCommitError != nil
|
||||
var postCommitLastError pqtype.NullRawMessage
|
||||
var postCommitMessage string
|
||||
if commitHooks.PostCommitError != nil {
|
||||
classified := chaterror.Classify(commitHooks.PostCommitError)
|
||||
s.opts.Logger.Warn(ctx, "chat generation failed",
|
||||
slog.F("chat_id", input.ChatID),
|
||||
slog.F("worker_id", input.WorkerID),
|
||||
slog.F("generation_attempt", input.GenerationAttempt),
|
||||
slog.F("error_kind", classified.Kind),
|
||||
slog.F("provider", classified.Provider),
|
||||
slog.F("status_code", classified.StatusCode),
|
||||
slog.F("retryable", classified.Retryable),
|
||||
slog.Error(commitHooks.PostCommitError),
|
||||
)
|
||||
postCommitLastError, postCommitMessage = generationLastError(commitHooks.PostCommitError)
|
||||
}
|
||||
var committed database.Chat
|
||||
insertedMessages := []runnerActionMessage{}
|
||||
err := machine.Update(ctx, func(tx *chatstate.Tx, store database.Store) error {
|
||||
@@ -895,19 +1136,43 @@ func (s *taskStarter) commitGenerationStep(
|
||||
if err != nil {
|
||||
return xerrors.Errorf("tx.CommitStep: %w", err)
|
||||
}
|
||||
insertedMessages = make([]runnerActionMessage, 0, len(commitResult.InsertedMessages))
|
||||
for _, msg := range commitResult.InsertedMessages {
|
||||
inserted := commitResult.InsertedMessages
|
||||
// The fail-closed hook error must commit atomically with the
|
||||
// step; a separate commit races the runner and can be dropped
|
||||
// on crash.
|
||||
if failClosed {
|
||||
if _, err := tx.FinishError(chatstate.FinishErrorInput{LastError: postCommitLastError}); err != nil {
|
||||
return xerrors.Errorf("tx.FinishError: %w", err)
|
||||
}
|
||||
}
|
||||
insertedMessages = make([]runnerActionMessage, 0, len(inserted))
|
||||
for _, msg := range inserted {
|
||||
insertedMessages = append(insertedMessages, runnerActionMessage{ID: msg.ID, Role: codersdk.ChatMessageRole(msg.Role)})
|
||||
}
|
||||
committed, err = store.GetChatByID(ctx, input.ChatID)
|
||||
loadedChat, err := store.GetChatByID(ctx, input.ChatID)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("load committed chat: %w", err)
|
||||
}
|
||||
committed = loadedChat
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return normalizeTaskTransitionError(err, "commit generation step")
|
||||
}
|
||||
if failClosed {
|
||||
input.DebugTurn.RecordOutcome(chatdebug.StatusError)
|
||||
postCommitCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), postCommitWatchPublishTimeout)
|
||||
defer cancel()
|
||||
if err := s.publishWatchAndRoute(postCommitCtx, committed, codersdk.ChatWatchEventKindStatusChange); err != nil {
|
||||
return xerrors.Errorf("publish watch and route: %w", err)
|
||||
}
|
||||
return s.afterGenerationOutcome(postCommitCtx, generationOutcome{
|
||||
Chat: committed,
|
||||
Kind: runnerActionKindFinishError,
|
||||
WatchEventKind: codersdk.ChatWatchEventKindStatusChange,
|
||||
LastError: postCommitMessage,
|
||||
})
|
||||
}
|
||||
s.routeStateHint(ctx, stateUpdateFromChat(committed))
|
||||
return s.afterGenerationOutcome(ctx, generationOutcome{
|
||||
Chat: committed,
|
||||
@@ -997,7 +1262,32 @@ func recordGenerationFinishFailure(turn *runnerDebugTurn, err error) {
|
||||
turn.RecordOutcome(chatdebug.StatusError)
|
||||
}
|
||||
|
||||
func (s *taskStarter) finishGenerationTurn(
|
||||
func (s *taskStarter) completeGenerationTurn(
|
||||
ctx context.Context,
|
||||
input chatWorkerTaskStartInput,
|
||||
committed database.Chat,
|
||||
promotedMessageID int64,
|
||||
) error {
|
||||
input.StopNudges.reset()
|
||||
input.DebugTurn.RecordOutcome(chatdebug.StatusCompleted)
|
||||
watchCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), postCommitWatchPublishTimeout)
|
||||
defer cancel()
|
||||
if err := s.publishWatchWithRetry(watchCtx, committed, codersdk.ChatWatchEventKindStatusChange); err != nil {
|
||||
return xerrors.Errorf("publish watch and route: %w", err)
|
||||
}
|
||||
if err := s.afterGenerationOutcome(ctx, generationOutcome{
|
||||
Chat: committed,
|
||||
Kind: runnerActionKindFinishTurn,
|
||||
WatchEventKind: codersdk.ChatWatchEventKindStatusChange,
|
||||
PromotedMessageID: promotedMessageID,
|
||||
}); err != nil {
|
||||
return xerrors.Errorf("after generation outcome: %w", err)
|
||||
}
|
||||
s.routeStateHint(ctx, stateUpdateFromChat(committed))
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *taskStarter) finishGenerationTurnWithoutHook(
|
||||
ctx context.Context,
|
||||
machine *chatstate.ChatMachine,
|
||||
input chatWorkerTaskStartInput,
|
||||
@@ -1024,22 +1314,97 @@ func (s *taskStarter) finishGenerationTurn(
|
||||
recordGenerationFinishFailure(input.DebugTurn, err)
|
||||
return err
|
||||
}
|
||||
input.DebugTurn.RecordOutcome(chatdebug.StatusCompleted)
|
||||
watchCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), postCommitWatchPublishTimeout)
|
||||
defer cancel()
|
||||
if err := s.publishWatchWithRetry(watchCtx, committed, codersdk.ChatWatchEventKindStatusChange); err != nil {
|
||||
return xerrors.Errorf("publish watch and route: %w", err)
|
||||
return s.completeGenerationTurn(ctx, input, committed, decision.promotedMessageID)
|
||||
}
|
||||
|
||||
func (s *taskStarter) finishGenerationTurn(
|
||||
ctx context.Context,
|
||||
machine *chatstate.ChatMachine,
|
||||
input chatWorkerTaskStartInput,
|
||||
decision generationDecision,
|
||||
fence generationAttemptFence,
|
||||
) error {
|
||||
if !s.server.hooks.Enabled() {
|
||||
return s.finishGenerationTurnWithoutHook(ctx, machine, input, decision, fence)
|
||||
}
|
||||
if err := s.afterGenerationOutcome(ctx, generationOutcome{
|
||||
Chat: committed,
|
||||
Kind: runnerActionKindFinishTurn,
|
||||
WatchEventKind: codersdk.ChatWatchEventKindStatusChange,
|
||||
PromotedMessageID: decision.promotedMessageID,
|
||||
}); err != nil {
|
||||
return xerrors.Errorf("after generation outcome: %w", err)
|
||||
var chat database.Chat
|
||||
var messages []database.ChatMessage
|
||||
err := machine.ReadLock(ctx, func(store database.Store) error {
|
||||
loadedChat, err := loadChatForGeneration(ctx, store, input, fence)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("load chat for stop hook: %w", err)
|
||||
}
|
||||
loadedMessages, err := store.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: input.ChatID,
|
||||
AfterID: 0,
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("load messages for stop hook: %w", err)
|
||||
}
|
||||
chat = loadedChat
|
||||
messages = loadedMessages
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return normalizeTaskTransitionError(err, "load stop hook state")
|
||||
}
|
||||
s.routeStateHint(ctx, stateUpdateFromChat(committed))
|
||||
return nil
|
||||
response, err := s.server.hooks.Trigger(ctx, chathooks.ChatFor(chat, input.hookTurnID()), chathooks.Message{}, agenthooks.EventStop)
|
||||
if err != nil {
|
||||
return s.finishGenerationError(ctx, machine, input, chathooks.GenerationDispatchError(agenthooks.EventStop, err), fence)
|
||||
}
|
||||
stopMessages, err := chathooks.EventMessages(response, chat.LastModelConfigID)
|
||||
if err != nil {
|
||||
return s.finishGenerationError(ctx, machine, input, err, fence)
|
||||
}
|
||||
nudgeKey := stopNudgeKey(messages)
|
||||
// Prompt conversion drops whitespace-only text parts, so a blank
|
||||
// model context would buy a continuation that nudges nothing.
|
||||
continueTurn := strings.TrimSpace(response.GetModelContext()) != "" && input.StopNudges.claim(nudgeKey)
|
||||
|
||||
var committed database.Chat
|
||||
err = machine.Update(ctx, func(tx *chatstate.Tx, store database.Store) error {
|
||||
if _, err := loadChatForGeneration(ctx, store, input, fence); err != nil {
|
||||
return xerrors.Errorf("load chat for generation: %w", err)
|
||||
}
|
||||
if len(stopMessages) > 0 {
|
||||
if _, err := tx.CommitStep(chatstate.CommitStepInput{Messages: stopMessages}); err != nil {
|
||||
return xerrors.Errorf("commit stop hook messages: %w", err)
|
||||
}
|
||||
}
|
||||
if !continueTurn {
|
||||
finishResult, err := tx.FinishTurn(chatstate.FinishTurnInput{})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("tx.FinishTurn: %w", err)
|
||||
}
|
||||
if finishResult.PromotedMessage != nil {
|
||||
decision.promotedMessageID = finishResult.PromotedMessage.ID
|
||||
}
|
||||
committed = finishResult.Chat
|
||||
return nil
|
||||
}
|
||||
loadedChat, err := store.GetChatByID(ctx, input.ChatID)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("load committed chat: %w", err)
|
||||
}
|
||||
committed = loadedChat
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
if continueTurn {
|
||||
input.StopNudges.cancel(nudgeKey)
|
||||
}
|
||||
err := normalizeTaskTransitionError(err, "finish generation turn")
|
||||
recordGenerationFinishFailure(input.DebugTurn, err)
|
||||
return err
|
||||
}
|
||||
if continueTurn {
|
||||
s.routeStateHint(ctx, stateUpdateFromChat(committed))
|
||||
return s.afterGenerationOutcome(ctx, generationOutcome{
|
||||
Chat: committed,
|
||||
Kind: runnerActionKind(generationActionGenerateAssistant),
|
||||
})
|
||||
}
|
||||
return s.completeGenerationTurn(ctx, input, committed, decision.promotedMessageID)
|
||||
}
|
||||
|
||||
func (s *taskStarter) finishGenerationError(
|
||||
|
||||
@@ -3,6 +3,7 @@ package chatd //nolint:testpackage // Exercises unexported generation helpers.
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
@@ -13,6 +14,39 @@ import (
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
func TestExclusiveBatchRejected(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
call := func(name string) fantasy.ToolCallContent {
|
||||
return fantasy.ToolCallContent{ToolCallID: "call_" + name, ToolName: name}
|
||||
}
|
||||
exclusive := map[string]bool{"advisor": true}
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
toolCalls []fantasy.ToolCallContent
|
||||
exclusives map[string]bool
|
||||
want bool
|
||||
}{
|
||||
{name: "ExclusiveAlone", toolCalls: []fantasy.ToolCallContent{call("advisor")}, exclusives: exclusive},
|
||||
{name: "NoExclusive", toolCalls: []fantasy.ToolCallContent{call("execute"), call("read_file")}, exclusives: exclusive},
|
||||
{name: "NoExclusiveNames", toolCalls: []fantasy.ToolCallContent{call("advisor"), call("execute")}},
|
||||
{
|
||||
name: "ExclusiveMixed",
|
||||
toolCalls: []fantasy.ToolCallContent{call("advisor"), call("execute")},
|
||||
exclusives: exclusive,
|
||||
want: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range cases {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
require.Equal(t, test.want, exclusiveBatchRejected(test.toolCalls, test.exclusives))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompactionMetricIdentity(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -0,0 +1,187 @@
|
||||
package chatd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chathooks"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
)
|
||||
|
||||
// applyHookResultMessages inserts hook event rows before the step's
|
||||
// own rows so injected model context precedes the assistant content it
|
||||
// steers; providers require tool results to directly follow the
|
||||
// assistant tool calls.
|
||||
func applyHookResultMessages(
|
||||
messages stepMessagesForCommit,
|
||||
results []*chathooks.Result,
|
||||
modelConfigID uuid.UUID,
|
||||
) (stepMessagesForCommit, error) {
|
||||
return insertHookResultMessages(messages, results, modelConfigID, hookRowsBeforeStep)
|
||||
}
|
||||
|
||||
func appendHookResultMessages(
|
||||
messages stepMessagesForCommit,
|
||||
results []*chathooks.Result,
|
||||
modelConfigID uuid.UUID,
|
||||
) (stepMessagesForCommit, error) {
|
||||
return insertHookResultMessages(messages, results, modelConfigID, hookRowsAfterStep)
|
||||
}
|
||||
|
||||
type hookRowPlacement int
|
||||
|
||||
const (
|
||||
hookRowsBeforeStep hookRowPlacement = iota
|
||||
hookRowsAfterStep
|
||||
)
|
||||
|
||||
func insertHookResultMessages(
|
||||
messages stepMessagesForCommit,
|
||||
results []*chathooks.Result,
|
||||
modelConfigID uuid.UUID,
|
||||
placement hookRowPlacement,
|
||||
) (stepMessagesForCommit, error) {
|
||||
rows, err := chathooks.EventMessagesForResults(results, modelConfigID)
|
||||
if err != nil {
|
||||
return stepMessagesForCommit{}, err
|
||||
}
|
||||
if len(rows) > 0 {
|
||||
if placement == hookRowsBeforeStep {
|
||||
messages.Messages = append(rows, messages.Messages...)
|
||||
} else {
|
||||
messages.Messages = append(messages.Messages, rows...)
|
||||
}
|
||||
messages.VisibleIndexes = visibleMessageIndexes(messages.Messages)
|
||||
}
|
||||
return messages, nil
|
||||
}
|
||||
|
||||
func (p *Server) handleUserPromptDispatchError(ctx context.Context, chatID uuid.UUID, dispatchErr error) error {
|
||||
return p.handleAPIDispatchError(ctx, chatID, agenthooks.EventUserPromptSubmit, dispatchErr)
|
||||
}
|
||||
|
||||
func (p *Server) handleAPIDispatchError(ctx context.Context, chatID uuid.UUID, eventType agenthooks.EventType, dispatchErr error) error {
|
||||
lastError, ok := chathooks.DispatchErrorMessage(eventType, dispatchErr)
|
||||
if !ok {
|
||||
return dispatchErr
|
||||
}
|
||||
encoded, marshalErr := json.Marshal(codersdk.ChatError{
|
||||
Message: lastError,
|
||||
Kind: codersdk.ChatErrorKindHookDispatchFailed,
|
||||
})
|
||||
if marshalErr != nil {
|
||||
return errors.Join(dispatchErr, xerrors.Errorf("encode hook dispatch error: %w", marshalErr))
|
||||
}
|
||||
var failedChat database.Chat
|
||||
machine := p.newChatMachine(chatID)
|
||||
err := machine.Update(ctx, func(tx *chatstate.Tx, store database.Store) error {
|
||||
current, err := store.GetChatByID(ctx, chatID)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("load chat for hook failure: %w", err)
|
||||
}
|
||||
// Park only idle chats. FinishError is also allowed from running
|
||||
// states, but a running chat keeps its active turn and the
|
||||
// request error alone surfaces to the caller.
|
||||
if current.Status != database.ChatStatusWaiting {
|
||||
return chatstate.ErrTransitionNotAllowed
|
||||
}
|
||||
if _, err := tx.FinishError(chatstate.FinishErrorInput{
|
||||
LastError: pqtype.NullRawMessage{RawMessage: encoded, Valid: true},
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
chat, err := store.GetChatByID(ctx, chatID)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("reload chat after hook failure: %w", err)
|
||||
}
|
||||
failedChat = chat
|
||||
return nil
|
||||
})
|
||||
if errors.Is(err, chatstate.ErrTransitionNotAllowed) {
|
||||
return dispatchErr
|
||||
}
|
||||
if err != nil {
|
||||
return errors.Join(dispatchErr, xerrors.Errorf("fail idle chat after hook dispatch: %w", err))
|
||||
}
|
||||
p.publishChatPubsubEvent(failedChat, codersdk.ChatWatchEventKindStatusChange, nil)
|
||||
return dispatchErr
|
||||
}
|
||||
|
||||
type dynamicPostToolUseState struct {
|
||||
chat database.Chat
|
||||
modelConfigID uuid.UUID
|
||||
toolNames map[string]string
|
||||
}
|
||||
|
||||
func loadDynamicPostToolUseState(
|
||||
ctx context.Context,
|
||||
machine *chatstate.ChatMachine,
|
||||
opts SubmitToolResultsOptions,
|
||||
) (dynamicPostToolUseState, error) {
|
||||
var state dynamicPostToolUseState
|
||||
err := machine.ReadLock(ctx, func(store database.Store) error {
|
||||
chat, err := store.GetChatByID(ctx, opts.ChatID)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("load chat: %w", err)
|
||||
}
|
||||
if chat.Archived {
|
||||
return ErrChatArchived
|
||||
}
|
||||
if chat.Status != database.ChatStatusRequiresAction {
|
||||
return &ToolResultStatusConflictError{ActualStatus: chat.Status}
|
||||
}
|
||||
messages, err := store.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: opts.ChatID,
|
||||
AfterID: 0,
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("load chat messages: %w", err)
|
||||
}
|
||||
_, pending, err := unresolvedToolCallsFromHistory(messages, dynamicToolNamesFromChat(chat))
|
||||
if err != nil {
|
||||
return xerrors.Errorf("load pending dynamic tool calls: %w", err)
|
||||
}
|
||||
toolNames := make(map[string]string, len(pending))
|
||||
for _, call := range pending {
|
||||
toolNames[call.ToolCallID] = call.ToolName
|
||||
}
|
||||
if err := validateSubmittedToolResults(opts.Results, toolNames); err != nil {
|
||||
return err
|
||||
}
|
||||
modelConfigID := opts.ModelConfigID
|
||||
if modelConfigID == uuid.Nil {
|
||||
modelConfigID = chat.LastModelConfigID
|
||||
}
|
||||
state = dynamicPostToolUseState{
|
||||
chat: chat,
|
||||
modelConfigID: modelConfigID,
|
||||
toolNames: toolNames,
|
||||
}
|
||||
return nil
|
||||
})
|
||||
return state, err
|
||||
}
|
||||
|
||||
// validateSubmittedToolResults rejects invalid results before hook dispatch,
|
||||
// using the same rules as CompleteRequiresAction.
|
||||
func validateSubmittedToolResults(results []codersdk.ToolResult, toolNames map[string]string) error {
|
||||
inputs := make([]chatstate.ToolResultInput, 0, len(results))
|
||||
for _, result := range results {
|
||||
inputs = append(inputs, chatstate.ToolResultInput{
|
||||
ToolCallID: result.ToolCallID,
|
||||
Output: result.Output,
|
||||
})
|
||||
}
|
||||
if invalid := chatstate.ValidateToolResults(inputs, toolNames); invalid != nil {
|
||||
return translateToolResultValidationError(invalid)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,156 @@
|
||||
package chatd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog/v3/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/x/agenthooks/dispatch"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chathooks"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
func TestSessionStartTrackerRetriesIncompleteDispatch(t *testing.T) {
|
||||
t.Parallel()
|
||||
tracker := &sessionStartTracker{}
|
||||
claimed, complete, err := tracker.claim(t.Context())
|
||||
require.NoError(t, err)
|
||||
require.True(t, claimed)
|
||||
|
||||
canceled, cancel := context.WithCancel(t.Context())
|
||||
cancel()
|
||||
_, _, err = tracker.claim(canceled)
|
||||
require.ErrorIs(t, err, context.Canceled)
|
||||
complete(false)
|
||||
|
||||
claimed, complete, err = tracker.claim(t.Context())
|
||||
require.NoError(t, err)
|
||||
require.True(t, claimed)
|
||||
complete(true)
|
||||
claimed, _, err = tracker.claim(t.Context())
|
||||
require.NoError(t, err)
|
||||
require.False(t, claimed)
|
||||
}
|
||||
|
||||
func TestSessionStartDispatchFailureFinishesGeneration(t *testing.T) {
|
||||
t.Parallel()
|
||||
f := newTaskTestFixture(t)
|
||||
chat := f.createRunningChat(t)
|
||||
workerID := uuid.New()
|
||||
runnerID := uuid.New()
|
||||
chat = f.acquireChat(t, chat.ID, workerID, runnerID)
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
dispatcher := dispatch.New(
|
||||
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
|
||||
consumer.Client(),
|
||||
consumer.URL,
|
||||
"test-hook-secret-32-bytes-minimum!!",
|
||||
time.Second,
|
||||
"test-deployment",
|
||||
"test-version",
|
||||
prometheus.NewRegistry(),
|
||||
)
|
||||
starter := newTestTaskStarter(t, f, newTaskSideEffectRecorder())
|
||||
starter.server.hooks = chathooks.NewTrigger(dispatcher)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
debugTurn := newRunnerDebugTurn(ctx, starter.opts.Logger)
|
||||
defer debugTurn.Finalize(ctx)
|
||||
err := starter.StartGeneration(ctx, chatWorkerTaskStartInput{
|
||||
ChatID: chat.ID,
|
||||
WorkerID: workerID,
|
||||
RunnerID: runnerID,
|
||||
HistoryVersion: chat.HistoryVersion,
|
||||
GenerationAttempt: chat.GenerationAttempt,
|
||||
Status: database.ChatStatusRunning,
|
||||
DebugTurn: debugTurn,
|
||||
SessionStart: &sessionStartTracker{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
updated, err := f.db.GetChatByID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, database.ChatStatusError, updated.Status)
|
||||
var chatErr codersdk.ChatError
|
||||
require.NoError(t, json.Unmarshal(updated.LastError.RawMessage, &chatErr))
|
||||
require.Equal(t, codersdk.ChatErrorKindHookDispatchFailed, chatErr.Kind)
|
||||
require.Contains(t, chatErr.Message, "hook dispatch failed: session_start: http_error (dispatch ")
|
||||
require.False(t, chatErr.Retryable)
|
||||
}
|
||||
|
||||
func TestApplySessionStartResponse(t *testing.T) {
|
||||
t.Parallel()
|
||||
f := newTaskTestFixture(t)
|
||||
chat := f.createRunningChat(t)
|
||||
workerID := uuid.New()
|
||||
runnerID := uuid.New()
|
||||
chat = f.acquireChat(t, chat.ID, workerID, runnerID)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
input := chatWorkerTaskStartInput{
|
||||
ChatID: chat.ID,
|
||||
WorkerID: workerID,
|
||||
RunnerID: runnerID,
|
||||
HistoryVersion: chat.HistoryVersion,
|
||||
Status: database.ChatStatusRunning,
|
||||
}
|
||||
_, err := applySessionStartResponse(
|
||||
ctx,
|
||||
chatstate.NewChatMachine(f.db, f.pubsub, chat.ID),
|
||||
input,
|
||||
chat,
|
||||
&chathooks.Result{
|
||||
ModelContext: "model context",
|
||||
UserMessage: "user notice",
|
||||
},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
promptRows, err := f.db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "model context", hookMessageTextInternal(t, promptRows[len(promptRows)-1]))
|
||||
require.Equal(t, database.ChatMessageVisibilityModel, promptRows[len(promptRows)-1].Visibility)
|
||||
allRows, err := f.db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ChatID: chat.ID})
|
||||
require.NoError(t, err)
|
||||
userNotice := allRows[len(allRows)-1]
|
||||
require.Equal(t, database.ChatMessageRoleSystem, userNotice.Role)
|
||||
require.Equal(t, database.ChatMessageVisibilityUser, userNotice.Visibility)
|
||||
require.Equal(t, "user notice", hookMessageTextInternal(t, userNotice))
|
||||
}
|
||||
|
||||
func TestApplySessionStartResponseNoOp(t *testing.T) {
|
||||
t.Parallel()
|
||||
f := newTaskTestFixture(t)
|
||||
chat := f.createRunningChat(t)
|
||||
f.pubsub.clear()
|
||||
result, err := applySessionStartResponse(
|
||||
testutil.Context(t, testutil.WaitLong),
|
||||
chatstate.NewChatMachine(f.db, f.pubsub, chat.ID),
|
||||
chatWorkerTaskStartInput{},
|
||||
chat,
|
||||
nil,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, chat.SnapshotVersion, result.Chat.SnapshotVersion)
|
||||
require.Empty(t, f.pubsub.events())
|
||||
}
|
||||
|
||||
func hookMessageTextInternal(t *testing.T, message database.ChatMessage) string {
|
||||
t.Helper()
|
||||
parts, err := chatprompt.ParseContent(message)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, parts, 1)
|
||||
return parts[0].Text
|
||||
}
|
||||
@@ -0,0 +1,683 @@
|
||||
package chatd_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"slices"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog/v3/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/x/agenthooks/dispatch"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chathooks"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
func TestSendMessageUserPromptSubmitHook(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("override", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
})
|
||||
|
||||
submitted := []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("before"),
|
||||
codersdk.ChatMessageFileReference("main.go", 1, 3, "package main"),
|
||||
}
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
data := decodeHookData[agenthooks.UserPromptSubmitData](t, request)
|
||||
require.Equal(t, "before", data.Prompt)
|
||||
var hookParts []codersdk.ChatMessagePart
|
||||
require.NoError(t, json.Unmarshal(data.Parts, &hookParts))
|
||||
require.Equal(t, submitted, hookParts, "hook payload must carry non-text parts")
|
||||
require.NotNil(t, request.Meta.TurnID)
|
||||
_, err := w.Write([]byte(`{"permission":{"decision":"allow","input_override":{"prompt":"after"}},"model_context":"model only","user_message":"user only"}`))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
server := newHookTestServer(t, db, ps, consumer)
|
||||
result, err := server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: user.ID,
|
||||
Content: submitted,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
parts, err := chatprompt.ParseContent(result.Message)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("after"),
|
||||
{Type: codersdk.ChatMessagePartTypeHookContext, Text: "model only"},
|
||||
{Type: codersdk.ChatMessagePartTypeHookNotice, Text: "user only"},
|
||||
}, parts)
|
||||
require.Len(t, result.InsertedMessages, 1)
|
||||
require.Equal(t, result.Message.ID, result.InsertedMessages[0].ID)
|
||||
})
|
||||
|
||||
t.Run("deny", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
})
|
||||
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, err := w.Write([]byte(`{"permission":{"decision":"deny"},"user_message":"blocked"}`))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
server := newHookTestServer(t, db, ps, consumer)
|
||||
_, err := server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: user.ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("blocked prompt")},
|
||||
})
|
||||
var denied *chathooks.UserPromptDeniedError
|
||||
require.ErrorAs(t, err, &denied)
|
||||
require.Equal(t, "blocked", denied.UserMessage)
|
||||
|
||||
messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ChatID: chat.ID})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, messages)
|
||||
})
|
||||
}
|
||||
|
||||
func newHookDispatcher(t *testing.T, _ database.Store, consumer *httptest.Server) *dispatch.Dispatcher {
|
||||
t.Helper()
|
||||
return dispatch.New(
|
||||
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
|
||||
consumer.Client(),
|
||||
consumer.URL,
|
||||
"test-hook-secret-32-bytes-minimum!!",
|
||||
time.Second,
|
||||
"test-deployment",
|
||||
"test-version",
|
||||
prometheus.NewRegistry(),
|
||||
)
|
||||
}
|
||||
|
||||
func newHookTestServer(t *testing.T, db database.Store, ps dbpubsub.Pubsub, consumer *httptest.Server) *chatd.Server {
|
||||
t.Helper()
|
||||
return newTestServer(t, db, ps, uuid.New(), func(cfg *chatd.Config) {
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
})
|
||||
}
|
||||
|
||||
func TestHookDispatcherRequiresExperiment(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
})
|
||||
|
||||
var hookRequests atomic.Int32
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
hookRequests.Add(1)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
server := newTestServer(t, db, ps, uuid.New(), func(cfg *chatd.Config) {
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
cfg.Experiments = slices.DeleteFunc(
|
||||
slices.Clone(codersdk.ExperimentsKnown),
|
||||
func(e codersdk.Experiment) bool { return e == codersdk.ExperimentAgentLifecycleHooks },
|
||||
)
|
||||
})
|
||||
result, err := server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: user.ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
parts, err := chatprompt.ParseContent(result.Message)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, parts)
|
||||
|
||||
require.Zero(t, hookRequests.Load())
|
||||
}
|
||||
|
||||
func TestSendMessageUserPromptSubmitPassthrough(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
})
|
||||
var received agenthooks.Request
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&received))
|
||||
_, err := w.Write([]byte(`{}`))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
server := newHookTestServer(t, db, ps, consumer)
|
||||
result, err := server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("passthrough")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "passthrough", hookMessageText(t, result.Message))
|
||||
require.Equal(t, agenthooks.EventUserPromptSubmit, received.Type)
|
||||
promptData := decodeHookData[agenthooks.UserPromptSubmitData](t, received)
|
||||
require.Equal(t, "passthrough", promptData.Prompt)
|
||||
// The persisted content is jsonb-normalized, so compare JSON
|
||||
// semantics rather than raw bytes.
|
||||
require.JSONEq(t, string(result.Message.Content.RawMessage), string(promptData.Parts))
|
||||
}
|
||||
|
||||
func TestSendMessageUserPromptSubmitQueue(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
chat, err := newTestServer(t, db, ps, uuid.New()).CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "queued hook",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("running")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
var received agenthooks.Request
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&received))
|
||||
_, err := w.Write([]byte(`{"permission":{"decision":"allow","input_override":{"prompt":"queued override"}},"model_context":"queued context","user_message":"queued notice"}`))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
server := newHookTestServer(t, db, ps, consumer)
|
||||
|
||||
result, err := server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued original")},
|
||||
BusyBehavior: chatd.SendMessageBusyBehaviorQueue,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.True(t, result.Queued)
|
||||
require.NotNil(t, result.QueuedMessage)
|
||||
queuedParts, err := chatprompt.ParseContent(database.ChatMessage{
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: pqtype.NullRawMessage{RawMessage: result.QueuedMessage.Content, Valid: true},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
wantQueuedParts := []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("queued override"),
|
||||
{Type: codersdk.ChatMessagePartTypeHookContext, Text: "queued context"},
|
||||
{Type: codersdk.ChatMessagePartTypeHookNotice, Text: "queued notice"},
|
||||
}
|
||||
require.Equal(t, wantQueuedParts, queuedParts)
|
||||
queued, err := db.GetChatQueuedMessages(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, queued, 1)
|
||||
persistedParts, err := chatprompt.ParseContent(database.ChatMessage{
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: pqtype.NullRawMessage{RawMessage: queued[0].Content, Valid: true},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, wantQueuedParts, persistedParts)
|
||||
require.Equal(t, agenthooks.EventUserPromptSubmit, received.Type)
|
||||
require.Equal(t, "queued original", decodeHookData[agenthooks.UserPromptSubmitData](t, received).Prompt)
|
||||
}
|
||||
|
||||
func TestSendMessageUserPromptSubmitQueuedRejections(t *testing.T) {
|
||||
t.Parallel()
|
||||
tests := []struct {
|
||||
name string
|
||||
statusCode int
|
||||
response string
|
||||
assertErr func(*testing.T, error)
|
||||
}{
|
||||
{
|
||||
name: "deny",
|
||||
statusCode: http.StatusOK,
|
||||
response: `{"permission":{"decision":"deny"},"user_message":"blocked"}`,
|
||||
assertErr: func(t *testing.T, err error) {
|
||||
var denied *chathooks.UserPromptDeniedError
|
||||
require.ErrorAs(t, err, &denied)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "dispatch failure",
|
||||
statusCode: http.StatusInternalServerError,
|
||||
assertErr: func(t *testing.T, err error) {
|
||||
var dispatchErr *dispatch.Error
|
||||
require.ErrorAs(t, err, &dispatchErr)
|
||||
},
|
||||
},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
chat, err := newTestServer(t, db, ps, uuid.New()).CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "queued rejection",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("running")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(test.statusCode)
|
||||
if test.response != "" {
|
||||
_, err := w.Write([]byte(test.response))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
server := newHookTestServer(t, db, ps, consumer)
|
||||
_, err = server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued")},
|
||||
BusyBehavior: chatd.SendMessageBusyBehaviorQueue,
|
||||
})
|
||||
test.assertErr(t, err)
|
||||
queued, err := db.GetChatQueuedMessages(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, queued)
|
||||
updated, err := db.GetChatByID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, database.ChatStatusRunning, updated.Status)
|
||||
require.False(t, updated.LastError.Valid)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSubagentSpawnHookDispatchFailureFailsTurn(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
chunk := chattest.OpenAIToolCallChunk("spawn_agent", `{"type":"general","prompt":"child admission prompt"}`)
|
||||
chunk.Choices[0].ToolCalls[0].ID = "call_spawn"
|
||||
return chattest.OpenAIStreamingResponse(chunk)
|
||||
})
|
||||
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
|
||||
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
if request.Type == agenthooks.EventUserPromptSubmit {
|
||||
data := decodeHookData[agenthooks.UserPromptSubmitData](t, request)
|
||||
if data.Prompt == "child admission prompt" {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
}
|
||||
_, err := w.Write([]byte(`{}`))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
|
||||
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
})
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "spawn-hook-failure",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("spawn a child"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
failed := waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusError)
|
||||
require.Contains(t, chatLastErrorMessage(failed.LastError), "hook dispatch failed: user_prompt_submit: http_error")
|
||||
|
||||
messages := chatMessages(ctx, t, db, chat.ID)
|
||||
require.Len(t, messages, 3)
|
||||
require.Equal(t, database.ChatMessageRoleUser, messages[0].Role)
|
||||
require.Equal(t, database.ChatMessageRoleAssistant, messages[1].Role)
|
||||
require.Equal(t, database.ChatMessageRoleTool, messages[2].Role)
|
||||
require.Contains(t, string(messages[2].Content.RawMessage), "lifecycle hook returned HTTP status 500")
|
||||
|
||||
chats, err := db.GetChats(ctx, database.GetChatsParams{
|
||||
OwnedOnly: true,
|
||||
ViewerID: user.ID,
|
||||
AfterID: uuid.Nil,
|
||||
OffsetOpt: 0,
|
||||
LimitOpt: 100,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, chats, 1)
|
||||
}
|
||||
|
||||
func TestSendMessageUserPromptSubmitDispatchFailure(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
})
|
||||
var received agenthooks.Request
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&received))
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
server := newHookTestServer(t, db, ps, consumer)
|
||||
|
||||
_, err := server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("fails")},
|
||||
})
|
||||
var dispatchErr *dispatch.Error
|
||||
require.ErrorAs(t, err, &dispatchErr)
|
||||
require.Equal(t, dispatch.ResultHTTPError, dispatchErr.Class)
|
||||
updated, err := db.GetChatByID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, database.ChatStatusError, updated.Status)
|
||||
var chatErr codersdk.ChatError
|
||||
require.NoError(t, json.Unmarshal(updated.LastError.RawMessage, &chatErr))
|
||||
require.Equal(t, "hook dispatch failed: user_prompt_submit: http_error (dispatch "+dispatchErr.DispatchID.String()+")", chatErr.Message)
|
||||
require.Equal(t, agenthooks.EventUserPromptSubmit, received.Type)
|
||||
prompt := decodeHookData[agenthooks.UserPromptSubmitData](t, received)
|
||||
require.Equal(t, "fails", prompt.Prompt)
|
||||
messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ChatID: chat.ID})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, messages)
|
||||
queued, err := db.GetChatQueuedMessages(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, queued)
|
||||
}
|
||||
|
||||
func TestEditMessageUserPromptSubmitHook(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
})
|
||||
content, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{codersdk.ChatMessageText("original")})
|
||||
require.NoError(t, err)
|
||||
inserted, err := db.InsertChatMessages(ctx, chatd.BuildSingleChatMessageInsertParams(
|
||||
chat.ID, database.ChatMessageRoleUser, content, database.ChatMessageVisibilityBoth, model.ID, chatprompt.CurrentContentVersion, user.ID,
|
||||
))
|
||||
require.NoError(t, err)
|
||||
require.Len(t, inserted, 1)
|
||||
type receivedHook struct {
|
||||
request agenthooks.Request
|
||||
claims agenthooks.Claims
|
||||
}
|
||||
var receivedMu sync.Mutex
|
||||
received := make([]receivedHook, 0, 2)
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
claims, err := agenthooks.Verify(r.Header.Get("Authorization"), []byte("test-hook-secret-32-bytes-minimum!!"))
|
||||
require.NoError(t, err)
|
||||
receivedMu.Lock()
|
||||
received = append(received, receivedHook{request: request, claims: claims})
|
||||
receivedMu.Unlock()
|
||||
response := `{"model_context":"clear context","user_message":"clear notice"}`
|
||||
if request.Type == agenthooks.EventUserPromptSubmit {
|
||||
response = `{"permission":{"decision":"allow","input_override":{"prompt":"edited override"}},"model_context":"edit context","user_message":"edit notice"}`
|
||||
}
|
||||
_, err = w.Write([]byte(response))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
server := newHookTestServer(t, db, ps, consumer)
|
||||
|
||||
result, err := server.EditMessage(ctx, chatd.EditMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: user.ID,
|
||||
EditedMessageID: inserted[0].ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited original")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
parts, err := chatprompt.ParseContent(result.Message)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("edited override"),
|
||||
{Type: codersdk.ChatMessagePartTypeHookContext, Text: "edit context"},
|
||||
{Type: codersdk.ChatMessagePartTypeHookNotice, Text: "edit notice"},
|
||||
}, parts)
|
||||
receivedMu.Lock()
|
||||
received = slices.Clone(received)
|
||||
receivedMu.Unlock()
|
||||
require.Len(t, received, 2)
|
||||
require.Equal(t, agenthooks.EventSessionStart, received[0].request.Type)
|
||||
require.Equal(t, agenthooks.SessionStartData{Source: "clear"}, decodeHookData[agenthooks.SessionStartData](t, received[0].request))
|
||||
require.Equal(t, received[0].request.Meta.DispatchID, received[0].claims.JTI)
|
||||
require.Equal(t, agenthooks.EventUserPromptSubmit, received[1].request.Type)
|
||||
prompt := decodeHookData[agenthooks.UserPromptSubmitData](t, received[1].request)
|
||||
require.Equal(t, "edited original", prompt.Prompt)
|
||||
require.NotNil(t, received[0].request.Meta.TurnID)
|
||||
require.Equal(t, received[0].request.Meta.TurnID, received[1].request.Meta.TurnID)
|
||||
require.Equal(t, received[1].request.Meta.DispatchID, received[1].claims.JTI)
|
||||
require.NotEqual(t, received[0].claims.JTI, received[1].claims.JTI)
|
||||
rows, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ChatID: chat.ID})
|
||||
require.NoError(t, err)
|
||||
var foundNotice bool
|
||||
for _, row := range rows {
|
||||
if row.Role == database.ChatMessageRoleSystem && row.Visibility == database.ChatMessageVisibilityUser && hookMessageText(t, row) == "clear notice" {
|
||||
foundNotice = true
|
||||
}
|
||||
}
|
||||
require.True(t, foundNotice)
|
||||
promptRows, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
var foundContext bool
|
||||
for _, row := range promptRows {
|
||||
if row.Visibility == database.ChatMessageVisibilityModel && hookMessageText(t, row) == "clear context" {
|
||||
foundContext = true
|
||||
}
|
||||
}
|
||||
require.True(t, foundContext)
|
||||
}
|
||||
|
||||
func TestEditMessageInvalidTargetSkipsHooks(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
})
|
||||
var dispatched atomic.Int32
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
dispatched.Add(1)
|
||||
_, err := w.Write([]byte(`{}`))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
server := newHookTestServer(t, db, ps, consumer)
|
||||
|
||||
_, err := server.EditMessage(ctx, chatd.EditMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: user.ID,
|
||||
EditedMessageID: 999999,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edit of nothing")},
|
||||
})
|
||||
require.ErrorIs(t, err, chatd.ErrEditedMessageNotFound)
|
||||
|
||||
// dbgen.Chat ignores seed.Archived; archive explicitly.
|
||||
archived := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
})
|
||||
_, err = db.ArchiveChatByID(ctx, archived.ID)
|
||||
require.NoError(t, err)
|
||||
_, err = server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: archived.ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("send to archived")},
|
||||
})
|
||||
require.ErrorIs(t, err, chatd.ErrChatArchived)
|
||||
_, err = server.EditMessage(ctx, chatd.EditMessageOptions{
|
||||
ChatID: archived.ID,
|
||||
CreatedBy: user.ID,
|
||||
EditedMessageID: 1,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edit archived")},
|
||||
})
|
||||
require.ErrorIs(t, err, chatd.ErrChatArchived)
|
||||
|
||||
require.Zero(t, dispatched.Load(), "invalid targets must not dispatch hooks")
|
||||
}
|
||||
|
||||
func TestPromptHooksAdmissionPreflight(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
received := make(chan agenthooks.Request, 8)
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
received <- request
|
||||
_, err := w.Write([]byte(`{}`))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
server := newHookTestServer(t, db, ps, consumer)
|
||||
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
})
|
||||
_, err := server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("bad model")},
|
||||
ModelConfigID: uuid.New(),
|
||||
})
|
||||
require.ErrorIs(t, err, chatd.ErrInvalidModelConfigID)
|
||||
|
||||
content, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{codersdk.ChatMessageText("original")})
|
||||
require.NoError(t, err)
|
||||
inserted, err := db.InsertChatMessages(ctx, chatd.BuildSingleChatMessageInsertParams(
|
||||
chat.ID, database.ChatMessageRoleUser, content, database.ChatMessageVisibilityBoth, model.ID, chatprompt.CurrentContentVersion, user.ID,
|
||||
))
|
||||
require.NoError(t, err)
|
||||
require.Len(t, inserted, 1)
|
||||
_, err = server.EditMessage(ctx, chatd.EditMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: user.ID,
|
||||
EditedMessageID: inserted[0].ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("bad model edit")},
|
||||
ModelConfigID: uuid.New(),
|
||||
})
|
||||
require.ErrorIs(t, err, chatd.ErrInvalidModelConfigID)
|
||||
|
||||
busy := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
Status: database.ChatStatusRunning,
|
||||
})
|
||||
queuedContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{codersdk.ChatMessageText("queued")})
|
||||
require.NoError(t, err)
|
||||
for range chatstate.MaxQueueSize {
|
||||
_, err = db.InsertChatQueuedMessageWithCreator(ctx, database.InsertChatQueuedMessageWithCreatorParams{
|
||||
ChatID: busy.ID,
|
||||
Content: queuedContent.RawMessage,
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
CreatedBy: user.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
_, err = server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: busy.ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queue full")},
|
||||
})
|
||||
require.ErrorIs(t, err, chatstate.ErrMessageQueueFull)
|
||||
|
||||
select {
|
||||
case request := <-received:
|
||||
t.Fatalf("admission-rejected prompt dispatched %s", request.Type)
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendMessageHooksDisabled(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
})
|
||||
server := newTestServer(t, db, ps, uuid.New())
|
||||
result, err := server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("unchanged")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "unchanged", hookMessageText(t, result.Message))
|
||||
}
|
||||
|
||||
func hookMessageText(t *testing.T, message database.ChatMessage) string {
|
||||
t.Helper()
|
||||
parts, err := chatprompt.ParseContent(message)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, parts, 1)
|
||||
return parts[0].Text
|
||||
}
|
||||
|
||||
func decodeHookData[T any](t *testing.T, request agenthooks.Request) T {
|
||||
t.Helper()
|
||||
var data T
|
||||
require.NoError(t, json.Unmarshal(request.Data, &data))
|
||||
return data
|
||||
}
|
||||
@@ -336,19 +336,28 @@ func buildCompactionMessages(input buildCompactionMessagesInput) (compactionMess
|
||||
return compactionMessagesForCommit{Messages: messages, HiddenCount: 1}, nil
|
||||
}
|
||||
|
||||
func currentTurnStepCount(messages []database.ChatMessage) int {
|
||||
latestUser := -1
|
||||
// Hook model-context messages use the user role but must not reset
|
||||
// per-turn guards.
|
||||
func lastUserPromptIndex(messages []database.ChatMessage) int {
|
||||
index := -1
|
||||
for i, msg := range messages {
|
||||
if msg.Deleted || msg.Compressed {
|
||||
continue
|
||||
}
|
||||
if msg.Role == database.ChatMessageRoleUser {
|
||||
latestUser = i
|
||||
if msg.Role == database.ChatMessageRoleUser && msg.Visibility != database.ChatMessageVisibilityModel {
|
||||
index = i
|
||||
}
|
||||
}
|
||||
return index
|
||||
}
|
||||
|
||||
func currentTurnStartIndex(messages []database.ChatMessage) int {
|
||||
return lastUserPromptIndex(messages) + 1
|
||||
}
|
||||
|
||||
func currentTurnStepCount(messages []database.ChatMessage) int {
|
||||
count := 0
|
||||
for i := latestUser + 1; i < len(messages); i++ {
|
||||
msg := messages[i]
|
||||
for _, msg := range messages[currentTurnStartIndex(messages):] {
|
||||
if msg.Deleted || msg.Compressed {
|
||||
continue
|
||||
}
|
||||
@@ -477,16 +486,7 @@ func historyHasStopAfterToolResult(messages []database.ChatMessage, stopAfterToo
|
||||
if len(stopAfterTools) == 0 {
|
||||
return false, nil
|
||||
}
|
||||
start := 0
|
||||
for i, msg := range messages {
|
||||
if msg.Deleted || msg.Compressed {
|
||||
continue
|
||||
}
|
||||
if msg.Role == database.ChatMessageRoleUser {
|
||||
start = i + 1
|
||||
}
|
||||
}
|
||||
for _, msg := range messages[start:] {
|
||||
for _, msg := range messages[currentTurnStartIndex(messages):] {
|
||||
if msg.Deleted || msg.Compressed || msg.Role != database.ChatMessageRoleTool {
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -271,6 +271,22 @@ func TestCurrentTurnStepCount_CountsAssistantMessagesAfterLatestUser(t *testing.
|
||||
require.Equal(t, 2, got)
|
||||
}
|
||||
|
||||
func TestCurrentTurnStepCount_IgnoresHookModelContext(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
hookContext := dbMessage(t, 4, database.ChatMessageRoleUser, false, codersdk.ChatMessageText("hook context"))
|
||||
hookContext.Visibility = database.ChatMessageVisibilityModel
|
||||
messages := []database.ChatMessage{
|
||||
dbMessage(t, 1, database.ChatMessageRoleUser, false, codersdk.ChatMessageText("prompt")),
|
||||
dbMessage(t, 2, database.ChatMessageRoleAssistant, false, codersdk.ChatMessageText("one")),
|
||||
dbMessage(t, 3, database.ChatMessageRoleTool, false, codersdk.ChatMessageToolResult("call", "tool", json.RawMessage(`{}`), false, false)),
|
||||
hookContext,
|
||||
dbMessage(t, 5, database.ChatMessageRoleAssistant, false, codersdk.ChatMessageText("two")),
|
||||
}
|
||||
got := currentTurnStepCount(messages)
|
||||
require.Equal(t, 2, got)
|
||||
}
|
||||
|
||||
func TestDecisionCompactsAgainAfterPostCompactionTurn(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -545,6 +561,22 @@ func TestDecisionDetectsStopAfterToolFromCommittedHistory(t *testing.T) {
|
||||
require.False(t, got)
|
||||
}
|
||||
|
||||
func TestDecisionDetectsStopAfterToolAcrossHookContext(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
hookContext := dbMessage(t, 4, database.ChatMessageRoleUser, false, codersdk.ChatMessageText("hook context"))
|
||||
hookContext.Visibility = database.ChatMessageVisibilityModel
|
||||
messages := []database.ChatMessage{
|
||||
dbMessage(t, 1, database.ChatMessageRoleUser, false, codersdk.ChatMessageText("plan")),
|
||||
dbMessage(t, 2, database.ChatMessageRoleAssistant, false, codersdk.ChatMessageToolCall("plan-1", "propose_plan", json.RawMessage(`{}`))),
|
||||
dbMessage(t, 3, database.ChatMessageRoleTool, false, codersdk.ChatMessageToolResult("plan-1", "propose_plan", json.RawMessage(`{"ok":true}`), false, false)),
|
||||
hookContext,
|
||||
}
|
||||
got, err := historyHasStopAfterToolResult(messages, map[string]struct{}{"propose_plan": {}})
|
||||
require.NoError(t, err)
|
||||
require.True(t, got)
|
||||
}
|
||||
|
||||
func TestDecisionDetectsCurrentHistoryCompletion(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
+115
-2
@@ -3,6 +3,7 @@ package chatd
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
@@ -53,8 +54,12 @@ type chatWorkerTaskStarter interface {
|
||||
|
||||
// chatWorkerTaskStartInput describes one runner task invocation.
|
||||
type chatWorkerTaskStartInput struct {
|
||||
TaskID uuid.UUID
|
||||
ChatID uuid.UUID
|
||||
TaskID uuid.UUID
|
||||
ChatID uuid.UUID
|
||||
// TurnID is a process-local correlation ID minted per generation
|
||||
// task run. It groups the run's hook events; it is best-effort only
|
||||
// and never persisted.
|
||||
TurnID uuid.UUID
|
||||
WorkerID uuid.UUID
|
||||
RunnerID uuid.UUID
|
||||
HistoryVersion int64
|
||||
@@ -62,6 +67,114 @@ type chatWorkerTaskStartInput struct {
|
||||
Status database.ChatStatus
|
||||
RequiresActionDeadlineAt sql.NullTime
|
||||
DebugTurn *runnerDebugTurn
|
||||
SessionStart *sessionStartTracker
|
||||
StopNudges *stopNudgeTracker
|
||||
}
|
||||
|
||||
func (i chatWorkerTaskStartInput) hookTurnID() *uuid.UUID {
|
||||
if i.TurnID == uuid.Nil {
|
||||
return nil
|
||||
}
|
||||
turnID := i.TurnID
|
||||
return &turnID
|
||||
}
|
||||
|
||||
// stopNudgeTracker allows at most one stop-hook nudge continuation per
|
||||
// turn. Turns are keyed by the last user prompt's message ID so the
|
||||
// claim survives task restarts, which mint fresh process-local turn
|
||||
// IDs.
|
||||
type stopNudgeTracker struct {
|
||||
mu sync.Mutex
|
||||
turnKey int64
|
||||
claimed bool
|
||||
pending bool
|
||||
}
|
||||
|
||||
// stopNudgeKey identifies the current turn by its prompt row. Model
|
||||
// visibility user rows are hook context, not prompts.
|
||||
func stopNudgeKey(messages []database.ChatMessage) int64 {
|
||||
index := lastUserPromptIndex(messages)
|
||||
if index == -1 {
|
||||
return 0
|
||||
}
|
||||
return messages[index].ID
|
||||
}
|
||||
|
||||
func (t *stopNudgeTracker) claim(turnKey int64) bool {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
if t.turnKey != turnKey {
|
||||
t.turnKey = turnKey
|
||||
t.claimed = false
|
||||
}
|
||||
if t.claimed {
|
||||
return false
|
||||
}
|
||||
t.claimed = true
|
||||
t.pending = true
|
||||
return true
|
||||
}
|
||||
|
||||
func (t *stopNudgeTracker) consume(turnKey int64) bool {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
if t.turnKey != turnKey || !t.pending {
|
||||
return false
|
||||
}
|
||||
t.pending = false
|
||||
return true
|
||||
}
|
||||
|
||||
func (t *stopNudgeTracker) cancel(turnKey int64) {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
if t.turnKey != turnKey || !t.pending {
|
||||
return
|
||||
}
|
||||
t.pending = false
|
||||
t.claimed = false
|
||||
}
|
||||
|
||||
func (t *stopNudgeTracker) reset() {
|
||||
t.mu.Lock()
|
||||
t.turnKey = 0
|
||||
t.claimed = false
|
||||
t.pending = false
|
||||
t.mu.Unlock()
|
||||
}
|
||||
|
||||
type sessionStartTracker struct {
|
||||
mu sync.Mutex
|
||||
completed bool
|
||||
inFlight chan struct{}
|
||||
}
|
||||
|
||||
func (t *sessionStartTracker) claim(ctx context.Context) (bool, func(bool), error) {
|
||||
for {
|
||||
t.mu.Lock()
|
||||
if t.completed {
|
||||
t.mu.Unlock()
|
||||
return false, nil, nil
|
||||
}
|
||||
if t.inFlight == nil {
|
||||
t.inFlight = make(chan struct{})
|
||||
t.mu.Unlock()
|
||||
return true, func(completed bool) {
|
||||
t.mu.Lock()
|
||||
t.completed = completed
|
||||
close(t.inFlight)
|
||||
t.inFlight = nil
|
||||
t.mu.Unlock()
|
||||
}, nil
|
||||
}
|
||||
inFlight := t.inFlight
|
||||
t.mu.Unlock()
|
||||
select {
|
||||
case <-inFlight:
|
||||
case <-ctx.Done():
|
||||
return false, nil, ctx.Err()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// chatWorkerOptions configures a chatWorker.
|
||||
|
||||
@@ -0,0 +1,606 @@
|
||||
package chatd_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/x/agenthooks/dispatch"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk/agentconnmock"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
func TestPostToolUseHookResponsesCommitWithResults(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
var modelCalls atomic.Int32
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
if modelCalls.Add(1) == 1 {
|
||||
first := chattest.OpenAIToolCallChunk("read_file", `{"path":"/tmp/first.txt"}`)
|
||||
first.Choices[0].ToolCalls[0].ID = "call_first"
|
||||
second := chattest.OpenAIToolCallChunk("read_file", `{"path":"/tmp/second.txt"}`).Choices[0].ToolCalls[0]
|
||||
second.ID = "call_second"
|
||||
second.Index = 1
|
||||
first.Choices[0].ToolCalls = append(first.Choices[0].ToolCalls, second)
|
||||
return chattest.OpenAIStreamingResponse(first)
|
||||
}
|
||||
toolResultIndex := -1
|
||||
contextIndex := -1
|
||||
for i, message := range req.Messages {
|
||||
if message.Role == "tool" && strings.Contains(message.Content, "data") {
|
||||
toolResultIndex = i
|
||||
}
|
||||
if strings.Contains(message.Content, "lint feedback") {
|
||||
contextIndex = i
|
||||
}
|
||||
}
|
||||
require.NotEqual(t, -1, toolResultIndex)
|
||||
require.NotEqual(t, -1, contextIndex)
|
||||
require.Less(t, toolResultIndex, contextIndex)
|
||||
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("done")...)
|
||||
})
|
||||
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
|
||||
ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID)
|
||||
|
||||
var mu sync.Mutex
|
||||
var received []agenthooks.PostToolUseData
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
if request.Type != agenthooks.EventPostToolUse {
|
||||
_, err := w.Write([]byte(`{}`))
|
||||
require.NoError(t, err)
|
||||
return
|
||||
}
|
||||
data := decodeHookData[agenthooks.PostToolUseData](t, request)
|
||||
mu.Lock()
|
||||
received = append(received, data)
|
||||
index := len(received)
|
||||
mu.Unlock()
|
||||
|
||||
messages := chatMessages(ctx, t, db, request.Meta.ChatID)
|
||||
for _, message := range messages {
|
||||
require.NotEqual(t, database.ChatMessageRoleTool, message.Role)
|
||||
}
|
||||
var err error
|
||||
if index == 1 {
|
||||
_, err = w.Write([]byte(`{"model_context":"lint feedback","user_message":"tool notice"}`))
|
||||
} else {
|
||||
_, err = w.Write([]byte(`{}`))
|
||||
}
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
mockConn := agentconnmock.NewMockAgentConn(ctrl)
|
||||
setupToolExecutionAgentConn(t, mockConn)
|
||||
mockConn.EXPECT().ReadFileLines(gomock.Any(), gomock.Any(), int64(1), int64(0), gomock.Any()).
|
||||
Return(workspacesdk.ReadFileLinesResponse{
|
||||
Success: true, FileSize: 4, TotalLines: 1, LinesRead: 1, Content: "data",
|
||||
}, nil).
|
||||
Times(2)
|
||||
|
||||
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
|
||||
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
|
||||
require.Equal(t, dbAgent.ID, agentID)
|
||||
return mockConn, func() {}, nil
|
||||
}
|
||||
})
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true},
|
||||
AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true},
|
||||
Title: "post-tool-use-responses",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("read both files"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
|
||||
|
||||
mu.Lock()
|
||||
receivedSnapshot := append([]agenthooks.PostToolUseData(nil), received...)
|
||||
mu.Unlock()
|
||||
require.Len(t, receivedSnapshot, 2)
|
||||
require.Equal(t, "call_first", receivedSnapshot[0].ToolUseID)
|
||||
require.Equal(t, "call_second", receivedSnapshot[1].ToolUseID)
|
||||
require.Equal(t, "read_file", receivedSnapshot[0].ToolName)
|
||||
require.Empty(t, receivedSnapshot[0].ToolError)
|
||||
require.Contains(t, string(receivedSnapshot[0].ToolResponse), "data")
|
||||
|
||||
var toolResults, userMessages int
|
||||
for _, message := range chatMessages(ctx, t, db, chat.ID) {
|
||||
parts, err := chatprompt.ParseContent(message)
|
||||
require.NoError(t, err)
|
||||
if message.Role == database.ChatMessageRoleTool {
|
||||
toolResults++
|
||||
}
|
||||
if len(parts) == 1 && parts[0].Text == "tool notice" {
|
||||
userMessages++
|
||||
require.Equal(t, database.ChatMessageVisibilityUser, message.Visibility)
|
||||
}
|
||||
}
|
||||
promptMessages, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
var modelContexts int
|
||||
for _, message := range promptMessages {
|
||||
parts, err := chatprompt.ParseContent(message)
|
||||
require.NoError(t, err)
|
||||
if len(parts) == 1 && parts[0].Text == "lint feedback" {
|
||||
modelContexts++
|
||||
require.Equal(t, database.ChatMessageVisibilityModel, message.Visibility)
|
||||
}
|
||||
}
|
||||
require.Equal(t, 2, toolResults)
|
||||
require.Equal(t, 1, modelContexts)
|
||||
require.Equal(t, 1, userMessages)
|
||||
}
|
||||
|
||||
func TestPostToolUseHookDynamicResult(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
var modelCalls atomic.Int32
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
if modelCalls.Add(1) == 1 {
|
||||
chunk := chattest.OpenAIToolCallChunk("my_dynamic_tool", `{"query":"value"}`)
|
||||
chunk.Choices[0].ToolCalls[0].ID = "call_dynamic_result"
|
||||
return chattest.OpenAIStreamingResponse(chunk)
|
||||
}
|
||||
resultIndex := -1
|
||||
contextIndex := -1
|
||||
for i, message := range req.Messages {
|
||||
if message.Role == "tool" && strings.Contains(message.Content, "answer") {
|
||||
resultIndex = i
|
||||
}
|
||||
if strings.Contains(message.Content, "dynamic feedback") {
|
||||
contextIndex = i
|
||||
}
|
||||
}
|
||||
require.NotEqual(t, -1, resultIndex)
|
||||
require.NotEqual(t, -1, contextIndex)
|
||||
require.Less(t, resultIndex, contextIndex)
|
||||
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("done")...)
|
||||
})
|
||||
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
|
||||
var postCalls atomic.Int32
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
if request.Type != agenthooks.EventPostToolUse {
|
||||
_, err := w.Write([]byte(`{}`))
|
||||
require.NoError(t, err)
|
||||
return
|
||||
}
|
||||
postCalls.Add(1)
|
||||
data := decodeHookData[agenthooks.PostToolUseData](t, request)
|
||||
require.Equal(t, "call_dynamic_result", data.ToolUseID)
|
||||
require.Equal(t, "my_dynamic_tool", data.ToolName)
|
||||
require.JSONEq(t, `{"answer":42}`, string(data.ToolResponse))
|
||||
_, err := w.Write([]byte(`{"model_context":"dynamic feedback","user_message":"dynamic notice"}`))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
|
||||
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
})
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "post-tool-use-dynamic",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("call the dynamic tool"),
|
||||
},
|
||||
DynamicTools: dynamicToolJSON(t, "my_dynamic_tool"),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
testutil.Eventually(ctx, t, func(context.Context) bool {
|
||||
updated, err := db.GetChatByID(ctx, chat.ID)
|
||||
return err == nil && updated.Status == database.ChatStatusRequiresAction
|
||||
}, testutil.IntervalFast)
|
||||
|
||||
err = server.SubmitToolResults(ctx, chatd.SubmitToolResultsOptions{
|
||||
ChatID: chat.ID,
|
||||
UserID: user.ID,
|
||||
ModelConfigID: model.ID,
|
||||
Results: []codersdk.ToolResult{{
|
||||
ToolCallID: "call_dynamic_result",
|
||||
Output: json.RawMessage(`{"answer":42}`),
|
||||
}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
|
||||
require.Equal(t, int32(1), postCalls.Load())
|
||||
|
||||
err = server.SubmitToolResults(ctx, chatd.SubmitToolResultsOptions{
|
||||
ChatID: chat.ID,
|
||||
UserID: user.ID,
|
||||
ModelConfigID: model.ID,
|
||||
Results: []codersdk.ToolResult{{
|
||||
ToolCallID: "call_dynamic_result",
|
||||
Output: json.RawMessage(`{"answer":42}`),
|
||||
}},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Equal(t, int32(1), postCalls.Load())
|
||||
var notices int
|
||||
for _, message := range chatMessages(ctx, t, db, chat.ID) {
|
||||
if hookMessageText(t, message) == "dynamic notice" {
|
||||
notices++
|
||||
}
|
||||
}
|
||||
require.Equal(t, 1, notices)
|
||||
}
|
||||
|
||||
func TestPostToolUseHookDynamicFailureRejectsSubmission(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
chunk := chattest.OpenAIToolCallChunk("my_dynamic_tool", `{}`)
|
||||
chunk.Choices[0].ToolCalls[0].ID = "call_dynamic_failure"
|
||||
return chattest.OpenAIStreamingResponse(chunk)
|
||||
})
|
||||
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
|
||||
var postCalls atomic.Int32
|
||||
var failPostToolUse atomic.Bool
|
||||
failPostToolUse.Store(true)
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
if request.Type == agenthooks.EventPostToolUse {
|
||||
postCalls.Add(1)
|
||||
if failPostToolUse.Load() {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
}
|
||||
_, err := w.Write([]byte(`{}`))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
|
||||
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
})
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "post-tool-use-dynamic-failure",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("call the dynamic tool"),
|
||||
},
|
||||
DynamicTools: dynamicToolJSON(t, "my_dynamic_tool"),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
testutil.Eventually(ctx, t, func(context.Context) bool {
|
||||
updated, err := db.GetChatByID(ctx, chat.ID)
|
||||
return err == nil && updated.Status == database.ChatStatusRequiresAction
|
||||
}, testutil.IntervalFast)
|
||||
|
||||
results := []codersdk.ToolResult{{
|
||||
ToolCallID: "call_dynamic_failure",
|
||||
Output: json.RawMessage(`{"answer":42}`),
|
||||
}}
|
||||
err = server.SubmitToolResults(ctx, chatd.SubmitToolResultsOptions{
|
||||
ChatID: chat.ID,
|
||||
UserID: user.ID,
|
||||
ModelConfigID: model.ID,
|
||||
Results: results,
|
||||
})
|
||||
var dispatchErr *dispatch.Error
|
||||
require.ErrorAs(t, err, &dispatchErr)
|
||||
|
||||
unchanged, err := db.GetChatByID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, database.ChatStatusRequiresAction, unchanged.Status)
|
||||
require.False(t, unchanged.LastError.Valid)
|
||||
for _, part := range chatToolParts(ctx, t, db, chat.ID) {
|
||||
require.NotEqual(t, codersdk.ChatMessagePartTypeToolResult, part.Type,
|
||||
"rejected submission must not commit tool results")
|
||||
}
|
||||
require.Equal(t, int32(1), postCalls.Load())
|
||||
|
||||
failPostToolUse.Store(false)
|
||||
require.NoError(t, server.SubmitToolResults(ctx, chatd.SubmitToolResultsOptions{
|
||||
ChatID: chat.ID,
|
||||
UserID: user.ID,
|
||||
ModelConfigID: model.ID,
|
||||
Results: results,
|
||||
}))
|
||||
result := requireToolResultPart(t, chatToolParts(ctx, t, db, chat.ID), "my_dynamic_tool")
|
||||
require.JSONEq(t, `{"answer":42}`, string(result.Result))
|
||||
require.Equal(t, int32(2), postCalls.Load())
|
||||
}
|
||||
|
||||
func TestPostToolUseHookFailureCommitsResultThenErrors(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
chunk := chattest.OpenAIToolCallChunk("read_file", `{"path":"/tmp/file.txt"}`)
|
||||
chunk.Choices[0].ToolCalls[0].ID = "call_failure"
|
||||
return chattest.OpenAIStreamingResponse(chunk)
|
||||
})
|
||||
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
|
||||
ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID)
|
||||
var postCalls atomic.Int32
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
if request.Type == agenthooks.EventPostToolUse {
|
||||
postCalls.Add(1)
|
||||
data := decodeHookData[agenthooks.PostToolUseData](t, request)
|
||||
require.Equal(t, "call_failure", data.ToolUseID)
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
_, err := w.Write([]byte(`{}`))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
mockConn := agentconnmock.NewMockAgentConn(ctrl)
|
||||
setupToolExecutionAgentConn(t, mockConn)
|
||||
mockConn.EXPECT().ReadFileLines(gomock.Any(), "/tmp/file.txt", int64(1), int64(0), gomock.Any()).
|
||||
Return(workspacesdk.ReadFileLinesResponse{
|
||||
Success: true, FileSize: 4, TotalLines: 1, LinesRead: 1, Content: "data",
|
||||
}, nil)
|
||||
|
||||
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
|
||||
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
|
||||
require.Equal(t, dbAgent.ID, agentID)
|
||||
return mockConn, func() {}, nil
|
||||
}
|
||||
})
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true},
|
||||
AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true},
|
||||
Title: "post-tool-use-failure",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("read the file"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
failed := waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusError)
|
||||
|
||||
result := requireToolResultPart(t, chatToolParts(ctx, t, db, chat.ID), "read_file")
|
||||
require.Contains(t, string(result.Result), "data")
|
||||
require.Equal(t, int32(1), postCalls.Load())
|
||||
lastError := chatLastErrorMessage(failed.LastError)
|
||||
require.Contains(t, lastError, "hook dispatch failed: post_tool_use: http_error")
|
||||
}
|
||||
|
||||
func TestSubagentSpawnHookDispatchFailureCommitsSiblingResults(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
chunk := chattest.OpenAIToolCallChunk("read_file", `{"path":"/tmp/sibling.txt"}`)
|
||||
chunk.Choices[0].ToolCalls[0].ID = "call_sibling"
|
||||
spawn := chattest.OpenAIToolCallChunk("spawn_agent", `{"type":"general","prompt":"child admission prompt"}`).Choices[0].ToolCalls[0]
|
||||
spawn.ID = "call_spawn"
|
||||
spawn.Index = 1
|
||||
chunk.Choices[0].ToolCalls = append(chunk.Choices[0].ToolCalls, spawn)
|
||||
return chattest.OpenAIStreamingResponse(chunk)
|
||||
})
|
||||
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
|
||||
ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID)
|
||||
|
||||
var mu sync.Mutex
|
||||
var postToolUseIDs []string
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
switch request.Type {
|
||||
case agenthooks.EventUserPromptSubmit:
|
||||
data := decodeHookData[agenthooks.UserPromptSubmitData](t, request)
|
||||
if data.Prompt == "child admission prompt" {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
case agenthooks.EventPostToolUse:
|
||||
data := decodeHookData[agenthooks.PostToolUseData](t, request)
|
||||
mu.Lock()
|
||||
postToolUseIDs = append(postToolUseIDs, data.ToolUseID)
|
||||
mu.Unlock()
|
||||
}
|
||||
_, err := w.Write([]byte(`{}`))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
mockConn := agentconnmock.NewMockAgentConn(ctrl)
|
||||
setupToolExecutionAgentConn(t, mockConn)
|
||||
mockConn.EXPECT().ReadFileLines(gomock.Any(), gomock.Any(), int64(1), int64(0), gomock.Any()).
|
||||
Return(workspacesdk.ReadFileLinesResponse{
|
||||
Success: true, FileSize: 7, TotalLines: 1, LinesRead: 1, Content: "sibling",
|
||||
}, nil)
|
||||
|
||||
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
|
||||
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
|
||||
require.Equal(t, dbAgent.ID, agentID)
|
||||
return mockConn, func() {}, nil
|
||||
}
|
||||
})
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true},
|
||||
AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true},
|
||||
Title: "spawn-hook-failure-siblings",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("read a file and spawn a child"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
failed := waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusError)
|
||||
require.Contains(t, chatLastErrorMessage(failed.LastError), "hook dispatch failed: user_prompt_submit: http_error")
|
||||
|
||||
messages := chatMessages(ctx, t, db, chat.ID)
|
||||
require.Len(t, messages, 4)
|
||||
require.Equal(t, database.ChatMessageRoleTool, messages[2].Role)
|
||||
require.Equal(t, database.ChatMessageRoleTool, messages[3].Role)
|
||||
sibling := string(messages[2].Content.RawMessage)
|
||||
require.Contains(t, sibling, "call_sibling")
|
||||
require.Contains(t, sibling, "sibling")
|
||||
spawn := string(messages[3].Content.RawMessage)
|
||||
require.Contains(t, spawn, "call_spawn")
|
||||
require.Contains(t, spawn, "lifecycle hook returned HTTP status 500")
|
||||
|
||||
mu.Lock()
|
||||
dispatched := slices.Clone(postToolUseIDs)
|
||||
mu.Unlock()
|
||||
require.Equal(t, []string{"call_sibling"}, dispatched,
|
||||
"the spawn tool never ran, so it has no use to post-process")
|
||||
}
|
||||
|
||||
func TestPostToolUseHookFailureDispatchesRemainingResults(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
first := chattest.OpenAIToolCallChunk("read_file", `{"path":"/tmp/first.txt"}`)
|
||||
first.Choices[0].ToolCalls[0].ID = "call_first"
|
||||
second := chattest.OpenAIToolCallChunk("read_file", `{"path":"/tmp/second.txt"}`).Choices[0].ToolCalls[0]
|
||||
second.ID = "call_second"
|
||||
second.Index = 1
|
||||
first.Choices[0].ToolCalls = append(first.Choices[0].ToolCalls, second)
|
||||
return chattest.OpenAIStreamingResponse(first)
|
||||
})
|
||||
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
|
||||
ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID)
|
||||
|
||||
var mu sync.Mutex
|
||||
results := map[string]string{}
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
if request.Type != agenthooks.EventPostToolUse {
|
||||
_, err := w.Write([]byte(`{}`))
|
||||
require.NoError(t, err)
|
||||
return
|
||||
}
|
||||
data := decodeHookData[agenthooks.PostToolUseData](t, request)
|
||||
result := "ok"
|
||||
if data.ToolUseID == "call_first" {
|
||||
result = "http_error"
|
||||
}
|
||||
mu.Lock()
|
||||
results[data.ToolUseID] = result
|
||||
mu.Unlock()
|
||||
if result == "http_error" {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
_, err := w.Write([]byte(`{}`))
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
mockConn := agentconnmock.NewMockAgentConn(ctrl)
|
||||
setupToolExecutionAgentConn(t, mockConn)
|
||||
mockConn.EXPECT().ReadFileLines(gomock.Any(), gomock.Any(), int64(1), int64(0), gomock.Any()).
|
||||
Return(workspacesdk.ReadFileLinesResponse{
|
||||
Success: true, FileSize: 4, TotalLines: 1, LinesRead: 1, Content: "data",
|
||||
}, nil).
|
||||
Times(2)
|
||||
|
||||
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
|
||||
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
|
||||
require.Equal(t, dbAgent.ID, agentID)
|
||||
return mockConn, func() {}, nil
|
||||
}
|
||||
})
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true},
|
||||
AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true},
|
||||
Title: "post-tool-use-continue",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("read both files"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
failed := waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusError)
|
||||
|
||||
mu.Lock()
|
||||
received := make(map[string]string, len(results))
|
||||
for toolUseID, result := range results {
|
||||
received[toolUseID] = result
|
||||
}
|
||||
mu.Unlock()
|
||||
require.Equal(t, map[string]string{
|
||||
"call_first": "http_error",
|
||||
"call_second": "ok",
|
||||
}, received)
|
||||
require.Contains(t, chatLastErrorMessage(failed.LastError), "hook dispatch failed: post_tool_use: http_error")
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -57,6 +57,8 @@ type runner struct {
|
||||
tasksByIndex map[taskIndexKey]taskInstanceID
|
||||
localLocks *localLockSet
|
||||
debugTurn *runnerDebugTurn
|
||||
sessionStart sessionStartTracker
|
||||
stopNudges stopNudgeTracker
|
||||
}
|
||||
|
||||
func newRunner(ctx context.Context, mgr *runnerManager, rec *runnerRecord, opts chatWorkerOptions) *runner {
|
||||
@@ -227,6 +229,8 @@ func (r *runner) spawnTaskIfNeeded(kind taskKind, state runnerStateUpdate) {
|
||||
Status: state.Status,
|
||||
RequiresActionDeadlineAt: state.RequiresActionDeadlineAt,
|
||||
DebugTurn: r.debugTurn,
|
||||
SessionStart: &r.sessionStart,
|
||||
StopNudges: &r.stopNudges,
|
||||
}
|
||||
go r.runTask(taskCtx, kind, key, input, done)
|
||||
}
|
||||
|
||||
@@ -42,6 +42,7 @@ func TestRunner_CancelsActiveTaskWhenHistoryChanges(t *testing.T) {
|
||||
require.NotErrorIs(t, context.Cause(first.ctx), errTaskTimeout)
|
||||
second := starter.waitCall(t, taskKindGeneration, chat.ID)
|
||||
require.Equal(t, updated.HistoryVersion, second.input.HistoryVersion)
|
||||
require.Same(t, first.input.SessionStart, second.input.SessionStart)
|
||||
}
|
||||
|
||||
func TestRunner_CancelsActiveTaskWhenStatusChanges(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,181 @@
|
||||
package chatd_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
func TestStopHookNoOpFinishesTurn(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("done")...)
|
||||
})
|
||||
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
|
||||
var stopCalls atomic.Int32
|
||||
consumer := stopConsumer(t, func() (int, string) {
|
||||
stopCalls.Add(1)
|
||||
return http.StatusOK, `{}`
|
||||
})
|
||||
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
|
||||
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
})
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "stop-noop",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("finish normally"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
testutil.Eventually(ctx, t, func(context.Context) bool {
|
||||
return stopCalls.Load() == 1
|
||||
}, testutil.IntervalFast)
|
||||
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
|
||||
require.Equal(t, int32(1), stopCalls.Load())
|
||||
}
|
||||
|
||||
func TestStopHookNudgeContinuesOnce(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
var modelCalls atomic.Int32
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
switch modelCalls.Add(1) {
|
||||
case 1:
|
||||
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("first answer")...)
|
||||
case 2:
|
||||
var found bool
|
||||
for _, message := range req.Messages {
|
||||
found = found || strings.Contains(message.Content, "continue please")
|
||||
}
|
||||
require.True(t, found)
|
||||
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("second answer")...)
|
||||
default:
|
||||
require.FailNow(t, "stop nudge exceeded continuation cap")
|
||||
return chattest.OpenAIStreamingResponse()
|
||||
}
|
||||
})
|
||||
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
|
||||
var stopCalls atomic.Int32
|
||||
consumer := stopConsumer(t, func() (int, string) {
|
||||
stopCalls.Add(1)
|
||||
return http.StatusOK, `{"model_context":"continue please"}`
|
||||
})
|
||||
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
|
||||
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
})
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "stop-nudge",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("continue once"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
testutil.Eventually(ctx, t, func(context.Context) bool {
|
||||
return stopCalls.Load() == 2
|
||||
}, testutil.IntervalFast)
|
||||
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
|
||||
require.Equal(t, int32(2), modelCalls.Load())
|
||||
require.Equal(t, int32(2), stopCalls.Load())
|
||||
|
||||
promptMessages, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
var contextRows int
|
||||
for _, message := range promptMessages {
|
||||
parts, err := chatprompt.ParseContent(message)
|
||||
require.NoError(t, err)
|
||||
if len(parts) == 1 && parts[0].Text == "continue please" {
|
||||
contextRows++
|
||||
require.Equal(t, database.ChatMessageVisibilityModel, message.Visibility)
|
||||
}
|
||||
}
|
||||
require.Equal(t, 2, contextRows)
|
||||
}
|
||||
|
||||
func TestStopHookDispatchFailureErrorsChat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("done")...)
|
||||
})
|
||||
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
|
||||
consumer := stopConsumer(t, func() (int, string) {
|
||||
return http.StatusInternalServerError, ""
|
||||
})
|
||||
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
|
||||
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
|
||||
cfg.HookDispatcher = newHookDispatcher(t, db, consumer)
|
||||
})
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "stop-failure",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("fail on stop"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
failed := waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusError)
|
||||
lastError := chatLastErrorMessage(failed.LastError)
|
||||
require.Contains(t, lastError, "hook dispatch failed: stop: http_error")
|
||||
}
|
||||
|
||||
func stopConsumer(t *testing.T, response func() (int, string)) *httptest.Server {
|
||||
t.Helper()
|
||||
consumer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request agenthooks.Request
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
if request.Type != agenthooks.EventStop {
|
||||
_, err := w.Write([]byte(`{}`))
|
||||
require.NoError(t, err)
|
||||
return
|
||||
}
|
||||
status, body := response()
|
||||
w.WriteHeader(status)
|
||||
if body != "" {
|
||||
_, err := w.Write([]byte(body))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(consumer.Close)
|
||||
return consumer
|
||||
}
|
||||
@@ -20,11 +20,14 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/database/dbauthz"
|
||||
dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub"
|
||||
coderdpubsub "github.com/coder/coder/v2/coderd/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/x/agenthooks/dispatch"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chathooks"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk"
|
||||
"github.com/coder/coder/v2/codersdk/x/agenthooks"
|
||||
)
|
||||
|
||||
var ErrSubagentNotDescendant = xerrors.New("target chat is not a descendant of current chat")
|
||||
@@ -768,6 +771,13 @@ func (p *Server) subagentTools(
|
||||
options,
|
||||
)
|
||||
if err != nil {
|
||||
// A failed hook dispatch must fail closed instead of
|
||||
// degrading into a tool error the model can ignore.
|
||||
if _, ok := errors.AsType[*dispatch.Error](err); ok {
|
||||
return fantasy.ToolResponse{}, err
|
||||
}
|
||||
// chathooks.UserPromptDeniedError.Error() carries the user-facing
|
||||
// denial message, so the model can adjust its prompt.
|
||||
return fantasy.NewTextErrorResponse(err.Error()), nil
|
||||
}
|
||||
|
||||
@@ -1153,7 +1163,6 @@ func (p *Server) loadSubagentSpawnParentChat(
|
||||
if err := validateSubagentSpawnParent(parent); err != nil {
|
||||
return database.Chat{}, err
|
||||
}
|
||||
|
||||
return parent, nil
|
||||
}
|
||||
|
||||
@@ -1245,9 +1254,6 @@ func (p *Server) createChildSubagentChatWithOptions(
|
||||
}
|
||||
|
||||
title = strings.TrimSpace(title)
|
||||
if title == "" {
|
||||
title = subagentFallbackChatTitle(prompt)
|
||||
}
|
||||
|
||||
rootChatID := parent.ID
|
||||
if parent.RootChatID.Valid {
|
||||
@@ -1291,6 +1297,39 @@ func (p *Server) createChildSubagentChatWithOptions(
|
||||
return database.Chat{}, limitErr
|
||||
}
|
||||
|
||||
// Review before persistence so spawned chats cannot bypass prompt policy.
|
||||
childChatID := uuid.New()
|
||||
var promptResult *chathooks.Result
|
||||
if p.hooks.Enabled() {
|
||||
mintedTurnID := uuid.New()
|
||||
promptMessage, err := chathooks.UserPromptMessage([]codersdk.ChatMessagePart{codersdk.ChatMessageText(prompt)})
|
||||
if err != nil {
|
||||
return database.Chat{}, err
|
||||
}
|
||||
promptResult, err = p.hooks.Trigger(ctx, chathooks.Chat{
|
||||
ID: childChatID,
|
||||
OwnerID: parent.OwnerID,
|
||||
WorkspaceID: parent.WorkspaceID,
|
||||
ParentChatID: uuid.NullUUID{UUID: parent.ID, Valid: true},
|
||||
RootChatID: uuid.NullUUID{UUID: rootChatID, Valid: true},
|
||||
TurnID: &mintedTurnID,
|
||||
}, promptMessage, agenthooks.EventUserPromptSubmit)
|
||||
if err != nil {
|
||||
return database.Chat{}, chathooks.UserPromptDenial(err)
|
||||
}
|
||||
override, overridden, overrideErr := chathooks.UserPromptOverride(promptResult)
|
||||
if overrideErr != nil {
|
||||
return database.Chat{}, overrideErr
|
||||
}
|
||||
if overridden {
|
||||
// The overridden prompt also feeds the fallback title below.
|
||||
prompt = override
|
||||
}
|
||||
}
|
||||
if title == "" {
|
||||
title = subagentFallbackChatTitle(prompt)
|
||||
}
|
||||
|
||||
workspaceAwareness := workspaceDetachedNoCreateAwareness
|
||||
if parent.WorkspaceID.Valid {
|
||||
workspaceAwareness = workspaceAttachedAwareness
|
||||
@@ -1301,7 +1340,9 @@ func (p *Server) createChildSubagentChatWithOptions(
|
||||
if err != nil {
|
||||
return database.Chat{}, xerrors.Errorf("marshal workspace awareness: %w", err)
|
||||
}
|
||||
userContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{codersdk.ChatMessageText(prompt)})
|
||||
childUserParts := []codersdk.ChatMessagePart{codersdk.ChatMessageText(prompt)}
|
||||
childUserParts = append(childUserParts, chathooks.UserPromptParts(promptResult)...)
|
||||
userContent, err := chatprompt.MarshalParts(childUserParts)
|
||||
if err != nil {
|
||||
return database.Chat{}, xerrors.Errorf("marshal initial user content: %w", err)
|
||||
}
|
||||
@@ -1337,7 +1378,7 @@ func (p *Server) createChildSubagentChatWithOptions(
|
||||
if publisher == nil {
|
||||
publisher = dbpubsub.NewInMemory()
|
||||
}
|
||||
result, err := chatstate.CreateChat(ctx, p.db, publisher, chatstate.CreateChatInput{
|
||||
result, err := chatstate.CreateChatWithID(ctx, p.db, publisher, childChatID, chatstate.CreateChatInput{
|
||||
OrganizationID: parent.OrganizationID,
|
||||
OwnerID: parent.OwnerID,
|
||||
WorkspaceID: parent.WorkspaceID,
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
|
||||
"charm.land/fantasy"
|
||||
"github.com/google/uuid"
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -30,6 +31,8 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/database/pubsub"
|
||||
coderdpubsub "github.com/coder/coder/v2/coderd/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/util/ptr"
|
||||
"github.com/coder/coder/v2/coderd/x/agenthooks/dispatch"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chathooks"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatloop"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
|
||||
@@ -267,6 +270,154 @@ func insertInternalAIProvider(
|
||||
})
|
||||
}
|
||||
|
||||
func TestCreateChildSubagentChatDispatchesUserPromptSubmit(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
newFixture := func(t *testing.T, handler http.HandlerFunc) (context.Context, database.Store, database.Chat, *Server) {
|
||||
t.Helper()
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
db, _ := dbtestutil.NewDB(t)
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{UserID: user.ID, OrganizationID: org.ID})
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{})
|
||||
parent := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
})
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID})
|
||||
ctx = aibridge.WithDelegatedAPIKeyID(ctx, apiKey.ID)
|
||||
|
||||
consumer := httptest.NewServer(handler)
|
||||
t.Cleanup(consumer.Close)
|
||||
server := &Server{
|
||||
db: db,
|
||||
logger: slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
|
||||
hooks: chathooks.NewTrigger(dispatch.New(
|
||||
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
|
||||
consumer.Client(),
|
||||
consumer.URL,
|
||||
"test-hook-secret-32-bytes-minimum!!",
|
||||
time.Second,
|
||||
"test-deployment",
|
||||
"test-version",
|
||||
prometheus.NewRegistry(),
|
||||
)),
|
||||
}
|
||||
return ctx, db, parent, server
|
||||
}
|
||||
|
||||
t.Run("Rewrite", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var meta struct {
|
||||
sync.Mutex
|
||||
parentChatID string
|
||||
prompt string
|
||||
}
|
||||
ctx, db, parent, server := newFixture(t, func(rw http.ResponseWriter, r *http.Request) {
|
||||
var request struct {
|
||||
Type string `json:"type"`
|
||||
Meta struct {
|
||||
ParentChatID *uuid.UUID `json:"parent_chat_id"`
|
||||
} `json:"meta"`
|
||||
Data struct {
|
||||
Prompt string `json:"prompt"`
|
||||
} `json:"data"`
|
||||
}
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
|
||||
require.Equal(t, "user_prompt_submit", request.Type)
|
||||
meta.Lock()
|
||||
if request.Meta.ParentChatID != nil {
|
||||
meta.parentChatID = request.Meta.ParentChatID.String()
|
||||
}
|
||||
meta.prompt = request.Data.Prompt
|
||||
meta.Unlock()
|
||||
rw.Header().Set("Content-Type", "application/json")
|
||||
_, _ = rw.Write([]byte(`{
|
||||
"permission": {"decision": "allow", "input_override": {"prompt": "REVIEWED: inspect"}}
|
||||
}`))
|
||||
})
|
||||
|
||||
child, err := server.createChildSubagentChatWithOptions(ctx, parent, "inspect the workspace", "", childSubagentChatOptions{})
|
||||
require.NoError(t, err)
|
||||
|
||||
meta.Lock()
|
||||
require.Equal(t, parent.ID.String(), meta.parentChatID, "spawn dispatch must identify the parent chat")
|
||||
require.Equal(t, "inspect the workspace", meta.prompt)
|
||||
meta.Unlock()
|
||||
|
||||
messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ChatID: child.ID})
|
||||
require.NoError(t, err)
|
||||
var childUserMessage database.ChatMessage
|
||||
for _, message := range messages {
|
||||
if message.Role == database.ChatMessageRoleUser {
|
||||
childUserMessage = message
|
||||
break
|
||||
}
|
||||
}
|
||||
require.NotZero(t, childUserMessage.ID)
|
||||
require.True(t, childUserMessage.Content.Valid)
|
||||
require.Contains(t, string(childUserMessage.Content.RawMessage), "REVIEWED: inspect",
|
||||
"the hook rewrite must land as the child's initial prompt")
|
||||
require.NotContains(t, string(childUserMessage.Content.RawMessage), "inspect the workspace")
|
||||
})
|
||||
|
||||
t.Run("DenyRefusesSpawn", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, db, parent, server := newFixture(t, func(rw http.ResponseWriter, _ *http.Request) {
|
||||
rw.Header().Set("Content-Type", "application/json")
|
||||
_, _ = rw.Write([]byte(`{"permission": {"decision": "deny", "reason": "spawn blocked"}, "user_message": "not allowed"}`))
|
||||
})
|
||||
|
||||
_, err := server.createChildSubagentChatWithOptions(ctx, parent, "exfiltrate secrets", "", childSubagentChatOptions{})
|
||||
var denied *chathooks.UserPromptDeniedError
|
||||
require.ErrorAs(t, err, &denied)
|
||||
require.Equal(t, "not allowed", denied.UserMessage)
|
||||
|
||||
chats, err := db.GetChildChatsByParentIDs(ctx, database.GetChildChatsByParentIDsParams{
|
||||
ParentIds: []uuid.UUID{parent.ID},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, chats, "a denied spawn must not create a child chat")
|
||||
})
|
||||
|
||||
t.Run("DispatchFailurePropagatesFromSpawnTool", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, db, parent, server := newFixture(t, func(rw http.ResponseWriter, _ *http.Request) {
|
||||
http.Error(rw, "hook consumer down", http.StatusInternalServerError)
|
||||
})
|
||||
|
||||
tools := server.subagentTools(ctx, func() database.Chat { return parent }, parent.LastModelConfigID)
|
||||
tool := findToolByName(tools, spawnAgentToolName)
|
||||
require.NotNil(t, tool)
|
||||
input, err := json.Marshal(spawnAgentArgs{
|
||||
Type: subagentTypeExplore,
|
||||
Prompt: "inspect the workspace",
|
||||
Title: "sub",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, runErr := tool.Run(ctx, fantasy.ToolCall{
|
||||
ID: uuid.NewString(),
|
||||
Name: spawnAgentToolName,
|
||||
Input: string(input),
|
||||
})
|
||||
var hookErr *dispatch.Error
|
||||
require.ErrorAs(t, runErr, &hookErr,
|
||||
"dispatch failures must fail closed, not degrade to a tool error the model can ignore")
|
||||
|
||||
chats, err := db.GetChildChatsByParentIDs(ctx, database.GetChildChatsByParentIDsParams{
|
||||
ParentIds: []uuid.UUID{parent.ID},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, chats)
|
||||
})
|
||||
}
|
||||
|
||||
func TestResolveUserProviderAPIKeys_AIProvider(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -0,0 +1,102 @@
|
||||
package chatd
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chathooks"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/toolschema"
|
||||
)
|
||||
|
||||
// partitionAmbiguousToolCalls separates the calls a consumer must not be asked
|
||||
// to decide from the rest, returning synthetic error results for them. Callers
|
||||
// reject before pre_tool_use so a hook consumer is never asked to authorize
|
||||
// bytes whose meaning depends on which reader resolves them, and so input that
|
||||
// cannot be carried in a hook payload fails as a retryable tool error instead
|
||||
// of a dispatch failure.
|
||||
func partitionAmbiguousToolCalls(
|
||||
prepared generationPrepared,
|
||||
toolCalls []fantasy.ToolCallContent,
|
||||
) ([]fantasy.ToolCallContent, []fantasy.ToolResultContent) {
|
||||
var (
|
||||
allowed []fantasy.ToolCallContent
|
||||
rejected []fantasy.ToolResultContent
|
||||
)
|
||||
for _, toolCall := range toolCalls {
|
||||
if !json.Valid([]byte(toolCall.Input)) {
|
||||
rejected = append(rejected, malformedToolResult(toolCall))
|
||||
continue
|
||||
}
|
||||
if err := validateBuiltinToolInput(prepared, toolCall.ToolName, []byte(toolCall.Input)); err != nil {
|
||||
rejected = append(rejected, ambiguousToolResult(toolCall, err))
|
||||
continue
|
||||
}
|
||||
allowed = append(allowed, toolCall)
|
||||
}
|
||||
return allowed, rejected
|
||||
}
|
||||
|
||||
// validateOverriddenToolInputs rechecks the inputs a pre_tool_use consumer
|
||||
// replaced. The model cannot fix an ambiguous override, so the turn fails
|
||||
// closed instead of executing it.
|
||||
func validateOverriddenToolInputs(prepared generationPrepared, preflight chathooks.PreToolUseExecutionResult) error {
|
||||
for _, toolCall := range preflight.Allowed {
|
||||
if _, overridden := preflight.Overrides[toolCall.ToolCallID]; !overridden {
|
||||
continue
|
||||
}
|
||||
if err := validateBuiltinToolInput(prepared, toolCall.ToolName, []byte(toolCall.Input)); err != nil {
|
||||
return xerrors.Errorf("hook input override for tool %s: %w", toolCall.ToolName, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateBuiltinToolInput only guards builtin tools, whose input coderd
|
||||
// decodes itself. Dynamic calls are executed by the client and MCP calls by
|
||||
// their own server, and a dynamic tool cannot shadow a builtin name.
|
||||
func validateBuiltinToolInput(prepared generationPrepared, toolName string, input []byte) error {
|
||||
// Execution resolves a deprecated alias to its canonical tool, so
|
||||
// skipping that here would let the old name bypass validation.
|
||||
if canonical, aliased := subagentToolNameAliases[toolName]; aliased {
|
||||
toolName = canonical
|
||||
}
|
||||
if !prepared.BuiltinToolNames[toolName] {
|
||||
return nil
|
||||
}
|
||||
for _, tool := range prepared.Tools {
|
||||
info := tool.Info()
|
||||
if info.Name != toolName {
|
||||
continue
|
||||
}
|
||||
return toolschema.ValidateUnambiguous(info.Parameters, input)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// malformedToolResult reports input the tool decoder would reject anyway. It
|
||||
// is produced here because a hook payload carries the input as JSON, so
|
||||
// invalid bytes would otherwise surface as a dispatch failure and end the
|
||||
// turn instead of letting the model correct the call.
|
||||
func malformedToolResult(toolCall fantasy.ToolCallContent) fantasy.ToolResultContent {
|
||||
return fantasy.ToolResultContent{
|
||||
ToolCallID: toolCall.ToolCallID,
|
||||
ToolName: toolCall.ToolName,
|
||||
Result: fantasy.ToolResultOutputContentError{
|
||||
Error: xerrors.New("This tool call was not executed because its input is not valid JSON. Retry with a well-formed JSON object matching the tool schema."),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func ambiguousToolResult(toolCall fantasy.ToolCallContent, err error) fantasy.ToolResultContent {
|
||||
message := "This tool call was not executed because its input is ambiguous: " + err.Error() +
|
||||
". Retry with the exact property names from the tool schema, each key used once."
|
||||
return fantasy.ToolResultContent{
|
||||
ToolCallID: toolCall.ToolCallID,
|
||||
ToolName: toolCall.ToolName,
|
||||
Result: fantasy.ToolResultOutputContentError{
|
||||
Error: xerrors.New(message),
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,199 @@
|
||||
package chatd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chathooks"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
|
||||
)
|
||||
|
||||
func TestPartitionAmbiguousToolCallsGatesOnBuiltins(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fetch := fetchToolStub()
|
||||
ambiguous := fantasy.ToolCallContent{
|
||||
ToolCallID: "call_ambiguous",
|
||||
ToolName: "fetch",
|
||||
Input: `{"URL":"https://example.test","url":"https://other.test"}`,
|
||||
}
|
||||
clean := fantasy.ToolCallContent{
|
||||
ToolCallID: "call_clean",
|
||||
ToolName: "fetch",
|
||||
Input: `{"url":"https://example.test"}`,
|
||||
}
|
||||
|
||||
t.Run("builtin", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
prepared := generationPrepared{
|
||||
Tools: []fantasy.AgentTool{fetch},
|
||||
BuiltinToolNames: map[string]bool{"fetch": true},
|
||||
}
|
||||
allowed, rejected := partitionAmbiguousToolCalls(prepared, []fantasy.ToolCallContent{ambiguous, clean})
|
||||
require.Len(t, rejected, 1)
|
||||
require.Equal(t, "call_ambiguous", rejected[0].ToolCallID)
|
||||
require.Len(t, allowed, 1)
|
||||
require.Equal(t, "call_clean", allowed[0].ToolCallID)
|
||||
})
|
||||
|
||||
t.Run("non-builtin", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
prepared := generationPrepared{Tools: []fantasy.AgentTool{fetch}}
|
||||
allowed, rejected := partitionAmbiguousToolCalls(prepared, []fantasy.ToolCallContent{ambiguous, clean})
|
||||
require.Empty(t, rejected)
|
||||
require.Len(t, allowed, 2)
|
||||
})
|
||||
|
||||
// Execution resolves a deprecated name to its canonical tool, so
|
||||
// validation has to resolve it too.
|
||||
t.Run("deprecated alias", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var alias, canonical string
|
||||
for a, name := range subagentToolNameAliases {
|
||||
alias, canonical = a, name
|
||||
break
|
||||
}
|
||||
require.NotEmpty(t, canonical)
|
||||
|
||||
type input struct {
|
||||
ChatID string `json:"chat_id"`
|
||||
}
|
||||
tool := fantasy.NewAgentTool(canonical, "",
|
||||
func(context.Context, input, fantasy.ToolCall) (fantasy.ToolResponse, error) {
|
||||
return fantasy.ToolResponse{}, nil
|
||||
})
|
||||
aliased := fantasy.ToolCallContent{
|
||||
ToolCallID: "call_aliased",
|
||||
ToolName: alias,
|
||||
Input: `{"chat_id":"a","CHAT_ID":"b"}`,
|
||||
}
|
||||
|
||||
prepared := generationPrepared{
|
||||
Tools: []fantasy.AgentTool{tool},
|
||||
BuiltinToolNames: map[string]bool{canonical: true},
|
||||
}
|
||||
_, rejected := partitionAmbiguousToolCalls(prepared, []fantasy.ToolCallContent{aliased})
|
||||
require.Len(t, rejected, 1)
|
||||
})
|
||||
}
|
||||
|
||||
func TestValidateOverriddenToolInputs(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
prepared := generationPrepared{
|
||||
Tools: []fantasy.AgentTool{fetchToolStub()},
|
||||
BuiltinToolNames: map[string]bool{"fetch": true},
|
||||
}
|
||||
overridden := chathooks.PreToolUseExecutionResult{
|
||||
Allowed: []fantasy.ToolCallContent{{
|
||||
ToolCallID: "call_overridden",
|
||||
ToolName: "fetch",
|
||||
Input: `{"URL":"https://other.test"}`,
|
||||
}},
|
||||
Overrides: map[string]json.RawMessage{
|
||||
"call_overridden": json.RawMessage(`{"URL":"https://other.test"}`),
|
||||
},
|
||||
}
|
||||
require.ErrorContains(t, validateOverriddenToolInputs(prepared, overridden),
|
||||
`hook input override for tool fetch: input key "URL" differs from schema property "url" only by case`)
|
||||
|
||||
// The same input is left alone when no consumer replaced it, because
|
||||
// the model-authored batch is checked before the dispatch instead.
|
||||
untouched := chathooks.PreToolUseExecutionResult{Allowed: overridden.Allowed}
|
||||
require.NoError(t, validateOverriddenToolInputs(prepared, untouched))
|
||||
}
|
||||
|
||||
// TestBuiltinToolSchemasDescribeTheirInputs guards the validator's reach: it
|
||||
// cannot detect a case-variant key for a builtin that declares no properties.
|
||||
func TestBuiltinToolSchemasDescribeTheirInputs(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := chatdTestContext(t)
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{
|
||||
UserID: user.ID,
|
||||
OrganizationID: org.ID,
|
||||
})
|
||||
provider := dbgen.AIProviderWithOptionalKey(t, db, database.AIProvider{
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
}, "test-key")
|
||||
modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Model: "gpt-4o-mini",
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
}, func(p *database.InsertChatModelConfigParams) {
|
||||
p.Enabled = true
|
||||
})
|
||||
created, err := chatstate.CreateChat(ctx, db, ps, chatstate.CreateChatInput{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: modelConfig.ID,
|
||||
Title: "builtin tool schemas",
|
||||
ClientType: database.ChatClientTypeApi,
|
||||
InitialMessages: []chatstate.Message{{
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: mustMarshalText(t, "hello"),
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
server := newInternalTestServer(
|
||||
t,
|
||||
db,
|
||||
ps,
|
||||
chatprovider.ProviderAPIKeys{},
|
||||
withInternalTestServerTransportFactory(&aibridgeTestFactory{}),
|
||||
)
|
||||
prepared, err := server.prepareGeneration(ctx, generationPrepareInput{
|
||||
Chat: created.Chat,
|
||||
Messages: created.InitialMessages,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(prepared.Cleanup)
|
||||
|
||||
// These take an empty struct, so they carry no keys to validate.
|
||||
noInput := map[string]bool{
|
||||
"process_list": true,
|
||||
"stop_workspace": true,
|
||||
"list_subagent_models": true,
|
||||
}
|
||||
var unvalidated []string
|
||||
require.NotEmpty(t, prepared.BuiltinToolNames)
|
||||
for _, tool := range prepared.Tools {
|
||||
info := tool.Info()
|
||||
if !prepared.BuiltinToolNames[info.Name] || len(info.Parameters) > 0 || noInput[info.Name] {
|
||||
continue
|
||||
}
|
||||
unvalidated = append(unvalidated, info.Name)
|
||||
}
|
||||
require.Empty(t, unvalidated,
|
||||
"these builtin tools declare no schema properties, so their input is not validated")
|
||||
}
|
||||
|
||||
func fetchToolStub() fantasy.AgentTool {
|
||||
type input struct {
|
||||
URL string `json:"url"`
|
||||
}
|
||||
return fantasy.NewAgentTool("fetch", "",
|
||||
func(context.Context, input, fantasy.ToolCall) (fantasy.ToolResponse, error) {
|
||||
return fantasy.ToolResponse{}, nil
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,117 @@
|
||||
// Package toolschema rejects tool inputs whose object keys the Go decoder
|
||||
// and a case-sensitive reader resolve differently.
|
||||
package toolschema
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"maps"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/xerrors"
|
||||
)
|
||||
|
||||
// freeFormPropertyName is the property name fantasy generates for
|
||||
// map[string]T inputs. Its keys are data, so they are checked against the
|
||||
// value schema behind this name rather than against a fixed property set.
|
||||
const freeFormPropertyName = "*"
|
||||
|
||||
// ValidateUnambiguous reports an error when input holds an object key that
|
||||
// encoding/json folds into a declared property but a case-sensitive reader
|
||||
// treats as distinct, or when one object repeats a key. Either lets code
|
||||
// inspecting the raw input read one value while the tool executes another.
|
||||
//
|
||||
// properties is a fantasy ToolInfo.Parameters map, keyed by property name.
|
||||
// Keys matching no property are ignored because a generated struct decoder
|
||||
// drops them. A tool with a hand-written decoder that reads undeclared keys
|
||||
// has to reject ambiguous spellings of those keys itself.
|
||||
func ValidateUnambiguous(properties map[string]any, input []byte) error {
|
||||
decoder := json.NewDecoder(bytes.NewReader(input))
|
||||
token, err := decoder.Token()
|
||||
// Input that does not parse here does not decode for the tool either,
|
||||
// so its own decode reports the failure.
|
||||
if err != nil || token != json.Delim('{') {
|
||||
return nil
|
||||
}
|
||||
return validateObject(properties, decoder, "")
|
||||
}
|
||||
|
||||
func validateObject(properties map[string]any, decoder *json.Decoder, parent string) error {
|
||||
seen := make(map[string]struct{})
|
||||
for {
|
||||
token, err := decoder.Token()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
if token == json.Delim('}') {
|
||||
return nil
|
||||
}
|
||||
key, ok := token.(string)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
path := key
|
||||
if parent != "" {
|
||||
path = parent + "." + key
|
||||
}
|
||||
if _, duplicate := seen[key]; duplicate {
|
||||
return xerrors.Errorf("input repeats the key %q", path)
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
if err := checkKeyCase(properties, key, path); err != nil {
|
||||
return err
|
||||
}
|
||||
value, err := decoder.Token()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
if err := validateValue(childSchema(properties, key), decoder, value, path); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func validateValue(schema map[string]any, decoder *json.Decoder, token json.Token, path string) error {
|
||||
switch token {
|
||||
case json.Delim('{'):
|
||||
properties, _ := schema["properties"].(map[string]any)
|
||||
return validateObject(properties, decoder, path)
|
||||
case json.Delim('['):
|
||||
items, _ := schema["items"].(map[string]any)
|
||||
for {
|
||||
next, err := decoder.Token()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
if next == json.Delim(']') {
|
||||
return nil
|
||||
}
|
||||
if err := validateValue(items, decoder, next, path+"[]"); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func checkKeyCase(properties map[string]any, key, path string) error {
|
||||
if _, exact := properties[key]; exact {
|
||||
return nil
|
||||
}
|
||||
for _, name := range slices.Sorted(maps.Keys(properties)) {
|
||||
if strings.EqualFold(name, key) {
|
||||
return xerrors.Errorf("input key %q differs from schema property %q only by case", path, name)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func childSchema(properties map[string]any, key string) map[string]any {
|
||||
name := key
|
||||
if _, freeForm := properties[freeFormPropertyName]; freeForm {
|
||||
name = freeFormPropertyName
|
||||
}
|
||||
schema, _ := properties[name].(map[string]any)
|
||||
return schema
|
||||
}
|
||||
@@ -0,0 +1,140 @@
|
||||
package toolschema_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/toolschema"
|
||||
)
|
||||
|
||||
func TestValidateUnambiguous(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
readFile := chattool.ReadFile(chattool.ReadFileOptions{})
|
||||
editFiles := chattool.EditFiles(chattool.EditFilesOptions{})
|
||||
createWorkspace := chattool.CreateWorkspace(nil, uuid.Nil, uuid.Nil, chattool.CreateWorkspaceOptions{})
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
tool fantasy.AgentTool
|
||||
input string
|
||||
wantErr string
|
||||
}{
|
||||
{
|
||||
name: "case variant beside the canonical key",
|
||||
tool: readFile,
|
||||
input: `{"path":"/allowed","PATH":"/secret"}`,
|
||||
wantErr: `input key "PATH" differs from schema property "path" only by case`,
|
||||
},
|
||||
{
|
||||
name: "case variant alone",
|
||||
tool: readFile,
|
||||
input: `{"PATH":"/secret"}`,
|
||||
wantErr: `input key "PATH" differs from schema property "path" only by case`,
|
||||
},
|
||||
{
|
||||
name: "repeated key",
|
||||
tool: readFile,
|
||||
input: `{"path":"/allowed","path":"/secret"}`,
|
||||
wantErr: `input repeats the key "path"`,
|
||||
},
|
||||
{
|
||||
name: "case variant nested in an array element",
|
||||
tool: editFiles,
|
||||
input: `{"files":[{"path":"a","PATH":"b","edits":[{"old_text":"x","new_text":"y"}]}]}`,
|
||||
wantErr: `input key "files[].PATH" differs from schema property "path" only by case`,
|
||||
},
|
||||
{
|
||||
name: "case variant nested in an array element object",
|
||||
tool: editFiles,
|
||||
input: `{"files":[{"path":"a","edits":[{"old_text":"x","NEW_TEXT":"y"}]}]}`,
|
||||
wantErr: `input key "files[].edits[].NEW_TEXT" differs from schema property "new_text" only by case`,
|
||||
},
|
||||
{
|
||||
name: "free-form map keys differing by case",
|
||||
tool: createWorkspace,
|
||||
input: `{"template_id":"t","parameters":{"foo":"1","FOO":"2"}}`,
|
||||
},
|
||||
{
|
||||
name: "free-form map repeating a key",
|
||||
tool: createWorkspace,
|
||||
input: `{"parameters":{"foo":"1","foo":"2"}}`,
|
||||
wantErr: `input repeats the key "parameters.foo"`,
|
||||
},
|
||||
{
|
||||
name: "case variant inside a free-form map value",
|
||||
tool: freeFormValueTool(),
|
||||
input: `{"targets":{"first":{"PATH":"/secret"}}}`,
|
||||
wantErr: `input key "targets.first.PATH" differs from schema property "path" only by case`,
|
||||
},
|
||||
{
|
||||
name: "key matching no property",
|
||||
tool: readFile,
|
||||
input: `{"path":"/allowed","xyzzy":"b"}`,
|
||||
},
|
||||
{
|
||||
name: "exact keys",
|
||||
tool: editFiles,
|
||||
input: `{"files":[{"path":"a","edits":[{"old_text":"x","new_text":"y","replace_all":true}]}]}`,
|
||||
},
|
||||
{
|
||||
name: "input the tool cannot decode either",
|
||||
tool: readFile,
|
||||
input: `{"path":`,
|
||||
},
|
||||
{
|
||||
name: "empty input",
|
||||
tool: readFile,
|
||||
input: ``,
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
err := toolschema.ValidateUnambiguous(tt.tool.Info().Parameters, []byte(tt.input))
|
||||
if tt.wantErr == "" {
|
||||
require.NoError(t, err)
|
||||
return
|
||||
}
|
||||
require.ErrorContains(t, err, tt.wantErr)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestCaseVariantKeyDecodesCaseInsensitively pins the decoder behavior the
|
||||
// validator exists for: a case-sensitive reader of these bytes finds no
|
||||
// "path" key in the first case and "/allowed" in the second, while the
|
||||
// tool's own arguments resolve to "/secret" in both.
|
||||
func TestCaseVariantKeyDecodesCaseInsensitively(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var args chattool.ReadFileArgs
|
||||
require.NoError(t, json.Unmarshal([]byte(`{"PATH":"/secret"}`), &args))
|
||||
require.Equal(t, "/secret", args.Path)
|
||||
|
||||
require.NoError(t, json.Unmarshal([]byte(`{"path":"/allowed","PATH":"/secret"}`), &args))
|
||||
require.Equal(t, "/secret", args.Path)
|
||||
}
|
||||
|
||||
// freeFormValueTool builds a tool whose input nests a fixed property set
|
||||
// inside a free-form map, so the value schema fantasy renders under "*" has
|
||||
// to be carried into the map's values.
|
||||
func freeFormValueTool() fantasy.AgentTool {
|
||||
type target struct {
|
||||
Path string `json:"path"`
|
||||
}
|
||||
type input struct {
|
||||
Targets map[string]target `json:"targets"`
|
||||
}
|
||||
return fantasy.NewAgentTool("free_form_value", "",
|
||||
func(context.Context, input, fantasy.ToolCall) (fantasy.ToolResponse, error) {
|
||||
return fantasy.ToolResponse{}, nil
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user