diff --git a/docs/filters/llmisvc_model_provider_resolver.md b/docs/filters/llmisvc_model_provider_resolver.md new file mode 100644 index 0000000000..cb5c6f8a4a --- /dev/null +++ b/docs/filters/llmisvc_model_provider_resolver.md @@ -0,0 +1,28 @@ + + + +# `llmisvc_model_provider_resolver` + +Ports the `LLMISvc` / `KServe` BBR body-rewrite branch from IPP's `model-provider-resolver`. + +## 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/reference.md b/docs/filters/reference.md index 5a4213ea33..73966d92c1 100644 --- a/docs/filters/reference.md +++ b/docs/filters/reference.md @@ -64,6 +64,7 @@ see the [Praxis core filter reference][core-ref]. | Filter | Description | |--------|-------------| +| [`llmisvc_model_provider_resolver`](llmisvc_model_provider_resolver.md) | Ports the `LLMISvc` / `KServe` BBR body-rewrite branch from IPP's `model-provider-resolver`. | | [`model_to_header`](model_to_header.md) | Promotes the JSON `"model"` field from the request body to a request header. | ### Prompt Enrich diff --git a/examples/README.md b/examples/README.md index 6265b04f78..1aae70df21 100644 --- a/examples/README.md +++ b/examples/README.md @@ -29,6 +29,7 @@ before sending requests. | [intelligent-route-mcp.yaml](configs/intelligent-route-mcp.yaml) | Routes MCP `tools/call` requests to the cluster that owns the requested tool, using the `mcp.name` metadata set by the `mcp` filter | | [intelligent-route-overlay.yaml](configs/intelligent-route-overlay.yaml) | Routes requests using a routing overlay file (`routing-overlay.json`) instead of inline YAML candidates | | [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 | +| [llmisvc-model-provider-resolver.yaml](configs/llmisvc-model-provider-resolver.yaml) | Ports the LLMISvc / KServe BBR body-rewrite path: when the model name is a publisher ID (`publishers/{ns}/models/{name}`), rewrite the JSON body `"model"` field to `{name}` so vLLM receives the short name | | [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..37bcd5109f --- /dev/null +++ b/examples/configs/llmisvc-model-provider-resolver.yaml @@ -0,0 +1,46 @@ +# LLMISvc Model Provider Resolver +# +# Build: +# cargo build -p praxis-ai-proxy +# +# Ports the LLMISvc / KServe BBR body-rewrite path: when the model +# name is a publisher ID (`publishers/{ns}/models/{name}`), rewrite +# the JSON body `"model"` field to `{name}` so vLLM receives the +# short name. The routing header (default `X-Model`) is left +# unchanged so KServe can still route on 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" 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..75ecfcb846 --- /dev/null +++ b/filters/src/inference/llmisvc_model_provider_resolver.rs @@ -0,0 +1,593 @@ +// SPDX-License-Identifier: MIT +// Copyright (c) 2026 Praxis Contributors + +//! `LLMISvc` model-provider resolver: rewrites `KServe` publisher-ID body +//! `model` values to the short model name while leaving the routing +//! header untouched. + +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 +// ----------------------------------------------------------------------------- + +/// Ports the `LLMISvc` / `KServe` BBR body-rewrite branch from IPP's +/// `model-provider-resolver`. +/// +/// 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 { + // Require the routing header (usually from `model_to_header`); + // missing header means no-op. + 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); + }; + + let Some(raw) = body.as_ref() else { + return Ok(FilterAction::Continue); + }; + + let mut value: serde_json::Value = match serde_json::from_slice(raw) { + Ok(v) => v, + Err(_) => return Ok(FilterAction::Continue), + }; + + let Some(obj) = value.as_object_mut() else { + return Ok(FilterAction::Continue); + }; + + // Only rewrite an existing body field -- never invent `"model"` + // when the publisher ID came solely from the routing header. + if !obj.contains_key("model") { + return Ok(FilterAction::Continue); + } + + // Stash the original publisher ID for later metering. + ctx.set_metadata(META_PUBLISHER_ID, model_name.as_str()); + + if obj.get("model").and_then(serde_json::Value::as_str) == Some(short_name) { + return Ok(FilterAction::Continue); + } + + obj.insert("model".to_owned(), serde_json::Value::String(short_name.to_owned())); + + replace_json_body(body, &value, self.name(), "model").map_err(|e| -> FilterError { + format!("{}: failed to re-serialize rewritten request body: {e}", self.name()).into() + })?; + + debug!( + original = %model_name, + rewritten = %short_name, + "LLMISvc BBR: rewrote body model field" + ); + + 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 +// ----------------------------------------------------------------------------- + +/// 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()); + } + } + + 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. +/// +/// Mirrors IPP: 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 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 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 113d9b1cde..194facb182 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/lib.rs b/filters/src/lib.rs index e8aaca3bac..50e8e320ae 100644 --- a/filters/src/lib.rs +++ b/filters/src/lib.rs @@ -21,7 +21,7 @@ mod token_usage; pub use agentic::{a2a::A2aFilter, mcp::McpFilter}; pub use guardrails::AiGuardrailsFilter; -pub use inference::ModelToHeaderFilter; +pub use inference::{LlmisvcModelProviderResolverFilter, ModelToHeaderFilter}; pub use prompt_enrich::PromptEnrichFilter; pub use register::{build_ai_registry, register_ai_filters}; pub use routing::IntelligentRouteFilter; diff --git a/filters/src/register.rs b/filters/src/register.rs index 6f04506d55..def7c3a6e7 100644 --- a/filters/src/register.rs +++ b/filters/src/register.rs @@ -7,8 +7,8 @@ use praxis_core::subrequest::SubRequestClient; use praxis_filter::FilterRegistry; use crate::{ - A2aFilter, AiGuardrailsFilter, IntelligentRouteFilter, McpFilter, ModelToHeaderFilter, PromptEnrichFilter, - TimeToFirstTokenFilter, TokenCountFilter, TokenUsageHeadersFilter, + A2aFilter, AiGuardrailsFilter, IntelligentRouteFilter, LlmisvcModelProviderResolverFilter, McpFilter, + ModelToHeaderFilter, PromptEnrichFilter, TimeToFirstTokenFilter, TokenCountFilter, TokenUsageHeadersFilter, }; /// Register all in-tree AI HTTP filters into `registry`. @@ -77,6 +77,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 @@ -333,31 +337,18 @@ mod tests { fn build_ai_registry_includes_ai_and_builtin_filters() { let registry = build_ai_registry(); let names = registry.available_filters(); - assert!(names.contains(&"ai_guardrails"), "expected ai_guardrails in registry"); - assert!( - names.contains(&"openai_responses_validate"), - "expected openai_responses_validate in registry" - ); - assert!( - names.contains(&"responses_to_chat_completions"), - "expected responses_to_chat_completions in registry" - ); - assert!(names.contains(&"a2a"), "expected agentic filter a2a in registry"); - assert!( - names.contains(&"intelligent_route"), - "expected intelligent_route in registry" - ); - assert!( - names.contains(&"anthropic_validate"), - "expected anthropic filter in registry" - ); - assert!( - names.contains(&"anthropic_web_search"), - "expected anthropic_web_search in registry" - ); - assert!( - names.contains(&"request_id"), - "expected core builtin request_id in registry" - ); + for expected in [ + "ai_guardrails", + "llmisvc_model_provider_resolver", + "openai_responses_validate", + "responses_to_chat_completions", + "a2a", + "intelligent_route", + "anthropic_validate", + "anthropic_web_search", + "request_id", + ] { + assert!(names.contains(&expected), "expected {expected} in registry"); + } } } 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..1304835f49 --- /dev/null +++ b/tests/integration/tests/suite/examples/llmisvc_model_provider_resolver.rs @@ -0,0 +1,127 @@ +// SPDX-License-Identifier: MIT +// 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 0c665deab2..4980919860 100644 --- a/tests/integration/tests/suite/examples/mod.rs +++ b/tests/integration/tests/suite/examples/mod.rs @@ -17,6 +17,7 @@ mod full_flow; mod full_flow_agentic; mod guardrails; mod inference_fallback; +mod llmisvc_model_provider_resolver; mod mcp_broker; mod model_to_header; mod openai_agentic_loop;