From 8ea2586189be2d9723c34e5d46a9ffc0a1bec8ad Mon Sep 17 00:00:00 2001 From: Michael Suchacz <203725896+ibetitsmike@users.noreply.github.com> Date: Tue, 28 Jul 2026 13:59:37 +0200 Subject: [PATCH] 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/` 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. --- coderd/util/xnet/xnet.go | 62 ++ coderd/util/xnet/xnet_test.go | 89 +++ coderd/x/agenthooks/dispatch/dispatcher.go | 712 ++++++++++++++++++ .../dispatch/dispatcher_internal_test.go | 701 +++++++++++++++++ codersdk/x/agenthooks/agenthooks_test.go | 469 ++++++++++++ codersdk/x/agenthooks/http.go | 189 +++++ codersdk/x/agenthooks/jwt.go | 137 ++++ codersdk/x/agenthooks/types.go | 154 ++++ docs/admin/integrations/prometheus.md | 5 + scripts/agenthooks-server/main.go | 323 ++++++++ scripts/apitypings/main.go | 39 + scripts/metricsdocgen/generated_metrics | 15 + site/src/api/typesGenerated.ts | 204 +++++ 13 files changed, 3099 insertions(+) create mode 100644 coderd/util/xnet/xnet.go create mode 100644 coderd/util/xnet/xnet_test.go create mode 100644 coderd/x/agenthooks/dispatch/dispatcher.go create mode 100644 coderd/x/agenthooks/dispatch/dispatcher_internal_test.go create mode 100644 codersdk/x/agenthooks/agenthooks_test.go create mode 100644 codersdk/x/agenthooks/http.go create mode 100644 codersdk/x/agenthooks/jwt.go create mode 100644 codersdk/x/agenthooks/types.go create mode 100644 scripts/agenthooks-server/main.go diff --git a/coderd/util/xnet/xnet.go b/coderd/util/xnet/xnet.go new file mode 100644 index 0000000000..d0595e6c25 --- /dev/null +++ b/coderd/util/xnet/xnet.go @@ -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 + } +} diff --git a/coderd/util/xnet/xnet_test.go b/coderd/util/xnet/xnet_test.go new file mode 100644 index 0000000000..43beb0da83 --- /dev/null +++ b/coderd/util/xnet/xnet_test.go @@ -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) + }) + } +} diff --git a/coderd/x/agenthooks/dispatch/dispatcher.go b/coderd/x/agenthooks/dispatch/dispatcher.go new file mode 100644 index 0000000000..db7e1566e1 --- /dev/null +++ b/coderd/x/agenthooks/dispatch/dispatcher.go @@ -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() + } +} diff --git a/coderd/x/agenthooks/dispatch/dispatcher_internal_test.go b/coderd/x/agenthooks/dispatch/dispatcher_internal_test.go new file mode 100644 index 0000000000..320fd33b14 --- /dev/null +++ b/coderd/x/agenthooks/dispatch/dispatcher_internal_test.go @@ -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) +} diff --git a/codersdk/x/agenthooks/agenthooks_test.go b/codersdk/x/agenthooks/agenthooks_test.go new file mode 100644 index 0000000000..89d8f7a69f --- /dev/null +++ b/codersdk/x/agenthooks/agenthooks_test.go @@ -0,0 +1,469 @@ +package agenthooks_test + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/go-jose/go-jose/v4" + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/codersdk/x/agenthooks" +) + +var testSecret = []byte("0123456789abcdef0123456789abcdef") + +// testAudience is the URL the handler is configured to accept, kept +// independent of the test server address it is reached on. +const testAudience = "https://hooks.example.com" + +func TestSignClaimsVerify(t *testing.T) { + t.Parallel() + + claims := validClaims(t, "https://hooks.example.com/coder", agenthooks.EventPreToolUse, nil) + token, err := agenthooks.SignClaims(testSecret, claims) + require.NoError(t, err) + + got, err := agenthooks.Verify("Bearer "+token, testSecret) + require.NoError(t, err) + require.Equal(t, claims, got) +} + +func TestShortSecretRejected(t *testing.T) { + t.Parallel() + + claims := validClaims(t, "https://hooks.example.com/coder", agenthooks.EventPreToolUse, nil) + shortSecret := testSecret[:agenthooks.MinSecretLen-1] + + _, err := agenthooks.SignClaims(shortSecret, claims) + require.ErrorContains(t, err, "secret must be at least") + + token, err := agenthooks.SignClaims(testSecret, claims) + require.NoError(t, err) + _, err = agenthooks.Verify("Bearer "+token, shortSecret) + require.ErrorContains(t, err, "secret must be at least") + _, err = agenthooks.Verify("Bearer "+token, nil) + require.ErrorContains(t, err, "secret must be at least") +} + +func TestVerifyRejectsAlgorithmConfusion(t *testing.T) { + t.Parallel() + + claims := validClaims(t, "https://hooks.example.com/coder", agenthooks.EventStop, nil) + signer, err := jose.NewSigner( + jose.SigningKey{Algorithm: jose.HS512, Key: bytes.Repeat([]byte{1}, 64)}, + new(jose.SignerOptions).WithType("JWT"), + ) + require.NoError(t, err) + payload, err := json.Marshal(claims) + require.NoError(t, err) + signed, err := signer.Sign(payload) + require.NoError(t, err) + token, err := signed.CompactSerialize() + require.NoError(t, err) + + _, err = agenthooks.Verify("Bearer "+token, testSecret) + var algErr *jose.ErrUnexpectedSignatureAlgorithm + require.ErrorAs(t, err, &algErr) + require.Equal(t, jose.HS512, algErr.Got) +} + +func TestVerifyTimeBounds(t *testing.T) { + t.Parallel() + + now := time.Now() + tests := []struct { + name string + update func(*agenthooks.Claims) + wantErr string + }{ + { + name: "expired", + update: func(claims *agenthooks.Claims) { + claims.IssuedAt = now.Add(-2 * time.Minute).Unix() + claims.NotBefore = now.Add(-2 * time.Minute).Unix() + claims.Expires = now.Add(-time.Minute).Unix() + }, + wantErr: "token has expired", + }, + { + name: "not before", + update: func(claims *agenthooks.Claims) { + claims.NotBefore = now.Add(time.Minute).Unix() + claims.Expires = now.Add(2 * time.Minute).Unix() + }, + wantErr: "token is not valid yet", + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + claims := validClaims(t, "https://hooks.example.com/coder", agenthooks.EventStop, nil) + test.update(&claims) + token, err := agenthooks.SignClaims(testSecret, claims) + require.NoError(t, err) + + _, err = agenthooks.Verify("Bearer "+token, testSecret) + require.ErrorContains(t, err, test.wantErr) + }) + } +} + +func TestHTTPHandlerRoutesEvents(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + event agenthooks.EventType + data any + install func(*testing.T, *agenthooks.Hooks) + }{ + { + name: "session start", + event: agenthooks.EventSessionStart, + data: agenthooks.SessionStartData{Source: "startup"}, + install: func(t *testing.T, h *agenthooks.Hooks) { + h.SessionStart = func(_ context.Context, _ agenthooks.Meta, data agenthooks.SessionStartData) (agenthooks.Response, error) { + assert.Equal(t, "startup", data.Source) + return agenthooks.Response{UserMessage: "session start"}, nil + } + }, + }, + { + name: "user prompt submit", + event: agenthooks.EventUserPromptSubmit, + data: agenthooks.UserPromptSubmitData{Prompt: "hello"}, + install: func(t *testing.T, h *agenthooks.Hooks) { + h.UserPromptSubmit = func(_ context.Context, _ agenthooks.Meta, data agenthooks.UserPromptSubmitData) (agenthooks.Response, error) { + assert.Equal(t, "hello", data.Prompt) + return agenthooks.Response{UserMessage: "user prompt submit"}, nil + } + }, + }, + { + name: "pre tool use", + event: agenthooks.EventPreToolUse, + data: agenthooks.PreToolUseData{ + ToolUseID: "call_" + uuid.NewString(), + ToolName: "execute", + ToolInput: json.RawMessage(`{"command":"pwd"}`), + }, + install: func(t *testing.T, h *agenthooks.Hooks) { + h.PreToolUse = func(_ context.Context, _ agenthooks.Meta, data agenthooks.PreToolUseData) (agenthooks.Response, error) { + assert.Equal(t, "execute", data.ToolName) + return agenthooks.Response{UserMessage: "pre tool use"}, nil + } + }, + }, + { + name: "post tool use", + event: agenthooks.EventPostToolUse, + data: agenthooks.PostToolUseData{ + ToolUseID: "call_" + uuid.NewString(), + ToolName: "execute", + ToolResponse: json.RawMessage(`{"output":"ok"}`), + }, + install: func(t *testing.T, h *agenthooks.Hooks) { + h.PostToolUse = func(_ context.Context, _ agenthooks.Meta, data agenthooks.PostToolUseData) (agenthooks.Response, error) { + assert.Equal(t, "execute", data.ToolName) + return agenthooks.Response{UserMessage: "post tool use"}, nil + } + }, + }, + { + name: "pre compact", + event: agenthooks.EventPreCompact, + data: agenthooks.PreCompactData{}, + install: func(_ *testing.T, h *agenthooks.Hooks) { + h.PreCompact = func(context.Context, agenthooks.Meta, agenthooks.PreCompactData) (agenthooks.Response, error) { + return agenthooks.Response{UserMessage: "pre compact"}, nil + } + }, + }, + { + name: "post compact", + event: agenthooks.EventPostCompact, + data: agenthooks.PostCompactData{}, + install: func(_ *testing.T, h *agenthooks.Hooks) { + h.PostCompact = func(context.Context, agenthooks.Meta, agenthooks.PostCompactData) (agenthooks.Response, error) { + return agenthooks.Response{UserMessage: "post compact"}, nil + } + }, + }, + { + name: "stop", + event: agenthooks.EventStop, + data: agenthooks.StopData{}, + install: func(_ *testing.T, h *agenthooks.Hooks) { + h.Stop = func(context.Context, agenthooks.Meta, agenthooks.StopData) (agenthooks.Response, error) { + return agenthooks.Response{UserMessage: "stop"}, nil + } + }, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + var h agenthooks.Hooks + test.install(t, &h) + server := httptest.NewServer(agenthooks.NewHTTPHandler(testSecret, testAudience, h)) + t.Cleanup(server.Close) + + response := postEvent(t, server.URL, test.event, test.data, nil, nil) + defer response.Body.Close() + require.Equal(t, http.StatusOK, response.StatusCode) + var got agenthooks.Response + require.NoError(t, json.NewDecoder(response.Body).Decode(&got)) + require.Equal(t, test.name, got.UserMessage) + }) + } +} + +func TestHTTPHandlerUnencodableResponseFailsClosed(t *testing.T) { + t.Parallel() + + // An empty 200 reads as allow, so a response that cannot be marshaled + // must not reach the dispatcher as one. + server := httptest.NewServer(agenthooks.NewHTTPHandler(testSecret, testAudience, agenthooks.Hooks{ + Stop: func(context.Context, agenthooks.Meta, agenthooks.StopData) (agenthooks.Response, error) { + return agenthooks.Response{ + Permission: &agenthooks.Permission{ + Decision: agenthooks.PermissionDeny, + InputOverride: json.RawMessage(`{invalid`), + }, + }, nil + }, + })) + t.Cleanup(server.Close) + + response := postEvent(t, server.URL, agenthooks.EventStop, agenthooks.StopData{}, nil, nil) + defer response.Body.Close() + require.Equal(t, http.StatusInternalServerError, response.StatusCode) +} + +func TestHTTPHandlerNoOpHookDoesNotDecodeData(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(agenthooks.NewHTTPHandler(testSecret, testAudience, agenthooks.Hooks{})) + t.Cleanup(server.Close) + response := postEvent(t, server.URL, agenthooks.EventStop, "unused", nil, nil) + defer response.Body.Close() + require.Equal(t, http.StatusOK, response.StatusCode) + var got agenthooks.Response + require.NoError(t, json.NewDecoder(response.Body).Decode(&got)) + require.Equal(t, agenthooks.Response{}, got) +} + +func TestHTTPHandlerRejectsMismatches(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + updateRequest func(*agenthooks.Request) + updateClaims func(*agenthooks.Claims) + }{ + { + name: "dispatch ID", + updateRequest: func(request *agenthooks.Request) { + request.Meta.DispatchID = uuid.New() + }, + }, + { + name: "event type", + updateRequest: func(request *agenthooks.Request) { + request.Type = agenthooks.EventPreCompact + }, + }, + { + name: "chat ID", + updateRequest: func(request *agenthooks.Request) { + request.Meta.ChatID = uuid.New() + }, + }, + { + name: "audience", + updateClaims: func(claims *agenthooks.Claims) { + claims.Audience = "https://hooks.example.com/other" + }, + }, + { + name: "body SHA-256", + updateClaims: func(claims *agenthooks.Claims) { + claims.BodySHA256 = hex.EncodeToString(bytes.Repeat([]byte{1}, sha256.Size)) + }, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(agenthooks.NewHTTPHandler(testSecret, testAudience, agenthooks.Hooks{})) + t.Cleanup(server.Close) + response := postEvent(t, server.URL, agenthooks.EventStop, agenthooks.StopData{}, test.updateRequest, test.updateClaims) + defer response.Body.Close() + require.Equal(t, http.StatusBadRequest, response.StatusCode) + }) + } +} + +func TestHTTPHandlerExpectedIssuer(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(agenthooks.NewHTTPHandler( + testSecret, + testAudience, + agenthooks.Hooks{}, + agenthooks.WithExpectedIssuer("deployment-a"), + )) + t.Cleanup(server.Close) + + matching := postEvent(t, server.URL, agenthooks.EventStop, agenthooks.StopData{}, nil, func(claims *agenthooks.Claims) { + claims.Issuer = "deployment-a" + }) + defer matching.Body.Close() + require.Equal(t, http.StatusOK, matching.StatusCode) + + mismatched := postEvent(t, server.URL, agenthooks.EventStop, agenthooks.StopData{}, nil, func(claims *agenthooks.Claims) { + claims.Issuer = "deployment-b" + }) + defer mismatched.Body.Close() + require.Equal(t, http.StatusUnauthorized, mismatched.StatusCode) +} + +func TestHTTPHandlerAcceptsTrailingSlashAudience(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(agenthooks.NewHTTPHandler(testSecret, testAudience, agenthooks.Hooks{})) + t.Cleanup(server.Close) + response := postEvent(t, server.URL, agenthooks.EventStop, agenthooks.StopData{}, nil, func(claims *agenthooks.Claims) { + claims.Audience = testAudience + "/" + }) + defer response.Body.Close() + require.Equal(t, http.StatusOK, response.StatusCode) +} + +func TestHTTPHandlerRejectsRequestDerivedAudience(t *testing.T) { + t.Parallel() + + // Every request-controlled source of an audience names an attacker host, so + // a token minted for another listener sharing this secret would pass an + // audience check that reads any of them. + const spoofed = "https://hooks.attacker.example" + handler := agenthooks.NewHTTPHandler(testSecret, testAudience, agenthooks.Hooks{}) + body, token := signedEvent(t, spoofed, agenthooks.EventStop, agenthooks.StopData{}, nil, nil) + request, err := http.NewRequestWithContext(t.Context(), http.MethodPost, spoofed, bytes.NewReader(body)) + require.NoError(t, err) + request.Host = "hooks.attacker.example" + request.Header.Set("Authorization", "Bearer "+token) + request.Header.Set("X-Forwarded-Proto", "https") + request.Header.Set("X-Forwarded-Host", "hooks.attacker.example") + + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, request) + require.Equal(t, http.StatusBadRequest, recorder.Code) +} + +func TestHTTPHandlerWithoutAudienceRejectsEveryRequest(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(agenthooks.NewHTTPHandler(testSecret, "", agenthooks.Hooks{})) + t.Cleanup(server.Close) + response := postEvent(t, server.URL, agenthooks.EventStop, agenthooks.StopData{}, nil, nil) + defer response.Body.Close() + require.Equal(t, http.StatusInternalServerError, response.StatusCode) +} + +func postEvent(t *testing.T, target string, eventType agenthooks.EventType, data any, updateRequest func(*agenthooks.Request), updateClaims func(*agenthooks.Claims)) *http.Response { + t.Helper() + + body, token := signedEvent(t, testAudience, eventType, data, updateRequest, updateClaims) + httpRequest, err := http.NewRequestWithContext(t.Context(), http.MethodPost, target, bytes.NewReader(body)) + require.NoError(t, err) + httpRequest.Header.Set("Authorization", "Bearer "+token) + response, err := http.DefaultClient.Do(httpRequest) + require.NoError(t, err) + return response +} + +func signedEvent(t *testing.T, audience string, eventType agenthooks.EventType, data any, updateRequest func(*agenthooks.Request), updateClaims func(*agenthooks.Claims)) ([]byte, string) { + t.Helper() + + dataJSON, err := json.Marshal(data) + require.NoError(t, err) + request := agenthooks.Request{ + Type: eventType, + Meta: agenthooks.Meta{ + DispatchID: uuid.New(), + SchemaVersion: agenthooks.SchemaVersion, + ChatRef: agenthooks.ChatRef{ + ChatID: uuid.New(), + OwnerID: uuid.New(), + }, + }, + Data: dataJSON, + } + claims := validClaims(t, audience, eventType, &request) + if updateRequest != nil { + updateRequest(&request) + } + body, err := json.Marshal(request) + require.NoError(t, err) + digest := sha256.Sum256(body) + claims.BodySHA256 = hex.EncodeToString(digest[:]) + if updateClaims != nil { + updateClaims(&claims) + } + token, err := agenthooks.SignClaims(testSecret, claims) + require.NoError(t, err) + return body, token +} + +func validClaims(t *testing.T, audience string, eventType agenthooks.EventType, request *agenthooks.Request) agenthooks.Claims { + t.Helper() + + now := time.Now() + claims := agenthooks.Claims{ + Issuer: uuid.NewString(), + Subject: "coder:chat:" + uuid.NewString(), + Audience: audience, + IssuedAt: now.Unix(), + NotBefore: now.Add(-time.Second).Unix(), + Expires: now.Add(time.Minute).Unix(), + JTI: uuid.New(), + Type: eventType, + BodySHA256: hex.EncodeToString(make([]byte, sha256.Size)), + } + if request != nil { + claims.Subject = "coder:chat:" + request.Meta.ChatID.String() + claims.JTI = request.Meta.DispatchID + } + return claims +} + +func TestHTTPHandlerRejectsOversizedBody(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(agenthooks.NewHTTPHandler(testSecret, testAudience, agenthooks.Hooks{})) + t.Cleanup(server.Close) + + // A correctly signed body over the limit must be rejected by size + // before it is hashed or decoded. + huge, err := json.Marshal(strings.Repeat("a", int(agenthooks.MaxRequestBodyBytes))) + require.NoError(t, err) + response := postEvent(t, server.URL, agenthooks.EventStop, agenthooks.StopData{}, func(request *agenthooks.Request) { + request.Data = huge + }, nil) + defer response.Body.Close() + require.Equal(t, http.StatusRequestEntityTooLarge, response.StatusCode) +} diff --git a/codersdk/x/agenthooks/http.go b/codersdk/x/agenthooks/http.go new file mode 100644 index 0000000000..12322a395b --- /dev/null +++ b/codersdk/x/agenthooks/http.go @@ -0,0 +1,189 @@ +package agenthooks + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "io" + "net/http" + "net/url" + + "golang.org/x/xerrors" +) + +// MaxRequestBodyBytes limits memory used to verify hook requests. +const MaxRequestBodyBytes int64 = 10 << 20 // 10 MiB + +// Hooks lets a consumer implement only the lifecycle events it uses. +type Hooks struct { + SessionStart func(context.Context, Meta, SessionStartData) (Response, error) + UserPromptSubmit func(context.Context, Meta, UserPromptSubmitData) (Response, error) + PreToolUse func(context.Context, Meta, PreToolUseData) (Response, error) + PostToolUse func(context.Context, Meta, PostToolUseData) (Response, error) + PreCompact func(context.Context, Meta, PreCompactData) (Response, error) + PostCompact func(context.Context, Meta, PostCompactData) (Response, error) + Stop func(context.Context, Meta, StopData) (Response, error) +} + +// HandlerOption configures NewHTTPHandler. +type HandlerOption func(*handlerOptions) + +type handlerOptions struct { + expectedIssuer string +} + +// WithExpectedIssuer requires the verified iss claim to match issuer. +// If omitted, NewHTTPHandler accepts any non-empty issuer signed +// with the secret. +func WithExpectedIssuer(issuer string) HandlerOption { + return func(options *handlerOptions) { + options.expectedIssuer = issuer + } +} + +// NewHTTPHandler verifies hook POSTs, binds their claims to each request, +// and routes events to their configured callbacks. expectedAudience must be +// the URL Coder dispatches to, which is the value it signs into the aud +// claim. Deriving it from the request instead would let a caller replay a +// token minted for a different listener, because the request URL, the Host +// header, and any forwarding headers are all caller-controlled. A handler +// built with an empty audience rejects every request. +func NewHTTPHandler(secret []byte, expectedAudience string, hooks Hooks, opts ...HandlerOption) http.Handler { + var options handlerOptions + for _, opt := range opts { + opt(&options) + } + if expectedAudience == "" { + return http.HandlerFunc(func(rw http.ResponseWriter, _ *http.Request) { + http.Error(rw, "hook audience is not configured", http.StatusInternalServerError) + }) + } + audience := canonicalAudience(expectedAudience) + return http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + rw.Header().Set("Allow", http.MethodPost) + http.Error(rw, "method not allowed", http.StatusMethodNotAllowed) + return + } + + claims, err := Verify(r.Header.Get("Authorization"), secret) + if err != nil { + http.Error(rw, err.Error(), http.StatusUnauthorized) + return + } + if options.expectedIssuer != "" && claims.Issuer != options.expectedIssuer { + http.Error(rw, "unexpected issuer", http.StatusUnauthorized) + return + } + r.Body = http.MaxBytesReader(rw, r.Body, MaxRequestBodyBytes) + body, err := io.ReadAll(r.Body) + if err != nil { + var maxBytesErr *http.MaxBytesError + if errors.As(err, &maxBytesErr) { + http.Error(rw, "request body too large", http.StatusRequestEntityTooLarge) + return + } + http.Error(rw, "read request body", http.StatusBadRequest) + return + } + var request Request + if err := json.Unmarshal(body, &request); err != nil { + http.Error(rw, "decode request body", http.StatusBadRequest) + return + } + if err := verifyBody(body, claims, request, audience); err != nil { + http.Error(rw, err.Error(), http.StatusBadRequest) + return + } + response, err := dispatch(r.Context(), hooks, request) + if err != nil { + http.Error(rw, err.Error(), http.StatusInternalServerError) + return + } + + // Marshal before writing: streaming the encode would emit a 200 + // with a truncated body on failure, and the dispatcher reads an + // empty 200 as allow, so a malformed deny would fail open. + encoded, err := json.Marshal(response) + if err != nil { + http.Error(rw, "encode response", http.StatusInternalServerError) + return + } + rw.Header().Set("Content-Type", "application/json") + if _, err := rw.Write(encoded); err != nil { + return + } + }) +} + +func verifyBody(body []byte, claims Claims, request Request, audience string) error { + digest := sha256.Sum256(body) + if claims.BodySHA256 != hex.EncodeToString(digest[:]) { + return xerrors.New("request body does not match body_sha256 claim") + } + if canonicalAudience(claims.Audience) != audience { + return xerrors.New("audience claim does not match the configured audience") + } + if request.Meta.SchemaVersion != SchemaVersion { + return xerrors.New("unsupported schema version") + } + if request.Meta.DispatchID != claims.JTI { + return xerrors.New("dispatch ID does not match JWT ID") + } + if request.Type != claims.Type { + return xerrors.New("request type does not match type claim") + } + chatID, err := claims.ChatID() + if err != nil { + return err + } + if request.Meta.ChatID != chatID { + return xerrors.New("chat ID does not match subject claim") + } + return nil +} + +func canonicalAudience(audience string) string { + parsed, err := url.Parse(audience) + if err != nil { + return audience + } + if parsed.Path == "/" && parsed.RawPath == "" { + parsed.Path = "" + } + return parsed.String() +} + +func dispatch(ctx context.Context, hooks Hooks, request Request) (Response, error) { + switch request.Type { + case EventSessionStart: + return dispatchHook(ctx, request, hooks.SessionStart) + case EventUserPromptSubmit: + return dispatchHook(ctx, request, hooks.UserPromptSubmit) + case EventPreToolUse: + return dispatchHook(ctx, request, hooks.PreToolUse) + case EventPostToolUse: + return dispatchHook(ctx, request, hooks.PostToolUse) + case EventPreCompact: + return dispatchHook(ctx, request, hooks.PreCompact) + case EventPostCompact: + return dispatchHook(ctx, request, hooks.PostCompact) + case EventStop: + return dispatchHook(ctx, request, hooks.Stop) + default: + return Response{}, xerrors.Errorf("unknown event type %q", request.Type) + } +} + +func dispatchHook[T any](ctx context.Context, request Request, hook func(context.Context, Meta, T) (Response, error)) (Response, error) { + if hook == nil { + return Response{}, nil + } + var data T + if err := json.Unmarshal(request.Data, &data); err != nil { + return Response{}, xerrors.Errorf("decode %q event data: %w", request.Type, err) + } + return hook(ctx, request.Meta, data) +} diff --git a/codersdk/x/agenthooks/jwt.go b/codersdk/x/agenthooks/jwt.go new file mode 100644 index 0000000000..71d32a2a4f --- /dev/null +++ b/codersdk/x/agenthooks/jwt.go @@ -0,0 +1,137 @@ +package agenthooks + +import ( + "encoding/hex" + "encoding/json" + "strings" + "time" + + "github.com/go-jose/go-jose/v4" + "github.com/google/uuid" + "golang.org/x/xerrors" +) + +const jwtType = "JWT" + +// MinSecretLen is the minimum HS256 secret length in bytes. go-jose +// accepts shorter keys, so signing and verification enforce it to fail +// closed on missing or weak secrets. +const MinSecretLen = 32 + +// SignClaims signs claims with the shared secret using HS256. +func SignClaims(secret []byte, claims Claims) (string, error) { + if len(secret) < MinSecretLen { + return "", xerrors.Errorf("secret must be at least %d bytes", MinSecretLen) + } + signer, err := jose.NewSigner( + jose.SigningKey{Algorithm: jose.HS256, Key: secret}, + new(jose.SignerOptions).WithType(jwtType), + ) + if err != nil { + return "", xerrors.Errorf("create signer: %w", err) + } + + payload, err := json.Marshal(claims) + if err != nil { + return "", xerrors.Errorf("marshal claims: %w", err) + } + signed, err := signer.Sign(payload) + if err != nil { + return "", xerrors.Errorf("sign claims: %w", err) + } + token, err := signed.CompactSerialize() + if err != nil { + return "", xerrors.Errorf("serialize token: %w", err) + } + return token, nil +} + +// Verify authenticates an HS256 bearer token and validates its JWT header, +// required claims, and validity window. Request binding remains the caller's +// responsibility; see NewHTTPHandler. +func Verify(authzHeader string, secret []byte) (Claims, error) { + if len(secret) < MinSecretLen { + return Claims{}, xerrors.Errorf("secret must be at least %d bytes", MinSecretLen) + } + const bearerPrefix = "Bearer " + token, ok := strings.CutPrefix(authzHeader, bearerPrefix) + if !ok || token == "" || strings.ContainsAny(token, " \t\r\n") { + return Claims{}, xerrors.New("authorization header must contain one Bearer token") + } + + object, err := jose.ParseSigned(token, []jose.SignatureAlgorithm{jose.HS256}) + if err != nil { + return Claims{}, xerrors.Errorf("parse token: %w", err) + } + if len(object.Signatures) != 1 { + return Claims{}, xerrors.New("token must contain one signature") + } + header := object.Signatures[0].Header + typ, ok := header.ExtraHeaders[jose.HeaderType].(string) + if !ok || typ != jwtType { + return Claims{}, xerrors.Errorf("token type must be %q", jwtType) + } + + payload, err := object.Verify(secret) + if err != nil { + return Claims{}, xerrors.Errorf("verify token: %w", err) + } + var claims Claims + if err := json.Unmarshal(payload, &claims); err != nil { + return Claims{}, xerrors.Errorf("decode claims: %w", err) + } + if err := validateClaims(claims, time.Now()); err != nil { + return Claims{}, err + } + return claims, nil +} + +func validateClaims(claims Claims, now time.Time) error { + switch { + case claims.Issuer == "": + return xerrors.New("issuer is required") + case claims.Subject == "": + return xerrors.New("subject is required") + case claims.Audience == "": + return xerrors.New("audience is required") + case claims.IssuedAt == 0: + return xerrors.New("issued at is required") + case claims.NotBefore == 0: + return xerrors.New("not before is required") + case claims.Expires == 0: + return xerrors.New("expiry is required") + case claims.JTI == uuid.Nil: + return xerrors.New("JWT ID is required") + case !validEventType(claims.Type): + return xerrors.Errorf("invalid event type %q", claims.Type) + case !validSHA256(claims.BodySHA256): + return xerrors.New("body SHA-256 must be a hexadecimal SHA-256 digest") + case claims.NotBefore > claims.Expires: + return xerrors.New("not before must not be after expiry") + case claims.IssuedAt > claims.Expires: + return xerrors.New("issued at must not be after expiry") + case now.Unix() < claims.NotBefore: + return xerrors.New("token is not valid yet") + case now.Unix() >= claims.Expires: + return xerrors.New("token has expired") + } + if _, err := claims.ChatID(); err != nil { + return err + } + return nil +} + +func validEventType(eventType EventType) bool { + switch eventType { + case EventSessionStart, EventUserPromptSubmit, EventPreToolUse, + EventPostToolUse, EventPreCompact, EventPostCompact, EventStop: + return true + default: + return false + } +} + +func validSHA256(value string) bool { + decoded, err := hex.DecodeString(value) + return err == nil && len(decoded) == 32 +} diff --git a/codersdk/x/agenthooks/types.go b/codersdk/x/agenthooks/types.go new file mode 100644 index 0000000000..5277e15bca --- /dev/null +++ b/codersdk/x/agenthooks/types.go @@ -0,0 +1,154 @@ +// Package agenthooks defines the experimental wire protocol for Coder agent +// lifecycle hooks. The protocol, including SchemaVersion 1, has no +// backward-compatibility guarantee. +// +// Coder persists no hook dispatch state, so delivery is best-effort and may +// duplicate. A failed dispatch is never queued for redelivery; hooks are fail +// closed, so the operation that raised the event fails instead. Consumers must +// therefore tolerate duplicates without assuming every event arrives. +// +// A retried HTTP attempt reuses its Meta.DispatchID, so consumers deduplicate +// transport retries by that ID. A repeated logical event gets a new ID, so +// keep side effects keyed on the event's own identifiers, such as tool_use_id, +// or make them safe to repeat. +package agenthooks + +import ( + "encoding/json" + "strings" + + "github.com/google/uuid" + "golang.org/x/xerrors" +) + +// SchemaVersion is the current lifecycle hook request schema version. +const SchemaVersion = 1 + +// EventType names a lifecycle event carried by a hook request. +type EventType string + +const ( + EventSessionStart EventType = "session_start" + EventUserPromptSubmit EventType = "user_prompt_submit" + EventPreToolUse EventType = "pre_tool_use" + EventPostToolUse EventType = "post_tool_use" + EventPreCompact EventType = "pre_compact" + EventPostCompact EventType = "post_compact" + EventStop EventType = "stop" +) + +// Request is the body coderd posts to the configured lifecycle hook URL. +type Request struct { + Type EventType `json:"type"` + Meta Meta `json:"meta"` + Data json.RawMessage `json:"data"` +} + +// Meta identifies a hook dispatch and its chat. +type Meta struct { + DispatchID uuid.UUID `json:"dispatch_id"` + SchemaVersion int `json:"schema_version"` + ChatRef +} + +// ChatRef identifies the chat a lifecycle hook event refers to. +type ChatRef struct { + ChatID uuid.UUID `json:"chat_id"` + OwnerID uuid.UUID `json:"owner_id"` + WorkspaceID *uuid.UUID `json:"workspace_id,omitempty"` + TurnID *uuid.UUID `json:"turn_id,omitempty"` + ParentChatID *uuid.UUID `json:"parent_chat_id,omitempty"` + // RootChatID identifies the user-facing root of the chat tree. + RootChatID *uuid.UUID `json:"root_chat_id,omitempty"` +} + +// SessionStartData reports why a chat session started. Source is +// "startup", "resume", or "clear". +type SessionStartData struct { + Source string `json:"source"` +} + +// UserPromptSubmitData includes concatenated text and persisted parts. +// Inspect Parts when structure matters. +type UserPromptSubmitData struct { + Prompt string `json:"prompt"` + Parts json.RawMessage `json:"parts,omitempty"` +} + +// PreToolUseData describes a tool call before execution. +type PreToolUseData struct { + ToolUseID string `json:"tool_use_id"` + ToolName string `json:"tool_name"` + ToolInput json.RawMessage `json:"tool_input"` +} + +// PostToolUseData describes a completed tool call, carrying either +// ToolResponse or ToolError. +type PostToolUseData struct { + ToolUseID string `json:"tool_use_id"` + ToolName string `json:"tool_name"` + ToolResponse json.RawMessage `json:"tool_response,omitempty"` + ToolError string `json:"tool_error,omitempty"` +} + +// PreCompactData is empty; Meta identifies the chat being compacted. +type PreCompactData struct{} + +// PostCompactData is empty; Meta identifies the compacted chat. +type PostCompactData struct{} + +// StopData is empty; Meta identifies the chat that stopped. +type StopData struct{} + +// Response carries a consumer's decision and optional injected content. +// Permission is honored for user_prompt_submit and pre_tool_use only. +// user_prompt_submit folds injected content into the submitted message. +// A denied pre_tool_use yields a synthetic tool result carrying only the +// policy text and any Reason; ModelContext persists separately as +// model-only transcript content that never reaches clients. +type Response struct { + Permission *Permission `json:"permission,omitempty"` + ModelContext string `json:"model_context,omitempty"` + UserMessage string `json:"user_message,omitempty"` +} + +// Permission controls whether mutable hook input may proceed. +type Permission struct { + Decision PermissionDecision `json:"decision"` + Reason string `json:"reason,omitempty"` + InputOverride json.RawMessage `json:"input_override,omitempty"` +} + +// PermissionDecision is a consumer's verdict on mutable hook input. +type PermissionDecision string + +const ( + PermissionAllow PermissionDecision = "allow" + PermissionDeny PermissionDecision = "deny" +) + +// Claims describes the JWT minted by coderd for a lifecycle hook dispatch. +type Claims struct { + Issuer string `json:"iss"` + Subject string `json:"sub"` + Audience string `json:"aud"` + IssuedAt int64 `json:"iat"` + NotBefore int64 `json:"nbf"` + Expires int64 `json:"exp"` + JTI uuid.UUID `json:"jti"` + Type EventType `json:"type"` + BodySHA256 string `json:"body_sha256"` +} + +// ChatID returns the chat ID encoded in the "coder:chat:" subject. +func (c Claims) ChatID() (uuid.UUID, error) { + value, ok := strings.CutPrefix(c.Subject, "coder:chat:") + if !ok { + return uuid.Nil, xerrors.Errorf("invalid subject %q", c.Subject) + } + chatID, err := uuid.Parse(value) + if err != nil { + return uuid.Nil, xerrors.Errorf("parse chat ID: %w", err) + } + return chatID, nil +} diff --git a/docs/admin/integrations/prometheus.md b/docs/admin/integrations/prometheus.md index 26c16449ca..fc913277ae 100644 --- a/docs/admin/integrations/prometheus.md +++ b/docs/admin/integrations/prometheus.md @@ -229,6 +229,11 @@ deployment. They will always be available from the agent. | `coderd_chat_auto_archive_records_archived_total` | counter | Total number of chats archived by the auto-archive job (counting both roots and cascaded children). | | | `coderd_chatd_chats` | gauge | Number of chats being processed, by state. | `state` | | `coderd_chatd_compaction_total` | counter | Total compaction outcomes (only recorded when compaction was triggered or failed). | `model` `provider` `result` | +| `coderd_chatd_hook_context_size_bytes` | histogram | Lifecycle hook model context response size in bytes. | `event` | +| `coderd_chatd_hook_decisions_total` | counter | Total lifecycle hook permission decisions by event and decision. | `decision` `event` | +| `coderd_chatd_hook_dispatch_seconds` | histogram | Lifecycle hook dispatch duration in seconds. | `event` | +| `coderd_chatd_hook_dispatches_total` | counter | Total lifecycle hook dispatches by event and result. | `event` `result` | +| `coderd_chatd_hook_input_overrides_total` | counter | Total lifecycle hook input overrides by event. | `event` | | `coderd_chatd_message_count` | histogram | Number of messages in the prompt per LLM request. | `model` `provider` | | `coderd_chatd_prompt_size_bytes` | histogram | Estimated byte size of the prompt per LLM request. | `model` `provider` | | `coderd_chatd_steps_total` | counter | Total agentic loop steps across all chats. | `model` `provider` | diff --git a/scripts/agenthooks-server/main.go b/scripts/agenthooks-server/main.go new file mode 100644 index 0000000000..21c5cb6045 --- /dev/null +++ b/scripts/agenthooks-server/main.go @@ -0,0 +1,323 @@ +// agenthooks-server is a reference consumer that logs verified lifecycle +// events as JSON. +package main + +import ( + "context" + "encoding/json" + "errors" + "flag" + "fmt" + "net/http" + "os" + "os/signal" + "regexp" + "strconv" + "sync" + "syscall" + "time" + + "golang.org/x/xerrors" + + "github.com/coder/coder/v2/codersdk/x/agenthooks" +) + +type config struct { + listen string + audience string + secret string + issuer string + tlsCert string + tlsKey string + logOnly bool + denyToolPattern string + redactPrompt string +} + +type eventLog struct { + Event agenthooks.EventType `json:"event"` + DispatchID string `json:"dispatch_id"` + ChatID string `json:"chat_id"` + TurnID string `json:"turn_id,omitempty"` + ToolUseID string `json:"tool_use_id,omitempty"` + ToolName string `json:"tool_name,omitempty"` + Source string `json:"source,omitempty"` + Prompt string `json:"prompt,omitempty"` + ToolInput json.RawMessage `json:"tool_input,omitempty"` + ToolOutput json.RawMessage `json:"tool_output,omitempty"` + ToolError string `json:"tool_error,omitempty"` + Duplicate bool `json:"duplicate,omitempty"` +} + +// consumerState demonstrates consumer-owned hook state. Coder persists no +// hook decisions and delivery is best-effort, so consumers that need memory +// keep it themselves, keyed by the stable payload identifiers: chat_id, the +// event type, and tool_use_id. +type consumerState struct { + mu sync.Mutex + // preToolDecisions reuses responses for duplicate (chat_id, tool_use_id) + // deliveries while they remain cached. + preToolDecisions map[string]agenthooks.Response + // blockedTools records tool names this consumer denied per chat, so + // the policy outlives any single dispatch. Evicted with preToolDecisions. + blockedTools map[string]map[string]struct{} +} + +const maxRememberedDecisions = 8192 + +func newConsumerState() *consumerState { + return &consumerState{ + preToolDecisions: make(map[string]agenthooks.Response), + blockedTools: make(map[string]map[string]struct{}), + } +} + +// decidePreToolUse resolves one pre_tool_use delivery and reports whether the +// response was already remembered. The lookup, the policy decision, and the +// store share one lock, so concurrent duplicate deliveries of the same +// tool_use_id cannot each decide and then overwrite each other. +func (s *consumerState) decidePreToolUse(chatID, toolUseID, toolName string, logOnly bool, denyTool *regexp.Regexp) (agenthooks.Response, bool) { + s.mu.Lock() + defer s.mu.Unlock() + key := chatID + "\x00" + toolUseID + if response, ok := s.preToolDecisions[key]; ok { + return response, true + } + + var response agenthooks.Response + deniedTool := "" + switch { + case logOnly: + case s.isBlockedLocked(chatID, toolName): + response = agenthooks.Response{Permission: &agenthooks.Permission{ + Decision: agenthooks.PermissionDeny, + Reason: "use of this tool is blocked for this chat", + }} + case denyTool != nil && denyTool.MatchString(toolName): + deniedTool = toolName + response = agenthooks.Response{Permission: &agenthooks.Permission{ + Decision: agenthooks.PermissionDeny, + Reason: "use of this tool is denied by this deployment's policy", + }} + } + + if len(s.preToolDecisions) >= maxRememberedDecisions { + // Both maps grow per chat, so evict them together to keep a + // long-running consumer bounded. + s.preToolDecisions = make(map[string]agenthooks.Response) + s.blockedTools = make(map[string]map[string]struct{}) + } + s.preToolDecisions[key] = response + if deniedTool != "" { + blocked := s.blockedTools[chatID] + if blocked == nil { + blocked = make(map[string]struct{}) + s.blockedTools[chatID] = blocked + } + blocked[deniedTool] = struct{}{} + } + return response, false +} + +func (s *consumerState) isBlockedLocked(chatID, toolName string) bool { + _, ok := s.blockedTools[chatID][toolName] + return ok +} + +func main() { + if err := run(); err != nil { + _, _ = fmt.Fprintf(os.Stderr, "Error: %v\n", err) + os.Exit(1) + } +} + +func run() error { + cfg, err := parseFlags() + if err != nil { + return err + } + if cfg.secret == "" { + return xerrors.New("secret is required through --secret or CODER_AGENTHOOKS_SECRET") + } + // The bind address is not the audience: it may be a wildcard or an + // ephemeral port, and behind a proxy the signed audience is the proxy URL. + if cfg.audience == "" { + return xerrors.New("audience is required through --audience or CODER_AGENTHOOKS_AUDIENCE") + } + if len(cfg.secret) < agenthooks.MinSecretLen { + return xerrors.Errorf("secret must be at least %d bytes", agenthooks.MinSecretLen) + } + if (cfg.tlsCert == "") != (cfg.tlsKey == "") { + return xerrors.New("TLS certificate and key must be configured together") + } + + var denyTool *regexp.Regexp + if cfg.denyToolPattern != "" { + denyTool, err = regexp.Compile(cfg.denyToolPattern) + if err != nil { + return xerrors.Errorf("compile deny tool pattern: %w", err) + } + } + var redactPrompt *regexp.Regexp + if cfg.redactPrompt != "" { + redactPrompt, err = regexp.Compile(cfg.redactPrompt) + if err != nil { + return xerrors.Errorf("compile redact prompt pattern: %w", err) + } + } + + state := newConsumerState() + var logMu sync.Mutex + encoder := json.NewEncoder(os.Stdout) + logEvent := func(event eventLog) error { + logMu.Lock() + defer logMu.Unlock() + if err := encoder.Encode(event); err != nil { + return xerrors.Errorf("encode event: %w", err) + } + return nil + } + baseEvent := func(event agenthooks.EventType, meta agenthooks.Meta) eventLog { + entry := eventLog{ + Event: event, + DispatchID: meta.DispatchID.String(), + ChatID: meta.ChatID.String(), + } + if meta.TurnID != nil { + entry.TurnID = meta.TurnID.String() + } + return entry + } + + consumerHooks := agenthooks.Hooks{ + SessionStart: func(_ context.Context, meta agenthooks.Meta, data agenthooks.SessionStartData) (agenthooks.Response, error) { + entry := baseEvent(agenthooks.EventSessionStart, meta) + entry.Source = data.Source + return agenthooks.Response{}, logEvent(entry) + }, + UserPromptSubmit: func(_ context.Context, meta agenthooks.Meta, data agenthooks.UserPromptSubmitData) (agenthooks.Response, error) { + entry := baseEvent(agenthooks.EventUserPromptSubmit, meta) + entry.Prompt = data.Prompt + matches := redactPrompt != nil && redactPrompt.MatchString(data.Prompt) + if matches { + entry.Prompt = redactPrompt.ReplaceAllString(data.Prompt, "[REDACTED]") + } + if err := logEvent(entry); err != nil { + return agenthooks.Response{}, err + } + // Log-only mode still redacts the log entry above; it only + // suppresses the prompt override response. + if cfg.logOnly || !matches { + return agenthooks.Response{}, nil + } + override, err := json.Marshal(map[string]string{"prompt": entry.Prompt}) + if err != nil { + return agenthooks.Response{}, xerrors.Errorf("marshal prompt override: %w", err) + } + return agenthooks.Response{Permission: &agenthooks.Permission{ + Decision: agenthooks.PermissionAllow, + InputOverride: override, + }}, nil + }, + PreToolUse: func(_ context.Context, meta agenthooks.Meta, data agenthooks.PreToolUseData) (agenthooks.Response, error) { + entry := baseEvent(agenthooks.EventPreToolUse, meta) + entry.ToolUseID = data.ToolUseID + entry.ToolName = data.ToolName + entry.ToolInput = data.ToolInput + response, duplicate := state.decidePreToolUse(entry.ChatID, data.ToolUseID, data.ToolName, cfg.logOnly, denyTool) + entry.Duplicate = duplicate + return response, logEvent(entry) + }, + PostToolUse: func(_ context.Context, meta agenthooks.Meta, data agenthooks.PostToolUseData) (agenthooks.Response, error) { + entry := baseEvent(agenthooks.EventPostToolUse, meta) + entry.ToolUseID = data.ToolUseID + entry.ToolName = data.ToolName + entry.ToolOutput = data.ToolResponse + entry.ToolError = data.ToolError + return agenthooks.Response{}, logEvent(entry) + }, + PreCompact: func(_ context.Context, meta agenthooks.Meta, _ agenthooks.PreCompactData) (agenthooks.Response, error) { + return agenthooks.Response{}, logEvent(baseEvent(agenthooks.EventPreCompact, meta)) + }, + PostCompact: func(_ context.Context, meta agenthooks.Meta, _ agenthooks.PostCompactData) (agenthooks.Response, error) { + return agenthooks.Response{}, logEvent(baseEvent(agenthooks.EventPostCompact, meta)) + }, + Stop: func(_ context.Context, meta agenthooks.Meta, _ agenthooks.StopData) (agenthooks.Response, error) { + return agenthooks.Response{}, logEvent(baseEvent(agenthooks.EventStop, meta)) + }, + } + + var handlerOpts []agenthooks.HandlerOption + if cfg.issuer != "" { + handlerOpts = append(handlerOpts, agenthooks.WithExpectedIssuer(cfg.issuer)) + } + handler := agenthooks.NewHTTPHandler([]byte(cfg.secret), cfg.audience, consumerHooks, handlerOpts...) + server := &http.Server{ + Addr: cfg.listen, + Handler: handler, + ReadHeaderTimeout: 10 * time.Second, + } + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + go func() { + <-ctx.Done() + _ = server.Close() + }() + + mode := "enforcing" + if cfg.logOnly { + mode = "log-only" + } + _, _ = fmt.Fprintf(os.Stderr, "Agent hooks server listening on %s in %s mode\n", cfg.listen, mode) + if cfg.logOnly && (cfg.denyToolPattern != "" || cfg.redactPrompt != "") { + _, _ = fmt.Fprintln(os.Stderr, "Warning: log-only mode ignores -deny-tool-pattern and -redact-prompt-pattern; pass -log-only=false to act on them") + } + if cfg.tlsCert != "" { + err = server.ListenAndServeTLS(cfg.tlsCert, cfg.tlsKey) + } else { + err = server.ListenAndServe() + } + if err != nil && !errors.Is(err, http.ErrServerClosed) { + return xerrors.Errorf("serve lifecycle hooks: %w", err) + } + return nil +} + +func parseFlags() (config, error) { + logOnly, err := envBool("CODER_AGENTHOOKS_LOG_ONLY", true) + if err != nil { + return config{}, err + } + var cfg config + cfg.logOnly = logOnly + flag.StringVar(&cfg.listen, "listen", envOrDefault("CODER_AGENTHOOKS_LISTEN", "127.0.0.1:8081"), "Listen address (CODER_AGENTHOOKS_LISTEN)") + flag.StringVar(&cfg.audience, "audience", os.Getenv("CODER_AGENTHOOKS_AUDIENCE"), "Expected aud claim, which is the deployment's CODER_CHAT_HOOK_URL, required (CODER_AGENTHOOKS_AUDIENCE)") + flag.StringVar(&cfg.secret, "secret", os.Getenv("CODER_AGENTHOOKS_SECRET"), "Shared HS256 secret, required (CODER_AGENTHOOKS_SECRET)") + flag.StringVar(&cfg.issuer, "issuer", os.Getenv("CODER_AGENTHOOKS_ISSUER"), "Expected iss claim, normally the Coder deployment ID (CODER_AGENTHOOKS_ISSUER)") + flag.StringVar(&cfg.tlsCert, "tls-cert", os.Getenv("CODER_AGENTHOOKS_TLS_CERT"), "TLS certificate path (CODER_AGENTHOOKS_TLS_CERT)") + flag.StringVar(&cfg.tlsKey, "tls-key", os.Getenv("CODER_AGENTHOOKS_TLS_KEY"), "TLS private key path (CODER_AGENTHOOKS_TLS_KEY)") + flag.BoolVar(&cfg.logOnly, "log-only", cfg.logOnly, "Return an empty response for every event, on by default (CODER_AGENTHOOKS_LOG_ONLY)") + flag.StringVar(&cfg.denyToolPattern, "deny-tool-pattern", os.Getenv("CODER_AGENTHOOKS_DENY_TOOL_PATTERN"), "Example regexp for denied tool names, requires -log-only=false (CODER_AGENTHOOKS_DENY_TOOL_PATTERN)") + flag.StringVar(&cfg.redactPrompt, "redact-prompt-pattern", os.Getenv("CODER_AGENTHOOKS_REDACT_PROMPT_PATTERN"), "Example regexp to redact in prompts, requires -log-only=false to override the prompt (CODER_AGENTHOOKS_REDACT_PROMPT_PATTERN)") + flag.Parse() + return cfg, nil +} + +func envOrDefault(name, fallback string) string { + if value := os.Getenv(name); value != "" { + return value + } + return fallback +} + +func envBool(name string, fallback bool) (bool, error) { + value := os.Getenv(name) + if value == "" { + return fallback, nil + } + parsed, err := strconv.ParseBool(value) + if err != nil { + return false, xerrors.Errorf("parse %s: %w", name, err) + } + return parsed, nil +} diff --git a/scripts/apitypings/main.go b/scripts/apitypings/main.go index 77c648a050..0910fad145 100644 --- a/scripts/apitypings/main.go +++ b/scripts/apitypings/main.go @@ -26,6 +26,7 @@ func main() { generateDirectories := map[string]string{ "github.com/coder/coder/v2/codersdk": "", "github.com/coder/coder/v2/coderd/healthcheck/health": "Health", + "github.com/coder/coder/v2/codersdk/x/agenthooks": "AgentHook", "github.com/coder/coder/v2/codersdk/healthsdk": "", } for dir, prefix := range generateDirectories { @@ -78,6 +79,7 @@ func TSMutations(ts *guts.Typescript) { config.NotNullMaps, FixSerpentStruct, DiscriminatedChatMessagePart, + AgentHookRawMessages, // Prefer enums as types config.EnumAsTypes, // Enum list generator @@ -146,6 +148,43 @@ func TypeMappings(gen *guts.GoParser) error { return nil } +// AgentHookRawMessages maps agent-hook raw JSON fields to unknown instead of +// the global object type. +func AgentHookRawMessages(ts *guts.Typescript) { + if _, ok := ts.Node("AgentHookRequest"); !ok { + return + } + unknown := bindings.KeywordUnknown + fields := map[string]string{ + "AgentHookRequest": "data", + "AgentHookUserPromptSubmitData": "parts", + "AgentHookPreToolUseData": "tool_input", + "AgentHookPostToolUseData": "tool_response", + "AgentHookPermission": "input_override", + } + for typeName, fieldName := range fields { + node, ok := ts.Node(typeName) + if !ok { + panic(fmt.Sprintf("agent hook type %q was not generated", typeName)) + } + iface, ok := node.(*bindings.Interface) + if !ok { + panic(fmt.Sprintf("agent hook type %q is not an interface", typeName)) + } + found := false + for _, field := range iface.Fields { + if field.Name == fieldName { + field.Type = &unknown + found = true + break + } + } + if !found { + panic(fmt.Sprintf("agent hook field %q.%s was not generated", typeName, fieldName)) + } + } +} + // DiscriminatedChatMessagePart splits the flat ChatMessagePart // interface into a discriminated union of per-type sub-interfaces. // Each sub-interface narrows the `type` field to a string literal diff --git a/scripts/metricsdocgen/generated_metrics b/scripts/metricsdocgen/generated_metrics index 65e613adf6..a4b45ef0f2 100644 --- a/scripts/metricsdocgen/generated_metrics +++ b/scripts/metricsdocgen/generated_metrics @@ -277,6 +277,21 @@ coderd_chatd_chats{state=""} 0 # HELP coderd_chatd_compaction_total Total compaction outcomes (only recorded when compaction was triggered or failed). # TYPE coderd_chatd_compaction_total counter coderd_chatd_compaction_total{provider="",model="",result=""} 0 +# HELP coderd_chatd_hook_context_size_bytes Lifecycle hook model context response size in bytes. +# TYPE coderd_chatd_hook_context_size_bytes histogram +coderd_chatd_hook_context_size_bytes{event=""} 0 +# HELP coderd_chatd_hook_decisions_total Total lifecycle hook permission decisions by event and decision. +# TYPE coderd_chatd_hook_decisions_total counter +coderd_chatd_hook_decisions_total{event="",decision=""} 0 +# HELP coderd_chatd_hook_dispatch_seconds Lifecycle hook dispatch duration in seconds. +# TYPE coderd_chatd_hook_dispatch_seconds histogram +coderd_chatd_hook_dispatch_seconds{event=""} 0 +# HELP coderd_chatd_hook_dispatches_total Total lifecycle hook dispatches by event and result. +# TYPE coderd_chatd_hook_dispatches_total counter +coderd_chatd_hook_dispatches_total{event="",result=""} 0 +# HELP coderd_chatd_hook_input_overrides_total Total lifecycle hook input overrides by event. +# TYPE coderd_chatd_hook_input_overrides_total counter +coderd_chatd_hook_input_overrides_total{event=""} 0 # HELP coderd_chatd_message_count Number of messages in the prompt per LLM request. # TYPE coderd_chatd_message_count histogram coderd_chatd_message_count{provider="",model=""} 0 diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index 543a5d9dd8..322bad1b6d 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -1217,6 +1217,210 @@ export interface AgentFirewallSessionLogsResponse { readonly results: readonly AgentFirewallLog[]; } +// From agenthooks/types.go +/** + * ChatRef identifies the chat a lifecycle hook event refers to. + */ +export interface AgentHookChatRef { + readonly chat_id: string; + readonly owner_id: string; + readonly workspace_id?: string; + readonly turn_id?: string; + readonly parent_chat_id?: string; + /** + * RootChatID identifies the user-facing root of the chat tree. + */ + readonly root_chat_id?: string; +} + +// From agenthooks/types.go +/** + * Claims describes the JWT minted by coderd for a lifecycle hook dispatch. + */ +export interface AgentHookClaims { + readonly iss: string; + readonly sub: string; + readonly aud: string; + readonly iat: number; + readonly nbf: number; + readonly exp: number; + readonly jti: string; + readonly type: AgentHookEventType; + readonly body_sha256: string; +} + +// From agenthooks/types.go +export type AgentHookEventType = + | "post_compact" + | "post_tool_use" + | "pre_compact" + | "pre_tool_use" + | "session_start" + | "stop" + | "user_prompt_submit"; + +export const AgentHookEventTypes: AgentHookEventType[] = [ + "post_compact", + "post_tool_use", + "pre_compact", + "pre_tool_use", + "session_start", + "stop", + "user_prompt_submit", +]; + +// From agenthooks/http.go +/** + * Hooks lets a consumer implement only the lifecycle events it uses. + */ +export interface AgentHookHooks { + // Function type detected, and unsupported. Leaving the type as unknown + readonly SessionStart: unknown; + // Function type detected, and unsupported. Leaving the type as unknown + readonly UserPromptSubmit: unknown; + // Function type detected, and unsupported. Leaving the type as unknown + readonly PreToolUse: unknown; + // Function type detected, and unsupported. Leaving the type as unknown + readonly PostToolUse: unknown; + // Function type detected, and unsupported. Leaving the type as unknown + readonly PreCompact: unknown; + // Function type detected, and unsupported. Leaving the type as unknown + readonly PostCompact: unknown; + // Function type detected, and unsupported. Leaving the type as unknown + readonly Stop: unknown; +} + +// From agenthooks/http.go +/** + * MaxRequestBodyBytes limits memory used to verify hook requests. + */ +export const AgentHookMaxRequestBodyBytes = 10485760; // 10 MiB + +// From agenthooks/types.go +/** + * Meta identifies a hook dispatch and its chat. + */ +export interface AgentHookMeta extends AgentHookChatRef { + readonly dispatch_id: string; + readonly schema_version: number; +} + +// From agenthooks/jwt.go +/** + * MinSecretLen is the minimum HS256 secret length in bytes. go-jose + * accepts shorter keys, so signing and verification enforce it to fail + * closed on missing or weak secrets. + */ +export const AgentHookMinSecretLen = 32; + +// From agenthooks/types.go +/** + * Permission controls whether mutable hook input may proceed. + */ +export interface AgentHookPermission { + readonly decision: AgentHookPermissionDecision; + readonly reason?: string; + readonly input_override?: unknown; +} + +// From agenthooks/types.go +export type AgentHookPermissionDecision = "allow" | "deny"; + +export const AgentHookPermissionDecisions: AgentHookPermissionDecision[] = [ + "allow", + "deny", +]; + +// From agenthooks/types.go +/** + * PostCompactData is empty; Meta identifies the compacted chat. + */ +export interface AgentHookPostCompactData {} + +// From agenthooks/types.go +/** + * PostToolUseData describes a completed tool call, carrying either + * ToolResponse or ToolError. + */ +export interface AgentHookPostToolUseData { + readonly tool_use_id: string; + readonly tool_name: string; + readonly tool_response?: unknown; + readonly tool_error?: string; +} + +// From agenthooks/types.go +/** + * PreCompactData is empty; Meta identifies the chat being compacted. + */ +export interface AgentHookPreCompactData {} + +// From agenthooks/types.go +/** + * PreToolUseData describes a tool call before execution. + */ +export interface AgentHookPreToolUseData { + readonly tool_use_id: string; + readonly tool_name: string; + readonly tool_input: unknown; +} + +// From agenthooks/types.go +/** + * Request is the body coderd posts to the configured lifecycle hook URL. + */ +export interface AgentHookRequest { + readonly type: AgentHookEventType; + readonly meta: AgentHookMeta; + readonly data: unknown; +} + +// From agenthooks/types.go +/** + * Response carries a consumer's decision and optional injected content. + * Permission is honored for user_prompt_submit and pre_tool_use only. + * user_prompt_submit folds injected content into the submitted message. + * A denied pre_tool_use yields a synthetic tool result carrying only the + * policy text and any Reason; ModelContext persists separately as + * model-only transcript content that never reaches clients. + */ +export interface AgentHookResponse { + readonly permission?: AgentHookPermission; + readonly model_context?: string; + readonly user_message?: string; +} + +// From agenthooks/types.go +/** + * SchemaVersion is the current lifecycle hook request schema version. + */ +export const AgentHookSchemaVersion = 1; + +// From agenthooks/types.go +/** + * SessionStartData reports why a chat session started. Source is + * "startup", "resume", or "clear". + */ +export interface AgentHookSessionStartData { + readonly source: string; +} + +// From agenthooks/types.go +/** + * StopData is empty; Meta identifies the chat that stopped. + */ +export interface AgentHookStopData {} + +// From agenthooks/types.go +/** + * UserPromptSubmitData includes concatenated text and persisted parts. + * Inspect Parts when structure matters. + */ +export interface AgentHookUserPromptSubmitData { + readonly prompt: string; + readonly parts?: unknown; +} + // From codersdk/workspacebuilds.go export interface AgentScriptTiming { readonly started_at: string;