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;