mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add chat lifecycle hook dispatch backend (#27401)
Adds the chat lifecycle hook wire contract and dispatch plumbing, first PR of the lifecycle hooks stack (followed by #27428, #27429, #27430). - `codersdk/x/agenthooks`: event and response wire types, JWT creation and verification with the shared secret (HS256, request body digest, expiry and not-before freshness checks), and an HTTP handler helper so consumers only implement the events they use. The `codersdk/x` location marks the consumer SDK as experimental. - `coderd/x/agenthooks/dispatch`: a stateless dispatcher that signs and posts hook events, enforces a concurrency cap under one configured timeout that bounds both the capacity wait and both post attempts, retries one connection failure with the same JWT, sends a distinctive `coderd-agenthooks/<version>` User-Agent, and records Prometheus metrics. Delivery is at least once; consumers own durable decision state, audit records, and deduplication keyed by the stable payload identifiers. Nothing is persisted by Coder. - Response bodies decode strictly: unknown fields, duplicate JSON keys (including inside `input_override`), and trailing data fail the dispatch closed as protocol errors instead of silently reading as allow. - `coderd/util/xnet`: shared timeout and connection error classification used by the dispatcher retry logic. Transient HTTP/2 stream aborts count as connection errors, so the documented single retry also applies to h2 consumers, which is the shape Go's default transport negotiates against any TLS consumer. Deterministic protocol failures stay terminal. Only the struct form of a stream error is matched, because `net/http` bundles its own HTTP/2 types and `h2_error.go` bridges only that shape. - `scripts/agenthooks-server`: a reference consumer that logs events and demonstrates consumer-owned pre-tool decision deduplication. It requires an explicitly configured JWT audience rather than deriving one from the request, and its startup output names the mode it is running in so an operator can see that the example policy flags need `-log-only=false`. - `scripts/apitypings`: generate TypeScript types for the hook wire contract. Dispatch failures log without the error's stack frames, since a failed dispatch is an expected, operator-visible condition. Nothing dispatches these events yet; chatd wiring lands in #27429. > This PR was written by Mux, an AI coding agent, on Mike's behalf.
This commit is contained in:
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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:<id>" 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
|
||||
}
|
||||
Reference in New Issue
Block a user