diff --git a/Cargo.lock b/Cargo.lock index 9f4ee17347..49151a3aad 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2723,6 +2723,7 @@ dependencies = [ "metrics", "metrics-util", "notify", + "percent-encoding", "praxis-ai-apis", "praxis-proxy-core", "praxis-proxy-filter", diff --git a/apis/src/lib.rs b/apis/src/lib.rs index 5857d6885d..72c72ae697 100644 --- a/apis/src/lib.rs +++ b/apis/src/lib.rs @@ -17,7 +17,7 @@ pub mod openai; pub mod promotion; #[cfg(feature = "store")] pub mod store; -pub(crate) mod subrequest; +pub mod subrequest; pub(crate) mod web_search; /// Whether a `Content-Type` header value indicates `text/event-stream`, diff --git a/apis/src/subrequest.rs b/apis/src/subrequest.rs index 902d06e5bd..4594ca6daa 100644 --- a/apis/src/subrequest.rs +++ b/apis/src/subrequest.rs @@ -13,7 +13,7 @@ use std::{future::Future, net::SocketAddr, time::Duration}; use pingora_core::upstreams::peer::HttpPeer; -pub(crate) use praxis_core::subrequest::{SubRequest, SubRequestClient, SubRequestError, SubResponse}; +pub use praxis_core::subrequest::{SubRequest, SubRequestClient, SubRequestError, SubResponse}; use tracing::debug; /// Parsed URL components needed to resolve and execute a request. @@ -103,7 +103,13 @@ async fn with_deadline( /// The configured timeout covers URL resolution and the complete HTTP /// exchange. All resolved addresses are tried in order when connecting, /// while the original URL authority is preserved in `Host`. -pub(crate) async fn execute_url( +/// +/// # Errors +/// +/// Returns [`SubRequestError`] when the URL is invalid, DNS resolution +/// fails, every resolved address refuses the connection, the exchange +/// fails, or the deadline expires. +pub async fn execute_url( client: &SubRequestClient, url: &str, request: SubRequest, diff --git a/docs/filters/external_metering.md b/docs/filters/external_metering.md new file mode 100644 index 0000000000..6948049d2d --- /dev/null +++ b/docs/filters/external_metering.md @@ -0,0 +1,39 @@ + + + +# `external_metering` + +Integrates with an external metering service for pre-request balance checks and post-response token usage reporting. + +## Configuration Notes + +Tenant identity is resolved from the highest-trust source available: verified `{prefix}*` metadata written by an authentication filter, then the `identity_header_guard` filter's namespaced `{namespace}.{prefix}*` metadata, then raw `{prefix}*` request headers. A higher tier always wins, so forged client headers can never override verified claims, and identity headers plus client credentials are always stripped before the request is forwarded. + +## Configuration + +| Field | Type | Required | Description | +|-------|------|---------|-------------| +| `metering_url` | string | yes | Base URL of the external metering service (required). | +| `timeout_seconds` | integer | no | HTTP timeout in seconds for all metering calls. | +| `feature_key` | string | no | Entitlement feature key used in balance check URL path. | +| `source` | string | no | `CloudEvents` `source` field value. | +| `fail_open` | bool | no | When `true` (default), requests proceed if the metering service is unavailable. When `false`, requests are rejected with 503. | +| `identity_header_prefix` | string | no | Prefix for tenant identity headers to capture and strip. Expected headers: `{prefix}username`, `{prefix}group`, `{prefix}subscription`, `{prefix}model`. | +| `identity_metadata_namespace` | string | no | Metadata namespace the `identity_header_guard` filter writes captured identity headers under. Must match that filter's `metadata_namespace` setting when both run in one pipeline. | +| `default_username` | string | no | Fallback username when no identity header is present. If set, requests without `{prefix}username` are still metered under this name. If unset, metering is skipped entirely. | +| `default_model` | string | no | Fallback model name when no identity model header is present. | + +## Example + +```yaml +filter: external_metering +metering_url: "http://metering-service:8080" +timeout_seconds: 5 +feature_key: "inference-tokens" +source: "ai-gateway" +fail_open: true +identity_header_prefix: "x-tenant-" +identity_metadata_namespace: "identity" +default_username: "anonymous" +default_model: "unknown" +``` diff --git a/docs/filters/reference.md b/docs/filters/reference.md index 8a95a5382b..5bab799e34 100644 --- a/docs/filters/reference.md +++ b/docs/filters/reference.md @@ -90,6 +90,12 @@ see the [Praxis core filter reference][core-ref]. |--------|-------------| | [`model_to_header`](model_to_header.md) | Promotes the JSON `"model"` field from the request body to a request header. | +### Metering + +| Filter | Description | +|--------|-------------| +| [`external_metering`](external_metering.md) | Integrates with an external metering service for pre-request balance checks and post-response token usage reporting. | + ### Prompt Enrich | Filter | Description | diff --git a/examples/README.md b/examples/README.md index 6849626ef0..e2f73c4170 100644 --- a/examples/README.md +++ b/examples/README.md @@ -26,6 +26,7 @@ before sending requests. | [aws-sigv4.yaml](configs/aws-sigv4.yaml) | Signs outbound requests to an AWS service (Bedrock, in this example) using Signature Version 4. Credentials are static, sourced from environment variables — see the module docs on Sigv4SignFilter for the planned OIDC/default-credential-chain follow-up | | [azure-ad.yaml](configs/azure-ad.yaml) | Acquires an Entra ID bearer token via the client-credentials grant and injects "Authorization: Bearer " on every proxied request to Azure OpenAI | | [credential-injection.yaml](configs/credential-injection.yaml) | Injects per-cluster API credentials into upstream requests and strips client-provided credentials to prevent forwarding | +| [external-metering.yaml](configs/external-metering.yaml) | Pre-request balance check and post-response token usage reporting against an external metering service | | [gcp-adc.yaml](configs/gcp-adc.yaml) | Establishes the gcp_adc filter's configuration surface and fail-closed behavior | | [intelligent-route-all-capabilities.yaml](configs/intelligent-route-all-capabilities.yaml) | Demonstrates every candidate capability and selection input handled by intelligent_route today | | [intelligent-route-inference.yaml](configs/intelligent-route-inference.yaml) | Routes requests to different upstream clusters based on the inference model name extracted from a configured request header. The header value is set by an earlier filter such as `json_body_field` | diff --git a/examples/configs/external-metering.yaml b/examples/configs/external-metering.yaml new file mode 100644 index 0000000000..d9fadcdab8 --- /dev/null +++ b/examples/configs/external-metering.yaml @@ -0,0 +1,69 @@ +# External Metering +# +# Pre-request balance check and post-response token usage reporting +# against an external metering service. Identity headers are captured +# from configurable tenant headers (default prefix: x-tenant-) and +# stripped before forwarding upstream. +# +# Usage: +# cargo run -p praxis-ai-proxy -- -c examples/configs/external-metering.yaml +# curl -X POST http://localhost:8080/v1/chat/completions \ +# -H "Content-Type: application/json" \ +# -H "x-tenant-username: alice" \ +# -H "x-tenant-group: engineering" \ +# -H "x-tenant-subscription: sub-42" \ +# -d '{"model":"gpt-4","messages":[{"role":"user","content":"hi"}]}' +# +# Pipeline ordering matters: external_metering is declared before +# token_count so that response hooks (which run in reverse order) +# execute token_count first, writing token.input / token.output / +# token.total to filter_metadata, then external_metering reads them +# and sends a CloudEvent to the metering service. +# +# Balance check endpoint: +# GET {metering_url}/api/v1/customers/{username}/entitlements/{feature_key}/value?model={model} +# +# Usage report endpoint: +# POST {metering_url}/api/v1/events (CloudEvents 1.0 JSON) + +listeners: + - name: default + address: "127.0.0.1:8080" + filter_chains: + - main + +filter_chains: + - name: main + filters: + - filter: router + routes: + - path_prefix: "/" + cluster: backend + + - filter: external_metering + metering_url: "http://127.0.0.1:9090" + timeout_seconds: 5 + feature_key: "inference-tokens" + source: "ai-gateway" + fail_open: true + identity_header_prefix: "x-tenant-" + # Namespace the identity_header_guard filter writes captured headers + # under; must match that filter's metadata_namespace when both run. + # identity_metadata_namespace: "identity" + # Optional fallbacks for deployments where an upstream authentication + # layer does not inject identity headers. Without default_username, + # requests carrying no identity header are not metered at all. + # default_username: "anonymous" + # default_model: "unknown" + + - filter: token_count + provider: openai + + - filter: load_balancer + clusters: + - name: backend + endpoints: + - "127.0.0.1:3000" + +insecure_options: + allow_private_endpoints: true # example proxies to local backends diff --git a/filters/Cargo.toml b/filters/Cargo.toml index 9be2fa9879..0aa30114ce 100644 --- a/filters/Cargo.toml +++ b/filters/Cargo.toml @@ -56,6 +56,7 @@ futures = { workspace = true } http = { workspace = true } metrics = { workspace = true } notify = { workspace = true } +percent-encoding = { workspace = true } pingora-core.workspace = true praxis-ai-apis = { workspace = true } praxis-core = { workspace = true } diff --git a/filters/src/lib.rs b/filters/src/lib.rs index ae98cd3e05..165d955f35 100644 --- a/filters/src/lib.rs +++ b/filters/src/lib.rs @@ -18,6 +18,7 @@ pub mod callout; pub mod gcp; pub mod guardrails; pub mod inference; +pub mod metering; #[cfg(feature = "opentelemetry")] mod opentelemetry; pub mod prompt_enrich; @@ -38,6 +39,7 @@ pub use callout::HttpCalloutFilter; pub use gcp::GcpAdcFilter; pub use guardrails::AiGuardrailsFilter; pub use inference::ModelToHeaderFilter; +pub use metering::ExternalMeteringFilter; pub use prompt_enrich::PromptEnrichFilter; pub use register::{build_ai_registry, register_ai_filters}; pub use routing::{CredentialInjectFilter, IntelligentRouteFilter, ProviderRouteFilter}; diff --git a/filters/src/metering/config.rs b/filters/src/metering/config.rs new file mode 100644 index 0000000000..a1e2be3743 --- /dev/null +++ b/filters/src/metering/config.rs @@ -0,0 +1,143 @@ +// SPDX-License-Identifier: MIT +// Copyright (c) 2026 Praxis Contributors + +//! Deserialized YAML configuration types for the external metering filter. + +use praxis_filter::FilterError; +use serde::Deserialize; + +/// Default HTTP timeout for metering service calls (5 seconds). +const DEFAULT_TIMEOUT_SECONDS: u64 = 5; + +/// Default entitlement feature key for balance checks. +const DEFAULT_FEATURE_KEY: &str = "inference-tokens"; + +/// Default `CloudEvents` `source` field. +const DEFAULT_SOURCE: &str = "ai-gateway"; + +/// Default header prefix for tenant identity headers. +const DEFAULT_IDENTITY_HEADER_PREFIX: &str = "x-tenant-"; + +/// Default metadata namespace the `identity_header_guard` filter writes +/// captured identity headers under. +const DEFAULT_IDENTITY_METADATA_NAMESPACE: &str = "identity"; + +/// Deserialized YAML config for the `external_metering` filter. +/// +/// ```yaml +/// filter: external_metering +/// metering_url: "http://metering-service:8080" +/// timeout_seconds: 5 +/// feature_key: "inference-tokens" +/// source: "ai-gateway" +/// fail_open: true +/// identity_header_prefix: "x-tenant-" +/// identity_metadata_namespace: "identity" +/// default_username: "anonymous" +/// default_model: "unknown" +/// ``` +#[derive(Debug, Deserialize)] +#[serde(deny_unknown_fields)] +pub(super) struct ExternalMeteringConfig { + /// Base URL of the external metering service (required). + pub metering_url: String, + + /// HTTP timeout in seconds for all metering calls. + #[serde(default = "default_timeout_seconds")] + pub timeout_seconds: u64, + + /// Entitlement feature key used in balance check URL path. + #[serde(default = "default_feature_key")] + pub feature_key: String, + + /// `CloudEvents` `source` field value. + #[serde(default = "default_source")] + pub source: String, + + /// When `true` (default), requests proceed if the metering service + /// is unavailable. When `false`, requests are rejected with 503. + #[serde(default = "default_true")] + pub fail_open: bool, + + /// Prefix for tenant identity headers to capture and strip. + /// Expected headers: `{prefix}username`, `{prefix}group`, + /// `{prefix}subscription`, `{prefix}model`. + #[serde(default = "default_identity_header_prefix")] + pub identity_header_prefix: String, + + /// Metadata namespace the `identity_header_guard` filter writes captured + /// identity headers under. Must match that filter's + /// `metadata_namespace` setting when both run in one pipeline. + #[serde(default = "default_identity_metadata_namespace")] + pub identity_metadata_namespace: String, + + /// Fallback username when no identity header is present. + /// If set, requests without `{prefix}username` are still metered + /// under this name. If unset, metering is skipped entirely. + #[serde(default)] + pub default_username: Option, + + /// Fallback model name when no identity model header is present. + #[serde(default)] + pub default_model: Option, +} + +/// Validate config at construction time. +pub(super) fn validate_config(cfg: &ExternalMeteringConfig) -> Result<(), FilterError> { + if cfg.metering_url.is_empty() { + return Err("external_metering: metering_url must not be empty".into()); + } + + if cfg.timeout_seconds == 0 { + return Err("external_metering: timeout_seconds must be greater than 0".into()); + } + + if cfg.identity_header_prefix.is_empty() { + return Err("external_metering: identity_header_prefix must not be empty".into()); + } + + // A prefix with characters that cannot appear in an HTTP header name + // (e.g. a space) would silently match nothing, leaving tenant identity + // headers unstripped while the filter reports healthy. + if http::header::HeaderName::from_bytes(cfg.identity_header_prefix.as_bytes()).is_err() { + return Err( + "external_metering: identity_header_prefix must contain only valid HTTP header name characters".into(), + ); + } + + if cfg.identity_metadata_namespace.is_empty() { + return Err("external_metering: identity_metadata_namespace must not be empty".into()); + } + + Ok(()) +} + +/// Serde default for `timeout_seconds`. +fn default_timeout_seconds() -> u64 { + DEFAULT_TIMEOUT_SECONDS +} + +/// Serde default for `feature_key`. +fn default_feature_key() -> String { + DEFAULT_FEATURE_KEY.to_owned() +} + +/// Serde default for `source`. +fn default_source() -> String { + DEFAULT_SOURCE.to_owned() +} + +/// Serde default for `fail_open`. +fn default_true() -> bool { + true +} + +/// Serde default for `identity_header_prefix`. +fn default_identity_header_prefix() -> String { + DEFAULT_IDENTITY_HEADER_PREFIX.to_owned() +} + +/// Serde default for `identity_metadata_namespace`. +fn default_identity_metadata_namespace() -> String { + DEFAULT_IDENTITY_METADATA_NAMESPACE.to_owned() +} diff --git a/filters/src/metering/mod.rs b/filters/src/metering/mod.rs new file mode 100644 index 0000000000..14be0bb8c4 --- /dev/null +++ b/filters/src/metering/mod.rs @@ -0,0 +1,982 @@ +// SPDX-License-Identifier: MIT +// Copyright (c) 2026 Praxis Contributors + +//! External metering filter: pre-request balance checks and post-response +//! token usage reporting via [`CloudEvents`] to an external metering service. +//! +//! Reads token counts from [`filter_metadata`] keys set by the `token_count` +//! filter (`token.input`, `token.output`, `token.total`, and the prompt cache +//! breakdown `token.cache_read` / `token.cache_write`). The metering filter +//! must be declared *before* `token_count` in the YAML filter chain so that +//! response hooks (which run in reverse order) execute after token extraction. +//! +//! [`CloudEvents`]: https://github.com/cloudevents/spec/blob/v1.0.2/cloudevents/spec.md +//! [`filter_metadata`]: HttpFilterContext::filter_metadata + +mod config; + +#[cfg(test)] +#[expect(clippy::allow_attributes, reason = "blanket test suppressions")] +#[allow( + clippy::unwrap_used, + clippy::expect_used, + clippy::indexing_slicing, + clippy::needless_raw_strings, + clippy::needless_raw_string_hashes, + reason = "tests" +)] +mod tests; + +use std::time::Duration; + +use async_trait::async_trait; +use bytes::Bytes; +use http::header::HeaderName; +use metrics::counter; +use percent_encoding::{AsciiSet, CONTROLS, utf8_percent_encode}; +use praxis_ai_apis::subrequest::{self, SubRequest, SubRequestClient, SubRequestError, SubResponse}; +use praxis_core::subrequest::SubRequestConnector; +use praxis_filter::{ + BodyAccess, BodyMode, FilterAction, FilterError, HttpFilter, HttpFilterContext, Rejection, parse_filter_config, +}; +use serde::Deserialize; +use tracing::{debug, trace, warn}; + +use self::config::{ExternalMeteringConfig, validate_config}; + +// ----------------------------------------------------------------------------- +// Constants +// ----------------------------------------------------------------------------- + +/// `CloudEvents` spec version. +const CE_SPEC_VERSION: &str = "1.0"; + +/// `CloudEvent` type for provider error responses. +const CE_TYPE_ERROR: &str = "inference.request.error"; + +/// `CloudEvent` type for successful token usage. +const CE_TYPE_USAGE: &str = "inference.tokens.used"; + +/// Metadata key holding the resolved model name, written during the request +/// body phase and read during the response body phase. +const META_METERING_MODEL: &str = "metering.model"; + +/// Well-known `filter_metadata` key for input tokens (set by `token_count`). +const META_TOKEN_INPUT: &str = "token.input"; + +/// Well-known `filter_metadata` key for output tokens (set by `token_count`). +const META_TOKEN_OUTPUT: &str = "token.output"; + +/// Well-known `filter_metadata` key for total tokens (set by `token_count`). +const META_TOKEN_TOTAL: &str = "token.total"; + +/// Well-known `filter_metadata` key for prompt cache reads (set by +/// `token_count`). A breakdown of [`META_TOKEN_INPUT`], not an addition to it. +const META_TOKEN_CACHE_READ: &str = "token.cache_read"; + +/// Well-known `filter_metadata` key for prompt cache writes (set by +/// `token_count`). A breakdown of [`META_TOKEN_INPUT`], not an addition to it. +const META_TOKEN_CACHE_WRITE: &str = "token.cache_write"; + +/// Counter incremented whenever a usage or error event fails to reach the +/// metering service (transport failure or non-2xx acknowledgement). +const METRIC_REPORT_FAILURES: &str = "praxis_ai_metering_report_failures_total"; + +/// Upper bound on metering service response bodies. +/// +/// Balance and event-acknowledgement responses are small JSON documents; +/// anything larger indicates a misbehaving service and is cut off. +const MAX_CALLOUT_RESPONSE_BYTES: usize = 64 * 1024; + +/// Keepalive pool size for the private per-filter sub-request connector +/// created by [`ExternalMeteringFilter::from_config`]. +const PRIVATE_POOL_SIZE: usize = 4; + +/// Characters escaped when interpolating values into the balance check URL. +/// +/// Deliberately narrower than [`percent_encoding::NON_ALPHANUMERIC`]: feature +/// keys and model names routinely contain hyphens and dots, which are valid +/// path characters and must survive unescaped for the metering service to +/// match them. +const PATH_SEGMENT: &AsciiSet = &CONTROLS + .add(b' ') + .add(b'"') + .add(b'#') + .add(b'%') + .add(b'/') + .add(b':') + .add(b'<') + .add(b'>') + .add(b'?') + .add(b'@') + .add(b'`') + .add(b'{') + .add(b'}'); + +/// Status returned when the tenant has no remaining token budget. +const STATUS_BUDGET_EXHAUSTED: u16 = 429; + +/// Status returned when metering is unreachable and `fail_open` is disabled. +const STATUS_METERING_UNAVAILABLE: u16 = 503; + +// ----------------------------------------------------------------------------- +// ExternalMeteringFilter +// ----------------------------------------------------------------------------- + +/// Integrates with an external metering service for pre-request balance +/// checks and post-response token usage reporting. +/// +/// Tenant identity is resolved from the highest-trust source available: +/// verified `{prefix}*` metadata written by an authentication filter, +/// then the `identity_header_guard` filter's namespaced +/// `{namespace}.{prefix}*` metadata, then raw `{prefix}*` request +/// headers. A higher tier always wins, so forged client headers can +/// never override verified claims, and identity headers plus client +/// credentials are always stripped before the request is forwarded. +/// +/// # YAML +/// +/// ```yaml +/// filter: external_metering +/// metering_url: "http://metering-service:8080" +/// timeout_seconds: 5 +/// feature_key: "inference-tokens" +/// source: "ai-gateway" +/// fail_open: true +/// identity_header_prefix: "x-tenant-" +/// identity_metadata_namespace: "identity" +/// default_username: "anonymous" +/// default_model: "unknown" +/// ``` +pub struct ExternalMeteringFilter { + /// Model name reported when neither the identity header nor the request + /// body reveals one. + default_model: Option, + + /// Username reported when no identity header is present. When unset, + /// unidentified requests are not metered at all. + default_username: Option, + + /// Whether to admit requests when the metering service is unreachable. + fail_open: bool, + + /// Entitlement feature key used in the balance check path. + feature_key: String, + + /// Prefix of the tenant identity headers to capture and strip. + /// Lowercased once at construction. + identity_header_prefix: String, + + /// Metadata namespace the `identity_header_guard` filter writes + /// captured headers under (tier 2 of identity resolution). + identity_metadata_namespace: String, + + /// Base URL of the external metering service. + metering_url: String, + + /// `CloudEvents` `source` attribute for emitted events. + source: String, + + /// Shared HTTP client for balance checks and usage reports. + subrequest_client: SubRequestClient, + + /// HTTP timeout applied to each metering call. + timeout: Duration, +} + +impl ExternalMeteringFilter { + /// Create from parsed YAML config with a private sub-request client. + /// + /// The private connector uses a small keepalive pool. Prefer + /// [`from_config_with_client`] when a shared server-level client is + /// available. + /// + /// # Errors + /// + /// Returns [`FilterError`] if config parsing or validation fails. + /// + /// [`from_config_with_client`]: Self::from_config_with_client + pub fn from_config(config: &serde_yaml::Value) -> Result, FilterError> { + let client = SubRequestClient::new(SubRequestConnector::new(PRIVATE_POOL_SIZE, None)); + Ok(Box::new(Self::build(config, client)?)) + } + + /// Create from parsed YAML config using the shared [`SubRequestClient`]. + /// + /// The shared client inherits the server-level pool size and + /// connection limits from the runtime configuration. + /// + /// # Errors + /// + /// Returns [`FilterError`] if config parsing or validation fails. + pub fn from_config_with_client( + config: &serde_yaml::Value, + client: SubRequestClient, + ) -> Result, FilterError> { + Ok(Box::new(Self::build(config, client)?)) + } + + /// Build the concrete filter from parsed YAML config. + fn build(config: &serde_yaml::Value, subrequest_client: SubRequestClient) -> Result { + let cfg: ExternalMeteringConfig = parse_filter_config("external_metering", config)?; + validate_config(&cfg)?; + + Ok(Self { + default_model: cfg.default_model, + default_username: cfg.default_username, + fail_open: cfg.fail_open, + feature_key: cfg.feature_key, + identity_header_prefix: cfg.identity_header_prefix.to_ascii_lowercase(), + identity_metadata_namespace: cfg.identity_metadata_namespace, + metering_url: cfg.metering_url, + source: cfg.source, + subrequest_client, + timeout: Duration::from_secs(cfg.timeout_seconds), + }) + } + + /// Ask the metering service whether the tenant may spend more tokens. + /// + /// The metering service expresses denial only through a `2xx` response + /// whose body carries `hasAccess: false`. Any non-`2xx` status means the + /// service itself is misbehaving (bad route, internal error), so it is + /// handled by the availability policy (`fail_open`) rather than treated + /// as a denial. + async fn check_balance(&self, state: &MeteringState) -> FilterAction { + let url = build_balance_url(&self.metering_url, &state.username, &self.feature_key, &state.model); + let request = SubRequest { + method: http::Method::GET, + uri: http::Uri::default(), + headers: http::HeaderMap::new(), + body: Bytes::new(), + }; + + let result = subrequest::execute_url( + &self.subrequest_client, + &url, + request, + MAX_CALLOUT_RESPONSE_BYTES, + self.timeout, + ) + .await; + + match result { + Ok(resp) if is_success(resp.status) => parse_balance_result(&resp.body, self.fail_open), + Ok(resp) => { + warn!(status = resp.status, "balance check returned a non-2xx status"); + self.on_metering_unavailable() + }, + Err(error) => { + warn!(%error, "balance check unreachable"); + self.on_metering_unavailable() + }, + } + } + + /// Resolve the action to take when the balance check cannot be completed. + fn on_metering_unavailable(&self) -> FilterAction { + if self.fail_open { + trace!("admitting request (fail-open)"); + FilterAction::Continue + } else { + reject_unavailable() + } + } + + /// Resolve the model to report: the captured model, else the + /// body/header metadata, else the configured default. + fn resolve_report_model(&self, ctx: &HttpFilterContext<'_>, state: &MeteringState) -> String { + if !state.model.is_empty() { + return state.model.clone(); + } + if let Some(model) = ctx.filter_metadata.get(META_METERING_MODEL) { + return model.clone(); + } + self.default_model.clone().unwrap_or_default() + } + + /// Emit the terminal usage or error event for a completed request. + fn report(&self, ctx: &HttpFilterContext<'_>, mut state: MeteringState) { + state.model = self.resolve_report_model(ctx, &state); + + let request_id = ctx.id_generator.generate(ctx.time_source); + let provider = ctx.cluster_name().unwrap_or_default().to_owned(); + let event_ctx = EventContext { + duration_ms: u64::try_from(state.request_start.elapsed().as_millis()).unwrap_or(u64::MAX), + event_id: &request_id, + provider: &provider, + source: &self.source, + state: &state, + }; + + let event = if state.is_error { + build_error_event(&event_ctx) + } else { + build_usage_event(&event_ctx, &TokenCounts::read(ctx)) + }; + + spawn_usage_report(self.subrequest_client.clone(), &self.metering_url, self.timeout, &event); + } +} + +#[async_trait] +impl HttpFilter for ExternalMeteringFilter { + fn name(&self) -> &'static str { + "external_metering" + } + + async fn on_request(&self, ctx: &mut HttpFilterContext<'_>) -> Result { + let mut state = capture_identity(ctx, &self.identity_header_prefix, &self.identity_metadata_namespace); + + if state.username.is_empty() { + let Some(fallback) = self.default_username.as_ref() else { + trace!("no tenant identity header, skipping metering"); + return Ok(FilterAction::Continue); + }; + state.username.clone_from(fallback); + } + + if !state.model.is_empty() { + ctx.filter_metadata + .insert(META_METERING_MODEL.to_owned(), state.model.clone()); + } + + let action = self.check_balance(&state).await; + store_state(ctx, state); + + Ok(action) + } + + fn request_body_access(&self) -> BodyAccess { + BodyAccess::ReadOnly + } + + fn request_body_mode(&self) -> BodyMode { + BodyMode::Stream + } + + async fn on_request_body( + &self, + ctx: &mut HttpFilterContext<'_>, + body: &mut Option, + _end_of_stream: bool, + ) -> Result { + // The identity header wins when present; only fall back to the body so + // that clients that do not send a model header are still attributed. + let unresolved = !ctx.filter_metadata.contains_key(META_METERING_MODEL); + if let Some(model) = body + .as_ref() + .filter(|_| unresolved) + .and_then(|chunk| extract_model_from_bytes(chunk)) + { + ctx.filter_metadata.insert(META_METERING_MODEL.to_owned(), model); + } + + Ok(FilterAction::Release) + } + + async fn on_response(&self, ctx: &mut HttpFilterContext<'_>) -> Result { + let status = ctx.response_header.as_ref().map_or(0, |r| r.status.as_u16()); + let is_error = ctx.response_header.as_ref().is_some_and(|r| !r.status.is_success()); + + if let Some(state) = ctx + .filter_state + .get_mut(&filter_state_key(ctx.current_filter_id)) + .and_then(|s| s.downcast_mut::()) + { + state.is_error = is_error; + state.response_status = status; + } + + Ok(FilterAction::Continue) + } + + fn response_body_access(&self) -> BodyAccess { + BodyAccess::ReadOnly + } + + fn response_body_mode(&self) -> BodyMode { + BodyMode::Stream + } + + fn on_response_body( + &self, + ctx: &mut HttpFilterContext<'_>, + _body: &mut Option, + end_of_stream: bool, + ) -> Result { + if !end_of_stream { + return Ok(FilterAction::Continue); + } + + let Some(state) = ctx + .filter_state + .remove(&filter_state_key(ctx.current_filter_id)) + .and_then(|s| s.downcast::().ok()) + else { + return Ok(FilterAction::Continue); + }; + + if !state.username.is_empty() { + self.report(ctx, *state); + } + + Ok(FilterAction::Continue) + } +} + +// ----------------------------------------------------------------------------- +// Per-Request State +// ----------------------------------------------------------------------------- + +/// Identity and timing state captured during the request phase and consumed +/// during the response phase. +struct MeteringState { + /// Tenant group, falling back to the subscription when absent. + group: String, + + /// Whether the upstream returned a non-success status. + is_error: bool, + + /// Model attributed to this request. + model: String, + + /// Start of the request, used to derive the reported duration. + request_start: std::time::Instant, + + /// Upstream response status, reported on error events. + response_status: u16, + + /// Tenant subscription identifier. + subscription: String, + + /// Client user agent, reported for attribution. + user_agent: String, + + /// Tenant username; empty means the request is not metered. + username: String, +} + +/// Token counts published by the `token_count` filter. +struct TokenCounts { + /// Completion tokens. + output: u64, + + /// Prompt tokens. + input: u64, + + /// Total tokens billed. + total: u64, + + /// Prompt tokens served from the provider's cache; a subset of `input`. + cache_read: u64, + + /// Prompt tokens written to the provider's cache; a subset of `input`. + cache_write: u64, +} + +impl TokenCounts { + /// Read the counts the `token_count` filter left in `filter_metadata`. + fn read(ctx: &HttpFilterContext<'_>) -> Self { + Self { + input: read_token_meta(ctx, META_TOKEN_INPUT), + output: read_token_meta(ctx, META_TOKEN_OUTPUT), + total: read_token_meta(ctx, META_TOKEN_TOTAL), + cache_read: read_token_meta(ctx, META_TOKEN_CACHE_READ), + cache_write: read_token_meta(ctx, META_TOKEN_CACHE_WRITE), + } + } +} + +/// Fields shared by the usage and error `CloudEvents`. +struct EventContext<'a> { + /// Wall-clock duration of the proxied request. + duration_ms: u64, + + /// Unique `CloudEvents` `id`. + event_id: &'a str, + + /// Upstream cluster the request was routed to. + provider: &'a str, + + /// Configured `CloudEvents` `source`. + source: &'a str, + + /// Captured tenant identity. + state: &'a MeteringState, +} + +/// JSON response from the metering service balance check endpoint. +#[derive(Debug, Deserialize)] +#[serde(rename_all = "camelCase")] +struct BalanceResponse { + /// Whether the tenant may spend more tokens. + has_access: bool, +} + +// ----------------------------------------------------------------------------- +// Identity Header Capture +// ----------------------------------------------------------------------------- + +/// Capture tenant identity from the configured headers and mark those headers, +/// along with client credentials, for removal before the request is forwarded. +/// +/// `prefix` must already be lowercase; [`ExternalMeteringFilter::build`] +/// lowercases it once at construction. +fn capture_identity(ctx: &mut HttpFilterContext<'_>, prefix: &str, identity_namespace: &str) -> MeteringState { + let user_agent = ctx + .request + .headers + .get("user-agent") + .and_then(|v| v.to_str().ok()) + .unwrap_or_default() + .to_owned(); + + let mut identity = read_identity_headers(ctx, prefix, identity_namespace); + strip_client_credentials(ctx); + + if identity.group.is_empty() { + identity.group.clone_from(&identity.subscription); + } + + MeteringState { + group: identity.group, + is_error: false, + model: identity.model, + request_start: std::time::Instant::now(), + response_status: 0, + subscription: identity.subscription, + user_agent, + username: identity.username, + } +} + +/// Tenant identity as carried on the request headers. +#[derive(Default)] +struct Identity { + /// Value of `{prefix}group`. + group: String, + + /// Value of `{prefix}model`. + model: String, + + /// Value of `{prefix}subscription`. + subscription: String, + + /// Value of `{prefix}username`. + username: String, +} + +/// Resolve tenant identity from the highest-trust source available. +/// +/// Three tiers, most trusted first: +/// +/// 1. Unnamespaced `{prefix}*` metadata keys, written by an authentication filter from verified credentials (e.g. JWT +/// claims or a validated API key). When any of these are present, every lower tier is ignored entirely so a client +/// cannot spoof the remaining fields via forged headers alongside valid credentials. +/// 2. Namespaced `{namespace}.{prefix}*` metadata keys, written by the `identity_header_guard` filter from captured +/// request headers. +/// 3. Raw `{prefix}*` request headers, set by a trusted upstream auth layer when neither metadata tier is populated. +/// +/// Identity headers are always marked for removal so tenant identity +/// never leaks to the upstream provider, regardless of which tier +/// supplied the identity. +fn read_identity_headers(ctx: &mut HttpFilterContext<'_>, prefix: &str, identity_namespace: &str) -> Identity { + let mut identity = Identity::default(); + read_metadata_identity(ctx, prefix, "", &mut identity); + + // Any verified field blocks the lower tiers entirely: an auth filter + // may map only some claims (e.g. group without username), and a + // partially verified identity must not be extended by forgeable + // sources. + let has_verified_identity = !(identity.username.is_empty() + && identity.group.is_empty() + && identity.subscription.is_empty() + && identity.model.is_empty()); + + if !has_verified_identity { + read_metadata_identity(ctx, prefix, &format!("{identity_namespace}."), &mut identity); + } + + let has_metadata_identity = has_verified_identity || !identity.username.is_empty(); + if has_metadata_identity { + strip_identity_headers(ctx, prefix); + } else { + read_header_identity(ctx, prefix, &mut identity); + } + + identity +} + +/// Read the `{namespace}{prefix}*` identity keys from `filter_metadata`, +/// overwriting only the fields the namespace carries. +fn read_metadata_identity(ctx: &HttpFilterContext<'_>, prefix_lower: &str, namespace: &str, identity: &mut Identity) { + if let Some(val) = ctx.filter_metadata.get(&format!("{namespace}{prefix_lower}group")) { + identity.group.clone_from(val); + } + if let Some(val) = ctx.filter_metadata.get(&format!("{namespace}{prefix_lower}model")) { + identity.model.clone_from(val); + } + if let Some(val) = ctx + .filter_metadata + .get(&format!("{namespace}{prefix_lower}subscription")) + { + identity.subscription.clone_from(val); + } + if let Some(val) = ctx.filter_metadata.get(&format!("{namespace}{prefix_lower}username")) { + identity.username.clone_from(val); + } +} + +/// Read identity from raw `{prefix}*` request headers, marking each one +/// for removal. +/// +/// [`HeaderName::as_str`] already returns the lowercased name, so the +/// pre-lowercased prefix compares directly without per-header allocation. +fn read_header_identity(ctx: &mut HttpFilterContext<'_>, prefix_lower: &str, identity: &mut Identity) { + for (key, value) in &ctx.request.headers { + let Some(suffix) = key.as_str().strip_prefix(prefix_lower) else { + continue; + }; + let val = value.to_str().unwrap_or_default(); + + match suffix { + "group" if identity.group.is_empty() => val.clone_into(&mut identity.group), + "model" if identity.model.is_empty() => val.clone_into(&mut identity.model), + "subscription" if identity.subscription.is_empty() => val.clone_into(&mut identity.subscription), + "username" if identity.username.is_empty() => val.clone_into(&mut identity.username), + _ => {}, + } + + ctx.request_headers_to_remove.push(key.clone()); + } +} + +/// Mark every `{prefix}*` header for removal without reading it, so +/// unused identity headers still never reach the upstream provider. +fn strip_identity_headers(ctx: &mut HttpFilterContext<'_>, prefix_lower: &str) { + for (key, _) in &ctx.request.headers { + if key.as_str().starts_with(prefix_lower) { + ctx.request_headers_to_remove.push(key.clone()); + } + } +} + +/// Mark client-supplied credentials for removal. +/// +/// `accept-encoding` is stripped alongside them so the upstream returns an +/// uncompressed body: `token_count` parses the response inline and cannot read +/// usage out of a compressed stream. +fn strip_client_credentials(ctx: &mut HttpFilterContext<'_>) { + ctx.request_headers_to_remove.push(http::header::AUTHORIZATION); + ctx.request_headers_to_remove.push(http::header::ACCEPT_ENCODING); + ctx.request_headers_to_remove.push(HeaderName::from_static("x-api-key")); +} + +// ----------------------------------------------------------------------------- +// Balance Check +// ----------------------------------------------------------------------------- + +/// Interpret a balance check response body. +fn parse_balance_result(body: &[u8], fail_open: bool) -> FilterAction { + let Some(balance) = decode_balance(body) else { + return admit_or_reject(fail_open); + }; + + if balance.has_access { + trace!("balance check passed"); + FilterAction::Continue + } else { + reject_budget_exhausted() + } +} + +/// Decode a balance payload, logging why an unusable one was discarded. +fn decode_balance(body: &[u8]) -> Option { + if body.is_empty() { + debug!("balance check returned an empty body"); + return None; + } + + match serde_json::from_slice::(body) { + Ok(balance) => Some(balance), + Err(e) => { + debug!("balance response parse error: {e}"); + None + }, + } +} + +/// Reject with a 429 because the tenant has no remaining token budget. +fn reject_budget_exhausted() -> FilterAction { + debug!("token budget exhausted"); + FilterAction::Reject( + Rejection::status(STATUS_BUDGET_EXHAUSTED).with_body(Bytes::from_static(b"token budget exhausted")), + ) +} + +/// Admit the request when configured to fail open, reject it otherwise. +fn admit_or_reject(fail_open: bool) -> FilterAction { + if fail_open { + FilterAction::Continue + } else { + reject_unavailable() + } +} + +/// Reject with a 503 because metering could not authorize the request. +fn reject_unavailable() -> FilterAction { + FilterAction::Reject( + Rejection::status(STATUS_METERING_UNAVAILABLE).with_body(Bytes::from_static(b"metering system unavailable")), + ) +} + +// ----------------------------------------------------------------------------- +// Usage Reporting (fire-and-forget) +// ----------------------------------------------------------------------------- + +/// Deliver an event without blocking the response. +/// +/// Metering is an observer: a slow or failing metering service must never add +/// latency to, or fail, a request the upstream already answered. +fn spawn_usage_report(client: SubRequestClient, metering_url: &str, timeout: Duration, event: &serde_json::Value) { + let url = format!("{}/api/v1/events", metering_url.trim_end_matches('/')); + + let body = match serde_json::to_vec(event) { + Ok(b) => b, + Err(e) => { + warn!("failed to serialize metering event: {e}"); + return; + }, + }; + + tokio::spawn(async move { + let mut headers = http::HeaderMap::new(); + headers.insert( + http::header::CONTENT_TYPE, + http::HeaderValue::from_static("application/json"), + ); + let request = SubRequest { + method: http::Method::POST, + uri: http::Uri::default(), + headers, + body: Bytes::from(body), + }; + + report_delivery(subrequest::execute_url(&client, &url, request, MAX_CALLOUT_RESPONSE_BYTES, timeout).await); + }); +} + +/// Log the outcome of a usage report delivery and count failures. +/// +/// Emits [`METRIC_REPORT_FAILURES`] so dropped billing events are visible +/// on a dashboard rather than only in debug logs. +fn report_delivery(result: Result) { + match result { + Ok(resp) if is_success(resp.status) => { + trace!(status = resp.status, "metering usage report sent"); + }, + Ok(resp) => { + counter!(METRIC_REPORT_FAILURES).increment(1); + warn!(status = resp.status, "metering usage report rejected"); + }, + Err(error) => { + counter!(METRIC_REPORT_FAILURES).increment(1); + warn!(%error, "metering usage report failed"); + }, + } +} + +/// Whether an HTTP status code is in the 2xx range. +fn is_success(status: u16) -> bool { + (200..300).contains(&status) +} + +// ----------------------------------------------------------------------------- +// URL Construction +// ----------------------------------------------------------------------------- + +/// Build the entitlement balance check URL for a tenant and model. +fn build_balance_url(base_url: &str, customer_id: &str, feature_key: &str, model: &str) -> String { + let base = base_url.trim_end_matches('/'); + let customer = utf8_percent_encode(customer_id, PATH_SEGMENT); + let feature = utf8_percent_encode(feature_key, PATH_SEGMENT); + let model = utf8_percent_encode(model, PATH_SEGMENT); + + format!("{base}/api/v1/customers/{customer}/entitlements/{feature}/value?model={model}") +} + +// ----------------------------------------------------------------------------- +// CloudEvent Construction +// ----------------------------------------------------------------------------- + +/// Build an `inference.tokens.used` event. +fn build_usage_event(ctx: &EventContext<'_>, tokens: &TokenCounts) -> serde_json::Value { + let mut event = build_envelope(ctx, CE_TYPE_USAGE); + + if let Some(data) = event.get_mut("data").and_then(serde_json::Value::as_object_mut) { + data.insert("prompt_tokens".to_owned(), tokens.input.into()); + data.insert("completion_tokens".to_owned(), tokens.output.into()); + data.insert("total_tokens".to_owned(), tokens.total.into()); + data.insert("cached_input_tokens".to_owned(), tokens.cache_read.into()); + data.insert("cache_creation_tokens".to_owned(), tokens.cache_write.into()); + } + + event +} + +/// Build an `inference.request.error` event. +fn build_error_event(ctx: &EventContext<'_>) -> serde_json::Value { + let mut event = build_envelope(ctx, CE_TYPE_ERROR); + + if let Some(data) = event.get_mut("data").and_then(serde_json::Value::as_object_mut) { + data.insert("status_code".to_owned(), ctx.state.response_status.into()); + } + + event +} + +/// Build the `CloudEvents` envelope and the attribution fields both event +/// types carry. +fn build_envelope(ctx: &EventContext<'_>, event_type: &str) -> serde_json::Value { + let state = ctx.state; + + serde_json::json!({ + "specversion": CE_SPEC_VERSION, + "id": ctx.event_id, + "source": ctx.source, + "type": event_type, + "subject": state.username, + "time": chrono::Utc::now().to_rfc3339(), + "datacontenttype": "application/json", + "data": { + "user": state.username, + "group": state.group, + "subscription": state.subscription, + "provider": ctx.provider, + "model": state.model, + "duration_ms": ctx.duration_ms, + "user_agent": state.user_agent, + } + }) +} + +// ----------------------------------------------------------------------------- +// Helpers +// ----------------------------------------------------------------------------- + +/// Extract the top-level `model` field from a JSON body fragment. +/// +/// Request bodies arrive in chunks and can be megabytes long, so this scans +/// with a string- and depth-aware state machine instead of buffering and +/// deserializing the whole document. Only a `"model"` key belonging to the +/// top-level object is matched, so `"model"` occurrences nested inside +/// `messages` content can never misattribute the request. Returns `None` +/// when the chunk does not contain the complete top-level `"model": "..."` +/// pair. +fn extract_model_from_bytes(bytes: &[u8]) -> Option { + let text = std::str::from_utf8(bytes).ok()?; + let mut scanner = TopLevelKeyScanner::new(text); + + while let Some(key) = scanner.next_top_level_key() { + if key == "model" { + return scanner.string_value().map(str::to_owned); + } + } + + None +} + +/// Incremental scanner over the top-level keys of a JSON object fragment. +/// +/// Tracks nesting depth and string boundaries (including escapes) so keys +/// inside nested objects, arrays, or string values are never surfaced. +struct TopLevelKeyScanner<'a> { + /// Remaining unscanned input. + rest: &'a str, + + /// Current object/array nesting depth; the document object is depth 1. + depth: u32, +} + +impl<'a> TopLevelKeyScanner<'a> { + /// Start scanning at the beginning of a JSON document fragment. + fn new(text: &'a str) -> Self { + Self { rest: text, depth: 0 } + } + + /// Advance to the next key of the top-level object and return it. + fn next_top_level_key(&mut self) -> Option<&'a str> { + loop { + let mut chars = self.rest.char_indices(); + let (pos, ch) = chars.next()?; + match ch { + '{' | '[' => { + self.depth = self.depth.checked_add(1)?; + self.rest = self.rest.get(pos + 1..)?; + }, + '}' | ']' => { + self.depth = self.depth.checked_sub(1)?; + self.rest = self.rest.get(pos + 1..)?; + }, + '"' => { + let start = pos + 1; + let end = find_string_end(self.rest, start)?; + let content = self.rest.get(start..end)?; + let after = self.rest.get(end + 1..)?; + let is_key = after.trim_start().starts_with(':'); + self.rest = after; + if self.depth == 1 && is_key { + return Some(content); + } + }, + _ => { + self.rest = self.rest.get(pos + ch.len_utf8()..)?; + }, + } + } + } + + /// Read the string value following the key just returned. + /// + /// Returns `None` when the value is not a string (e.g. `null`) or the + /// fragment is cut off before the closing quote. + fn string_value(&self) -> Option<&'a str> { + let after_colon = self.rest.trim_start().strip_prefix(':')?; + let value = after_colon.trim_start().strip_prefix('"')?; + let end = find_string_end(value, 0)?; + value.get(..end) + } +} + +/// Find the byte offset of the unescaped closing quote for the string +/// starting at `from` (which must point just past the opening quote). +fn find_string_end(text: &str, from: usize) -> Option { + let mut escaped = false; + for (offset, ch) in text.get(from..)?.char_indices() { + if escaped { + escaped = false; + } else if ch == '\\' { + escaped = true; + } else if ch == '"' { + return Some(from + offset); + } + } + None +} + +/// Read a numeric `filter_metadata` value, defaulting to zero. +fn read_token_meta(ctx: &HttpFilterContext<'_>, key: &str) -> u64 { + ctx.filter_metadata.get(key).and_then(|v| v.parse().ok()).unwrap_or(0) +} + +/// Key under which this filter instance stores its per-request state. +fn filter_state_key(id: Option) -> usize { + id.unwrap_or(0) +} + +/// Store per-request state for retrieval during the response phase. +fn store_state(ctx: &mut HttpFilterContext<'_>, state: MeteringState) { + let key = filter_state_key(ctx.current_filter_id); + ctx.filter_state.insert(key, Box::new(state)); +} diff --git a/filters/src/metering/tests.rs b/filters/src/metering/tests.rs new file mode 100644 index 0000000000..e6a766afe7 --- /dev/null +++ b/filters/src/metering/tests.rs @@ -0,0 +1,741 @@ +// SPDX-License-Identifier: MIT +// Copyright (c) 2026 Praxis Contributors + +use super::*; +use crate::test_utils::{make_filter_context, make_request}; + +/// Build the concrete filter with a private test client. +fn build_filter(yaml: &serde_yaml::Value) -> Result { + ExternalMeteringFilter::build(yaml, SubRequestClient::new(SubRequestConnector::new(1, None))) +} + +// ----------------------------------------------------------------------------- +// Config Parsing +// ----------------------------------------------------------------------------- + +#[test] +fn valid_config_parses() { + let yaml: serde_yaml::Value = serde_yaml::from_str( + r#" +metering_url: "http://metering:8080" +"#, + ) + .unwrap(); + let filter = ExternalMeteringFilter::from_config(&yaml).unwrap(); + assert_eq!(filter.name(), "external_metering"); +} + +#[test] +fn config_with_all_fields_parses() { + let yaml: serde_yaml::Value = serde_yaml::from_str( + r#" +metering_url: "http://metering:8080" +timeout_seconds: 10 +feature_key: "custom-tokens" +source: "my-gateway" +fail_open: false +identity_header_prefix: "x-custom-" +"#, + ) + .unwrap(); + let filter = ExternalMeteringFilter::from_config(&yaml).unwrap(); + assert_eq!(filter.name(), "external_metering"); +} + +#[test] +fn config_with_fallbacks_parses() { + let yaml: serde_yaml::Value = serde_yaml::from_str( + r#" +metering_url: "http://metering:8080" +default_username: "anonymous" +default_model: "unknown" +"#, + ) + .unwrap(); + let filter = build_filter(&yaml).unwrap(); + + assert_eq!(filter.default_username.as_deref(), Some("anonymous")); + assert_eq!(filter.default_model.as_deref(), Some("unknown")); +} + +#[test] +fn config_without_fallbacks_leaves_them_unset() { + let yaml: serde_yaml::Value = serde_yaml::from_str( + r#" +metering_url: "http://metering:8080" +"#, + ) + .unwrap(); + let filter = build_filter(&yaml).unwrap(); + + assert!(filter.default_username.is_none()); + assert!(filter.default_model.is_none()); +} + +// ----------------------------------------------------------------------------- +// Fallback Runtime Behavior +// ----------------------------------------------------------------------------- + +#[tokio::test] +async fn default_username_meters_anonymous_request() { + // Unroutable metering URL makes the balance check fail instantly; + // fail-open (the default) still admits and stores the metering state. + let filter = filter_from_yaml("metering_url: \"http://127.0.0.1:1\"\ndefault_username: \"anonymous\"\n"); + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = make_filter_context(&req); + + let action = filter.on_request(&mut ctx).await.unwrap(); + + assert!(matches!(action, FilterAction::Continue)); + assert!( + !ctx.filter_state.is_empty(), + "default_username must cause metering state to be stored" + ); +} + +#[test] +fn default_model_is_used_when_no_model_was_captured() { + let filter = filter_from_yaml("metering_url: \"http://metering:8080\"\ndefault_model: \"unknown\"\n"); + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let ctx = make_filter_context(&req); + let state = state_for("alice", ""); + + assert_eq!(filter.resolve_report_model(&ctx, &state), "unknown"); +} + +#[test] +fn captured_model_wins_over_default_model() { + let filter = filter_from_yaml("metering_url: \"http://metering:8080\"\ndefault_model: \"unknown\"\n"); + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = make_filter_context(&req); + ctx.filter_metadata + .insert(META_METERING_MODEL.to_owned(), "gpt-4".to_owned()); + let state = state_for("alice", ""); + + assert_eq!(filter.resolve_report_model(&ctx, &state), "gpt-4"); +} + +#[test] +fn empty_model_without_default_stays_empty() { + let filter = filter_from_yaml("metering_url: \"http://metering:8080\"\n"); + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let ctx = make_filter_context(&req); + let state = state_for("alice", ""); + + assert_eq!(filter.resolve_report_model(&ctx, &state), ""); +} + +#[test] +fn config_empty_prefix_fails() { + let yaml: serde_yaml::Value = serde_yaml::from_str( + r#" +metering_url: "http://metering:8080" +identity_header_prefix: "" +"#, + ) + .unwrap(); + let result = ExternalMeteringFilter::from_config(&yaml); + assert!(result.is_err()); +} + +#[test] +fn config_missing_url_fails() { + let yaml: serde_yaml::Value = serde_yaml::from_str( + r#" +metering_url: "" +"#, + ) + .unwrap(); + let result = ExternalMeteringFilter::from_config(&yaml); + assert!(result.is_err()); +} + +#[test] +fn config_zero_timeout_fails() { + let yaml: serde_yaml::Value = serde_yaml::from_str( + r#" +metering_url: "http://metering:8080" +timeout_seconds: 0 +"#, + ) + .unwrap(); + let result = ExternalMeteringFilter::from_config(&yaml); + assert!(result.is_err()); +} + +#[test] +fn config_prefix_with_invalid_header_chars_fails() { + let yaml: serde_yaml::Value = serde_yaml::from_str( + r#" +metering_url: "http://metering:8080" +identity_header_prefix: "x tenant " +"#, + ) + .unwrap(); + let result = ExternalMeteringFilter::from_config(&yaml); + assert!(result.is_err(), "prefix with spaces must be rejected"); +} + +#[test] +fn config_empty_namespace_fails() { + let yaml: serde_yaml::Value = serde_yaml::from_str( + r#" +metering_url: "http://metering:8080" +identity_metadata_namespace: "" +"#, + ) + .unwrap(); + let result = ExternalMeteringFilter::from_config(&yaml); + assert!(result.is_err(), "empty namespace must be rejected"); +} + +#[test] +fn config_unknown_field_fails() { + let yaml: serde_yaml::Value = serde_yaml::from_str( + r#" +metering_url: "http://metering:8080" +unknown_field: true +"#, + ) + .unwrap(); + let result = ExternalMeteringFilter::from_config(&yaml); + assert!(result.is_err()); +} + +// ----------------------------------------------------------------------------- +// Balance Response Parsing +// ----------------------------------------------------------------------------- + +#[test] +fn balance_response_has_access_continues() { + let body = br#"{"hasAccess": true, "balance": 9000.0, "usage": 1000.0}"#; + let action = parse_balance_result(body, true); + assert!(matches!(action, FilterAction::Continue)); +} + +#[test] +fn balance_response_no_access_rejects() { + let body = br#"{"hasAccess": false, "balance": 0.0, "usage": 10000.0}"#; + let action = parse_balance_result(body, true); + assert!(matches!(action, FilterAction::Reject(_))); +} + +#[test] +fn balance_response_invalid_json_fail_open() { + let body = b"not json"; + let action = parse_balance_result(body, true); + assert!(matches!(action, FilterAction::Continue)); +} + +#[test] +fn balance_response_invalid_json_fail_closed() { + let body = b"not json"; + let action = parse_balance_result(body, false); + assert!(matches!(action, FilterAction::Reject(_))); +} + +#[test] +fn balance_response_empty_body_fail_open() { + let action = parse_balance_result(b"", true); + assert!(matches!(action, FilterAction::Continue)); +} + +#[test] +fn balance_response_empty_body_fail_closed() { + let action = parse_balance_result(b"", false); + assert!(matches!(action, FilterAction::Reject(_))); +} + +// ----------------------------------------------------------------------------- +// URL Construction +// ----------------------------------------------------------------------------- + +#[test] +fn balance_url_encodes_special_chars() { + let url = build_balance_url("http://metering:8080", "user@example.com", "inference-tokens", "gpt-4o"); + assert!(url.contains("user%40example.com"), "@ should be encoded: {url}"); + assert!(url.contains("inference-tokens"), "hyphens should not be encoded: {url}"); + assert!(url.contains("gpt-4o"), "hyphens in model should not be encoded: {url}"); +} + +#[test] +fn balance_url_strips_trailing_slash() { + let url = build_balance_url("http://metering:8080/", "testuser", "tokens", "llama"); + assert!(url.starts_with("http://metering:8080/api/")); + assert!(!url.contains("//api/")); +} + +// ----------------------------------------------------------------------------- +// CloudEvent Construction +// ----------------------------------------------------------------------------- + +#[test] +fn usage_event_has_correct_structure() { + let state = state_for("testuser", "gpt-4"); + let tokens = TokenCounts { + input: 100, + output: 50, + total: 150, + cache_read: 80, + cache_write: 20, + }; + + let event = build_usage_event(&event_ctx("evt-1", &state), &tokens); + + assert_eq!(event["specversion"], "1.0"); + assert_eq!(event["type"], CE_TYPE_USAGE); + assert_eq!(event["subject"], "testuser"); + assert_eq!(event["data"]["prompt_tokens"], 100); + assert_eq!(event["data"]["completion_tokens"], 50); + assert_eq!(event["data"]["total_tokens"], 150); + assert_eq!(event["data"]["cached_input_tokens"], 80); + assert_eq!(event["data"]["cache_creation_tokens"], 20); + assert_eq!(event["data"]["duration_ms"], 500); + assert_eq!(event["data"]["model"], "gpt-4"); +} + +#[test] +fn error_event_has_correct_structure() { + let mut state = state_for("testuser", "gpt-4"); + state.is_error = true; + state.response_status = 500; + + let event = build_error_event(&event_ctx("evt-2", &state)); + + assert_eq!(event["type"], CE_TYPE_ERROR); + assert_eq!(event["data"]["status_code"], 500); + assert_eq!(event["data"]["user"], "testuser"); +} + +// ----------------------------------------------------------------------------- +// Identity Header Capture +// ----------------------------------------------------------------------------- + +#[test] +fn captures_tenant_headers_with_default_prefix() { + let mut req = make_request(http::Method::POST, "/v1/chat/completions"); + req.headers.insert("x-tenant-username", "alice".parse().unwrap()); + req.headers.insert("x-tenant-group", "engineering".parse().unwrap()); + req.headers.insert("x-tenant-subscription", "sub-42".parse().unwrap()); + req.headers.insert("x-tenant-model", "gpt-4".parse().unwrap()); + + let mut ctx = make_filter_context(&req); + let state = capture_identity(&mut ctx, "x-tenant-", "identity"); + + assert_eq!(state.username, "alice"); + assert_eq!(state.group, "engineering"); + assert_eq!(state.subscription, "sub-42"); + assert_eq!(state.model, "gpt-4"); +} + +#[test] +fn strips_identity_and_auth_headers() { + let mut req = make_request(http::Method::POST, "/v1/chat/completions"); + req.headers.insert("x-tenant-username", "alice".parse().unwrap()); + req.headers.insert("authorization", "Bearer sk-test".parse().unwrap()); + + let mut ctx = make_filter_context(&req); + let _state = capture_identity(&mut ctx, "x-tenant-", "identity"); + + let removed: Vec<&str> = ctx.request_headers_to_remove.iter().map(HeaderName::as_str).collect(); + assert!(removed.contains(&"x-tenant-username")); + assert!(removed.contains(&"authorization")); + assert!(removed.contains(&"x-api-key")); +} + +#[test] +fn custom_prefix_captures_correctly() { + let mut req = make_request(http::Method::POST, "/v1/chat/completions"); + req.headers.insert("x-myco-username", "bob".parse().unwrap()); + + let mut ctx = make_filter_context(&req); + let state = capture_identity(&mut ctx, "x-myco-", "identity"); + + assert_eq!(state.username, "bob"); +} + +#[test] +fn missing_username_returns_empty() { + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = make_filter_context(&req); + let state = capture_identity(&mut ctx, "x-tenant-", "identity"); + + assert!(state.username.is_empty()); +} + +#[test] +fn verified_identity_ignores_forged_headers_and_guard_metadata() { + let mut req = make_request(http::Method::POST, "/v1/chat/completions"); + // Client-forged headers alongside a valid JWT. + req.headers + .insert("x-tenant-subscription", "sub-forged".parse().unwrap()); + req.headers.insert("x-tenant-model", "model-forged".parse().unwrap()); + + let mut ctx = make_filter_context(&req); + // Verified claims written by an authentication filter (unnamespaced). + ctx.filter_metadata + .insert("x-tenant-username".to_owned(), "alice".to_owned()); + ctx.filter_metadata + .insert("x-tenant-group".to_owned(), "engineering".to_owned()); + // Guard-captured copies of the forged headers (namespaced). + ctx.filter_metadata + .insert("identity.x-tenant-subscription".to_owned(), "sub-forged".to_owned()); + ctx.filter_metadata + .insert("identity.x-tenant-model".to_owned(), "model-forged".to_owned()); + + let state = capture_identity(&mut ctx, "x-tenant-", "identity"); + + assert_eq!(state.username, "alice"); + assert_eq!(state.group, "engineering"); + assert!( + state.subscription.is_empty(), + "forged subscription must be ignored when identity is verified: {}", + state.subscription + ); + assert!( + state.model.is_empty(), + "forged model must be ignored when identity is verified: {}", + state.model + ); +} + +#[test] +fn guard_metadata_supplies_identity_without_verified_claims() { + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = make_filter_context(&req); + ctx.filter_metadata + .insert("identity.x-tenant-username".to_owned(), "bob".to_owned()); + ctx.filter_metadata + .insert("identity.x-tenant-group".to_owned(), "ml".to_owned()); + ctx.filter_metadata + .insert("identity.x-tenant-subscription".to_owned(), "sub-7".to_owned()); + ctx.filter_metadata + .insert("identity.x-tenant-model".to_owned(), "claude-3".to_owned()); + + let state = capture_identity(&mut ctx, "x-tenant-", "identity"); + + assert_eq!(state.username, "bob"); + assert_eq!(state.group, "ml"); + assert_eq!(state.subscription, "sub-7"); + assert_eq!(state.model, "claude-3"); +} + +#[test] +fn partial_verified_identity_blocks_lower_tiers() { + let mut req = make_request(http::Method::POST, "/v1/chat/completions"); + req.headers + .insert("x-tenant-username", "header-mallory".parse().unwrap()); + + let mut ctx = make_filter_context(&req); + // An auth filter that maps only the group claim, no username. + ctx.filter_metadata + .insert("x-tenant-group".to_owned(), "engineering".to_owned()); + // Guard-captured copy of a forged subscription. + ctx.filter_metadata + .insert("identity.x-tenant-subscription".to_owned(), "sub-forged".to_owned()); + + let state = capture_identity(&mut ctx, "x-tenant-", "identity"); + + assert_eq!(state.group, "engineering", "verified group must survive"); + assert!( + state.subscription.is_empty(), + "guard metadata must not extend a partially verified identity: {}", + state.subscription + ); + assert!( + state.username.is_empty(), + "raw header must not extend a partially verified identity" + ); +} + +#[test] +fn custom_namespace_supplies_guard_identity() { + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = make_filter_context(&req); + ctx.filter_metadata + .insert("tenant-meta.x-tenant-username".to_owned(), "bob".to_owned()); + + let state = capture_identity(&mut ctx, "x-tenant-", "tenant-meta"); + + assert_eq!(state.username, "bob", "configured namespace must resolve tier 2"); +} + +#[test] +fn multi_value_identity_header_is_fully_stripped() { + let mut req = make_request(http::Method::POST, "/v1/chat/completions"); + req.headers.insert("x-tenant-username", "alice".parse().unwrap()); + req.headers.append("x-tenant-username", "mallory".parse().unwrap()); + + let mut ctx = make_filter_context(&req); + let state = capture_identity(&mut ctx, "x-tenant-", "identity"); + + assert_eq!(state.username, "alice", "first value wins during capture"); + let removed: Vec<&str> = ctx.request_headers_to_remove.iter().map(HeaderName::as_str).collect(); + assert!( + removed.contains(&"x-tenant-username"), + "the duplicated header name must be marked for removal (removing a name drops every value): {removed:?}" + ); +} + +#[test] +fn guard_identity_blocks_raw_header_fallback() { + let mut req = make_request(http::Method::POST, "/v1/chat/completions"); + // Raw header not captured by the guard — anomalous, must not + // be trusted once any metadata identity exists. + req.headers.insert("x-tenant-subscription", "raw-sub".parse().unwrap()); + + let mut ctx = make_filter_context(&req); + ctx.filter_metadata + .insert("identity.x-tenant-username".to_owned(), "bob".to_owned()); + + let state = capture_identity(&mut ctx, "x-tenant-", "identity"); + + assert_eq!(state.username, "bob"); + assert!( + state.subscription.is_empty(), + "raw header must be ignored when guard metadata identity exists: {}", + state.subscription + ); + let removed: Vec<&str> = ctx.request_headers_to_remove.iter().map(HeaderName::as_str).collect(); + assert!( + removed.contains(&"x-tenant-subscription"), + "unused identity header must still be stripped: {removed:?}" + ); +} + +#[test] +fn group_falls_back_to_subscription() { + let mut req = make_request(http::Method::POST, "/v1/chat/completions"); + req.headers.insert("x-tenant-username", "alice".parse().unwrap()); + req.headers.insert("x-tenant-subscription", "sub-99".parse().unwrap()); + + let mut ctx = make_filter_context(&req); + let state = capture_identity(&mut ctx, "x-tenant-", "identity"); + + assert_eq!(state.group, "sub-99"); +} + +#[test] +fn strips_accept_encoding_to_keep_response_readable() { + let mut req = make_request(http::Method::POST, "/v1/chat/completions"); + req.headers.insert("x-tenant-username", "alice".parse().unwrap()); + req.headers + .insert("accept-encoding", "gzip, deflate, br".parse().unwrap()); + + let mut ctx = make_filter_context(&req); + let _state = capture_identity(&mut ctx, "x-tenant-", "identity"); + + let removed: Vec<&str> = ctx.request_headers_to_remove.iter().map(HeaderName::as_str).collect(); + assert!(removed.contains(&"accept-encoding")); +} + +// ----------------------------------------------------------------------------- +// Model Extraction +// ----------------------------------------------------------------------------- + +#[test] +fn extracts_model_from_compact_json() { + let body = br#"{"model":"gpt-4","messages":[]}"#; + assert_eq!(extract_model_from_bytes(body).as_deref(), Some("gpt-4")); +} + +#[test] +fn extracts_model_with_whitespace_around_colon() { + let body = br#"{ "model" : "claude-sonnet-4" , "stream": true }"#; + assert_eq!(extract_model_from_bytes(body).as_deref(), Some("claude-sonnet-4")); +} + +#[test] +fn extracts_model_when_not_first_field() { + let body = br#"{"stream":true,"max_tokens":100,"model":"gpt-4o-mini"}"#; + assert_eq!(extract_model_from_bytes(body).as_deref(), Some("gpt-4o-mini")); +} + +#[test] +fn extract_model_ignores_nested_decoy() { + // A "model" key inside a message object must not win over the real + // top-level field, regardless of field order. + let body = br#"{"messages":[{"role":"user","content":"hi","model":"decoy"}],"model":"gpt-4"}"#; + assert_eq!( + extract_model_from_bytes(body).as_deref(), + Some("gpt-4"), + "nested decoy must not shadow the top-level model" + ); +} + +#[test] +fn extract_model_ignores_decoy_inside_string_value() { + let body = br#"{"prompt":"please say \"model\": \"decoy\" back","model":"gpt-4o"}"#; + assert_eq!( + extract_model_from_bytes(body).as_deref(), + Some("gpt-4o"), + "a quoted decoy inside a string value must be skipped" + ); +} + +#[test] +fn extract_model_returns_none_when_only_nested() { + let body = br#"{"messages":[{"model":"decoy","content":"hi"}]}"#; + assert!( + extract_model_from_bytes(body).is_none(), + "a nested-only model key must not be attributed" + ); +} + +#[test] +fn extract_model_returns_none_when_absent() { + let body = br#"{"messages":[{"role":"user","content":"hi"}]}"#; + assert!(extract_model_from_bytes(body).is_none()); +} + +#[test] +fn extract_model_returns_none_on_truncated_chunk() { + // A streamed first chunk may cut off mid-value. + let body = br#"{"model":"gpt-4"#; + assert!(extract_model_from_bytes(body).is_none()); +} + +#[test] +fn extract_model_returns_none_on_non_utf8() { + let body = &[0xFF_u8, 0xFE, 0x00, 0x01]; + assert!(extract_model_from_bytes(body).is_none()); +} + +#[test] +fn extract_model_returns_none_on_non_string_value() { + let body = br#"{"model":null}"#; + assert!(extract_model_from_bytes(body).is_none()); +} + +// ----------------------------------------------------------------------------- +// Token Metadata Reading +// ----------------------------------------------------------------------------- + +#[test] +fn reads_token_metadata() { + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = make_filter_context(&req); + ctx.filter_metadata.insert("token.input".to_owned(), "150".to_owned()); + ctx.filter_metadata.insert("token.output".to_owned(), "80".to_owned()); + ctx.filter_metadata.insert("token.total".to_owned(), "230".to_owned()); + + assert_eq!(read_token_meta(&ctx, META_TOKEN_INPUT), 150); + assert_eq!(read_token_meta(&ctx, META_TOKEN_OUTPUT), 80); + assert_eq!(read_token_meta(&ctx, META_TOKEN_TOTAL), 230); +} + +#[test] +fn missing_token_metadata_returns_zero() { + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let ctx = make_filter_context(&req); + + assert_eq!(read_token_meta(&ctx, META_TOKEN_INPUT), 0); +} + +// ----------------------------------------------------------------------------- +// Request Lifecycle +// ----------------------------------------------------------------------------- + +fn filter_from_yaml(yaml: &str) -> ExternalMeteringFilter { + let parsed: serde_yaml::Value = serde_yaml::from_str(yaml).unwrap(); + build_filter(&parsed).unwrap() +} + +#[tokio::test] +async fn skips_metering_when_no_identity_and_no_fallback() { + let filter = filter_from_yaml("metering_url: \"http://metering:8080\"\n"); + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = make_filter_context(&req); + + let action = filter.on_request(&mut ctx).await.unwrap(); + + assert!(matches!(action, FilterAction::Continue)); + assert!(ctx.filter_state.is_empty()); +} + +#[tokio::test] +async fn request_body_records_model_in_metadata() { + let filter = filter_from_yaml("metering_url: \"http://metering:8080\"\n"); + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = make_filter_context(&req); + let mut body = Some(Bytes::from_static(br#"{"model":"gpt-4","stream":true}"#)); + + let action = filter.on_request_body(&mut ctx, &mut body, false).await.unwrap(); + + assert!(matches!(action, FilterAction::Release)); + assert_eq!( + ctx.filter_metadata.get(META_METERING_MODEL).map(String::as_str), + Some("gpt-4") + ); +} + +#[tokio::test] +async fn request_body_does_not_override_identity_model() { + let filter = filter_from_yaml("metering_url: \"http://metering:8080\"\n"); + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = make_filter_context(&req); + ctx.filter_metadata + .insert(META_METERING_MODEL.to_owned(), "from-header".to_owned()); + let mut body = Some(Bytes::from_static(br#"{"model":"from-body"}"#)); + + let _action = filter.on_request_body(&mut ctx, &mut body, false).await.unwrap(); + + assert_eq!( + ctx.filter_metadata.get(META_METERING_MODEL).map(String::as_str), + Some("from-header") + ); +} + +// ----------------------------------------------------------------------------- +// Response Lifecycle +// ----------------------------------------------------------------------------- + +#[test] +fn response_body_is_noop_before_end_of_stream() { + let filter = filter_from_yaml("metering_url: \"http://metering:8080\"\n"); + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = make_filter_context(&req); + store_state(&mut ctx, state_for("alice", "gpt-4")); + let mut body = Some(Bytes::from_static(b"chunk")); + + let action = filter.on_response_body(&mut ctx, &mut body, false).unwrap(); + + assert!(matches!(action, FilterAction::Continue)); + // State survives so the terminal chunk can still report usage. + assert!(!ctx.filter_state.is_empty()); +} + +#[test] +fn response_body_without_state_is_noop() { + let filter = filter_from_yaml("metering_url: \"http://metering:8080\"\n"); + let req = make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = make_filter_context(&req); + let mut body = None; + + let action = filter.on_response_body(&mut ctx, &mut body, true).unwrap(); + + assert!(matches!(action, FilterAction::Continue)); +} + +fn event_ctx<'a>(event_id: &'a str, state: &'a MeteringState) -> EventContext<'a> { + EventContext { + duration_ms: 500, + event_id, + provider: "openai", + source: "gw", + state, + } +} + +fn state_for(username: &str, model: &str) -> MeteringState { + MeteringState { + username: username.into(), + group: "engineering".into(), + subscription: "sub-1".into(), + model: model.into(), + user_agent: "test/1.0".into(), + request_start: std::time::Instant::now(), + is_error: false, + response_status: 200, + } +} diff --git a/filters/src/register.rs b/filters/src/register.rs index ceb8f0bdf6..f5f575401b 100644 --- a/filters/src/register.rs +++ b/filters/src/register.rs @@ -15,16 +15,16 @@ use crate::HttpCalloutFilter; #[cfg(feature = "token-rate-limit-filter")] use crate::TokenRateLimitFilter; use crate::{ - A2aFilter, AiGuardrailsFilter, CredentialInjectFilter, IntelligentRouteFilter, McpFilter, ModelToHeaderFilter, - PromptEnrichFilter, ProviderRouteFilter, Sigv4SignFilter, TimeToFirstTokenFilter, TokenCountFilter, - TokenUsageHeadersFilter, + A2aFilter, AiGuardrailsFilter, CredentialInjectFilter, ExternalMeteringFilter, IntelligentRouteFilter, McpFilter, + ModelToHeaderFilter, PromptEnrichFilter, ProviderRouteFilter, Sigv4SignFilter, TimeToFirstTokenFilter, + TokenCountFilter, TokenUsageHeadersFilter, }; /// Register all in-tree AI HTTP filters into `registry`. /// /// When `subrequest_client` is provided, filters that make HTTP /// callouts (`openai_file_resolve`, `openai_web_search`, -/// `anthropic_web_search`) capture the +/// `anthropic_web_search`, `external_metering`) capture the /// shared client instead of creating isolated per-filter connectors. /// /// Does not call [`FilterRegistry::with_builtins`]. @@ -45,6 +45,7 @@ pub fn register_ai_filters(registry: &mut FilterRegistry, subrequest_client: Opt #[cfg(feature = "gcp-adc-filter")] register_gcp_filters(registry); register_general_ai_filters(registry); + register_external_metering(registry, subrequest_client); register_anthropic_filters(registry, subrequest_client); register_openai_filters(registry, subrequest_client); register_routing_filters(registry); @@ -141,6 +142,28 @@ fn register_token_filters(registry: &mut FilterRegistry) { ); } +/// Register the external metering filter, capturing the shared +/// sub-request client when one is available. +#[expect(clippy::panic, reason = "duplicate filter registration is a fatal configuration bug")] +fn register_external_metering(registry: &mut FilterRegistry, subrequest_client: Option<&SubRequestClient>) { + if let Some(client) = subrequest_client { + let client = client.clone(); + registry + .register( + "external_metering", + praxis_filter::FilterFactory::Http(std::sync::Arc::new(move |config| { + ExternalMeteringFilter::from_config_with_client(config, client.clone()) + })), + ) + .unwrap_or_else(|_| panic!("duplicate filter name: 'external_metering'")); + } else { + praxis_filter::register_filters!( + @register registry, + http "external_metering" => ExternalMeteringFilter::from_config + ); + } +} + /// Register intelligent routing filters. fn register_routing_filters(registry: &mut FilterRegistry) { praxis_filter::register_filters!( diff --git a/tests/integration/tests/suite/examples/external_metering.rs b/tests/integration/tests/suite/examples/external_metering.rs new file mode 100644 index 0000000000..933657c06d --- /dev/null +++ b/tests/integration/tests/suite/examples/external_metering.rs @@ -0,0 +1,187 @@ +// SPDX-License-Identifier: MIT +// Copyright (c) 2026 Praxis Contributors + +//! Tests for the external metering example configuration. + +use std::collections::HashMap; + +use praxis_test_utils::{RoutedBackend, free_port, http_send, parse_body, parse_status, start_header_echo_backend}; + +// ----------------------------------------------------------------------------- +// Config Parsing +// ----------------------------------------------------------------------------- + +#[test] +fn external_metering_config_parses() { + let metering_port = free_port(); + let config = super::load_example_config( + "external-metering.yaml", + 29800, + HashMap::from([("127.0.0.1:3000", 29801_u16), ("127.0.0.1:9090", metering_port)]), + ); + + assert_eq!(config.listeners.len(), 1, "should have 1 listener"); +} + +// ----------------------------------------------------------------------------- +// Balance Check — Access Granted +// ----------------------------------------------------------------------------- + +#[test] +fn external_metering_allows_request_when_balance_available() { + let backend_guard = start_header_echo_backend(); + let backend_port = backend_guard.port(); + let proxy_port = free_port(); + + let balance_body = r#"{"hasAccess": true, "balance": 9000.0}"#; + let metering_port = RoutedBackend::new() + .route("/api/v1/customers", 200, balance_body) + .route("/api/v1/events", 204, "") + .start(); + + let config = super::load_example_config( + "external-metering.yaml", + proxy_port, + HashMap::from([("127.0.0.1:3000", backend_port), ("127.0.0.1:9090", metering_port)]), + ); + + let proxy = praxis_test_utils::start_proxy(&config); + let raw = http_send( + proxy.addr(), + "POST /v1/chat/completions HTTP/1.1\r\n\ + Host: localhost\r\n\ + x-tenant-username: alice\r\n\ + x-tenant-group: engineering\r\n\ + Connection: close\r\n\r\n", + ); + + assert_eq!(parse_status(&raw), 200, "should proxy request when balance available"); +} + +// ----------------------------------------------------------------------------- +// Fail-Closed — Metering Unavailable +// ----------------------------------------------------------------------------- + +#[test] +fn external_metering_rejects_when_fail_closed_and_metering_down() { + let backend_guard = start_header_echo_backend(); + let backend_port = backend_guard.port(); + let proxy_port = free_port(); + + // Point metering_url at a port with nothing listening + let metering_port = free_port(); + + let path = praxis_test_utils::example_config_path("external-metering.yaml"); + let yaml = std::fs::read_to_string(&path).unwrap(); + let patched = praxis_test_utils::patch_yaml( + &yaml, + proxy_port, + &HashMap::from([("127.0.0.1:3000", backend_port), ("127.0.0.1:9090", metering_port)]), + ) + .replace("fail_open: true", "fail_open: false"); + let config = praxis_core::config::Config::from_yaml(&patched).unwrap(); + + let proxy = praxis_test_utils::start_proxy(&config); + let raw = http_send( + proxy.addr(), + "POST /v1/chat/completions HTTP/1.1\r\n\ + Host: localhost\r\n\ + x-tenant-username: alice\r\n\ + Connection: close\r\n\r\n", + ); + + assert_eq!( + parse_status(&raw), + 503, + "should reject with 503 when metering is unavailable and fail_open=false" + ); + let body = parse_body(&raw); + assert!( + body.contains("metering system unavailable"), + "rejection body should mention metering unavailable: {body}" + ); +} + +// ----------------------------------------------------------------------------- +// Identity Headers Stripped +// ----------------------------------------------------------------------------- + +#[test] +fn external_metering_strips_tenant_and_auth_headers() { + let backend_guard = start_header_echo_backend(); + let backend_port = backend_guard.port(); + let proxy_port = free_port(); + + let balance_body = r#"{"hasAccess": true, "balance": 9000.0}"#; + let metering_port = RoutedBackend::new() + .route("/api/v1/customers", 200, balance_body) + .route("/api/v1/events", 204, "") + .start(); + + let config = super::load_example_config( + "external-metering.yaml", + proxy_port, + HashMap::from([("127.0.0.1:3000", backend_port), ("127.0.0.1:9090", metering_port)]), + ); + + let proxy = praxis_test_utils::start_proxy(&config); + let raw = http_send( + proxy.addr(), + "POST /v1/chat/completions HTTP/1.1\r\n\ + Host: localhost\r\n\ + x-tenant-username: alice\r\n\ + x-tenant-group: engineering\r\n\ + Authorization: Bearer sk-client-secret\r\n\ + x-api-key: client-key\r\n\ + Connection: close\r\n\r\n", + ); + + assert_eq!(parse_status(&raw), 200, "should proxy successfully"); + let body = parse_body(&raw); + assert!( + !body.contains("x-tenant-username"), + "tenant header should be stripped from upstream: {body}" + ); + assert!( + !body.contains("sk-client-secret"), + "authorization should be stripped from upstream: {body}" + ); + assert!( + !body.contains("client-key"), + "x-api-key should be stripped from upstream: {body}" + ); +} + +// ----------------------------------------------------------------------------- +// No Identity — Metering Skipped +// ----------------------------------------------------------------------------- + +#[test] +fn external_metering_skips_when_no_identity() { + let backend_guard = start_header_echo_backend(); + let backend_port = backend_guard.port(); + let proxy_port = free_port(); + + // No metering mock needed — filter skips entirely without identity + let metering_port = free_port(); + + let config = super::load_example_config( + "external-metering.yaml", + proxy_port, + HashMap::from([("127.0.0.1:3000", backend_port), ("127.0.0.1:9090", metering_port)]), + ); + + let proxy = praxis_test_utils::start_proxy(&config); + let raw = http_send( + proxy.addr(), + "GET / HTTP/1.1\r\n\ + Host: localhost\r\n\ + Connection: close\r\n\r\n", + ); + + assert_eq!( + parse_status(&raw), + 200, + "should proxy without metering when no identity headers" + ); +} diff --git a/tests/integration/tests/suite/examples/mod.rs b/tests/integration/tests/suite/examples/mod.rs index 1c3427b716..0e9e6e50d3 100644 --- a/tests/integration/tests/suite/examples/mod.rs +++ b/tests/integration/tests/suite/examples/mod.rs @@ -15,6 +15,7 @@ mod aws_sigv4; mod azure_ad; mod compact; mod credential_injection; +mod external_metering; mod file_search_callout; mod full_flow; mod full_flow_agentic;