diff --git a/Cargo.lock b/Cargo.lock index 996124e574..e155723951 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -8,17 +8,6 @@ version = "2.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa" -[[package]] -name = "ahash" -version = "0.7.8" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "891477e0c6a8957309ee5c45a6368af3ae14bb510732d2684ffa19af310920f9" -dependencies = [ - "getrandom 0.2.17", - "once_cell", - "version_check", -] - [[package]] name = "ahash" version = "0.8.12" @@ -121,7 +110,7 @@ version = "1.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" dependencies = [ - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -132,7 +121,7 @@ checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" dependencies = [ "anstyle", "once_cell_polyfill", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -1184,7 +1173,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -1307,7 +1296,7 @@ version = "0.5.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cf5efcf77a4da27927d3ab0509dec5b0954bb3bc59da5a1de9e52642ebd4cdf9" dependencies = [ - "ahash 0.8.12", + "ahash", "num_cpus", "parking_lot", "seize", @@ -1560,9 +1549,6 @@ name = "hashbrown" version = "0.12.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8a9ee70c43aaf417c914396645a0fa852624801b24ebb7ae78fe8272889ac888" -dependencies = [ - "ahash 0.7.8", -] [[package]] name = "hashbrown" @@ -2089,7 +2075,7 @@ version = "0.54.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7070bf0681439ff992bb94947cb3a361f57c880561c1e9e90ddb5941c7ec2d04" dependencies = [ - "ahash 0.8.12", + "ahash", "bytecount", "data-encoding", "email_address", @@ -2126,7 +2112,7 @@ version = "0.54.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4309c6b52390b1ebdd2053e513e9f23d4bb18d0c3c2df8142eded3eb188a85bf" dependencies = [ - "ahash 0.8.12", + "ahash", "bytecount", "fraction", "getrandom 0.3.4", @@ -2541,7 +2527,7 @@ version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -3466,7 +3452,7 @@ dependencies = [ "once_cell", "socket2", "tracing", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -3475,7 +3461,7 @@ version = "0.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2a820a1a3e59644b38bc03606b638b41f1c5a474215efcec37f94b15871a7e1c" dependencies = [ - "ahash 0.8.12", + "ahash", "async-trait", "blake2", "bstr", @@ -3510,7 +3496,7 @@ version = "0.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d7d6d9e04414ba6819c51ec210eddf17865c505431a492c86c6356338c2b4f60" dependencies = [ - "ahash 0.8.12", + "ahash", "async-trait", "brotli", "bstr", @@ -3598,7 +3584,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ec55d3bcbd56e52b3424bc62739a237d43c14df8ed04aaf294ae187b60da348a" dependencies = [ "arrayvec", - "hashbrown 0.12.3", + "hashbrown 0.17.1", "parking_lot", "rand 0.8.7", ] @@ -3891,7 +3877,7 @@ version = "0.54.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0304d4734d208eaf24528093715e2cb5260c977b1be1eb81cd487f601e3ff35d" dependencies = [ - "ahash 0.8.12", + "ahash", "fluent-uri", "getrandom 0.3.4", "hashbrown 0.17.1", @@ -4095,14 +4081,14 @@ dependencies = [ "errno", "libc", "linux-raw-sys", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] name = "rustls" -version = "0.23.43" +version = "0.23.45" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0283386ce02abc0151e1761d08802dfe86c173b0b494af5cbc086574e453da06" +checksum = "0d41d731c7d2f962d1ccc364cec258de3c0e93b38c2fb3ba97ac74513048d634" dependencies = [ "aws-lc-rs", "log", @@ -4176,7 +4162,7 @@ dependencies = [ "security-framework 3.7.0", "security-framework-sys", "webpki-root-certs", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -4581,7 +4567,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c3d1e2c7f27f8d4cb10542a02c49005dbd6e93095799d6f3be745fae9f8fedd4" dependencies = [ "libc", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -4894,10 +4880,10 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd" dependencies = [ "fastrand", - "getrandom 0.3.4", + "getrandom 0.4.3", "once_cell", "rustix", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -5713,7 +5699,7 @@ version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" dependencies = [ - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 126cdbaea3..0cddbe0800 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -48,7 +48,7 @@ redis = { version = "1.5.0", default-features = false, features = ["tokio-comp"] regex = "1.13.1" reqwest = { version = "0.13.4", default-features = false, features = ["rustls", "json", "stream"] } rmcp = { version = "3.1.4", default-features = false, features = ["client", "transport-streamable-http-client-reqwest", "reqwest"] } -rustls = "0.23.43" +rustls = "0.23.45" schemars = "1.2.2" rustls-pemfile = "2.2.0" secrecy = { version = "0.10.3", features = ["serde"] } diff --git a/docs/filters/llmisvc_model_provider_resolver.md b/docs/filters/llmisvc_model_provider_resolver.md new file mode 100644 index 0000000000..a825942035 --- /dev/null +++ b/docs/filters/llmisvc_model_provider_resolver.md @@ -0,0 +1,28 @@ + + + +# `llmisvc_model_provider_resolver` + +Rewrites publisher-ID body `model` values to the short model name for `LLMISvc` / `KServe` routing; the routing header is left unchanged. + +## Configuration Notes + +Reads the model name from the configured request header (default `X-Model`, typically set by an earlier `model_to_header`). When that value is a publisher ID (`publishers/.../models/`), rewrite the body `"model"` field to `` only. The routing header is never modified -- `KServe` routes on the publisher ID. + +If the header is absent/empty, or the body has no `"model"` field, this filter is a no-op (it does not invent a body `"model"`). + +Does **not** resolve `ExternalModel` / `ExternalProvider` CRDs, perform weighted provider selection, rewrite `Host`, or inject credentials. + +## Configuration + +| Field | Type | Required | Description | +|-------|------|---------|-------------| +| `header` | string | no | Request header that carries the publisher ID for `KServe` routing. Defaults to `X-Model` (same as `model_to_header`). | +| `max_body_bytes` | integer | no | Maximum request body size to buffer before parsing. | + +## Example + +```yaml +filter: llmisvc_model_provider_resolver +header: X-Model # optional, defaults to X-Model +``` diff --git a/docs/filters/model_to_header.md b/docs/filters/model_to_header.md index f3b70f0e58..d772c2f017 100644 --- a/docs/filters/model_to_header.md +++ b/docs/filters/model_to_header.md @@ -5,6 +5,10 @@ Promotes the JSON `"model"` field from the request body to a request header. +## Configuration Notes + +Promotion is deferred until end-of-stream so a later body-writing filter (for example `llmisvc_model_provider_resolver`) can observe the pending header in the same `StreamBuffer` pre-read pass. + ## Configuration | Field | Type | Required | Description | diff --git a/docs/filters/reference.md b/docs/filters/reference.md index d6557f4f71..ab1d228d6a 100644 --- a/docs/filters/reference.md +++ b/docs/filters/reference.md @@ -100,6 +100,7 @@ see the [Praxis core filter reference][core-ref]. | Filter | Description | |--------|-------------| +| [`llmisvc_model_provider_resolver`](llmisvc_model_provider_resolver.md) | Rewrites publisher-ID body `model` values to the short model name for `LLMISvc` / `KServe` routing; the routing header is left unchanged. | | [`model_to_header`](model_to_header.md) | Promotes the JSON `"model"` field from the request body to a request header. | ### Metering diff --git a/examples/README.md b/examples/README.md index d4c59fe6f8..f364114689 100644 --- a/examples/README.md +++ b/examples/README.md @@ -36,6 +36,7 @@ before sending requests. | [json-rpc-routing.yaml](configs/json-rpc-routing.yaml) | Routes JSON-RPC 2.0 requests to different backends based on the "method" field in the JSON request body | | [lakera-guard.yaml](configs/lakera-guard.yaml) | Screens every request body through Lakera Guard for content moderation before forwarding to the upstream | | [llmd-ext-proc-routing.yaml](configs/llmd-ext-proc-routing.yaml) | A real llm-d EPP or test processor returns the trusted x-gateway-destination-endpoint header | +| [llmisvc-model-provider-resolver.yaml](configs/llmisvc-model-provider-resolver.yaml) | Rewrites publisher-ID body `model` values to the short model name for LLMISvc / KServe routing; the routing header (default `X-Model`) is left unchanged so routing can still use the publisher ID | | [mcp-classifier-routing.yaml](configs/mcp-classifier-routing.yaml) | Routes MCP requests by body-derived method and tool name | | [mcp-stateless-broker.yaml](configs/mcp-stateless-broker.yaml) | Configurable stateless MCP broker using the final MCP 2026-07-28 stateless profile | | [model-to-header-routing.yaml](configs/model-to-header-routing.yaml) | Routes LLM API requests to different backends based on the "model" field in the JSON request body | diff --git a/examples/configs/llmisvc-model-provider-resolver.yaml b/examples/configs/llmisvc-model-provider-resolver.yaml new file mode 100644 index 0000000000..70386a4209 --- /dev/null +++ b/examples/configs/llmisvc-model-provider-resolver.yaml @@ -0,0 +1,47 @@ +# LLMISvc Model Provider Resolver +# +# Build: +# cargo build -p praxis-ai-proxy +# +# Rewrites publisher-ID body `model` values to the short model name for +# LLMISvc / KServe routing; the routing header (default `X-Model`) is +# left unchanged so routing can still use the publisher ID. +# +# Typical chain: `model_to_header` promotes the body model to +# `X-Model`, then `llmisvc_model_provider_resolver` strips the +# publisher prefix from the body only. +# +# Example request: +# +# curl -X POST http://localhost:8080/v1/chat/completions \ +# -H "Content-Type: application/json" \ +# -d '{"model":"publishers/rhoai/models/granite-3.1-8b","messages":[{"role":"user","content":"hi"}]}' +# +# Upstream body model becomes `granite-3.1-8b` while `X-Model` +# remains `publishers/rhoai/models/granite-3.1-8b`. + +listeners: + - name: llmisvc-gateway + address: "0.0.0.0:8080" # dev value; binds all interfaces + filter_chains: + - rewrite-and-route + +filter_chains: + - name: rewrite-and-route + filters: + - filter: model_to_header + header: X-Model + - filter: llmisvc_model_provider_resolver + header: X-Model + - filter: router + routes: + - path_prefix: "/" + cluster: provider + - filter: load_balancer + clusters: + - name: provider + endpoints: + - "127.0.0.1:3000" + +insecure_options: + allow_private_endpoints: true # example proxies to local backends diff --git a/filters/src/inference/llmisvc_model_provider_resolver.rs b/filters/src/inference/llmisvc_model_provider_resolver.rs new file mode 100644 index 0000000000..c37ff93f83 --- /dev/null +++ b/filters/src/inference/llmisvc_model_provider_resolver.rs @@ -0,0 +1,673 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 Praxis Contributors + +//! `LLMISvc` model-provider resolver: rewrites publisher-ID body `model` +//! values to the short model name for `LLMISvc` / `KServe` routing while +//! leaving the routing header unchanged. + +use async_trait::async_trait; +use bytes::Bytes; +use http::HeaderName; +use praxis_ai_apis::json_body::replace_json_body; +use praxis_filter::{ + BodyAccess, BodyMode, FilterAction, FilterError, HttpFilter, HttpFilterContext, PendingHeaderResult, + body::DEFAULT_JSON_BODY_MAX_BYTES, builtins::http::payload_processing::config_validation::validate_max_body_bytes, + parse_filter_config, +}; +use serde::Deserialize; +use tracing::debug; + +// ----------------------------------------------------------------------------- +// Constants +// ----------------------------------------------------------------------------- + +/// Default header name for the routing model value (aligned with +/// [`super::ModelToHeaderFilter`]). +const DEFAULT_HEADER: &str = "X-Model"; + +/// Filter metadata key for the original publisher ID (for metering). +const META_PUBLISHER_ID: &str = "llmisvc_model_provider_resolver.publisher_id"; + +/// Prefix that identifies a `KServe` / `LLMISvc` publisher model ID. +const PUBLISHERS_PREFIX: &str = "publishers/"; + +/// Separator between the publisher path and the short model name. +const MODELS_SEPARATOR: &str = "/models/"; + +// ----------------------------------------------------------------------------- +// Config +// ----------------------------------------------------------------------------- + +/// Deserialized YAML config for the `LLMISvc` model-provider resolver. +#[derive(Debug, Deserialize)] +#[serde(deny_unknown_fields)] +struct LlmisvcModelProviderResolverConfig { + /// Request header that carries the publisher ID for `KServe` routing. + /// + /// Defaults to `X-Model` (same as `model_to_header`). + #[serde(default = "default_header")] + header: String, + + /// Maximum request body size to buffer before parsing. + #[serde(default = "default_max_body_bytes")] + max_body_bytes: usize, +} + +/// Default header name. +fn default_header() -> String { + DEFAULT_HEADER.to_owned() +} + +/// Default for `max_body_bytes`. +fn default_max_body_bytes() -> usize { + DEFAULT_JSON_BODY_MAX_BYTES +} + +// ----------------------------------------------------------------------------- +// LlmisvcModelProviderResolverFilter +// ----------------------------------------------------------------------------- + +/// Rewrites publisher-ID body `model` values to the short model name for +/// `LLMISvc` / `KServe` routing; the routing header is left unchanged. +/// +/// Reads the model name from the configured request header (default +/// `X-Model`, typically set by an earlier `model_to_header`). When that +/// value is a publisher ID (`publishers/.../models/`), rewrite the +/// body `"model"` field to `` only. The routing header is never +/// modified -- `KServe` routes on the publisher ID. +/// +/// If the header is absent/empty, or the body has no `"model"` field, +/// this filter is a no-op (it does not invent a body `"model"`). +/// +/// Does **not** resolve `ExternalModel` / `ExternalProvider` CRDs, perform +/// weighted provider selection, rewrite `Host`, or inject credentials. +/// +/// # YAML configuration +/// +/// ```yaml +/// filter: llmisvc_model_provider_resolver +/// header: X-Model # optional, defaults to X-Model +/// ``` +/// +/// # Example +/// +/// ```ignore +/// use praxis_ai_filters::LlmisvcModelProviderResolverFilter; +/// +/// let yaml = serde_yaml::Value::Null; +/// let filter = LlmisvcModelProviderResolverFilter::from_config(&yaml).unwrap(); +/// assert_eq!(filter.name(), "llmisvc_model_provider_resolver"); +/// ``` +pub struct LlmisvcModelProviderResolverFilter { + /// Header that carries the publisher ID used for `KServe` routing. + header: HeaderName, + + /// Maximum request body size to buffer. + max_body_bytes: usize, +} + +impl LlmisvcModelProviderResolverFilter { + /// Create from parsed YAML config. + /// + /// Accepts an optional `header` field (defaults to `X-Model`) and + /// optional `max_body_bytes`. + /// + /// # Errors + /// + /// Returns [`FilterError`] if config parsing fails, `header` is + /// empty/invalid, or `max_body_bytes` is invalid. + /// + /// [`FilterError`]: praxis_filter::FilterError + /// + /// ```ignore + /// use praxis_ai_filters::LlmisvcModelProviderResolverFilter; + /// + /// let yaml: serde_yaml::Value = serde_yaml::from_str("header: X-AI-Model").unwrap(); + /// let filter = LlmisvcModelProviderResolverFilter::from_config(&yaml).unwrap(); + /// assert_eq!(filter.name(), "llmisvc_model_provider_resolver"); + /// ``` + pub fn from_config(config: &serde_yaml::Value) -> Result, FilterError> { + let cfg: LlmisvcModelProviderResolverConfig = parse_filter_config("llmisvc_model_provider_resolver", config)?; + + let header = cfg.header.trim(); + if header.is_empty() { + return Err("llmisvc_model_provider_resolver: 'header' must not be empty".into()); + } + let header: HeaderName = header + .parse() + .map_err(|e| format!("llmisvc_model_provider_resolver: invalid 'header' name: {e}"))?; + validate_max_body_bytes("llmisvc_model_provider_resolver", cfg.max_body_bytes)?; + + Ok(Box::new(Self { + header, + max_body_bytes: cfg.max_body_bytes, + })) + } + + /// Resolve model name from the routing header, rewrite publisher-ID + /// body field when needed. + fn rewrite_body( + &self, + ctx: &mut HttpFilterContext<'_>, + body: &mut Option, + ) -> Result { + let Some(model_name) = header_model_name(ctx, &self.header) else { + return Ok(FilterAction::Continue); + }; + let Some(short_name) = llmisvc_short_model_name(&model_name) else { + return Ok(FilterAction::Continue); + }; + + match rewrite_publisher_body_model(body, short_name, self.name())? { + PublisherBodyRewrite::Noop => {}, + PublisherBodyRewrite::Matched { rewritten } => { + ctx.set_metadata(META_PUBLISHER_ID, model_name.as_str()); + if rewritten { + debug!( + original = %model_name, + rewritten = %short_name, + "LLMISvc: rewrote body model field for publisher ID" + ); + } + }, + } + + Ok(FilterAction::Continue) + } +} + +#[async_trait] +impl HttpFilter for LlmisvcModelProviderResolverFilter { + fn name(&self) -> &'static str { + "llmisvc_model_provider_resolver" + } + + async fn on_request(&self, _ctx: &mut HttpFilterContext<'_>) -> Result { + Ok(FilterAction::Continue) + } + + fn request_body_access(&self) -> BodyAccess { + BodyAccess::ReadWrite + } + + fn request_body_mode(&self) -> BodyMode { + BodyMode::StreamBuffer { + max_bytes: Some(self.max_body_bytes), + } + } + + async fn on_request_body( + &self, + ctx: &mut HttpFilterContext<'_>, + body: &mut Option, + end_of_stream: bool, + ) -> Result { + if !end_of_stream { + return Ok(FilterAction::Continue); + } + + self.rewrite_body(ctx, body) + } +} + +// ----------------------------------------------------------------------------- +// Private Utilities +// ----------------------------------------------------------------------------- + +/// Result of attempting to rewrite a publisher-ID body `"model"` field. +enum PublisherBodyRewrite { + /// Body was not a JSON object with a `"model"` field. + Noop, + /// Body `"model"` matched the publisher ID; `rewritten` is `false` + /// when the short name was already present. + Matched { + /// Whether the serialized body was updated. + rewritten: bool, + }, +} + +/// Rewrite an existing body `"model"` field to `short_name` when present. +fn rewrite_publisher_body_model( + body: &mut Option, + short_name: &str, + filter_name: &'static str, +) -> Result { + let Some(raw) = body.as_ref() else { + return Ok(PublisherBodyRewrite::Noop); + }; + + let mut value: serde_json::Value = match serde_json::from_slice(raw) { + Ok(v) => v, + Err(_) => return Ok(PublisherBodyRewrite::Noop), + }; + + let Some(obj) = value.as_object_mut() else { + return Ok(PublisherBodyRewrite::Noop); + }; + + if !obj.contains_key("model") { + return Ok(PublisherBodyRewrite::Noop); + } + + if obj.get("model").and_then(serde_json::Value::as_str) == Some(short_name) { + return Ok(PublisherBodyRewrite::Matched { rewritten: false }); + } + + obj.insert("model".to_owned(), serde_json::Value::String(short_name.to_owned())); + + replace_json_body(body, &value, filter_name, "model").map_err(|e| -> FilterError { + format!("{filter_name}: failed to re-serialize rewritten request body: {e}").into() + })?; + + Ok(PublisherBodyRewrite::Matched { rewritten: true }) +} + +/// Read a non-empty model name from the request headers or pending +/// mutations (e.g. `extra_request_headers` from an earlier +/// `model_to_header`). +fn header_model_name(ctx: &HttpFilterContext<'_>, header: &HeaderName) -> Option { + if let Some(value) = ctx.request.headers.get(header) + && let Ok(s) = value.to_str() + { + let s = s.trim(); + if !s.is_empty() { + return Some(s.to_owned()); + } + } + + if let Ok(Some(value)) = ctx.resolve_trusted_header(header) { + let trimmed = value.trim(); + if !trimmed.is_empty() { + return Some(trimmed.to_owned()); + } + } + + match ctx.pending_header_value(header) { + Ok(PendingHeaderResult::Value(value)) => { + let trimmed = value.trim(); + (!trimmed.is_empty()).then(|| trimmed.to_owned()) + }, + Ok(PendingHeaderResult::Absent | PendingHeaderResult::Removed) | Err(_) => None, + } +} + +/// Extract the short model name from a `KServe` publisher ID. +/// +/// Require `publishers/` prefix, then take the segment after the first +/// `/models/` when non-empty. +fn llmisvc_short_model_name(model_name: &str) -> Option<&str> { + if !model_name.starts_with(PUBLISHERS_PREFIX) { + return None; + } + let (_, short_name) = model_name.split_once(MODELS_SEPARATOR)?; + if short_name.is_empty() { None } else { Some(short_name) } +} + +// ----------------------------------------------------------------------------- +// Tests +// ----------------------------------------------------------------------------- + +#[cfg(test)] +#[expect(clippy::allow_attributes, reason = "blanket test suppressions")] +#[allow( + clippy::unwrap_used, + clippy::expect_used, + clippy::indexing_slicing, + clippy::panic, + reason = "tests" +)] +mod tests { + use std::borrow::Cow; + + use http::HeaderValue; + + use super::*; + + fn filter_default() -> Box { + LlmisvcModelProviderResolverFilter::from_config(&serde_yaml::Value::Null).unwrap() + } + + #[test] + fn from_config_default_header() { + let filter = filter_default(); + assert_eq!( + filter.name(), + "llmisvc_model_provider_resolver", + "default config should produce llmisvc_model_provider_resolver" + ); + } + + #[test] + fn from_config_custom_header() { + let yaml: serde_yaml::Value = serde_yaml::from_str("header: X-AI-Model").unwrap(); + let filter = LlmisvcModelProviderResolverFilter::from_config(&yaml).unwrap(); + assert_eq!( + filter.name(), + "llmisvc_model_provider_resolver", + "custom header config should parse" + ); + } + + #[test] + fn from_config_rejects_empty_header() { + let yaml: serde_yaml::Value = serde_yaml::from_str("header: \"\"").unwrap(); + match LlmisvcModelProviderResolverFilter::from_config(&yaml) { + Err(err) => assert!( + err.to_string().contains("header"), + "empty header should be rejected: {err}" + ), + Ok(_) => panic!("empty header should be rejected"), + } + } + + #[test] + fn from_config_rejects_zero_max_body_bytes() { + let yaml: serde_yaml::Value = serde_yaml::from_str("max_body_bytes: 0").unwrap(); + match LlmisvcModelProviderResolverFilter::from_config(&yaml) { + Err(err) => assert!( + err.to_string().contains("max_body_bytes"), + "zero max_body_bytes should be rejected: {err}" + ), + Ok(_) => panic!("zero max_body_bytes should be rejected"), + } + } + + #[test] + fn from_config_rejects_unknown_fields() { + let yaml: serde_yaml::Value = serde_yaml::from_str( + r#" +header: "X-Model" +bogus_field: true +"#, + ) + .unwrap(); + assert!( + LlmisvcModelProviderResolverFilter::from_config(&yaml).is_err(), + "unknown fields should be rejected" + ); + } + + #[test] + fn from_config_rejects_invalid_header_name() { + let yaml: serde_yaml::Value = serde_yaml::from_str("header: \"@\"").unwrap(); + match LlmisvcModelProviderResolverFilter::from_config(&yaml) { + Err(err) => assert!( + err.to_string().contains("header"), + "invalid header name should be rejected: {err}" + ), + Ok(_) => panic!("invalid header name should be rejected"), + } + } + + #[test] + fn body_access_is_read_write_stream_buffer() { + let filter = filter_default(); + assert_eq!( + filter.request_body_access(), + BodyAccess::ReadWrite, + "must mutate the request body" + ); + assert!( + matches!( + filter.request_body_mode(), + BodyMode::StreamBuffer { + max_bytes: Some(limit) + } if limit > 0 + ), + "body mode should be StreamBuffer with a default size limit" + ); + } + + #[test] + fn short_model_name_extracts_after_models() { + assert_eq!( + llmisvc_short_model_name("publishers/ns/models/granite-3.1-8b"), + Some("granite-3.1-8b") + ); + assert_eq!( + llmisvc_short_model_name("publishers/ns/models/a/b"), + Some("a/b"), + "split_once keeps remainder after first /models/" + ); + assert_eq!(llmisvc_short_model_name("publishers/ns/models/"), None); + assert_eq!(llmisvc_short_model_name("publishers/ns/foo"), None); + assert_eq!(llmisvc_short_model_name("granite-3.1-8b"), None); + assert_eq!(llmisvc_short_model_name("other/models/foo"), None); + } + + #[tokio::test] + async fn rewrites_body_model_from_header_publisher_id() { + let filter = filter_default(); + let mut req = crate::test_utils::make_request(http::Method::POST, "/v1/chat/completions"); + req.headers.insert( + "X-Model", + HeaderValue::from_static("publishers/rhoai/models/granite-3.1-8b"), + ); + let mut ctx = crate::test_utils::make_filter_context(&req); + let mut body = Some(Bytes::from_static( + br#"{"model":"publishers/rhoai/models/granite-3.1-8b","messages":[]}"#, + )); + + let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + assert!(matches!(action, FilterAction::Continue), "rewrite should continue"); + + let parsed: serde_json::Value = serde_json::from_slice(body.as_ref().unwrap()).unwrap(); + assert_eq!(parsed["model"].as_str(), Some("granite-3.1-8b")); + assert_eq!( + ctx.filter_metadata.get(META_PUBLISHER_ID).map(String::as_str), + Some("publishers/rhoai/models/granite-3.1-8b"), + ); + assert!(ctx.extra_request_headers.is_empty()); + assert!(ctx.request_headers_to_set.is_empty()); + assert!(ctx.request_headers_to_remove.is_empty()); + } + + #[tokio::test] + async fn noops_when_header_absent() { + let filter = filter_default(); + let req = crate::test_utils::make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = crate::test_utils::make_filter_context(&req); + + let json = br#"{"model":"publishers/ns/models/mistral","prompt":"hi"}"#; + let mut body = Some(Bytes::from_static(json)); + let original = body.clone(); + + let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + assert!(matches!(action, FilterAction::Continue)); + assert_eq!(body, original, "missing routing header must not rewrite body"); + assert!( + !ctx.filter_metadata.contains_key(META_PUBLISHER_ID), + "no rewrite means no publisher metadata" + ); + } + + #[tokio::test] + async fn leaves_body_unchanged_when_header_publisher_id_but_no_body_model() { + let filter = filter_default(); + let mut req = crate::test_utils::make_request(http::Method::POST, "/v1/chat/completions"); + req.headers.insert( + "X-Model", + HeaderValue::from_static("publishers/rhoai/models/granite-3.1-8b"), + ); + let mut ctx = crate::test_utils::make_filter_context(&req); + + let json = br#"{"messages":[]}"#; + let mut body = Some(Bytes::from_static(json)); + let original = body.clone(); + + let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + assert!(matches!(action, FilterAction::Continue)); + assert_eq!(body, original, "must not invent a body model field"); + assert!( + !ctx.filter_metadata.contains_key(META_PUBLISHER_ID), + "no rewrite means no publisher metadata" + ); + } + + #[tokio::test] + async fn prefers_header_over_body_model() { + let filter = filter_default(); + let mut req = crate::test_utils::make_request(http::Method::POST, "/v1/chat/completions"); + req.headers + .insert("X-Model", HeaderValue::from_static("publishers/ns/models/from-header")); + let mut ctx = crate::test_utils::make_filter_context(&req); + + let json = br#"{"model":"publishers/ns/models/from-body"}"#; + let mut body = Some(Bytes::from_static(json)); + + let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + assert!(matches!(action, FilterAction::Continue)); + + let parsed: serde_json::Value = serde_json::from_slice(body.as_ref().unwrap()).unwrap(); + assert_eq!(parsed["model"].as_str(), Some("from-header")); + assert_eq!( + ctx.filter_metadata.get(META_PUBLISHER_ID).map(String::as_str), + Some("publishers/ns/models/from-header"), + ); + } + + #[tokio::test] + async fn composes_with_model_to_header_at_end_of_stream() { + use crate::inference::ModelToHeaderFilter; + + let model_to_header = ModelToHeaderFilter::from_config(&serde_yaml::Value::Null).unwrap(); + let llmisvc = filter_default(); + let req = crate::test_utils::make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = crate::test_utils::make_filter_context(&req); + + let json = br#"{"model":"publishers/rhoai/models/granite-3.1-8b","messages":[]}"#; + let mut body = Some(Bytes::from_static(json)); + + let action = model_to_header + .on_request_body(&mut ctx, &mut body, true) + .await + .unwrap(); + assert!(matches!(action, FilterAction::BodyDone)); + + let action = llmisvc.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + assert!(matches!(action, FilterAction::Continue)); + + let parsed: serde_json::Value = serde_json::from_slice(body.as_ref().unwrap()).unwrap(); + assert_eq!(parsed["model"].as_str(), Some("granite-3.1-8b")); + } + + #[tokio::test] + async fn reads_pending_extra_request_headers() { + let filter = filter_default(); + let req = crate::test_utils::make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = crate::test_utils::make_filter_context(&req); + ctx.extra_request_headers + .push((Cow::Borrowed("X-Model"), "publishers/ns/models/via-extra".to_owned())); + + let json = br#"{"model":"publishers/ns/models/via-extra"}"#; + let mut body = Some(Bytes::from_static(json)); + + let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + assert!(matches!(action, FilterAction::Continue)); + + let parsed: serde_json::Value = serde_json::from_slice(body.as_ref().unwrap()).unwrap(); + assert_eq!(parsed["model"].as_str(), Some("via-extra")); + assert_eq!( + ctx.extra_request_headers.len(), + 1, + "extra header from model_to_header must remain" + ); + assert_eq!(ctx.extra_request_headers[0].1, "publishers/ns/models/via-extra"); + } + + #[tokio::test] + async fn leaves_non_publisher_model_unchanged() { + let filter = filter_default(); + let mut req = crate::test_utils::make_request(http::Method::POST, "/v1/chat/completions"); + req.headers + .insert("X-Model", HeaderValue::from_static("mistral-large-latest")); + let mut ctx = crate::test_utils::make_filter_context(&req); + + let json = br#"{"model":"mistral-large-latest"}"#; + let mut body = Some(Bytes::from_static(json)); + let original = body.clone(); + + let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + assert!(matches!(action, FilterAction::Continue)); + assert_eq!(body, original, "non-publisher body must not be rewritten"); + assert!( + !ctx.filter_metadata.contains_key(META_PUBLISHER_ID), + "no publisher metadata for non-publisher models" + ); + } + + #[tokio::test] + async fn continues_when_model_absent() { + let filter = filter_default(); + let req = crate::test_utils::make_request(http::Method::POST, "/v1/chat/completions"); + let mut ctx = crate::test_utils::make_filter_context(&req); + + let json = br#"{"messages":[]}"#; + let mut body = Some(Bytes::from_static(json)); + let original = body.clone(); + + let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + assert!(matches!(action, FilterAction::Continue)); + assert_eq!(body, original); + } + + #[tokio::test] + async fn continues_on_invalid_json() { + let filter = filter_default(); + let mut req = crate::test_utils::make_request(http::Method::POST, "/v1/chat/completions"); + req.headers + .insert("X-Model", HeaderValue::from_static("publishers/ns/models/granite")); + let mut ctx = crate::test_utils::make_filter_context(&req); + + let mut body = Some(Bytes::from_static(b"not-json")); + let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + assert!(matches!(action, FilterAction::Continue)); + assert_eq!(body.as_deref(), Some(b"not-json".as_slice())); + } + + #[tokio::test] + async fn custom_header_name_used() { + let yaml: serde_yaml::Value = serde_yaml::from_str("header: X-AI-Model").unwrap(); + let filter = LlmisvcModelProviderResolverFilter::from_config(&yaml).unwrap(); + let mut req = crate::test_utils::make_request(http::Method::POST, "/v1/chat"); + req.headers + .insert("X-AI-Model", HeaderValue::from_static("publishers/ns/models/custom")); + let mut ctx = crate::test_utils::make_filter_context(&req); + + let json = br#"{"model":"publishers/ns/models/custom"}"#; + let mut body = Some(Bytes::from_static(json)); + + let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + assert!(matches!(action, FilterAction::Continue)); + + let parsed: serde_json::Value = serde_json::from_slice(body.as_ref().unwrap()).unwrap(); + assert_eq!(parsed["model"].as_str(), Some("custom")); + } + + #[tokio::test] + async fn on_request_is_noop() { + let filter = filter_default(); + let req = crate::test_utils::make_request(http::Method::POST, "/v1/chat"); + let mut ctx = crate::test_utils::make_filter_context(&req); + + let action = filter.on_request(&mut ctx).await.unwrap(); + assert!(matches!(action, FilterAction::Continue)); + } + + #[tokio::test] + async fn waits_for_end_of_stream() { + let filter = filter_default(); + let mut req = crate::test_utils::make_request(http::Method::POST, "/v1/chat"); + req.headers + .insert("X-Model", HeaderValue::from_static("publishers/ns/models/granite")); + let mut ctx = crate::test_utils::make_filter_context(&req); + + let json = br#"{"model":"publishers/ns/models/granite"}"#; + let mut body = Some(Bytes::from_static(json)); + let original = body.clone(); + + let action = filter.on_request_body(&mut ctx, &mut body, false).await.unwrap(); + assert!(matches!(action, FilterAction::Continue)); + assert_eq!(body, original, "must not rewrite before end_of_stream"); + } +} diff --git a/filters/src/inference/mod.rs b/filters/src/inference/mod.rs index 9965387141..c6bf58fc8c 100644 --- a/filters/src/inference/mod.rs +++ b/filters/src/inference/mod.rs @@ -3,6 +3,8 @@ //! AI inference proxy filters. +mod llmisvc_model_provider_resolver; mod model_to_header; +pub use llmisvc_model_provider_resolver::LlmisvcModelProviderResolverFilter; pub use model_to_header::ModelToHeaderFilter; diff --git a/filters/src/inference/model_to_header.rs b/filters/src/inference/model_to_header.rs index d1e4ad76aa..28c01d1ef1 100644 --- a/filters/src/inference/model_to_header.rs +++ b/filters/src/inference/model_to_header.rs @@ -45,6 +45,10 @@ fn default_header() -> String { /// Promotes the JSON `"model"` field from the request body to a request header. /// +/// Promotion is deferred until end-of-stream so a later body-writing filter +/// (for example `llmisvc_model_provider_resolver`) can observe the pending +/// header in the same `StreamBuffer` pre-read pass. +/// /// # YAML configuration /// /// ```yaml @@ -152,6 +156,10 @@ impl HttpFilter for ModelToHeaderFilter { body: &mut Option, end_of_stream: bool, ) -> Result { + if !end_of_stream { + return Ok(FilterAction::Continue); + } + self.inner.on_request_body(ctx, body, end_of_stream).await } @@ -241,6 +249,22 @@ mod tests { ); } + #[tokio::test] + async fn waits_for_end_of_stream() { + let filter = ModelToHeaderFilter::from_config(&serde_yaml::Value::Null).unwrap(); + let req = crate::test_utils::make_request(http::Method::POST, "/v1/chat"); + let mut ctx = crate::test_utils::make_filter_context(&req); + + let json = br#"{"model":"mistral-large-latest","prompt":"hello"}"#; + let mut body = Some(Bytes::from_static(json)); + let original = body.clone(); + + let action = filter.on_request_body(&mut ctx, &mut body, false).await.unwrap(); + assert!(matches!(action, FilterAction::Continue)); + assert_eq!(body, original, "must not promote before end_of_stream"); + assert!(ctx.extra_request_headers.is_empty()); + } + #[tokio::test] async fn extracts_model_field() { let filter = ModelToHeaderFilter::from_config(&serde_yaml::Value::Null).unwrap(); diff --git a/filters/src/lib.rs b/filters/src/lib.rs index bbe56f285a..6383023176 100644 --- a/filters/src/lib.rs +++ b/filters/src/lib.rs @@ -40,7 +40,7 @@ pub use callout::HttpCalloutFilter; pub use gcp::GcpAdcFilter; pub use guardrails::AiGuardrailsFilter; pub use identity_guard::IdentityHeaderGuardFilter; -pub use inference::ModelToHeaderFilter; +pub use inference::{LlmisvcModelProviderResolverFilter, ModelToHeaderFilter}; pub use metering::ExternalMeteringFilter; pub use prompt_enrich::PromptEnrichFilter; pub use register::{build_ai_registry, register_ai_filters}; diff --git a/filters/src/register.rs b/filters/src/register.rs index 3891311775..c909b01be0 100644 --- a/filters/src/register.rs +++ b/filters/src/register.rs @@ -16,8 +16,8 @@ use crate::HttpCalloutFilter; use crate::TokenRateLimitFilter; use crate::{ A2aFilter, AiGuardrailsFilter, CredentialInjectFilter, ExternalMeteringFilter, IdentityHeaderGuardFilter, - IntelligentRouteFilter, McpFilter, ModelToHeaderFilter, PromptEnrichFilter, ProviderRouteFilter, Sigv4SignFilter, - TimeToFirstTokenFilter, TokenCountFilter, TokenUsageHeadersFilter, + IntelligentRouteFilter, LlmisvcModelProviderResolverFilter, McpFilter, ModelToHeaderFilter, PromptEnrichFilter, + ProviderRouteFilter, Sigv4SignFilter, TimeToFirstTokenFilter, TokenCountFilter, TokenUsageHeadersFilter, }; /// Register all in-tree AI HTTP filters into `registry`. @@ -124,6 +124,10 @@ fn register_general_ai_filters(registry: &mut FilterRegistry) { @register registry, http "model_to_header" => ModelToHeaderFilter::from_config ); + praxis_filter::register_filters!( + @register registry, + http "llmisvc_model_provider_resolver" => LlmisvcModelProviderResolverFilter::from_config + ); praxis_filter::register_filters!( @register registry, http "prompt_enrich" => PromptEnrichFilter::from_config @@ -459,6 +463,7 @@ mod tests { let expected = [ "ai_guardrails", "identity_header_guard", + "llmisvc_model_provider_resolver", "openai_responses_validate", "responses_to_chat_completions", "a2a", diff --git a/tests/integration/tests/suite/examples/llmisvc_model_provider_resolver.rs b/tests/integration/tests/suite/examples/llmisvc_model_provider_resolver.rs new file mode 100644 index 0000000000..6f721ffec6 --- /dev/null +++ b/tests/integration/tests/suite/examples/llmisvc_model_provider_resolver.rs @@ -0,0 +1,127 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 Praxis Contributors + +//! Tests for the LLMISvc model-provider resolver example configuration. + +use std::collections::HashMap; + +use praxis_test_utils::{ + free_port, http_send, json_post, parse_body, parse_status, start_echo_backend, start_header_echo_backend, + start_proxy, +}; + +// ----------------------------------------------------------------------------- +// Tests +// ----------------------------------------------------------------------------- + +#[test] +fn llmisvc_model_provider_resolver_config_parses() { + let config = super::load_example_config( + "llmisvc-model-provider-resolver.yaml", + 29920, + HashMap::from([("127.0.0.1:3000", 29921_u16)]), + ); + + assert_eq!(config.listeners.len(), 1, "should have 1 listener"); + assert_eq!( + &*config.listeners[0].name, "llmisvc-gateway", + "listener name should be llmisvc-gateway" + ); +} + +#[test] +fn llmisvc_rewrites_publisher_id_body_model() { + let backend_guard = start_echo_backend(); + let backend_port = backend_guard.port(); + let proxy_port = free_port(); + + let config = super::load_example_config( + "llmisvc-model-provider-resolver.yaml", + proxy_port, + HashMap::from([("127.0.0.1:3000", backend_port)]), + ); + + let proxy = start_proxy(&config); + let raw = http_send( + proxy.addr(), + &json_post( + "/v1/chat/completions", + r#"{"model":"publishers/rhoai/models/granite-3.1-8b","messages":[{"role":"user","content":"hi"}]}"#, + ), + ); + + assert_eq!(parse_status(&raw), 200, "rewrite should return 200"); + let body = parse_body(&raw); + let parsed: serde_json::Value = serde_json::from_str(&body).expect("backend should echo valid JSON"); + assert_eq!( + parsed["model"].as_str(), + Some("granite-3.1-8b"), + "upstream body model should be the short name" + ); + assert_eq!( + parsed["messages"][0]["content"].as_str(), + Some("hi"), + "other body fields should be preserved" + ); +} + +#[test] +fn llmisvc_preserves_routing_header_publisher_id() { + let backend_guard = start_header_echo_backend(); + let backend_port = backend_guard.port(); + let proxy_port = free_port(); + + let config = super::load_example_config( + "llmisvc-model-provider-resolver.yaml", + proxy_port, + HashMap::from([("127.0.0.1:3000", backend_port)]), + ); + + let proxy = start_proxy(&config); + let raw = http_send( + proxy.addr(), + &json_post( + "/v1/chat/completions", + r#"{"model":"publishers/rhoai/models/granite-3.1-8b","messages":[]}"#, + ), + ); + + assert_eq!(parse_status(&raw), 200, "header echo should return 200"); + let headers = parse_body(&raw); + assert!( + headers + .lines() + .any(|line| line.eq_ignore_ascii_case("x-model: publishers/rhoai/models/granite-3.1-8b")), + "X-Model routing header must remain the publisher ID, got:\n{headers}" + ); +} + +#[test] +fn llmisvc_passes_non_publisher_model_through() { + let backend_guard = start_echo_backend(); + let backend_port = backend_guard.port(); + let proxy_port = free_port(); + + let config = super::load_example_config( + "llmisvc-model-provider-resolver.yaml", + proxy_port, + HashMap::from([("127.0.0.1:3000", backend_port)]), + ); + + let proxy = start_proxy(&config); + let raw = http_send( + proxy.addr(), + &json_post( + "/v1/chat/completions", + r#"{"model":"mistral-large-latest","messages":[]}"#, + ), + ); + + assert_eq!(parse_status(&raw), 200, "passthrough should return 200"); + let parsed: serde_json::Value = serde_json::from_str(&parse_body(&raw)).expect("backend should echo valid JSON"); + assert_eq!( + parsed["model"].as_str(), + Some("mistral-large-latest"), + "non-publisher model must not be rewritten" + ); +} diff --git a/tests/integration/tests/suite/examples/mod.rs b/tests/integration/tests/suite/examples/mod.rs index 2c1622f6e1..3432bb4577 100644 --- a/tests/integration/tests/suite/examples/mod.rs +++ b/tests/integration/tests/suite/examples/mod.rs @@ -32,6 +32,7 @@ mod irr_terminal_streaming; mod lakera_guard; #[cfg(feature = "llmd-ext-proc")] mod llmd_ext_proc; +mod llmisvc_model_provider_resolver; mod mcp_broker; mod model_to_header; mod openai_agentic_loop;