Skip to main content

systemprompt_api/services/
request_base_url.rs

1//! Request-derived base URL for OAuth discovery responses.
2//!
3//! RFC 9728 implementations identify themselves coherently from the host the
4//! client actually dialled. A single gateway reachable via both `127.0.0.1`
5//! and `localhost` must echo whichever the client used in every URL it
6//! returns (`issuer`, `authorization_endpoint`, `token_endpoint`, `resource`…),
7//! otherwise the client's RFC 8707 `resource` indicator won't round-trip
8//! against the configured `api_external_url` origin.
9//!
10//! [`RequestBaseUrl`] is an axum extractor that resolves
11//! `scheme://host[:port]` from the incoming request, validating the host
12//! against a small allowlist seeded from `api_external_url`. On allowlist
13//! miss or missing/invalid header it falls back to `api_external_url` — the
14//! gateway never advertises a hostname an attacker fabricated via Host
15//! header injection.
16//!
17//! Copyright (c) systemprompt.io — Business Source License 1.1.
18//! See <https://systemprompt.io> for licensing details.
19
20use axum::extract::FromRequestParts;
21use axum::response::{IntoResponse, Response};
22use http::header;
23use http::request::Parts;
24use systemprompt_manifest::Config;
25use systemprompt_models::api::ApiError;
26use systemprompt_models::errors::GlobalConfigError;
27
28#[derive(Debug, Clone)]
29pub struct RequestBaseUrl {
30    base: String,
31    origin: url::Origin,
32}
33
34impl RequestBaseUrl {
35    #[must_use]
36    pub fn as_str(&self) -> &str {
37        &self.base
38    }
39
40    #[must_use]
41    pub const fn origin(&self) -> &url::Origin {
42        &self.origin
43    }
44
45    #[must_use]
46    pub fn into_string(self) -> String {
47        self.base
48    }
49}
50
51fn is_loopback_host(host: &str) -> bool {
52    let bare = host.split(':').next().unwrap_or(host).to_ascii_lowercase();
53    bare == "localhost" || bare == "127.0.0.1" || bare == "[::1]" || bare == "::1"
54}
55
56fn host_in_allowlist(candidate_host: &str, configured: &url::Url) -> bool {
57    let candidate_bare = candidate_host
58        .rsplit_once(':')
59        .map_or(candidate_host, |(h, _)| h)
60        .to_ascii_lowercase();
61    let configured_host = configured.host_str().unwrap_or("").to_ascii_lowercase();
62
63    if candidate_bare == configured_host {
64        return true;
65    }
66    if is_loopback_host(&configured_host) && is_loopback_host(&candidate_bare) {
67        return true;
68    }
69    false
70}
71
72fn fallback_from_url(configured: &url::Url) -> RequestBaseUrl {
73    let trimmed = configured.as_str().trim_end_matches('/').to_owned();
74    RequestBaseUrl {
75        base: trimmed,
76        origin: configured.origin(),
77    }
78}
79
80#[must_use]
81pub fn resolve(raw_host: Option<&str>, configured: &url::Url) -> RequestBaseUrl {
82    if let Some(host) = raw_host.map(str::trim).filter(|s| !s.is_empty())
83        && let Ok(resolved) = build_from_host(host, configured)
84    {
85        return resolved;
86    }
87    fallback_from_url(configured)
88}
89
90fn build_from_host(raw_host: &str, configured: &url::Url) -> Result<RequestBaseUrl, &'static str> {
91    if raw_host.is_empty() || raw_host.contains('/') || raw_host.contains(' ') {
92        return Err("invalid host header");
93    }
94    if !host_in_allowlist(raw_host, configured) {
95        return Err("host not in allowlist");
96    }
97    let host_bare = raw_host
98        .rsplit_once(':')
99        .map_or(raw_host, |(h, _)| h)
100        .to_ascii_lowercase();
101    let scheme = if is_loopback_host(&host_bare) {
102        "http"
103    } else {
104        configured.scheme()
105    };
106    let base = format!("{scheme}://{raw_host}");
107    let parsed = url::Url::parse(&base).map_err(|_e| "host header did not parse as URL")?;
108    Ok(RequestBaseUrl {
109        base: base.trim_end_matches('/').to_owned(),
110        origin: parsed.origin(),
111    })
112}
113
114#[derive(Debug, thiserror::Error)]
115pub enum RequestBaseUrlError {
116    #[error("configuration unavailable")]
117    Config(#[from] GlobalConfigError),
118    #[error("api_external_url {url:?} is not a valid URL")]
119    InvalidExternalUrl {
120        url: String,
121        #[source]
122        source: url::ParseError,
123    },
124}
125
126impl IntoResponse for RequestBaseUrlError {
127    fn into_response(self) -> Response {
128        ApiError::internal("Request base URL unavailable", self).into_response()
129    }
130}
131
132impl<S: Send + Sync> FromRequestParts<S> for RequestBaseUrl {
133    type Rejection = RequestBaseUrlError;
134
135    #[expect(
136        clippy::unused_async_trait_impl,
137        reason = "async signature required by the FromRequestParts trait; this \
138                  extractor resolves the base URL synchronously"
139    )]
140    async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
141        let cfg = Config::get()?;
142        let configured = url::Url::parse(&cfg.api_external_url).map_err(|source| {
143            RequestBaseUrlError::InvalidExternalUrl {
144                url: cfg.api_external_url.clone(),
145                source,
146            }
147        })?;
148
149        let raw_host = parts
150            .headers
151            .get(header::HOST)
152            .and_then(|v| v.to_str().ok());
153        Ok(resolve(raw_host, &configured))
154    }
155}