diff --git a/docs/arch/03-transport-architecture.md b/docs/arch/03-transport-architecture.md index 4b0feb9a83..c85a7883dd 100644 --- a/docs/arch/03-transport-architecture.md +++ b/docs/arch/03-transport-architecture.md @@ -211,9 +211,9 @@ All proxy types integrate with the middleware chain: ```mermaid graph LR - Client[Client Request] --> MW1[Middleware 1
Auth] - MW1 --> MW2[Middleware 2
Parser] - MW2 --> MW3[Middleware 3
Audit] + Client[Client Request] --> MW1[Middleware 1
Audit] + MW1 --> MW2[Middleware 2
Auth] + MW2 --> MW3[Middleware 3
Parser] MW3 --> MW4[Middleware 4
Authz] MW4 --> Proxy[Proxy Handler] Proxy --> Container[MCP Server] diff --git a/docs/arch/10-virtual-mcp-architecture.md b/docs/arch/10-virtual-mcp-architecture.md index 883ce9e502..ae815141ab 100644 --- a/docs/arch/10-virtual-mcp-architecture.md +++ b/docs/arch/10-virtual-mcp-architecture.md @@ -426,8 +426,8 @@ Middleware is applied by wrapping handlers, so execution order is outer-to-inner | 1 | Recovery | Always | Catches panics, returns HTTP 500 | | 2 | WriteTimeout | Always | Clears the server `WriteTimeout` for qualifying SSE connections | | 3 | Header Validation | Always | Rejects GETs without `Accept: text/event-stream` before they reach the MCP handler | -| 4 | Authentication (+ MCP parsing) | Optional | Validates incoming credentials (OIDC/local/anonymous); MCP parsing is composed inside so downstream layers see `ParsedMCPRequest` | -| 5 | Audit | Optional | Logs request events for compliance | +| 4 | Audit | Optional | Logs every request outcome, including 401s from the auth middleware it wraps; identity and parsed MCP data flow back via the `auth.IdentityHolder` / `mcp.ParsedRequestHolder` carriers | +| 5 | Authentication (+ MCP parsing) | Optional | Validates incoming credentials (OIDC/local/anonymous); MCP parsing is composed inside so downstream layers see `ParsedMCPRequest` | | 6 | Discovery | Always | Aggregates backend capabilities per session | | 7 | Annotation Enrichment | Optional | Injects tool annotations into context for annotation-aware authz (only when Authorization is configured) | | 8 | Authorization | Optional | Evaluates Cedar policies after discovery and annotation enrichment | @@ -474,7 +474,7 @@ This enriches audit events with the backend name for better observability. The server wires them around discovery/annotation-enrichment so the effective execution order is: ``` -Authentication → MCP Parsing → Audit → Discovery → Annotation Enrichment → Authorization → Next Handler +Audit → Authentication → MCP Parsing → Discovery → Annotation Enrichment → Authorization → Next Handler ``` **Implementation**: `pkg/vmcp/server/server.go`, `pkg/vmcp/discovery/middleware.go`, `pkg/vmcp/auth/factory/` diff --git a/docs/middleware.md b/docs/middleware.md index 559d6b0cb2..619b816ef8 100644 --- a/docs/middleware.md +++ b/docs/middleware.md @@ -10,14 +10,14 @@ This document primarily covers the middleware system for `thv` and `thv-proxyrun The middleware chain consists of the following components: -1. **Authentication Middleware**: Validates JWT tokens and extracts client identity -2. **Upstream Token Swap Middleware**: Exchanges ToolHive JWTs for upstream IdP tokens (automatic with embedded auth server) -3. **Token Exchange Middleware**: Exchanges JWT tokens for external service tokens via OAuth 2.0 Token Exchange (optional) -4. **MCP Parsing Middleware**: Parses JSON-RPC MCP requests and extracts structured data -5. **Tool Mapping Middleware**: Enables tool filtering and override capabilities through two complementary middleware components that process outgoing `tools/list` responses and incoming `tools/call` requests (optional) -6. **Usage Metrics Middleware**: Collects anonymous usage metrics for ToolHive development (optional) -7. **Telemetry Middleware**: Instruments requests with OpenTelemetry (optional) -8. **Audit Middleware**: Logs request events for compliance and monitoring (optional) +1. **Audit Middleware**: Wraps the rest of the chain and logs every request outcome — including rejections from authentication, webhooks, and authorization (optional) +2. **Authentication Middleware**: Validates JWT tokens and extracts client identity +3. **Upstream Token Swap Middleware**: Exchanges ToolHive JWTs for upstream IdP tokens (automatic with embedded auth server) +4. **Token Exchange Middleware**: Exchanges JWT tokens for external service tokens via OAuth 2.0 Token Exchange (optional) +5. **MCP Parsing Middleware**: Parses JSON-RPC MCP requests and extracts structured data +6. **Tool Mapping Middleware**: Enables tool filtering and override capabilities through two complementary middleware components that process outgoing `tools/list` responses and incoming `tools/call` requests (optional) +7. **Usage Metrics Middleware**: Collects anonymous usage metrics for ToolHive development (optional) +8. **Telemetry Middleware**: Instruments requests with OpenTelemetry (optional) 9. **Authorization Middleware**: Evaluates Cedar policies to authorize requests (optional) 10. **Header Forward Middleware**: Injects custom headers into requests to remote MCP servers (optional) 11. **Recovery Middleware**: Catches panics and returns HTTP 500 errors (always present) @@ -33,12 +33,13 @@ Two webhook types are supported: When configured together, the effective order is: -1. Authentication -2. Token exchange and related auth middleware, when configured -3. MCP parsing -4. Mutating webhooks -5. Validating webhooks -6. Telemetry, audit, and authorization middleware +1. Audit (wraps everything below, so webhook denials are audited) +2. Authentication +3. Token exchange and related auth middleware, when configured +4. MCP parsing +5. Mutating webhooks +6. Validating webhooks +7. Telemetry and authorization middleware Multiple webhook definitions of the same type run in configuration order. When multiple `--webhook-config` files are provided, later files override earlier webhook definitions with the same `name`. @@ -59,10 +60,10 @@ Example config files: ```mermaid graph TD - A[Incoming MCP Request] --> B[Authentication Middleware] + A[Incoming MCP Request] --> E[Audit Middleware] + E --> B[Authentication Middleware] B --> C[MCP Parsing Middleware] - C --> E[Audit Middleware] - E --> D[Authorization Middleware] + C --> D[Authorization Middleware] D --> R[Recovery Middleware] R --> F[MCP Server Handler] @@ -101,29 +102,29 @@ graph TD ```mermaid sequenceDiagram participant Client + participant Audit as Audit participant Auth as Authentication participant Parser as MCP Parser - participant Audit as Audit participant Authz as Authorization participant Recovery as Recovery participant Server as MCP Server - Client->>Auth: HTTP Request with JWT + Client->>Audit: HTTP Request + Note over Audit: Injects identity / parsed-request holders,
logs the outcome after the inner chain returns Note over Recovery: Innermost wrapper: catches panics
from the handler and inner middleware + + Audit->>Auth: HTTP Request with JWT Auth->>Auth: Validate JWT Token Auth->>Auth: Extract Claims - Note over Auth: Add claims to context + Note over Auth: Add identity to context and holder Auth->>Parser: Request + JWT Claims Parser->>Parser: Parse JSON-RPC Parser->>Parser: Extract MCP Method Parser->>Parser: Extract Resource ID & Arguments - Note over Parser: Add parsed data to context - - Parser->>Audit: Request + Parsed MCP Data - Note over Audit: Wraps authorization so every
request is logged, including denials + Note over Parser: Add parsed data to context and holder - Audit->>Authz: Request + Parser->>Authz: Request + Parsed MCP Data Authz->>Authz: Get Parsed Data from Context Authz->>Authz: Create Cedar Entities Authz->>Authz: Evaluate Policies @@ -133,7 +134,11 @@ sequenceDiagram Server->>Audit: Response Audit->>Audit: Log Audit Event (outcome success) Audit->>Client: Response - else Unauthorized + else Authentication fails + Auth->>Audit: 401 Unauthorized + Audit->>Audit: Log Audit Event (outcome denied) + Audit->>Client: 401 Unauthorized + else Denied by policy Authz->>Audit: 403 Forbidden Audit->>Audit: Log Audit Event (outcome denied) Audit->>Client: 403 Forbidden @@ -391,7 +396,7 @@ thv config usage-metrics enable - Log structured audit events as JSON - Track request duration and outcome - Support file-based and stdout log destinations -- Wrap the authorization middleware so denied requests are still recorded (outcome `denied`) +- Wrap the rest of the chain (authentication, webhooks, authorization) so rejected requests are still recorded (outcome `denied`); the identity and parsed MCP data flow back from the inner middlewares via holder carriers **Event Types**: - `mcp_initialize` - Client initialization events @@ -446,6 +451,7 @@ thv run --transport sse --name my-server --audit-config audit.json my-image:late **Important Notes**: - `excludeEventTypes` takes precedence over `eventTypes` +- Requests rejected before the MCP parser runs (e.g. authentication failures) are typed `http_request`, not `mcp_*`. An `eventTypes` allowlist containing only `mcp_*` types will drop those rejection events — include `http_request` to keep them. - When `includeRequestData` or `includeResponseData` is enabled, **`maxDataSize` must be set** (non-zero) for data capture to work - Log files are created with restrictive permissions (0600) for security - Logs are written in newline-delimited JSON format for easy parsing @@ -618,10 +624,14 @@ The middleware chain uses Go's `context.Context` to pass data between components ```mermaid graph LR - A[Request Context] --> B[+ JWT Claims] + A[Request Context] --> E[+ Audit holder carriers] + E --> B[+ JWT Claims] B --> C[+ Parsed MCP Data] - C --> E[+ Audit Metadata] - E --> D[+ Authorization Result] + C --> D[+ Authorization Result] + + subgraph "Audit" + E + end subgraph "Authentication" B @@ -631,10 +641,6 @@ graph LR C end - subgraph "Audit" - E - end - subgraph "Authorization" D end @@ -661,9 +667,9 @@ thv run --transport sse --name my-server --audit-config audit.yaml my-image:late The middleware order is critical and enforced by the system: -1. **Authentication** - Must be first to establish client identity -2. **MCP Parsing** - Must come after authentication to access JWT context -3. **Audit** - Must wrap authorization so every request is logged, including policy denials (outcome `denied`) +1. **Audit** - Wraps the rest of the chain so every request outcome is logged, including authentication failures (401) and policy denials (403, outcome `denied`). The identity and parsed MCP data are published back to it by the inner middlewares via holder carriers. +2. **Authentication** - Establishes client identity +3. **MCP Parsing** - Must come after authentication to access JWT context 4. **Authorization** - Must come after parsing to access structured MCP data ## Error Handling @@ -995,28 +1001,30 @@ func CreateMiddleware(config *types.MiddlewareConfig, runner types.MiddlewareRun The middleware chain execution order is critical and controlled by the order in `PopulateMiddlewareConfigs()` in `pkg/runner/middleware.go`. -1. **Authentication Middleware** (always present) - Validates JWT tokens and extracts claims -2. **Upstream Token Swap Middleware** (if embedded auth server configured) - Swaps ToolHive JWT for upstream IdP token -3. **Token Exchange Middleware** (if enabled) - Exchanges JWT for external service tokens via OAuth 2.0 Token Exchange -4. **Tool Filter Middleware** (if enabled) - Filters available tools in list responses -5. **Tool Call Filter Middleware** (if enabled) - Filters tool call requests -6. **MCP Parser Middleware** (always present) - Parses JSON-RPC MCP requests -7. **Usage Metrics Middleware** (if enabled) - Tracks tool call counts -8. **Telemetry Middleware** (if enabled) - OpenTelemetry instrumentation -9. **Audit Middleware** (if enabled) - Request logging +1. **Audit Middleware** (if enabled) - Request logging; wraps everything below so every rejection is audited +2. **Authentication Middleware** (always present) - Validates JWT tokens and extracts claims +3. **Upstream Token Swap Middleware** (if embedded auth server configured) - Swaps ToolHive JWT for upstream IdP token +4. **Token Exchange Middleware** (if enabled) - Exchanges JWT for external service tokens via OAuth 2.0 Token Exchange +5. **Tool Filter Middleware** (if enabled) - Filters available tools in list responses +6. **Tool Call Filter Middleware** (if enabled) - Filters tool call requests +7. **MCP Parser Middleware** (always present) - Parses JSON-RPC MCP requests +8. **Usage Metrics Middleware** (if enabled) - Tracks tool call counts +9. **Telemetry Middleware** (if enabled) - OpenTelemetry instrumentation 10. **Authorization Middleware** (if enabled) - Cedar policy evaluation 11. **Header Forward Middleware** (if configured for remote servers) - Injects custom headers 12. **Recovery Middleware** (always present) - Catches panics **Important Ordering Rules**: -- Authentication must come first to establish client identity +- Audit wraps the whole chain (directly inside the body-size limit): every request that passes the size cap produces an audit event no matter which middleware rejects it. It does not need to run inside auth or the parser — those publish the identity and parsed MCP data back to it via `auth.IdentityHolder` and `mcp.ParsedRequestHolder`. + - Tradeoff: with `includeRequestData` enabled (default off), audit buffers the request body (up to `maxDataSize`, itself capped by the body-size limit) *before* authentication, so unauthenticated requests that will 401 also pay the buffer cost. This is inherent to auditing rejected requests — the body must be read before the outcome is known — and is bounded per-request. +- Authentication must come before the other middlewares to establish client identity - Upstream Token Swap must come after Authentication (requires `tsid` claim) and before Token Exchange (so it can read the original JWT) - Token Exchange must come after Upstream Swap if both are used (can further transform the upstream IdP token) - Tool filters should come before MCP Parser to operate on raw requests - MCP Parser must come before Authorization (provides structured MCP data) -- Audit must come before Authorization so it wraps it: policy denials (403) must still produce an audit event with outcome `denied` - Header Forward executes close to the backend handler (innermost position) - Recovery is always last in config, making it the innermost wrapper (the chain wraps in reverse config order, so the first entry is the outermost and runs first) +- Body-size limit and Origin validation stay OUTSIDE audit: oversized bodies must be rejected before audit buffers request data, and origin validation is a pre-auth DNS-rebind guard. Their rejections (413/403) are the only ones not audited. ### Custom Authorization Policies diff --git a/pkg/audit/auditor.go b/pkg/audit/auditor.go index dc4f164755..96e3134879 100644 --- a/pkg/audit/auditor.go +++ b/pkg/audit/auditor.go @@ -172,6 +172,12 @@ func (rw *responseWriter) Flush() { } } +// Unwrap exposes the underlying ResponseWriter so http.ResponseController +// can reach interfaces this wrapper does not re-implement (e.g. SetWriteDeadline). +func (rw *responseWriter) Unwrap() http.ResponseWriter { + return rw.ResponseWriter +} + // isMCPStreamOpenRequest returns true only for MCP "stream" opens: // - SSE transport's SSE endpoint (GET + Accept: text/event-stream) // - Streamable HTTP's GET stream (same header pattern) @@ -187,30 +193,53 @@ func (*Auditor) isMCPStreamOpenRequest(r *http.Request) bool { return strings.Contains(strings.ToLower(accept), "text/event-stream") } +// ensureAuditContext injects the mutable carriers the auditor reads after the +// inner chain returns: BackendInfo (backend routing), an auth.IdentityHolder +// (identity attached by an auth middleware running INSIDE audit), and an +// mcp.ParsedRequestHolder (parsed MCP data from a parser running INSIDE +// audit). Each is only injected when absent so nested auditors share carriers. +func ensureAuditContext(r *http.Request) *http.Request { + ctx := r.Context() + changed := false + if _, ok := BackendInfoFromContext(ctx); !ok { + ctx = WithBackendInfo(ctx, &BackendInfo{}) + changed = true + } + if _, ok := auth.IdentityHolderFromContext(ctx); !ok { + ctx = auth.WithIdentityHolder(ctx, &auth.IdentityHolder{}) + changed = true + } + if _, ok := mcp.ParsedRequestHolderFromContext(ctx); !ok { + ctx = mcp.WithParsedRequestHolder(ctx, &mcp.ParsedRequestHolder{}) + changed = true + } + if !changed { + return r + } + return r.WithContext(ctx) +} + // Middleware creates an HTTP middleware that logs audit events. func (a *Auditor) Middleware(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Handle SSE endpoints specially - log the connection event immediately - // since SSE connections are long-lived and don't follow normal request/response pattern - if a.isMCPStreamOpenRequest(r) { - // Log SSE connection event immediately - a.logSSEConnectionEvent(r) + r = ensureAuditContext(r) - // Pass through to SSE handler without waiting - next.ServeHTTP(w, r) + // Handle MCP stream opens (SSE endpoint, streamable GET) specially: + // these connections are long-lived, so instead of waiting for the + // response to complete, the connection event is logged on the FIRST + // write. By then any inner auth middleware has run, so the event + // carries the authenticated identity (or records the 401/403 denial). + if a.isMCPStreamOpenRequest(r) { + sw := &streamOpenWriter{ResponseWriter: w, auditor: a, req: r} + next.ServeHTTP(sw, r) + // Streams that end without a single write still get an event + // (net/http sends an implicit 200 in that case). + sw.logOnce(http.StatusOK) return } startTime := time.Now() - // Add BackendInfo to context if not already present - // (backend enrichment middleware may have already added it) - if _, ok := BackendInfoFromContext(r.Context()); !ok { - backendInfo := &BackendInfo{} - ctx := WithBackendInfo(r.Context(), backendInfo) - r = r.WithContext(ctx) - } - // Capture request data if configured var requestData []byte if a.config.IncludeRequestData && r.Body != nil { @@ -319,10 +348,35 @@ func (a *Auditor) logAuditEvent(r *http.Request, rw *responseWriter, requestData event.LogTo(r.Context(), a.auditLogger, LevelAudit) } +// mcpMethodFor returns the parsed MCP method for the request, whether the +// parser ran outside audit (context value) or inside it (holder filled by +// the parser and read back after the inner chain returns). +func mcpMethodFor(r *http.Request) string { + if m := mcp.GetMCPMethod(r.Context()); m != "" { + return m + } + if holder, ok := mcp.ParsedRequestHolderFromContext(r.Context()); ok && holder.Parsed != nil { + return holder.Parsed.Method + } + return "" +} + +// mcpResourceIDFor returns the parsed MCP resource ID for the request, with +// the same context-then-holder fallback as mcpMethodFor. +func mcpResourceIDFor(r *http.Request) string { + if id := mcp.GetMCPResourceID(r.Context()); id != "" { + return id + } + if holder, ok := mcp.ParsedRequestHolderFromContext(r.Context()); ok && holder.Parsed != nil { + return holder.Parsed.ResourceID + } + return "" +} + // determineEventType determines the event type based on the HTTP request. func (a *Auditor) determineEventType(r *http.Request) string { - // First, try to get the parsed MCP method from context - if mcpMethod := mcp.GetMCPMethod(r.Context()); mcpMethod != "" { + // First, try to get the parsed MCP method + if mcpMethod := mcpMethodFor(r); mcpMethod != "" { return a.mapMCPMethodToEventType(mcpMethod) } @@ -489,8 +543,17 @@ func extractSubjectsFromIdentity(identity *auth.Identity) map[string]string { func (*Auditor) extractSubjects(r *http.Request) map[string]string { subjects := make(map[string]string) - // Extract user information from Identity - if identity, ok := auth.IdentityFromContext(r.Context()); ok { + // Extract user information from Identity. The context value is present + // when an auth middleware runs OUTSIDE audit; the holder covers the + // audit-wraps-auth arrangement, where the identity attached for inner + // handlers is published back up via auth.WithIdentity. + identity, ok := auth.IdentityFromContext(r.Context()) + if !ok { + if holder, hok := auth.IdentityHolderFromContext(r.Context()); hok && holder.Identity != nil { + identity, ok = holder.Identity, true + } + } + if ok { subjects = extractSubjectsFromIdentity(identity) } @@ -521,12 +584,12 @@ func (*Auditor) extractTarget(r *http.Request, eventType string) map[string]stri target[TargetKeyMethod] = r.Method // Add MCP method if available from parsed data - if mcpMethod := mcp.GetMCPMethod(r.Context()); mcpMethod != "" { + if mcpMethod := mcpMethodFor(r); mcpMethod != "" { target[TargetKeyMethod] = mcpMethod } // Add resource ID if available from parsed data - if resourceID := mcp.GetMCPResourceID(r.Context()); resourceID != "" { + if resourceID := mcpResourceIDFor(r); resourceID != "" { target[TargetKeyName] = resourceID } @@ -612,8 +675,63 @@ func (a *Auditor) addEventData(event *AuditEvent, _ *http.Request, rw *responseW } } +// streamOpenWriter wraps the ResponseWriter for MCP stream-open requests +// (SSE endpoint, streamable GET). It logs the connection audit event exactly +// once, on the first WriteHeader/Write, so the event reflects the actual +// outcome (200 stream established, 401/403 denied by inner middleware) and +// carries the identity the inner auth middleware attached by that point. +type streamOpenWriter struct { + http.ResponseWriter + auditor *Auditor + req *http.Request + logged bool +} + +func (sw *streamOpenWriter) WriteHeader(statusCode int) { + // Informational (1xx) responses are not the final status — don't consume + // the one-shot connection event on them. + if statusCode >= http.StatusOK { + sw.logOnce(statusCode) + } + sw.ResponseWriter.WriteHeader(statusCode) +} + +func (sw *streamOpenWriter) Write(data []byte) (int, error) { + // An implicit WriteHeader(200) happens on first Write. + sw.logOnce(http.StatusOK) + return sw.ResponseWriter.Write(data) +} + +// Flush implements http.Flusher if the underlying ResponseWriter supports it. +func (sw *streamOpenWriter) Flush() { + if flusher, ok := sw.ResponseWriter.(http.Flusher); ok { + flusher.Flush() + } +} + +// Unwrap exposes the underlying ResponseWriter so http.ResponseController +// can reach interfaces this wrapper does not re-implement (e.g. SetWriteDeadline). +func (sw *streamOpenWriter) Unwrap() http.ResponseWriter { + return sw.ResponseWriter +} + +// logOnce logs the stream connection event with the given status on the first +// call; subsequent calls are no-ops. +func (sw *streamOpenWriter) logOnce(statusCode int) { + if sw.logged { + return + } + sw.logged = true + sw.auditor.logSSEConnectionEvent(sw.req, statusCode) +} + // logSSEConnectionEvent logs an audit event for SSE connection initiation. -func (a *Auditor) logSSEConnectionEvent(r *http.Request) { +func (a *Auditor) logSSEConnectionEvent(r *http.Request, statusCode int) { + // Honor the configured event-type filter, like logAuditEvent does. + if !a.config.ShouldAuditEvent(EventTypeSSEConnection) { + return + } + // Extract source information source := a.extractSource(r) @@ -624,7 +742,7 @@ func (a *Auditor) logSSEConnectionEvent(r *http.Request) { component := a.determineComponent(r) // Create the audit event for SSE connection - event := NewAuditEvent(EventTypeSSEConnection, source, OutcomeSuccess, subjects, component) + event := NewAuditEvent(EventTypeSSEConnection, source, a.determineOutcome(statusCode), subjects, component) // Add target information target := map[string]string{ diff --git a/pkg/audit/auditor_test.go b/pkg/audit/auditor_test.go index 1bda83cdda..633f3424d5 100644 --- a/pkg/audit/auditor_test.go +++ b/pkg/audit/auditor_test.go @@ -20,6 +20,7 @@ import ( "github.com/stretchr/testify/require" "github.com/stacklok/toolhive/pkg/auth" + "github.com/stacklok/toolhive/pkg/mcp" ) func TestNewAuditor(t *testing.T) { @@ -401,6 +402,7 @@ func TestDetermineOutcome(t *testing.T) { {403, OutcomeDenied}, {400, OutcomeFailure}, {404, OutcomeFailure}, + {429, OutcomeFailure}, // Rate limiting is a load condition, not an identity/policy denial {499, OutcomeFailure}, {500, OutcomeError}, {503, OutcomeError}, @@ -1040,3 +1042,186 @@ func TestAuditLoggerLevelFormat(t *testing.T) { assert.NotContains(t, logOutput, `"level":"AUDIT"`) }) } + +// newBufferAuditor returns an Auditor writing audit events to the returned +// buffer, for tests asserting on emitted events. +func newBufferAuditor(t *testing.T) (*Auditor, *bytes.Buffer) { + t.Helper() + auditor, err := NewAuditorWithTransport(&Config{Component: "test"}, "streamable-http") + require.NoError(t, err) + var logBuf bytes.Buffer + auditor.auditLogger = NewAuditLogger(&logBuf) + return auditor, &logBuf +} + +// decodeAuditEvents parses the newline-delimited JSON events in buf. +func decodeAuditEvents(t *testing.T, buf *bytes.Buffer) []map[string]any { + t.Helper() + var events []map[string]any + for _, line := range strings.Split(strings.TrimSpace(buf.String()), "\n") { + if line == "" { + continue + } + var event map[string]any + require.NoError(t, json.Unmarshal([]byte(line), &event), "audit log line is not JSON: %s", line) + events = append(events, event) + } + return events +} + +// newToolsCallRequest builds a POST tools/call request suitable for the parser. +func newToolsCallRequest() *http.Request { + req := httptest.NewRequest("POST", "/mcp", + strings.NewReader(`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"target_tool","arguments":{}}}`)) + req.Header.Set("Content-Type", "application/json") + return req +} + +// TestMiddlewareAuditsInnerChainOutcomes pins the audit-wraps-chain +// arrangement: auth and the MCP parser run INSIDE the audit middleware, and +// audit reads the identity and parsed MCP data back through the holder +// carriers it injects (auth.IdentityHolder, mcp.ParsedRequestHolder). This is +// what lets audit record rejections from any inner middleware — auth 401s, +// webhook denials, authz 403s — instead of only requests that reached its old +// position deep in the chain. +func TestMiddlewareAuditsInnerChainOutcomes(t *testing.T) { + t.Parallel() + + // innerAuth emulates an auth middleware running inside audit: it attaches + // the identity exactly as production middlewares do (via auth.WithIdentity, + // which also fills the IdentityHolder injected by audit). + innerAuth := func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + id := &auth.Identity{PrincipalInfo: auth.PrincipalInfo{ + Subject: "user-123", + Name: "Test User", + Claims: jwt.MapClaims{"sub": "user-123"}, + }} + next.ServeHTTP(w, r.WithContext(auth.WithIdentity(r.Context(), id))) + }) + } + + t.Run("identity attached by inner auth reaches the audit event", func(t *testing.T) { + t.Parallel() + auditor, logBuf := newBufferAuditor(t) + + handler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + }) + auditor.Middleware(innerAuth(handler)).ServeHTTP(httptest.NewRecorder(), newToolsCallRequest()) + + events := decodeAuditEvents(t, logBuf) + require.Len(t, events, 1) + assert.Equal(t, OutcomeSuccess, events[0]["outcome"]) + subjects, ok := events[0]["subjects"].(map[string]any) + require.True(t, ok, "event must carry subjects") + assert.Equal(t, "user-123", subjects[SubjectKeyUserID], + "identity attached by an inner auth middleware must reach the audit event via the holder") + }) + + t.Run("inner 401 rejection is audited as denied with anonymous subject", func(t *testing.T) { + t.Parallel() + auditor, logBuf := newBufferAuditor(t) + + reject := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + http.Error(w, "invalid token", http.StatusUnauthorized) + }) + auditor.Middleware(reject).ServeHTTP(httptest.NewRecorder(), newToolsCallRequest()) + + events := decodeAuditEvents(t, logBuf) + require.Len(t, events, 1, "an authentication failure must still produce an audit event") + assert.Equal(t, OutcomeDenied, events[0]["outcome"]) + subjects, ok := events[0]["subjects"].(map[string]any) + require.True(t, ok) + assert.Equal(t, "anonymous", subjects[SubjectKeyUser], + "no identity exists when authentication fails") + }) + + t.Run("event type comes from inner parser via holder", func(t *testing.T) { + t.Parallel() + auditor, logBuf := newBufferAuditor(t) + + handler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + }) + // The parser runs INSIDE audit, as in the real chains. + auditor.Middleware(mcp.ParsingMiddleware(handler)).ServeHTTP(httptest.NewRecorder(), newToolsCallRequest()) + + events := decodeAuditEvents(t, logBuf) + require.Len(t, events, 1) + assert.Equal(t, EventTypeMCPToolCall, events[0]["type"], + "parsed MCP data from an inner parser must drive the event type via the holder") + target, ok := events[0]["target"].(map[string]any) + require.True(t, ok) + assert.Equal(t, "target_tool", target[TargetKeyName]) + }) +} + +// TestStreamOpenAuditEvents pins the deferred stream-open logging: the +// connection event for SSE / streamable GET requests is logged on the FIRST +// response write, so it reflects the real outcome and the identity attached by +// inner middleware — instead of being logged on arrival with neither. +func TestStreamOpenAuditEvents(t *testing.T) { + t.Parallel() + + newStreamRequest := func() *http.Request { + req := httptest.NewRequest("GET", "/mcp", nil) + req.Header.Set("Accept", "text/event-stream") + return req + } + + t.Run("established stream logs one success event with identity", func(t *testing.T) { + t.Parallel() + auditor, logBuf := newBufferAuditor(t) + + handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Emulate inner auth attaching the identity before the stream starts. + // The returned context is intentionally discarded: only WithIdentity's + // side effect of filling the audit-injected IdentityHolder matters here. + id := &auth.Identity{PrincipalInfo: auth.PrincipalInfo{Subject: "user-123"}} + _ = auth.WithIdentity(r.Context(), id) + w.Header().Set("Content-Type", "text/event-stream") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("event: message\ndata: {}\n\n")) + _, _ = w.Write([]byte("event: message\ndata: {}\n\n")) + }) + auditor.Middleware(handler).ServeHTTP(httptest.NewRecorder(), newStreamRequest()) + + events := decodeAuditEvents(t, logBuf) + require.Len(t, events, 1, "the connection event must be logged exactly once") + assert.Equal(t, EventTypeSSEConnection, events[0]["type"]) + assert.Equal(t, OutcomeSuccess, events[0]["outcome"]) + subjects, ok := events[0]["subjects"].(map[string]any) + require.True(t, ok) + assert.Equal(t, "user-123", subjects[SubjectKeyUserID], + "deferring the log to first write makes the inner auth identity available") + }) + + t.Run("rejected stream open logs a denied event", func(t *testing.T) { + t.Parallel() + auditor, logBuf := newBufferAuditor(t) + + reject := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + http.Error(w, "invalid token", http.StatusUnauthorized) + }) + auditor.Middleware(reject).ServeHTTP(httptest.NewRecorder(), newStreamRequest()) + + events := decodeAuditEvents(t, logBuf) + require.Len(t, events, 1) + assert.Equal(t, EventTypeSSEConnection, events[0]["type"]) + assert.Equal(t, OutcomeDenied, events[0]["outcome"]) + }) + + t.Run("handler that never writes still produces one event", func(t *testing.T) { + t.Parallel() + auditor, logBuf := newBufferAuditor(t) + + silent := http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) {}) + auditor.Middleware(silent).ServeHTTP(httptest.NewRecorder(), newStreamRequest()) + + events := decodeAuditEvents(t, logBuf) + require.Len(t, events, 1) + assert.Equal(t, OutcomeSuccess, events[0]["outcome"], + "net/http sends an implicit 200 when the handler writes nothing") + }) +} diff --git a/pkg/auth/context.go b/pkg/auth/context.go index ea806db36b..9065858244 100644 --- a/pkg/auth/context.go +++ b/pkg/auth/context.go @@ -24,6 +24,16 @@ type IdentityContextKey struct{} // This function is typically called by authentication middleware after successful // authentication to make the identity available to downstream handlers. // +// Side effect: if the context carries an IdentityHolder (injected by the audit +// middleware wrapping this call), the identity is also published into it. Two +// invariants follow for callers on request-derived contexts: +// - Last write wins: the holder reports the most recently attached identity, +// so audit events carry that principal. Do not attach a different principal +// (e.g. a re-minted service identity) on a request-derived context unless +// that is the principal audit should report. +// - Call only on the request goroutine: the audit middleware reads the holder +// when the response is written, so off-goroutine writes would race. +// // Example: // // identity := &Identity{PrincipalInfo: PrincipalInfo{Subject: "user123", Name: "Alice"}} @@ -32,9 +42,46 @@ func WithIdentity(ctx context.Context, identity *Identity) context.Context { if identity == nil { return ctx } + // Also publish the identity to an IdentityHolder if one is present, so + // middleware wrapping the auth middleware (e.g. audit) can observe the + // identity even though the derived context only flows downstream. + if holder, ok := IdentityHolderFromContext(ctx); ok { + holder.Identity = identity + } return context.WithValue(ctx, IdentityContextKey{}, identity) } +// IdentityHolderContextKey is the key used to store an IdentityHolder in the +// request context. +type IdentityHolderContextKey struct{} + +// IdentityHolder is a mutable carrier that lets middleware running OUTSIDE the +// auth middleware observe the authenticated identity. Context values only flow +// downstream, so a wrapper such as the audit middleware cannot read the +// identity that auth attaches for inner handlers. The wrapper injects an empty +// holder via WithIdentityHolder before calling the inner chain; WithIdentity +// fills it when the identity is attached; the wrapper reads it after the inner +// chain returns. +// +// The holder is written and read by the single request goroutine (writes +// happen-before the wrapper's post-ServeHTTP read), mirroring the audit +// package's BackendInfo pattern, so no synchronization is needed. +type IdentityHolder struct { + Identity *Identity +} + +// WithIdentityHolder returns a new context carrying the given IdentityHolder. +func WithIdentityHolder(ctx context.Context, holder *IdentityHolder) context.Context { + return context.WithValue(ctx, IdentityHolderContextKey{}, holder) +} + +// IdentityHolderFromContext retrieves the IdentityHolder from the context. +// Returns (nil, false) if no holder is present. +func IdentityHolderFromContext(ctx context.Context) (*IdentityHolder, bool) { + holder, ok := ctx.Value(IdentityHolderContextKey{}).(*IdentityHolder) + return holder, ok && holder != nil +} + // IdentityFromContext retrieves an Identity from the context. // Returns the identity and true if a non-nil identity is present, nil and false otherwise. // A typed-nil *Identity stored directly in the context (bypassing WithIdentity) is diff --git a/pkg/mcp/parser.go b/pkg/mcp/parser.go index 59ebfc3e72..64aab470ab 100644 --- a/pkg/mcp/parser.go +++ b/pkg/mcp/parser.go @@ -79,10 +79,10 @@ type ParsedMCPRequest struct { // Example usage: // // middlewares := []types.Middleware{ -// authMiddleware, // Authentication first +// auditMiddleware, // Audit wraps the chain; parsed data flows back via ParsedRequestHolder +// authMiddleware, // Authentication // mcp.ParsingMiddleware, // MCP parsing after auth // authzMiddleware, // Authorization uses parsed data -// auditMiddleware, // Audit uses parsed data // } func ParsingMiddleware(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -126,6 +126,12 @@ func ParsingMiddleware(next http.Handler) http.Handler { if parsedRequest != nil { parsedRequest.MCPMethodHeader = r.Header.Get("Mcp-Method") parsedRequest.MCPNameHeader = r.Header.Get("Mcp-Name") + // Publish to a ParsedRequestHolder if one is present, so middleware + // wrapping the parser (e.g. audit) can observe the parsed request + // even though the derived context only flows downstream. + if holder, ok := ParsedRequestHolderFromContext(r.Context()); ok { + holder.Parsed = parsedRequest + } ctx := context.WithValue(r.Context(), MCPRequestContextKey, parsedRequest) r = r.WithContext(ctx) } @@ -135,6 +141,36 @@ func ParsingMiddleware(next http.Handler) http.Handler { }) } +// parsedRequestHolderContextKey is the context key for ParsedRequestHolder. +type parsedRequestHolderContextKey struct{} + +// ParsedRequestHolder is a mutable carrier that lets middleware running +// OUTSIDE the parsing middleware observe the parsed MCP request. Context +// values only flow downstream, so a wrapper such as the audit middleware +// cannot read the parsed request that the parser attaches for inner handlers. +// The wrapper injects an empty holder via WithParsedRequestHolder before +// calling the inner chain; ParsingMiddleware fills it; the wrapper reads it +// after the inner chain returns. +// +// The holder is written and read by the single request goroutine (writes +// happen-before the wrapper's post-ServeHTTP read), so no synchronization is +// needed. +type ParsedRequestHolder struct { + Parsed *ParsedMCPRequest +} + +// WithParsedRequestHolder returns a new context carrying the given holder. +func WithParsedRequestHolder(ctx context.Context, holder *ParsedRequestHolder) context.Context { + return context.WithValue(ctx, parsedRequestHolderContextKey{}, holder) +} + +// ParsedRequestHolderFromContext retrieves the ParsedRequestHolder from the +// context. Returns (nil, false) if no holder is present. +func ParsedRequestHolderFromContext(ctx context.Context) (*ParsedRequestHolder, bool) { + holder, ok := ctx.Value(parsedRequestHolderContextKey{}).(*ParsedRequestHolder) + return holder, ok && holder != nil +} + // GetParsedMCPRequest retrieves the parsed MCP request from the request context. // Returns nil if no parsed request is available. func GetParsedMCPRequest(ctx context.Context) *ParsedMCPRequest { diff --git a/pkg/runner/authz_audit_integration_test.go b/pkg/runner/authz_audit_integration_test.go index f78ceda731..bb45babc60 100644 --- a/pkg/runner/authz_audit_integration_test.go +++ b/pkg/runner/authz_audit_integration_test.go @@ -20,6 +20,7 @@ import ( "github.com/stacklok/toolhive/pkg/audit" "github.com/stacklok/toolhive/pkg/authz/authorizers" "github.com/stacklok/toolhive/pkg/authz/authorizers/cedar" + "github.com/stacklok/toolhive/pkg/webhook" statusesmocks "github.com/stacklok/toolhive/pkg/workloads/statuses/mocks" ) @@ -169,3 +170,64 @@ func TestAuthzDecisionIsAudited(t *testing.T) { }) } } + +// TestWebhookDenialIsAudited proves, through the full middleware chain built +// by PopulateMiddlewareConfigs, that a validating-webhook policy denial (403) +// still produces an audit event with outcome "denied". Webhooks run inside +// the audit middleware, so the rejection must be captured like any other. +func TestWebhookDenialIsAudited(t *testing.T) { + t.Parallel() + + denyServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req webhook.Request + require.NoError(t, json.NewDecoder(r.Body).Decode(&req)) + resp := webhook.Response{ + Version: webhook.APIVersion, + UID: req.UID, + Allowed: false, + Reason: "policy denied", + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(resp) + })) + t.Cleanup(denyServer.Close) + + auditLogPath := filepath.Join(t.TempDir(), "audit.log") + + runConfig := NewRunConfig() + runConfig.Name = "test-server" + runConfig.ValidatingWebhooks = []webhook.Config{{ + Name: "deny-all", + URL: denyServer.URL, + Timeout: webhook.DefaultTimeout, + FailurePolicy: webhook.FailurePolicyFail, + TLSConfig: &webhook.TLSConfig{InsecureSkipVerify: true}, + }} + runConfig.AuditConfig = &audit.Config{ + Component: "test-component", + LogFile: auditLogPath, + } + + handlerHit := false + backend := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + handlerHit = true + w.WriteHeader(http.StatusOK) + }) + + handler := buildRunnerMiddlewareChain(t, runConfig, backend) + + reqBody := `{"jsonrpc":"2.0","method":"tools/call","id":1,"params":{"name":"target_tool","arguments":{}}}` + req := httptest.NewRequest(http.MethodPost, "/", bytes.NewBufferString(reqBody)) + req.Header.Set("Content-Type", "application/json") + + rr := httptest.NewRecorder() + handler.ServeHTTP(rr, req) + + require.Equal(t, http.StatusForbidden, rr.Code, "response body: %s", rr.Body.String()) + assert.False(t, handlerHit, "a webhook-denied request must not reach the backend") + + events := readAuditEvents(t, auditLogPath) + require.Len(t, events, 1, "a webhook policy denial must produce exactly one audit event") + assert.Equal(t, "mcp_tool_call", events[0]["type"]) + assert.Equal(t, "denied", events[0]["outcome"]) +} diff --git a/pkg/runner/config_builder.go b/pkg/runner/config_builder.go index 1961b29f42..da563a4144 100644 --- a/pkg/runner/config_builder.go +++ b/pkg/runner/config_builder.go @@ -675,6 +675,15 @@ func WithMiddlewareFromFlags( // actual proxy determine the order of application of middlewares, since // the types of middleware are known at compile time. + // Audit middleware (if enabled) goes first so it is the outermost + // wrapper after the body-limit prepended in runner.Run: every request + // that passes the size cap produces an audit event no matter which + // middleware rejects it — authentication (401), webhook denials, and + // authorization (403, outcome "denied") included. Identity and parsed + // MCP data are read back from the inner auth/parser middlewares via + // the holder carriers (auth.IdentityHolder, mcp.ParsedRequestHolder). + middlewareConfigs = addAuditMiddleware(middlewareConfigs, enableAudit, auditConfigPath, serverName, transportType) + // Add tool filter middlewares middlewareConfigs = addToolFilterMiddlewares(middlewareConfigs, toolsFilter, toolsOverride) @@ -706,11 +715,10 @@ func WithMiddlewareFromFlags( return err } - // Add optional middlewares. Audit is added BEFORE authorization so it - // wraps it at request time: authorization denials (403) must still - // produce an audit event with outcome "denied". + // Add optional middlewares. Audit was added at the top of the chain + // so authorization denials (403) still produce an audit event with + // outcome "denied". middlewareConfigs = addTelemetryMiddleware(middlewareConfigs, telemetryConfig, serverName, transportType) - middlewareConfigs = addAuditMiddleware(middlewareConfigs, enableAudit, auditConfigPath, serverName, transportType) var authzErr error middlewareConfigs, authzErr = addAuthzMiddleware(middlewareConfigs, authzConfigPath, b.config.EmbeddedAuthServerConfig) if authzErr != nil { diff --git a/pkg/runner/config_builder_test.go b/pkg/runner/config_builder_test.go index 3b544ded28..0087d97ffd 100644 --- a/pkg/runner/config_builder_test.go +++ b/pkg/runner/config_builder_test.go @@ -1612,13 +1612,14 @@ func TestResolveRegistryServerName(t *testing.T) { } } -// TestWithMiddlewareFromFlags_AuditBeforeAuthz pins the same audit-wraps-authz +// TestWithMiddlewareFromFlags_AuditWrapsChain pins the same audit-wraps-chain // ordering invariant on the CLI flag path that -// TestPopulateMiddlewareConfigs_AuditBeforeAuthz pins on the operator path: -// audit must precede authorization in the config slice (earlier entries wrap -// later ones at request time) so authorization denials still produce an audit -// event with outcome "denied". -func TestWithMiddlewareFromFlags_AuditBeforeAuthz(t *testing.T) { +// TestPopulateMiddlewareConfigs_AuditWrapsChain pins on the operator path: +// audit is the first entry of the built slice (earlier entries wrap later +// ones at request time; runner.Run later prepends only body-limit/origin), so +// rejections from auth (401), webhooks, and authorization (403) all still +// produce audit events. +func TestWithMiddlewareFromFlags_AuditWrapsChain(t *testing.T) { t.Parallel() builder := &runConfigBuilder{config: NewRunConfig()} @@ -1651,12 +1652,14 @@ func TestWithMiddlewareFromFlags_AuditBeforeAuthz(t *testing.T) { parserIdx, ok := typeIndex[mcp.ParserMiddlewareType] require.True(t, ok, "MCP parser middleware must be present") + assert.Equal(t, 0, auditIdx, + "audit must be the outermost built entry so every rejection is audited") + assert.Less(t, auditIdx, authIdx, + "audit must wrap auth so authentication failures (401) are audited") + assert.Less(t, auditIdx, parserIdx, + "audit wraps the parser; parsed MCP data flows back via mcp.ParsedRequestHolder") assert.Less(t, auditIdx, authzIdx, - "audit must precede authz so authorization denials are audited") - assert.Less(t, authIdx, auditIdx, - "auth must precede audit so the identity is available to audit events") - assert.Less(t, parserIdx, auditIdx, - "MCP parser must precede audit so parsed MCP data is available to audit events") + "audit must wrap authz so authorization denials (403) are audited") } // TestWithAdditionalMiddlewareConfigs verifies the generic injected-middleware diff --git a/pkg/runner/middleware.go b/pkg/runner/middleware.go index 5e33ad5515..da116f1b18 100644 --- a/pkg/runner/middleware.go +++ b/pkg/runner/middleware.go @@ -76,6 +76,28 @@ func PopulateMiddlewareConfigs(config *RunConfig) error { return err } + // Audit middleware (if enabled). Added directly inside body-limit so it + // wraps the REST of the chain: every request that passes the size cap + // produces an audit event no matter which middleware rejects it. + // Authentication (401) and authorization/webhook denials (403) map to + // outcome "denied"; other rejections such as rate limiting (429) map to + // "failure" (see audit.determineOutcome). Identity and parsed MCP data are + // read back from the inner auth/parser middlewares via the holder + // carriers (see auth.IdentityHolder, mcp.ParsedRequestHolder). + if config.AuditConfig != nil { + auditParams := audit.MiddlewareParams{ + ConfigPath: config.AuditConfigPath, // Keep for backwards compatibility + ConfigData: config.AuditConfig, // Use the loaded config data + Component: config.AuditConfig.Component, + TransportType: config.Transport.String(), // Pass the actual transport type + } + auditConfig, err := types.NewMiddlewareConfig(audit.MiddlewareType, auditParams) + if err != nil { + return fmt.Errorf("failed to create audit middleware config: %w", err) + } + middlewareConfigs = append(middlewareConfigs, *auditConfig) + } + // Authentication middleware (always present) authParams := auth.MiddlewareParams{ OIDCConfig: config.OIDCConfig, @@ -151,7 +173,7 @@ func PopulateMiddlewareConfigs(config *RunConfig) error { // Mutating Webhooks middleware (if configured). // Must run BEFORE validating webhooks: - // MCP Parser -> [Mutating Webhooks] -> [Validating Webhooks] -> Audit -> Authz + // Audit -> ... -> MCP Parser -> [Mutating Webhooks] -> [Validating Webhooks] -> Authz middlewareConfigs, err = addMutatingWebhookMiddleware(middlewareConfigs, config) if err != nil { return err @@ -187,25 +209,6 @@ func PopulateMiddlewareConfigs(config *RunConfig) error { middlewareConfigs = append(middlewareConfigs, *telemetryConfig) } - // Audit middleware (if enabled) - // Added BEFORE authorization so it wraps it at request time: authorization - // denials (403) must still produce an audit event with outcome "denied". - // If audit ran inside authz, a deny would short-circuit before the auditor - // ever saw the request. - if config.AuditConfig != nil { - auditParams := audit.MiddlewareParams{ - ConfigPath: config.AuditConfigPath, // Keep for backwards compatibility - ConfigData: config.AuditConfig, // Use the loaded config data - Component: config.AuditConfig.Component, - TransportType: config.Transport.String(), // Pass the actual transport type - } - auditConfig, err := types.NewMiddlewareConfig(audit.MiddlewareType, auditParams) - if err != nil { - return fmt.Errorf("failed to create audit middleware config: %w", err) - } - middlewareConfigs = append(middlewareConfigs, *auditConfig) - } - // Authorization middleware (if enabled) if config.AuthzConfig != nil { authzCfgData, err := injectUpstreamProviderIfNeeded(config.AuthzConfig, config.EmbeddedAuthServerConfig) diff --git a/pkg/runner/middleware_test.go b/pkg/runner/middleware_test.go index 7cc8f8b99b..2a29e7c4d6 100644 --- a/pkg/runner/middleware_test.go +++ b/pkg/runner/middleware_test.go @@ -1354,13 +1354,16 @@ func TestPopulateMiddlewareConfigs_FullCoverage(t *testing.T) { assert.True(t, typeIndex[audit.MiddlewareType]) } -// TestPopulateMiddlewareConfigs_AuditBeforeAuthz pins the ordering invariant -// that the audit middleware precedes authorization in the config slice. -// Earlier entries wrap later ones at request time, so audit must wrap authz -// for an authorization denial (403) to still produce an audit event with -// outcome "denied". It must in turn come after auth and the MCP parser, which -// provide the identity and parsed MCP data the audit event is built from. -func TestPopulateMiddlewareConfigs_AuditBeforeAuthz(t *testing.T) { +// TestPopulateMiddlewareConfigs_AuditWrapsChain pins the ordering invariant +// that the audit middleware sits directly inside body-limit, wrapping the +// REST of the chain. Earlier entries wrap later ones at request time, so this +// placement is what guarantees every request that passes the size cap +// produces an audit event no matter which middleware rejects it — +// authentication 401s, webhook denials, and authorization 403s included. +// Identity and parsed MCP data flow back to audit from the inner auth/parser +// middlewares via the holder carriers, so audit no longer needs to run inside +// them. +func TestPopulateMiddlewareConfigs_AuditWrapsChain(t *testing.T) { t.Parallel() config := &RunConfig{ @@ -1383,13 +1386,17 @@ func TestPopulateMiddlewareConfigs_AuditBeforeAuthz(t *testing.T) { require.True(t, ok, "auth middleware must be present") parserIdx, ok := typeIndex[mcp.ParserMiddlewareType] require.True(t, ok, "MCP parser middleware must be present") - + bodyLimitIdx, ok := typeIndex[bodylimit.MiddlewareType] + require.True(t, ok, "body limit middleware must be present") + + assert.Less(t, bodyLimitIdx, auditIdx, + "body limit must stay outside audit so oversized bodies are rejected before audit buffers them") + assert.Less(t, auditIdx, authIdx, + "audit must wrap auth so authentication failures (401) are audited") + assert.Less(t, auditIdx, parserIdx, + "audit wraps the parser; parsed MCP data flows back via mcp.ParsedRequestHolder") assert.Less(t, auditIdx, authzIdx, - "audit must precede authz so authorization denials are audited") - assert.Less(t, authIdx, auditIdx, - "auth must precede audit so the identity is available to audit events") - assert.Less(t, parserIdx, auditIdx, - "MCP parser must precede audit so parsed MCP data is available to audit events") + "audit must wrap authz so authorization denials (403) are audited") } // TestPopulateMiddlewareConfigs_StripAuthOrdering pins the ordering invariant diff --git a/pkg/vmcp/server/authz_integration_test.go b/pkg/vmcp/server/authz_integration_test.go index 65025a52ef..a7580e8dd1 100644 --- a/pkg/vmcp/server/authz_integration_test.go +++ b/pkg/vmcp/server/authz_integration_test.go @@ -68,7 +68,7 @@ func parseRPCError(t *testing.T, body []byte) rpcErrorFields { // chain that replaced the legacy HTTP authz middleware on the Serve path. func newCedarAuthzTestServer(t *testing.T, backendURL string, policies ...string) *httptest.Server { t.Helper() - return buildCedarAuthzServer(t, backendURL, nil, nil, policies...) + return buildCedarAuthzServer(t, backendURL, nil, nil, nil, policies...) } // newCedarAuthzCodeModeServer is newCedarAuthzTestServer with code mode enabled, so @@ -77,15 +77,17 @@ func newCedarAuthzTestServer(t *testing.T, backendURL string, policies ...string // while a directly-denied tool still 403s. func newCedarAuthzCodeModeServer(t *testing.T, backendURL string, policies ...string) *httptest.Server { t.Helper() - return buildCedarAuthzServer(t, backendURL, &codemode.Config{}, nil, policies...) + return buildCedarAuthzServer(t, backendURL, &codemode.Config{}, nil, nil, policies...) } // buildCedarAuthzServer builds the vMCP test server. A non-nil codeModeCfg enables -// the codemode decorator; a non-nil auditCfg enables the audit middleware; a nil -// policies slice leaves Authz unset (allow-all, gate not installed) — used by the -// no-Authz parity guard. +// the codemode decorator; a non-nil auditCfg enables the audit middleware; a +// non-nil authMw replaces the default identity-injecting auth middleware (used by +// the auth-failure audit test); a nil policies slice leaves Authz unset +// (allow-all, gate not installed) — used by the no-Authz parity guard. func buildCedarAuthzServer( - t *testing.T, backendURL string, codeModeCfg *codemode.Config, auditCfg *audit.Config, policies ...string, + t *testing.T, backendURL string, codeModeCfg *codemode.Config, auditCfg *audit.Config, + authMw func(http.Handler) http.Handler, policies ...string, ) *httptest.Server { t.Helper() @@ -128,20 +130,23 @@ func buildCedarAuthzServer( // inside it, mirroring the production incoming-auth factory (see pkg/vmcp/auth/factory): // audit and authz read parsed MCP data from the request context, and the audit // middleware sits between auth and the parser applied in Handler. - identityMiddleware := func(next http.Handler) http.Handler { - withParser := mcpparser.ParsingMiddleware(next) - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - principal := r.Header.Get("X-Test-Principal") - if principal == "" { - principal = "user-123" - } - id := &auth.Identity{PrincipalInfo: auth.PrincipalInfo{ - Subject: principal, - Name: "Test User", - Claims: map[string]any{"sub": principal, "name": "Test User"}, - }} - withParser.ServeHTTP(w, r.WithContext(auth.WithIdentity(r.Context(), id))) - }) + identityMiddleware := authMw + if identityMiddleware == nil { + identityMiddleware = func(next http.Handler) http.Handler { + withParser := mcpparser.ParsingMiddleware(next) + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + principal := r.Header.Get("X-Test-Principal") + if principal == "" { + principal = "user-123" + } + id := &auth.Identity{PrincipalInfo: auth.PrincipalInfo{ + Subject: principal, + Name: "Test User", + Claims: map[string]any{"sub": principal, "name": "Test User"}, + }} + withParser.ServeHTTP(w, r.WithContext(auth.WithIdentity(r.Context(), id))) + }) + } } // A nil policies slice means "no authz": leave Config.Authz nil so the gate is not @@ -464,6 +469,7 @@ func TestIntegration_CedarAuthzDenialIsAudited(t *testing.T) { // Permit only an unrelated tool: "echo" is default-denied. ts := buildCedarAuthzServer(t, backendURL, nil, &audit.Config{Component: "vmcp-server", LogFile: auditLogPath}, + nil, `permit(principal, action == Action::"call_tool", resource == Tool::"unrelated");`) client := NewMCPTestClient(t, ts.URL) @@ -483,6 +489,46 @@ func TestIntegration_CedarAuthzDenialIsAudited(t *testing.T) { event := findToolCallAuditEvent(t, auditLogPath) assert.Equal(t, "denied", event["outcome"], "a policy-denied tools/call must be audited with outcome denied") + subjects, ok := event["subjects"].(map[string]any) + require.True(t, ok, "the event must carry subjects") + assert.Equal(t, "user-123", subjects["user_id"], + "audit wraps auth, so the identity must flow back via the auth.IdentityHolder carrier") +} + +// TestIntegration_AuthFailureIsAudited proves that an authentication failure +// (401 from the auth middleware) still produces an audit event: audit wraps +// auth on the vMCP Serve path, so rejected requests are recorded with outcome +// "denied" and an anonymous subject. +func TestIntegration_AuthFailureIsAudited(t *testing.T) { + t.Parallel() + + backendURL := startRealMCPBackend(t) + auditLogPath := filepath.Join(t.TempDir(), "audit.log") + rejectAll := func(http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + http.Error(w, "invalid token", http.StatusUnauthorized) + }) + } + ts := buildCedarAuthzServer(t, backendURL, nil, + &audit.Config{Component: "vmcp-server", LogFile: auditLogPath}, + rejectAll) + + resp, err := http.Post(ts.URL+"/mcp", "application/json", + strings.NewReader(`{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}`)) + require.NoError(t, err) + defer resp.Body.Close() + require.Equal(t, http.StatusUnauthorized, resp.StatusCode) + + require.Eventually(t, func() bool { + return findAuditEvent(t, auditLogPath, "http_request") != nil + }, 5*time.Second, 50*time.Millisecond, "an authentication failure must be audited") + + event := findAuditEvent(t, auditLogPath, "http_request") + assert.Equal(t, "denied", event["outcome"]) + subjects, ok := event["subjects"].(map[string]any) + require.True(t, ok) + assert.Equal(t, "anonymous", subjects["user"], + "no identity exists when authentication fails") } // TestIntegration_CedarAuthzDenial_ModernPath_IsAudited is the Modern (2026-07-28) @@ -502,7 +548,7 @@ func TestIntegration_CedarAuthzDenial_ModernPath_IsAudited(t *testing.T) { // Permit only an unrelated tool: "echo" is default-denied, so the re-homed gate // in dispatchModern rejects the call before it reaches the backend. ts := buildCedarAuthzServer(t, backendURL, nil, - &audit.Config{Component: "vmcp-server", LogFile: auditLogPath}, + &audit.Config{Component: "vmcp-server", LogFile: auditLogPath}, nil, `permit(principal, action == Action::"call_tool", resource == Tool::"unrelated");`) resp, decoded := postModern(t, ts.URL, "tools/call", map[string]any{ @@ -658,6 +704,13 @@ func TestRegression_PerSessionToolCall_DeniedForUnprivilegedPrincipal(t *testing // returns the first "mcp_tool_call" event, or nil if none is present yet. func findToolCallAuditEvent(t *testing.T, path string) map[string]any { t.Helper() + return findAuditEvent(t, path, "mcp_tool_call") +} + +// findAuditEvent reads the newline-delimited JSON audit log at path and returns +// the first event whose "type" matches eventType, or nil if none is present yet. +func findAuditEvent(t *testing.T, path string, eventType string) map[string]any { + t.Helper() data, err := os.ReadFile(path) if err != nil { @@ -671,7 +724,7 @@ func findToolCallAuditEvent(t *testing.T, path string) map[string]any { if err := json.Unmarshal([]byte(line), &event); err != nil { continue } - if event["type"] == "mcp_tool_call" { + if event["type"] == eventType { return event } } diff --git a/pkg/vmcp/server/server.go b/pkg/vmcp/server/server.go index fb800a5727..1062fa7776 100644 --- a/pkg/vmcp/server/server.go +++ b/pkg/vmcp/server/server.go @@ -557,8 +557,8 @@ func New( // This enables embedding the vmcp server inside another HTTP server or framework. // // The returned handler includes all routes (health, metrics, well-known, MCP) -// and the full middleware chain (recovery, body limit, header validation, auth, -// rate limit, audit, MCP parsing, telemetry). +// and the full middleware chain (recovery, body limit, header validation, +// audit, auth, MCP parsing, telemetry). // // Each call builds a fresh handler. The method is safe to call multiple times. // All returned handlers share the same underlying MCPServer and SessionManager, @@ -618,14 +618,18 @@ func (s *Server) Handler(_ context.Context) (http.Handler, error) { } // MCP endpoint - apply middleware chain (wrapping order, execution happens in reverse): - // Code wraps: auth → rate-limit → audit → MCP-parsing → telemetry → classification - // Execution order: recovery → body-limit → header-val → auth → - // rate-limit → audit → MCP-parsing → telemetry → classification → handler + // Code wraps: audit → auth → MCP-parsing → telemetry → classification + // Execution order: recovery → body-limit → header-val → audit → auth → + // MCP-parsing → telemetry → classification → handler // // Upstream token refresh failures are detected inside AuthMiddleware itself: // GetAllUpstreamCredentials returns a non-empty failed-provider slice when // any upstream refresh fails, and the middleware short-circuits with - // HTTP 401 + WWW-Authenticate before the request reaches any inner layer. + // HTTP 401 + WWW-Authenticate. Audit wraps auth, so those 401s (and every + // other rejection from the inner chain) still produce an audit event; the + // authenticated identity and parsed MCP data are read back through the + // holder carriers (auth.IdentityHolder, mcp.ParsedRequestHolder) that the + // inner auth/parser middlewares fill. // // The legacy HTTP authz, annotation-enrichment, and discovery layers have all been // removed: every caller now routes through Serve, so authorization is enforced by the @@ -664,7 +668,15 @@ func (s *Server) Handler(_ context.Context) (http.Handler, error) { // when auth middleware is nil. mcpHandler = mcpparser.ParsingMiddleware(mcpHandler) - // Apply audit middleware if configured (runs after auth, before discovery) + // Apply authentication middleware if configured + if s.config.AuthMiddleware != nil { + mcpHandler = s.config.AuthMiddleware(mcpHandler) + slog.Info("authentication middleware enabled for MCP endpoints") + } + + // Apply audit middleware if configured. It wraps authentication so auth + // failures (401) are audited too; the identity for successful requests is + // read back via the auth.IdentityHolder carrier. if s.config.AuditConfig != nil { if err := s.config.AuditConfig.Validate(); err != nil { return nil, fmt.Errorf("invalid audit configuration: %w", err) @@ -680,12 +692,6 @@ func (s *Server) Handler(_ context.Context) (http.Handler, error) { slog.Info("audit middleware enabled for MCP endpoints") } - // Apply authentication middleware if configured (runs first in chain) - if s.config.AuthMiddleware != nil { - mcpHandler = s.config.AuthMiddleware(mcpHandler) - slog.Info("authentication middleware enabled for MCP endpoints") - } - mcpHandler = s.applyForwardedHeaderCapture(mcpHandler) // Apply Accept header validation (rejects GET requests without Accept: text/event-stream)