feat: add chat lifecycle hook dispatch backend (#27401)

Adds the chat lifecycle hook wire contract and dispatch plumbing, first
PR of the lifecycle hooks stack (followed by #27428, #27429, #27430).

- `codersdk/x/agenthooks`: event and response wire types, JWT creation
and verification with the shared secret (HS256, request body digest,
expiry and not-before freshness checks), and an HTTP handler helper so
consumers only implement the events they use. The `codersdk/x` location
marks the consumer SDK as experimental.
- `coderd/x/agenthooks/dispatch`: a stateless dispatcher that signs and
posts hook events, enforces a concurrency cap under one configured
timeout that bounds both the capacity wait and both post attempts,
retries one connection failure with the same JWT, sends a distinctive
`coderd-agenthooks/<version>` User-Agent, and records Prometheus
metrics. Delivery is at least once; consumers own durable decision
state, audit records, and deduplication keyed by the stable payload
identifiers. Nothing is persisted by Coder.
- Response bodies decode strictly: unknown fields, duplicate JSON keys
(including inside `input_override`), and trailing data fail the dispatch
closed as protocol errors instead of silently reading as allow.
- `coderd/util/xnet`: shared timeout and connection error classification
used by the dispatcher retry logic. Transient HTTP/2 stream aborts count
as connection errors, so the documented single retry also applies to h2
consumers, which is the shape Go's default transport negotiates against
any TLS consumer. Deterministic protocol failures stay terminal. Only
the struct form of a stream error is matched, because `net/http` bundles
its own HTTP/2 types and `h2_error.go` bridges only that shape.
- `scripts/agenthooks-server`: a reference consumer that logs events and
demonstrates consumer-owned pre-tool decision deduplication. It requires
an explicitly configured JWT audience rather than deriving one from the
request, and its startup output names the mode it is running in so an
operator can see that the example policy flags need `-log-only=false`.
- `scripts/apitypings`: generate TypeScript types for the hook wire
contract.

Dispatch failures log without the error's stack frames, since a failed
dispatch is an expected, operator-visible condition.

Nothing dispatches these events yet; chatd wiring lands in #27429.

> This PR was written by Mux, an AI coding agent, on Mike's behalf.
This commit is contained in:
Michael Suchacz
2026-07-28 13:59:37 +02:00
committed by GitHub
parent bd5d640f1e
commit 8ea2586189
13 changed files with 3099 additions and 0 deletions
+62
View File
@@ -0,0 +1,62 @@
// Package xnet classifies network transport errors.
package xnet
import (
"context"
"errors"
"io"
"net"
"syscall"
"golang.org/x/net/http2"
)
// IsTimeoutError reports whether err indicates the peer did not respond in
// time, including context deadlines and net.Error timeouts.
func IsTimeoutError(err error) bool {
if err == nil {
return false
}
if errors.Is(err, context.DeadlineExceeded) || errors.Is(err, syscall.ETIMEDOUT) {
return true
}
var netErr net.Error
return errors.As(err, &netErr) && netErr.Timeout()
}
// IsConnectionError reports whether err indicates a failed or interrupted
// connection, such as refused dials, resets, and unexpected EOFs.
func IsConnectionError(err error) bool {
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) ||
errors.Is(err, net.ErrClosed) || errors.Is(err, syscall.ECONNRESET) ||
errors.Is(err, syscall.ECONNREFUSED) || errors.Is(err, syscall.EPIPE) {
return true
}
var opErr *net.OpError
if errors.As(err, &opErr) {
return true
}
// net/http bundles its own HTTP/2 implementation with unexported error
// types, and net/http/h2_error.go bridges only the struct form of a
// stream error to this package's type. A pointer target never matches,
// and GOAWAY errors have no such bridge.
var streamErr http2.StreamError
return errors.As(err, &streamErr) && isTransientHTTP2Error(streamErr.Code)
}
// isTransientHTTP2Error reports whether an HTTP/2 error code can result from a
// transient peer or transport condition. Deterministic protocol failures stay
// terminal so a malformed consumer response is not retried.
func isTransientHTTP2Error(code http2.ErrCode) bool {
switch code {
case http2.ErrCodeNo,
http2.ErrCodeInternal,
http2.ErrCodeRefusedStream,
http2.ErrCodeCancel,
http2.ErrCodeEnhanceYourCalm:
return true
default:
return false
}
}
+89
View File
@@ -0,0 +1,89 @@
package xnet_test
import (
"context"
"io"
"net"
"net/http"
"net/http/httptest"
"syscall"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/net/http2"
"golang.org/x/xerrors"
"github.com/coder/coder/v2/coderd/util/xnet"
)
func TestIsTimeoutError(t *testing.T) {
t.Parallel()
require.False(t, xnet.IsTimeoutError(nil))
require.False(t, xnet.IsTimeoutError(xerrors.New("other")))
require.True(t, xnet.IsTimeoutError(context.DeadlineExceeded))
require.True(t, xnet.IsTimeoutError(xerrors.Errorf("dial: %w", syscall.ETIMEDOUT)))
require.True(t, xnet.IsTimeoutError(&net.OpError{Op: "dial", Err: syscall.ETIMEDOUT}))
}
func TestIsConnectionError(t *testing.T) {
t.Parallel()
require.False(t, xnet.IsConnectionError(nil))
require.False(t, xnet.IsConnectionError(context.DeadlineExceeded))
require.True(t, xnet.IsConnectionError(io.EOF))
require.True(t, xnet.IsConnectionError(io.ErrUnexpectedEOF))
require.True(t, xnet.IsConnectionError(xerrors.Errorf("write: %w", syscall.EPIPE)))
require.True(t, xnet.IsConnectionError(&net.OpError{Op: "read", Err: syscall.ECONNRESET}))
for _, err := range []error{
http2.StreamError{Code: http2.ErrCodeNo},
http2.StreamError{Code: http2.ErrCodeInternal},
http2.StreamError{Code: http2.ErrCodeRefusedStream},
http2.StreamError{Code: http2.ErrCodeCancel},
http2.StreamError{Code: http2.ErrCodeEnhanceYourCalm},
xerrors.Errorf("read response: %w", http2.StreamError{Code: http2.ErrCodeInternal}),
} {
require.True(t, xnet.IsConnectionError(err), err)
}
require.False(t, xnet.IsConnectionError(http2.StreamError{Code: http2.ErrCodeProtocol}))
require.False(t, xnet.IsConnectionError(http2.StreamError{Code: http2.ErrCodeFlowControl}))
}
func TestIsConnectionErrorHTTPResponseAbort(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
enableHTTP2 bool
proto string
}{
{name: "http1", proto: "HTTP/1.1"},
{name: "http2", enableHTTP2: true, proto: "HTTP/2.0"},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, err := w.Write([]byte(`{"partial":`))
assert.NoError(t, err)
assert.NoError(t, http.NewResponseController(w).Flush())
panic(http.ErrAbortHandler)
}))
server.EnableHTTP2 = tc.enableHTTP2
server.StartTLS()
t.Cleanup(server.Close)
request, err := http.NewRequestWithContext(t.Context(), http.MethodGet, server.URL, nil)
require.NoError(t, err)
response, err := server.Client().Do(request)
require.NoError(t, err)
require.Equal(t, tc.proto, response.Proto)
_, err = io.ReadAll(response.Body)
require.Error(t, err)
require.NoError(t, response.Body.Close())
require.True(t, xnet.IsConnectionError(err), err)
})
}
}
+712
View File
@@ -0,0 +1,712 @@
// Package dispatch delivers chat lifecycle events to an external webhook.
package dispatch
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"io"
"net"
"net/http"
"net/url"
"time"
"unicode"
"unicode/utf8"
"github.com/google/uuid"
"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/promauto"
"golang.org/x/xerrors"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/coderd/util/xnet"
"github.com/coder/coder/v2/codersdk/x/agenthooks"
)
const (
maxConcurrentDispatches = 256
maxResponseBodyBytes = 1_048_576
maxModelContextBytes = 16_384
capacityWaitLimit = 250 * time.Millisecond
retryBackoff = 250 * time.Millisecond
clockSkewLeeway = 30 * time.Second
)
// Result classifies the terminal outcome of a dispatch attempt.
type Result string
const (
ResultOK Result = "ok"
ResultDenied Result = "denied"
ResultHTTPError Result = "http_error"
ResultProtocolError Result = "protocol_error"
ResultTimeout Result = "timeout"
ResultConnectionError Result = "connection_error"
ResultOverCapacity Result = "over_capacity"
ResultInternalError Result = "internal_error"
)
// Event carries the identities delivered with each dispatch attempt.
type Event struct {
Type agenthooks.EventType
agenthooks.ChatRef
Data any
}
// Error preserves the attempt ID and failure class.
type Error struct {
Class Result
DispatchID uuid.UUID
Err error
}
func (e *Error) Error() string {
return e.Err.Error()
}
func (e *Error) Unwrap() error {
return e.Err
}
func newError(class Result, dispatchID uuid.UUID, err error) error {
if err == nil {
return nil
}
return &Error{Class: class, DispatchID: dispatchID, Err: err}
}
// Dispatcher delivers lifecycle hook attempts. It keeps no delivery or
// decision state, so delivery is best-effort and consumers own both.
type Dispatcher struct {
logger slog.Logger
client *http.Client
hookURL string
hookURLErr error
secret []byte
timeout time.Duration
deploymentID string
userAgent string
semaphore chan struct{}
metrics *metrics
}
// validateHookURL requires HTTPS because hook traffic carries sensitive data
// and authorization tokens, and responses can control execution. Plain HTTP
// is allowed only for loopback development consumers.
func validateHookURL(raw string) error {
if raw == "" {
return nil
}
parsed, err := url.Parse(raw)
if err != nil {
return xerrors.Errorf("parse hook URL: %w", err)
}
// The raw URL is signed verbatim as the JWT audience, and neither
// component is ever transmitted, so a consumer configured with the URL
// it actually serves would never match.
if parsed.Fragment != "" {
return xerrors.New("chat hook URL must not contain a fragment")
}
if parsed.User != nil {
return xerrors.New("chat hook URL must not contain userinfo")
}
switch parsed.Scheme {
case "https":
if parsed.Hostname() == "" {
return xerrors.New("chat hook URL must include a host")
}
return nil
case "http":
host := parsed.Hostname()
if host == "" {
return xerrors.New("chat hook URL must include a host")
}
if host == "localhost" {
return nil
}
if ip := net.ParseIP(host); ip != nil && ip.IsLoopback() {
return nil
}
return xerrors.New("chat hook URL must use HTTPS; plain HTTP is allowed only for loopback addresses")
default:
return xerrors.Errorf("chat hook URL scheme %q is not supported", parsed.Scheme)
}
}
// New copies (or creates) the HTTP client and disables redirects for signed
// requests.
func New(
logger slog.Logger,
client *http.Client,
hookURL string,
secret string,
timeout time.Duration,
deploymentID string,
coderVersion string,
reg prometheus.Registerer,
) *Dispatcher {
if client == nil {
client = &http.Client{}
} else {
clientCopy := *client
client = &clientCopy
}
client.CheckRedirect = func(_ *http.Request, _ []*http.Request) error {
return http.ErrUseLastResponse
}
return &Dispatcher{
logger: logger.Named("chat_hook_dispatcher"),
client: client,
hookURL: hookURL,
hookURLErr: validateHookURL(hookURL),
secret: []byte(secret),
timeout: timeout,
deploymentID: deploymentID,
userAgent: "coderd-agenthooks/" + coderVersion,
semaphore: make(chan struct{}, maxConcurrentDispatches),
metrics: newMetrics(reg),
}
}
// Enabled reports whether a hook URL is configured. A nil dispatcher reads
// as disabled so callers can hold one unconditionally.
func (d *Dispatcher) Enabled() bool {
return d != nil && d.hookURL != ""
}
// Dispatch delivers one event. The returned ID correlates the attempt in logs
// and Error values; the dispatcher does not persist delivery state.
func (d *Dispatcher) Dispatch(ctx context.Context, event Event) (agenthooks.Response, uuid.UUID, error) {
if !d.Enabled() {
return agenthooks.Response{}, uuid.Nil, xerrors.New("chat hook dispatcher is not enabled")
}
if d.hookURLErr != nil {
return agenthooks.Response{}, uuid.Nil, xerrors.Errorf("chat hook URL rejected: %w", d.hookURLErr)
}
startedAt := time.Now()
dispatchID := uuid.New()
wait := max(min(d.timeout, capacityWaitLimit), 0)
capacityTimer := time.NewTimer(wait)
defer capacityTimer.Stop()
// The capacity wait runs against its own timer rather than a dispatch
// deadline, so a timeout shorter than capacityWaitLimit cannot make the
// over-capacity and caller-cancellation cases race.
select {
case d.semaphore <- struct{}{}:
defer func() { <-d.semaphore }()
case <-ctx.Done():
outcome := dispatchOutcome{result: ResultTimeout, err: ctx.Err()}
return agenthooks.Response{}, dispatchID, d.finish(ctx, event, dispatchID, startedAt, outcome)
case <-capacityTimer.C:
outcome := dispatchOutcome{result: ResultOverCapacity, err: context.DeadlineExceeded}
return agenthooks.Response{}, dispatchID, d.finish(ctx, event, dispatchID, startedAt, outcome)
}
// Both post attempts share whatever remains of the configured timeout so
// that waiting for capacity cannot extend the dispatch past it.
ctx, cancel := context.WithTimeout(ctx, d.timeout-time.Since(startedAt))
defer cancel()
outcome := d.prepareAndPost(ctx, event, dispatchID)
if err := d.finish(ctx, event, dispatchID, startedAt, outcome); err != nil {
return agenthooks.Response{}, dispatchID, err
}
return outcome.response, dispatchID, nil
}
func (d *Dispatcher) finish(
ctx context.Context,
event Event,
dispatchID uuid.UUID,
startedAt time.Time,
outcome dispatchOutcome,
) error {
if outcome.err != nil {
d.logger.Warn(context.WithoutCancel(ctx), "chat hook dispatch failed",
slog.F("dispatch_id", dispatchID),
slog.F("event", event.Type),
slog.F("result", outcome.result),
slog.F("error", outcome.err.Error()),
)
} else {
d.logger.Debug(context.WithoutCancel(ctx), "chat hook dispatched",
slog.F("dispatch_id", dispatchID),
slog.F("event", event.Type),
slog.F("duration", time.Since(startedAt)),
)
}
d.metrics.observe(event.Type, outcome.result, outcome.response, time.Since(startedAt))
return newError(outcome.result, dispatchID, outcome.err)
}
type dispatchOutcome struct {
result Result
response agenthooks.Response
err error
}
func (d *Dispatcher) prepareAndPost(ctx context.Context, event Event, dispatchID uuid.UUID) dispatchOutcome {
data, err := marshalEventData(event)
if err != nil {
return dispatchOutcome{result: ResultProtocolError, err: err}
}
request := agenthooks.Request{
Type: event.Type,
Meta: agenthooks.Meta{
DispatchID: dispatchID,
SchemaVersion: agenthooks.SchemaVersion,
ChatRef: event.ChatRef,
},
Data: data,
}
body, err := json.Marshal(request)
if err != nil {
return dispatchOutcome{result: ResultProtocolError, err: xerrors.Errorf("marshal request: %w", err)}
}
digest := sha256.Sum256(body)
now := time.Now()
token, err := agenthooks.SignClaims(d.secret, agenthooks.Claims{
Issuer: d.deploymentID,
Subject: "coder:chat:" + event.ChatID.String(),
Audience: d.hookURL,
IssuedAt: now.Unix(),
NotBefore: now.Add(-clockSkewLeeway).Unix(),
Expires: now.Add(d.timeout + clockSkewLeeway).Unix(),
JTI: dispatchID,
Type: event.Type,
BodySHA256: hex.EncodeToString(digest[:]),
})
if err != nil {
return dispatchOutcome{result: ResultProtocolError, err: xerrors.Errorf("sign request: %w", err)}
}
response, result, err := d.post(ctx, body, token)
outcome := dispatchOutcome{
result: result,
response: response,
err: err,
}
if err != nil {
return outcome
}
if err := validateResponse(event.Type, response); err != nil {
// Drop the rejected response so its decision, override, and context
// values are not observed as if they had been applied.
outcome.response = agenthooks.Response{}
outcome.result = ResultProtocolError
outcome.err = err
return outcome
}
if response.Permission != nil && response.Permission.Decision == agenthooks.PermissionDeny {
outcome.result = ResultDenied
}
return outcome
}
func (d *Dispatcher) post(
ctx context.Context,
body []byte,
token string,
) (response agenthooks.Response, result Result, err error) {
// The deadline set in Dispatch bounds both attempts so a retry cannot
// extend the dispatch past the configured timeout or the JWT lifetime.
for attempt := range 2 {
req, reqErr := http.NewRequestWithContext(ctx, http.MethodPost, d.hookURL, bytes.NewReader(body))
if reqErr != nil {
return agenthooks.Response{}, ResultProtocolError, xerrors.Errorf("create request: %w", reqErr)
}
req.Header.Set("Authorization", "Bearer "+token)
req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", d.userAgent)
httpResponse, requestErr := d.client.Do(req)
if requestErr != nil {
attemptErr := ctx.Err()
if xnet.IsTimeoutError(attemptErr) || xnet.IsTimeoutError(requestErr) || errors.Is(requestErr, context.Canceled) {
return agenthooks.Response{}, ResultTimeout, xerrors.Errorf("post lifecycle hook: %w", requestErr)
}
if !xnet.IsConnectionError(requestErr) {
return agenthooks.Response{}, ResultProtocolError, xerrors.Errorf("post lifecycle hook: %w", requestErr)
}
if attempt == 1 {
return agenthooks.Response{}, ResultConnectionError, xerrors.Errorf("post lifecycle hook: %w", requestErr)
}
backoff := time.NewTimer(retryBackoff)
select {
case <-ctx.Done():
backoff.Stop()
return agenthooks.Response{}, ResultTimeout, xerrors.Errorf("post lifecycle hook: %w", ctx.Err())
case <-backoff.C:
}
continue
}
if httpResponse.StatusCode < http.StatusOK || httpResponse.StatusCode >= http.StatusMultipleChoices {
_ = httpResponse.Body.Close()
return agenthooks.Response{}, ResultHTTPError, xerrors.Errorf("lifecycle hook returned HTTP status %d", httpResponse.StatusCode)
}
responseBody, readErr := io.ReadAll(io.LimitReader(httpResponse.Body, maxResponseBodyBytes+1))
attemptErr := ctx.Err()
_ = httpResponse.Body.Close()
if readErr != nil {
switch {
case xnet.IsTimeoutError(attemptErr), xnet.IsTimeoutError(readErr), errors.Is(readErr, context.Canceled):
return agenthooks.Response{}, ResultTimeout, xerrors.Errorf("read lifecycle hook response: %w", readErr)
case xnet.IsConnectionError(readErr):
if attempt == 1 {
return agenthooks.Response{}, ResultConnectionError, xerrors.Errorf("read lifecycle hook response: %w", readErr)
}
default:
return agenthooks.Response{}, ResultProtocolError, xerrors.Errorf("read lifecycle hook response: %w", readErr)
}
// Mid-body connection drops get the same single retry as dial
// failures, reusing the dispatch ID.
backoff := time.NewTimer(retryBackoff)
select {
case <-ctx.Done():
backoff.Stop()
return agenthooks.Response{}, ResultTimeout, xerrors.Errorf("read lifecycle hook response: %w", ctx.Err())
case <-backoff.C:
}
continue
}
if len(responseBody) > maxResponseBodyBytes {
return agenthooks.Response{}, ResultProtocolError, xerrors.New("lifecycle hook response exceeds 1 MiB")
}
trimmed := bytes.TrimSpace(responseBody)
if len(trimmed) == 0 {
return agenthooks.Response{}, ResultOK, nil
}
if bytes.Equal(trimmed, []byte("null")) {
return agenthooks.Response{}, ResultProtocolError, xerrors.New("lifecycle hook response must be a JSON object")
}
if err := decodeResponse(trimmed, &response); err != nil {
return agenthooks.Response{}, ResultProtocolError, xerrors.Errorf("decode lifecycle hook response: %w", err)
}
return response, ResultOK, nil
}
panic("unreachable")
}
// decodeResponse rejects response bodies that plain unmarshaling would
// silently misread as allow: unknown fields (misspelled keys), duplicate
// object keys (Go keeps the last value), and trailing JSON values.
func decodeResponse(trimmed []byte, response *agenthooks.Response) error {
if err := rejectDuplicateKeys(json.NewDecoder(bytes.NewReader(trimmed)), 0, keyMatchingFolded); err != nil {
return err
}
decoder := json.NewDecoder(bytes.NewReader(trimmed))
decoder.DisallowUnknownFields()
if err := decoder.Decode(response); err != nil {
return err
}
if err := decoder.Decode(&struct{}{}); err != io.EOF {
return xerrors.New("response must contain one JSON object")
}
return nil
}
const maxResponseJSONDepth = 128
// foldJSONName canonicalizes an object key the way encoding/json matches
// struct fields: ASCII is case-folded, and other runes collapse to the
// smallest member of their Unicode simple-fold orbit. strings.ToLower is
// not equivalent, so "permi\u017f\u017fion" would otherwise slip past the
// duplicate check and overwrite "permission".
func foldJSONName(name string) string {
folded := make([]byte, 0, len(name))
for _, r := range name {
if r < utf8.RuneSelf {
if 'a' <= r && r <= 'z' {
r -= 'a' - 'A'
}
folded = append(folded, byte(r))
continue
}
for {
next := unicode.SimpleFold(r)
if next <= r {
r = next
break
}
r = next
}
folded = utf8.AppendRune(folded, r)
}
return string(folded)
}
// foldedInputOverride is the canonical form of the one envelope field whose
// contents are opaque pass-through data.
var foldedInputOverride = foldJSONName("input_override")
// keyMatching selects how object keys are compared for duplicates.
type keyMatching int
const (
// keyMatchingFolded mirrors encoding/json struct-field matching, used
// for the typed response envelope.
keyMatchingFolded keyMatching = iota
// keyMatchingExact is used inside input_override, which is opaque data
// for a case-sensitive tool schema.
keyMatchingExact
)
// rejectDuplicateKeys consumes one JSON value and fails on duplicate object
// keys at any depth, including inside input_override, because a duplicated
// key such as {"permission":{"decision":"deny"},"permission":null} would
// otherwise drop the decision the consumer intended. Envelope keys also
// collide when they differ only by folding, matching how encoding/json
// resolves struct fields, so "Permission" cannot override "permission".
//
// CLEANUP: json/v2 removes the need for this. It matches names case-sensitively
// and RejectDuplicateNames rejects exact duplicates.
func rejectDuplicateKeys(decoder *json.Decoder, depth int, matching keyMatching) error {
if depth > maxResponseJSONDepth {
return xerrors.New("response JSON exceeds supported nesting depth")
}
token, err := decoder.Token()
if err != nil {
return err
}
delim, ok := token.(json.Delim)
if !ok {
return nil
}
switch delim {
case '{':
seen := make(map[string]struct{})
for decoder.More() {
keyToken, err := decoder.Token()
if err != nil {
return err
}
key, _ := keyToken.(string)
canonical := key
if matching == keyMatchingFolded {
canonical = foldJSONName(key)
}
if _, dup := seen[canonical]; dup {
return xerrors.Errorf("duplicate key %q", key)
}
seen[canonical] = struct{}{}
nested := matching
// The decoder binds this field case-insensitively, so the
// boundary must be detected the same way.
if matching == keyMatchingFolded && canonical == foldedInputOverride {
nested = keyMatchingExact
}
if err := rejectDuplicateKeys(decoder, depth+1, nested); err != nil {
return err
}
}
_, err = decoder.Token()
return err
case '[':
for decoder.More() {
if err := rejectDuplicateKeys(decoder, depth+1, matching); err != nil {
return err
}
}
_, err = decoder.Token()
return err
default:
return nil
}
}
func validateResponse(eventType agenthooks.EventType, response agenthooks.Response) error {
if len(response.ModelContext) > maxModelContextBytes {
return xerrors.New("model_context exceeds 16 KiB")
}
if response.Permission == nil {
return nil
}
if eventType != agenthooks.EventUserPromptSubmit && eventType != agenthooks.EventPreToolUse {
return xerrors.Errorf("permission is not valid for event %q", eventType)
}
switch response.Permission.Decision {
case agenthooks.PermissionAllow:
inputOverride := bytes.TrimSpace(response.Permission.InputOverride)
if len(inputOverride) == 0 || bytes.Equal(inputOverride, []byte("null")) {
return xerrors.New("allow decision requires input_override")
}
if eventType == agenthooks.EventUserPromptSubmit {
if err := validateUserPromptSubmitOverride(inputOverride); err != nil {
return err
}
}
case agenthooks.PermissionDeny:
// Denied input does not proceed, so reject overrides to surface consumer bugs.
inputOverride := bytes.TrimSpace(response.Permission.InputOverride)
if len(inputOverride) > 0 && !bytes.Equal(inputOverride, []byte("null")) {
return xerrors.New("deny decision must not include input_override")
}
default:
return xerrors.Errorf("invalid permission decision %q", response.Permission.Decision)
}
return nil
}
func validateUserPromptSubmitOverride(input json.RawMessage) error {
// This override is decoded into a struct rather than passed through, so
// it needs the same folded duplicate check as the response envelope.
if err := rejectDuplicateKeys(json.NewDecoder(bytes.NewReader(input)), 0, keyMatchingFolded); err != nil {
return xerrors.Errorf("user_prompt_submit input_override: %w", err)
}
var override struct {
Prompt *string `json:"prompt"`
}
decoder := json.NewDecoder(bytes.NewReader(input))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&override); err != nil {
return xerrors.Errorf("user_prompt_submit input_override must be {\"prompt\": string}: %w", err)
}
if override.Prompt == nil {
return xerrors.New("user_prompt_submit input_override must be {\"prompt\": string}")
}
if err := decoder.Decode(&struct{}{}); err != io.EOF {
return xerrors.New("user_prompt_submit input_override must contain one JSON object")
}
return nil
}
func marshalEventData(event Event) (json.RawMessage, error) {
switch event.Type {
case agenthooks.EventSessionStart:
if !isData[agenthooks.SessionStartData](event.Data) {
return nil, xerrors.New("session_start data has the wrong type")
}
case agenthooks.EventUserPromptSubmit:
if !isData[agenthooks.UserPromptSubmitData](event.Data) {
return nil, xerrors.New("user_prompt_submit data has the wrong type")
}
case agenthooks.EventPreToolUse:
value, ok := dataValue[agenthooks.PreToolUseData](event.Data)
if !ok {
return nil, xerrors.New("pre_tool_use data has the wrong type")
}
if value.ToolUseID == "" || value.ToolName == "" {
return nil, xerrors.New("pre_tool_use data requires tool_use_id and tool_name")
}
case agenthooks.EventPostToolUse:
value, ok := dataValue[agenthooks.PostToolUseData](event.Data)
if !ok {
return nil, xerrors.New("post_tool_use data has the wrong type")
}
if value.ToolUseID == "" || value.ToolName == "" {
return nil, xerrors.New("post_tool_use data requires tool_use_id and tool_name")
}
case agenthooks.EventPreCompact:
if !isData[agenthooks.PreCompactData](event.Data) {
return nil, xerrors.New("pre_compact data has the wrong type")
}
case agenthooks.EventPostCompact:
if !isData[agenthooks.PostCompactData](event.Data) {
return nil, xerrors.New("post_compact data has the wrong type")
}
case agenthooks.EventStop:
if !isData[agenthooks.StopData](event.Data) {
return nil, xerrors.New("stop data has the wrong type")
}
default:
return nil, xerrors.Errorf("unknown event type %q", event.Type)
}
encoded, marshalErr := json.Marshal(event.Data)
if marshalErr != nil {
return nil, xerrors.Errorf("marshal event data: %w", marshalErr)
}
return encoded, nil
}
func isData[T any](value any) bool {
_, ok := dataValue[T](value)
return ok
}
func dataValue[T any](value any) (T, bool) {
if typed, ok := value.(T); ok {
return typed, true
}
if typed, ok := value.(*T); ok && typed != nil {
return *typed, true
}
var zero T
return zero, false
}
type metrics struct {
dispatches *prometheus.CounterVec
duration *prometheus.HistogramVec
decisions *prometheus.CounterVec
contextSize *prometheus.HistogramVec
inputOverrides *prometheus.CounterVec
}
func newMetrics(reg prometheus.Registerer) *metrics {
if reg == nil {
reg = prometheus.NewRegistry()
}
factory := promauto.With(reg)
return &metrics{
dispatches: factory.NewCounterVec(prometheus.CounterOpts{
Namespace: "coderd",
Subsystem: "chatd",
Name: "hook_dispatches_total",
Help: "Total lifecycle hook dispatches by event and result.",
}, []string{"event", "result"}),
duration: factory.NewHistogramVec(prometheus.HistogramOpts{
Namespace: "coderd",
Subsystem: "chatd",
Name: "hook_dispatch_seconds",
Help: "Lifecycle hook dispatch duration in seconds.",
}, []string{"event"}),
decisions: factory.NewCounterVec(prometheus.CounterOpts{
Namespace: "coderd",
Subsystem: "chatd",
Name: "hook_decisions_total",
Help: "Total lifecycle hook permission decisions by event and decision.",
}, []string{"event", "decision"}),
contextSize: factory.NewHistogramVec(prometheus.HistogramOpts{
Namespace: "coderd",
Subsystem: "chatd",
Name: "hook_context_size_bytes",
Help: "Lifecycle hook model context response size in bytes.",
Buckets: prometheus.ExponentialBuckets(64, 2, 10),
}, []string{"event"}),
inputOverrides: factory.NewCounterVec(prometheus.CounterOpts{
Namespace: "coderd",
Subsystem: "chatd",
Name: "hook_input_overrides_total",
Help: "Total lifecycle hook input overrides by event.",
}, []string{"event"}),
}
}
func (m *metrics) observe(eventType agenthooks.EventType, result Result, response agenthooks.Response, duration time.Duration) {
event := string(eventType)
m.dispatches.WithLabelValues(event, string(result)).Inc()
m.duration.WithLabelValues(event).Observe(duration.Seconds())
if response.ModelContext != "" {
m.contextSize.WithLabelValues(event).Observe(float64(len(response.ModelContext)))
}
if response.Permission == nil {
return
}
switch response.Permission.Decision {
case agenthooks.PermissionAllow, agenthooks.PermissionDeny:
m.decisions.WithLabelValues(event, string(response.Permission.Decision)).Inc()
}
if response.Permission.Decision == agenthooks.PermissionAllow && response.Permission.InputOverride != nil {
m.inputOverrides.WithLabelValues(event).Inc()
}
}
@@ -0,0 +1,701 @@
package dispatch
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"sync/atomic"
"testing"
"time"
"github.com/google/uuid"
"github.com/prometheus/client_golang/prometheus"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/coder/coder/v2/codersdk/x/agenthooks"
"github.com/coder/coder/v2/testutil"
)
const (
testSecret = "test-hook-secret-32-bytes-minimum!!"
testDeploymentID = "test-deployment"
testVersion = "test-version"
)
func TestDispatcherSuccess(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventSessionStart, agenthooks.SessionStartData{Source: "new"})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
body, err := io.ReadAll(r.Body)
assert.NoError(t, err)
assert.Equal(t, "application/json", r.Header.Get("Content-Type"))
assert.Equal(t, "coderd-agenthooks/"+testVersion, r.Header.Get("User-Agent"))
claims, err := agenthooks.Verify(r.Header.Get("Authorization"), []byte(testSecret))
assert.NoError(t, err)
assert.Equal(t, testDeploymentID, claims.Issuer)
assert.Equal(t, serverURL(r), claims.Audience)
assert.Equal(t, event.Type, claims.Type)
chatID, err := claims.ChatID()
assert.NoError(t, err)
assert.Equal(t, event.ChatID, chatID)
digest := sha256.Sum256(body)
assert.Equal(t, hex.EncodeToString(digest[:]), claims.BodySHA256)
assert.Equal(t, claims.IssuedAt-int64(clockSkewLeeway/time.Second), claims.NotBefore)
var request agenthooks.Request
assert.NoError(t, json.Unmarshal(body, &request))
assert.Equal(t, claims.JTI, request.Meta.DispatchID)
assert.Equal(t, agenthooks.SchemaVersion, request.Meta.SchemaVersion)
var decoded agenthooks.SessionStartData
assert.NoError(t, json.Unmarshal(request.Data, &decoded))
assert.Equal(t, agenthooks.SessionStartData{Source: "new"}, decoded)
_, err = w.Write([]byte(`{}`))
assert.NoError(t, err)
}))
t.Cleanup(server.Close)
dispatcher := newTestDispatcher(t, server.Client(), server.URL, 2*time.Second)
response, _, err := dispatcher.Dispatch(testutil.Context(t, testutil.WaitLong), event)
require.NoError(t, err)
require.Equal(t, agenthooks.Response{}, response)
}
func TestDispatcherRejectsCleartextURL(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventPreToolUse, agenthooks.PreToolUseData{ToolUseID: "call_1", ToolName: "execute"})
dispatcher := newTestDispatcher(t, nil, "http://hooks.example.com/coder", 2*time.Second)
_, _, err := dispatcher.Dispatch(testutil.Context(t, testutil.WaitShort), event)
require.ErrorContains(t, err, "must use HTTPS")
require.NoError(t, validateHookURL(""))
require.NoError(t, validateHookURL("https://hooks.example.com/coder"))
require.NoError(t, validateHookURL("http://localhost:8080/hooks"))
require.NoError(t, validateHookURL("http://127.0.0.1:8080/hooks"))
require.NoError(t, validateHookURL("http://[::1]:8080/hooks"))
require.Error(t, validateHookURL("http://10.0.0.5/hooks"))
require.Error(t, validateHookURL("ftp://hooks.example.com/coder"))
require.ErrorContains(t, validateHookURL("https:///coder"), "must include a host")
require.ErrorContains(t, validateHookURL("https:hooks.example.com"), "must include a host")
require.ErrorContains(t, validateHookURL("http:///hooks"), "must include a host")
require.ErrorContains(t, validateHookURL("https://hooks.example.com/coder#frag"), "must not contain a fragment")
require.ErrorContains(t, validateHookURL("https://user:pass@hooks.example.com/coder"), "must not contain userinfo")
}
func TestDispatcherDeny(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventUserPromptSubmit, agenthooks.UserPromptSubmitData{Prompt: "delete everything"})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, err := w.Write([]byte(`{"permission":{"decision":"deny","reason":"blocked"},"user_message":"not allowed"}`))
assert.NoError(t, err)
}))
t.Cleanup(server.Close)
response, _, err := newTestDispatcher(t, server.Client(), server.URL, time.Second).Dispatch(
testutil.Context(t, testutil.WaitLong), event,
)
require.NoError(t, err)
require.NotNil(t, response.Permission)
require.Equal(t, agenthooks.PermissionDeny, response.Permission.Decision)
require.Equal(t, "not allowed", response.UserMessage)
}
func TestDispatcherDenyNullOverrideDecode(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventPreToolUse, agenthooks.PreToolUseData{
ToolUseID: "tool-use-1",
ToolName: "execute",
ToolInput: json.RawMessage(`{"cmd":"rm"}`),
})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, err := w.Write([]byte(`{"permission":{"decision":"deny","input_override":null}}`))
assert.NoError(t, err)
}))
t.Cleanup(server.Close)
response, _, err := newTestDispatcher(t, server.Client(), server.URL, time.Second).Dispatch(
testutil.Context(t, testutil.WaitLong), event,
)
require.NoError(t, err)
require.NotNil(t, response.Permission)
require.Equal(t, agenthooks.PermissionDeny, response.Permission.Decision)
require.JSONEq(t, `null`, string(response.Permission.InputOverride))
}
func TestDispatcherAllowInputOverride(t *testing.T) {
t.Parallel()
toolInput := json.RawMessage(`{"path":"before"}`)
toolUseID := "call_" + uuid.NewString()
event := newTestEvent(t, agenthooks.EventPreToolUse, agenthooks.PreToolUseData{
ToolUseID: toolUseID,
ToolName: "edit",
ToolInput: toolInput,
})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, err := w.Write([]byte(`{"permission":{"decision":"allow","input_override":{"path":"after"}}}`))
assert.NoError(t, err)
}))
t.Cleanup(server.Close)
response, _, err := newTestDispatcher(t, server.Client(), server.URL, time.Second).Dispatch(
testutil.Context(t, testutil.WaitLong), event,
)
require.NoError(t, err)
require.NotNil(t, response.Permission)
require.Equal(t, agenthooks.PermissionAllow, response.Permission.Decision)
require.JSONEq(t, `{"path":"after"}`, string(response.Permission.InputOverride))
}
func TestDispatcherAllowsCaseDistinctOverrideKeys(t *testing.T) {
t.Parallel()
// Tool schemas are case-sensitive, so "URL" and "url" are distinct
// properties even though the response envelope folds its own keys.
event := newTestEvent(t, agenthooks.EventPreToolUse, agenthooks.PreToolUseData{
ToolUseID: "call_" + uuid.NewString(),
ToolName: "fetch",
ToolInput: json.RawMessage(`{"url":"before"}`),
})
// The envelope key is folded by the decoder, so the boundary detection
// must fold too or the case-distinct override below is wrongly rejected.
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, err := w.Write([]byte(`{"permission":{"decision":"allow","INPUT_OVERRIDE":{"URL":"upper","url":"lower"}}}`))
assert.NoError(t, err)
}))
t.Cleanup(server.Close)
response, _, err := newTestDispatcher(t, server.Client(), server.URL, time.Second).Dispatch(
testutil.Context(t, testutil.WaitLong), event,
)
require.NoError(t, err)
require.NotNil(t, response.Permission)
require.JSONEq(t, `{"URL":"upper","url":"lower"}`, string(response.Permission.InputOverride))
}
func TestDispatcherTimeoutNoRetry(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventStop, agenthooks.StopData{})
var requests atomic.Int32
release := make(chan struct{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
requests.Add(1)
w.WriteHeader(http.StatusOK)
assert.NoError(t, http.NewResponseController(w).Flush())
<-release
}))
t.Cleanup(server.Close)
_, _, err := newTestDispatcher(t, server.Client(), server.URL, 50*time.Millisecond).Dispatch(
testutil.Context(t, testutil.WaitLong), event,
)
close(release)
assertDispatchErrorClass(t, err, ResultTimeout)
require.Equal(t, int32(1), requests.Load())
}
func TestDispatcherRetriesConnectionErrorWithSameJTI(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventPostCompact, agenthooks.PostCompactData{})
claimsCh := make(chan agenthooks.Claims, 2)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
claims, err := agenthooks.Verify(r.Header.Get("Authorization"), []byte(testSecret))
assert.NoError(t, err)
claimsCh <- claims
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(server.Close)
var attempts atomic.Int32
baseTransport := server.Client().Transport
client := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
claims, err := agenthooks.Verify(req.Header.Get("Authorization"), []byte(testSecret))
if err != nil {
return nil, err
}
if attempts.Add(1) == 1 {
claimsCh <- claims
_, err = io.Copy(io.Discard, req.Body)
if err != nil {
return nil, err
}
return nil, io.EOF
}
return baseTransport.RoundTrip(req)
})}
_, dispatchID, err := newTestDispatcher(t, client, server.URL, time.Second).Dispatch(
testutil.Context(t, testutil.WaitLong), event,
)
require.NoError(t, err)
require.Equal(t, int32(2), attempts.Load())
first := <-claimsCh
second := <-claimsCh
require.Equal(t, first.JTI, second.JTI)
require.Equal(t, first.JTI, dispatchID)
chatID, err := first.ChatID()
require.NoError(t, err)
require.Equal(t, event.ChatID, chatID)
}
func TestDispatcherRetriesHTTP2AbortWithSameDispatchID(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventPostCompact, agenthooks.PostCompactData{})
type attemptIdentity struct {
jti uuid.UUID
dispatchID uuid.UUID
}
identities := make(chan attemptIdentity, 2)
var attempts atomic.Int32
server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
assert.Equal(t, 2, r.ProtoMajor)
body, err := io.ReadAll(r.Body)
assert.NoError(t, err)
claims, err := agenthooks.Verify(r.Header.Get("Authorization"), []byte(testSecret))
assert.NoError(t, err)
var request agenthooks.Request
assert.NoError(t, json.Unmarshal(body, &request))
identities <- attemptIdentity{jti: claims.JTI, dispatchID: request.Meta.DispatchID}
if attempts.Add(1) == 1 {
_, err = w.Write([]byte(`{"partial":`))
assert.NoError(t, err)
assert.NoError(t, http.NewResponseController(w).Flush())
panic(http.ErrAbortHandler)
}
_, err = w.Write([]byte(`{}`))
assert.NoError(t, err)
}))
server.EnableHTTP2 = true
server.StartTLS()
t.Cleanup(server.Close)
_, dispatchID, err := newTestDispatcher(t, server.Client(), server.URL, time.Second).Dispatch(
testutil.Context(t, testutil.WaitLong), event,
)
require.NoError(t, err)
require.Equal(t, int32(2), attempts.Load())
first := <-identities
second := <-identities
require.Equal(t, first, second)
require.Equal(t, dispatchID, first.jti)
require.Equal(t, dispatchID, first.dispatchID)
}
func TestDispatcherRetriesMidBodyConnectionError(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventPostCompact, agenthooks.PostCompactData{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(server.Close)
var attempts atomic.Int32
baseTransport := server.Client().Transport
client := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
if attempts.Add(1) == 1 {
_, err := io.Copy(io.Discard, req.Body)
if err != nil {
return nil, err
}
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(errReader{err: io.ErrUnexpectedEOF}),
Header: http.Header{},
}, nil
}
return baseTransport.RoundTrip(req)
})}
_, _, err := newTestDispatcher(t, client, server.URL, time.Second).Dispatch(
testutil.Context(t, testutil.WaitLong), event,
)
require.NoError(t, err)
require.Equal(t, int32(2), attempts.Load())
}
type errReader struct{ err error }
func (r errReader) Read([]byte) (int, error) { return 0, r.err }
func TestDispatcherTLSFailureNoRetry(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventStop, agenthooks.StopData{})
server := httptest.NewTLSServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
t.Fatal("request with an untrusted certificate reached the handler")
}))
t.Cleanup(server.Close)
transport := http.DefaultTransport.(*http.Transport).Clone()
t.Cleanup(transport.CloseIdleConnections)
var attempts atomic.Int32
client := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
attempts.Add(1)
return transport.RoundTrip(req)
})}
_, _, err := newTestDispatcher(t, client, server.URL, time.Second).Dispatch(
testutil.Context(t, testutil.WaitLong), event,
)
assertDispatchErrorClass(t, err, ResultProtocolError)
require.Equal(t, int32(1), attempts.Load())
}
func TestDispatcherNon2xxNoRetry(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventPreCompact, agenthooks.PreCompactData{})
var requests atomic.Int32
release := make(chan struct{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
requests.Add(1)
w.WriteHeader(http.StatusServiceUnavailable)
assert.NoError(t, http.NewResponseController(w).Flush())
<-release
}))
t.Cleanup(server.Close)
_, _, err := newTestDispatcher(t, server.Client(), server.URL, time.Second).Dispatch(
testutil.Context(t, testutil.WaitLong), event,
)
close(release)
assertDispatchErrorClass(t, err, ResultHTTPError)
require.Equal(t, int32(1), requests.Load())
}
func TestDispatcherProtocolErrors(t *testing.T) {
t.Parallel()
tests := []struct {
name string
eventType agenthooks.EventType
data any
responseBody []byte
}{
{
name: "malformed JSON",
eventType: agenthooks.EventStop,
data: agenthooks.StopData{},
responseBody: []byte(`{"user_message":`),
},
{
name: "misspelled permission key",
eventType: agenthooks.EventPreToolUse,
data: agenthooks.PreToolUseData{ToolUseID: "call_typo", ToolName: "run_command", ToolInput: json.RawMessage(`{"cmd":"ls"}`)},
responseBody: []byte(`{"permision":{"decision":"deny"}}`),
},
{
name: "folded duplicate in prompt override",
eventType: agenthooks.EventUserPromptSubmit,
data: agenthooks.UserPromptSubmitData{Prompt: "original"},
// This override is struct-decoded, so folding applies to it.
responseBody: []byte(`{"permission":{"decision":"allow","input_override":{"prompt":"approved","Prompt":"different"}}}`),
},
{
name: "duplicate key inside input_override",
eventType: agenthooks.EventPreToolUse,
data: agenthooks.PreToolUseData{ToolUseID: "call_dup_override", ToolName: "run_command", ToolInput: json.RawMessage(`{"cmd":"ls"}`)},
// Exact duplicates stay rejected at any depth.
responseBody: []byte(`{"permission":{"decision":"allow","input_override":{"url":"a","url":"b"}}}`),
},
{
name: "unicode-folded duplicate permission key",
eventType: agenthooks.EventPreToolUse,
data: agenthooks.PreToolUseData{ToolUseID: "call_unicode", ToolName: "run_command", ToolInput: json.RawMessage(`{"cmd":"ls"}`)},
// encoding/json folds U+017F to "s", so this aliases "permission".
responseBody: []byte(`{"permission":{"decision":"deny"},"permiſſion":null}`),
},
{
name: "case-folded duplicate permission key",
eventType: agenthooks.EventPreToolUse,
data: agenthooks.PreToolUseData{ToolUseID: "call_fold", ToolName: "run_command", ToolInput: json.RawMessage(`{"cmd":"ls"}`)},
// encoding/json matches fields case-insensitively and keeps the
// last value, so this would otherwise decode as an empty allow.
responseBody: []byte(`{"permission":{"decision":"deny"},"Permission":null}`),
},
{
name: "duplicate permission key",
eventType: agenthooks.EventPreToolUse,
data: agenthooks.PreToolUseData{ToolUseID: "call_dup", ToolName: "run_command", ToolInput: json.RawMessage(`{"cmd":"ls"}`)},
responseBody: []byte(`{"permission":{"decision":"deny"},"permission":null}`),
},
{
name: "unknown permission field",
eventType: agenthooks.EventPreToolUse,
data: agenthooks.PreToolUseData{ToolUseID: "call_unknown", ToolName: "run_command", ToolInput: json.RawMessage(`{"cmd":"ls"}`)},
responseBody: []byte(`{"permission":{"decision":"deny","reasoning":"typo"}}`),
},
{
name: "duplicate key inside input_override",
eventType: agenthooks.EventPreToolUse,
data: agenthooks.PreToolUseData{ToolUseID: "call_override_dup", ToolName: "run_command", ToolInput: json.RawMessage(`{"cmd":"ls"}`)},
responseBody: []byte(`{"permission":{"decision":"allow","input_override":{"cmd":"ls","cmd":"rm -rf"}}}`),
},
{
name: "trailing JSON value",
eventType: agenthooks.EventStop,
data: agenthooks.StopData{},
responseBody: []byte(`{"user_message":"done"}{}`),
},
{
name: "oversized model context",
eventType: agenthooks.EventStop,
data: agenthooks.StopData{},
responseBody: mustJSON(t, agenthooks.Response{ModelContext: string(bytes.Repeat([]byte("x"), maxModelContextBytes+1))}),
},
{
name: "invalid user prompt override shape",
eventType: agenthooks.EventUserPromptSubmit,
data: agenthooks.UserPromptSubmitData{Prompt: "question"},
responseBody: mustJSON(t, agenthooks.Response{Permission: &agenthooks.Permission{
Decision: agenthooks.PermissionAllow,
InputOverride: json.RawMessage(`{"unexpected":"value"}`),
}}),
},
{
name: "deny with input override",
eventType: agenthooks.EventPreToolUse,
data: agenthooks.PreToolUseData{
ToolUseID: "call_deny_override",
ToolName: "edit",
ToolInput: json.RawMessage(`{"path":"a"}`),
},
responseBody: mustJSON(t, agenthooks.Response{Permission: &agenthooks.Permission{
Decision: agenthooks.PermissionDeny,
InputOverride: json.RawMessage(`{"path":"b"}`),
}}),
},
{
name: "unsupported ask decision",
eventType: agenthooks.EventUserPromptSubmit,
data: agenthooks.UserPromptSubmitData{Prompt: "question"},
responseBody: mustJSON(t, agenthooks.Response{Permission: &agenthooks.Permission{
Decision: agenthooks.PermissionDecision("ask"),
}}),
},
{
name: "pre_tool_use allow without input_override",
eventType: agenthooks.EventPreToolUse,
data: agenthooks.PreToolUseData{ToolUseID: "call_no_override", ToolName: "run_command", ToolInput: json.RawMessage(`{"cmd":"ls"}`)},
responseBody: mustJSON(t, agenthooks.Response{Permission: &agenthooks.Permission{
Decision: agenthooks.PermissionAllow,
}}),
},
{
name: "pre_tool_use without tool_use_id",
eventType: agenthooks.EventPreToolUse,
data: agenthooks.PreToolUseData{ToolName: "run_command", ToolInput: json.RawMessage(`{"cmd":"ls"}`)},
responseBody: []byte(`{}`),
},
{
name: "post_tool_use without tool_name",
eventType: agenthooks.EventPostToolUse,
data: agenthooks.PostToolUseData{ToolUseID: "call_no_name"},
responseBody: []byte(`{}`),
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
t.Parallel()
event := newTestEvent(t, test.eventType, test.data)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, err := w.Write(test.responseBody)
assert.NoError(t, err)
}))
t.Cleanup(server.Close)
_, _, err := newTestDispatcher(t, server.Client(), server.URL, time.Second).Dispatch(
testutil.Context(t, testutil.WaitLong), event,
)
assertDispatchErrorClass(t, err, ResultProtocolError)
})
}
}
func TestDispatcherOverCapacity(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventStop, agenthooks.StopData{})
dispatcher := newTestDispatcher(t, nil, "https://unused.test", 10*time.Millisecond)
for range maxConcurrentDispatches {
dispatcher.semaphore <- struct{}{}
}
defer func() {
for range maxConcurrentDispatches {
<-dispatcher.semaphore
}
}()
_, _, err := dispatcher.Dispatch(testutil.Context(t, testutil.WaitLong), event)
require.ErrorIs(t, err, context.DeadlineExceeded)
assertDispatchErrorClass(t, err, ResultOverCapacity)
}
func TestDispatcherRejectedResponseIsNotObserved(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventPreToolUse, agenthooks.PreToolUseData{
ToolUseID: "call_" + uuid.NewString(),
ToolName: "execute",
ToolInput: json.RawMessage(`{"cmd":"ls"}`),
})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, err := w.Write([]byte(`{"permission":{"decision":"deny","input_override":{"cmd":"rm"}},"model_context":"ctx"}`))
assert.NoError(t, err)
}))
t.Cleanup(server.Close)
registry := prometheus.NewRegistry()
dispatcher := New(
testutil.Logger(t), server.Client(), server.URL, testSecret, time.Second,
testDeploymentID, testVersion, registry,
)
_, _, err := dispatcher.Dispatch(testutil.Context(t, testutil.WaitLong), event)
assertDispatchErrorClass(t, err, ResultProtocolError)
families, err := registry.Gather()
require.NoError(t, err)
for _, family := range families {
switch family.GetName() {
case "coderd_chatd_hook_decisions_total",
"coderd_chatd_hook_input_overrides_total",
"coderd_chatd_hook_context_size_bytes":
require.Empty(t, family.GetMetric(), "rejected response observed in %s", family.GetName())
}
}
}
func TestDispatcherCanceledContextAtCapacityReturnsTimeout(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventStop, agenthooks.StopData{})
dispatcher := newTestDispatcher(t, nil, "https://unused.test", testutil.WaitLong)
for range maxConcurrentDispatches {
dispatcher.semaphore <- struct{}{}
}
defer func() {
for range maxConcurrentDispatches {
<-dispatcher.semaphore
}
}()
ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitLong))
cancel()
_, _, err := dispatcher.Dispatch(ctx, event)
require.ErrorIs(t, err, context.Canceled)
assertDispatchErrorClass(t, err, ResultTimeout)
}
func TestDispatcherCanceledContextReturnsTimeout(t *testing.T) {
t.Parallel()
event := newTestEvent(t, agenthooks.EventStop, agenthooks.StopData{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
}))
t.Cleanup(server.Close)
ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitLong))
cancel()
_, _, err := newTestDispatcher(t, server.Client(), server.URL, time.Second).Dispatch(ctx, event)
assertDispatchErrorClass(t, err, ResultTimeout)
}
func TestDispatcherInvalidToolInputReturnsProtocolError(t *testing.T) {
t.Parallel()
toolUseID := "call_" + uuid.NewString()
event := newTestEvent(t, agenthooks.EventPreToolUse, agenthooks.PreToolUseData{
ToolUseID: toolUseID,
ToolName: "edit",
ToolInput: json.RawMessage(`{"path":`),
})
var hookRequests atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
hookRequests.Add(1)
w.WriteHeader(http.StatusOK)
}))
t.Cleanup(server.Close)
_, _, err := newTestDispatcher(t, server.Client(), server.URL, time.Second).Dispatch(
testutil.Context(t, testutil.WaitLong), event,
)
assertDispatchErrorClass(t, err, ResultProtocolError)
require.Zero(t, hookRequests.Load())
}
func newTestDispatcher(
t *testing.T,
client *http.Client,
hookURL string,
timeout time.Duration,
) *Dispatcher {
t.Helper()
return New(
testutil.Logger(t),
client,
hookURL,
testSecret,
timeout,
testDeploymentID,
testVersion,
prometheus.NewRegistry(),
)
}
func newTestEvent(t *testing.T, eventType agenthooks.EventType, data any) Event {
t.Helper()
return Event{
Type: eventType,
ChatRef: agenthooks.ChatRef{
ChatID: uuid.New(),
OwnerID: uuid.New(),
},
Data: data,
}
}
func assertDispatchErrorClass(t *testing.T, err error, expected Result) {
t.Helper()
var dispatchErr *Error
require.ErrorAs(t, err, &dispatchErr)
require.Equal(t, expected, dispatchErr.Class)
require.NotEqual(t, uuid.Nil, dispatchErr.DispatchID)
}
func mustJSON(t *testing.T, value any) []byte {
t.Helper()
encoded, err := json.Marshal(value)
require.NoError(t, err)
return encoded
}
func serverURL(r *http.Request) string {
return "http://" + r.Host
}
type roundTripFunc func(*http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}