Skip to main content

renox_core/
domain.rs

1//! Routes for other hosts: `Routes::domain("admin.example.com", …)` and
2//! `Routes::domain("{account}.example.com", …)`.
3
4use std::collections::HashMap;
5use std::convert::Infallible;
6use std::sync::Arc;
7
8use axum::Router;
9use axum::extract::{FromRequestParts, Request};
10use axum::http::request::Parts;
11
12/// A host pattern: labels separated by dots, each literal or `{name}`
13/// (one label, e.g. `{account}` in `{account}.example.com`).
14#[derive(Debug, Clone, PartialEq, Eq)]
15pub(crate) struct DomainPattern {
16    text: String,
17    labels: Vec<Label>,
18}
19
20#[derive(Debug, Clone, PartialEq, Eq)]
21enum Label {
22    Literal(String),
23    Param(String),
24}
25
26impl DomainPattern {
27    pub(crate) fn parse(text: &str) -> anyhow::Result<Self> {
28        let text = text.trim().to_ascii_lowercase();
29        anyhow::ensure!(
30            !text.is_empty() && !text.contains(['/', ':']),
31            "Routes::domain(\"{text}\"): give a host like `admin.example.com` or \
32             `{{account}}.example.com`, without a scheme, port or path"
33        );
34        let labels = text
35            .split('.')
36            .map(|label| {
37                if let Some(name) = label.strip_prefix('{').and_then(|l| l.strip_suffix('}')) {
38                    anyhow::ensure!(
39                        !name.is_empty()
40                            && name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_'),
41                        "Routes::domain(\"{text}\"): `{label}` isn't a valid parameter"
42                    );
43                    Ok(Label::Param(name.to_owned()))
44                } else {
45                    anyhow::ensure!(
46                        !label.is_empty() && !label.contains(['{', '}']),
47                        "Routes::domain(\"{text}\"): `{label}` isn't a valid part of a host"
48                    );
49                    Ok(Label::Literal(label.to_owned()))
50                }
51            })
52            .collect::<anyhow::Result<_>>()?;
53        Ok(Self { text, labels })
54    }
55
56    pub(crate) fn as_str(&self) -> &str {
57        &self.text
58    }
59
60    /// The parameters when `host` (without its port) matches.
61    fn matches(&self, host: &str) -> Option<HashMap<String, String>> {
62        let host = host.to_ascii_lowercase();
63        let parts: Vec<&str> = host.split('.').collect();
64        if parts.len() != self.labels.len() {
65            return None;
66        }
67        let mut params = HashMap::new();
68        for (label, part) in self.labels.iter().zip(parts) {
69            match label {
70                Label::Literal(literal) if literal == part => {}
71                Label::Literal(_) => return None,
72                Label::Param(name) => {
73                    if part.is_empty() {
74                        return None;
75                    }
76                    params.insert(name.clone(), part.to_owned());
77                }
78            }
79        }
80        Some(params)
81    }
82}
83
84/// The request's host from `Host` (HTTP/1) or the URI's authority (HTTP/2),
85/// without a port.
86pub(crate) fn host(req: &Request) -> Option<String> {
87    let raw = req
88        .headers()
89        .get(axum::http::header::HOST)
90        .and_then(|v| v.to_str().ok())
91        .map(str::to_owned)
92        .or_else(|| req.uri().authority().map(|a| a.as_str().to_owned()))?;
93    // `[::1]:3000` and `example.com:3000`: drop the port.
94    let host = if raw.starts_with('[') {
95        raw.split(']').next().map(|h| format!("{h}]"))?
96    } else {
97        raw.split(':').next()?.to_owned()
98    };
99    Some(host)
100}
101
102/// Sends each request to the router of the first domain its host matches,
103/// else to `default`.
104pub(crate) fn dispatch(domains: Vec<(DomainPattern, Router)>, default: Router) -> Router {
105    let domains = Arc::new(domains);
106    Router::new().fallback_service(tower::service_fn(move |mut req: Request| {
107        let (domains, default) = (domains.clone(), default.clone());
108        async move {
109            if let Some(host) = host(&req) {
110                for (pattern, router) in domains.iter() {
111                    if let Some(params) = pattern.matches(&host) {
112                        req.extensions_mut().insert(DomainParams(Arc::new(params)));
113                        req.extensions_mut()
114                            .insert(MatchedDomain(Arc::from(pattern.as_str())));
115                        return tower::ServiceExt::oneshot(router.clone(), req).await;
116                    }
117                }
118            }
119            tower::ServiceExt::oneshot(default, req).await
120        }
121    }))
122}
123
124/// The `Routes::domain` pattern a request matched (for route names).
125#[derive(Debug, Clone)]
126pub(crate) struct MatchedDomain(pub Arc<str>);
127
128/// The parameters of the domain a request came to, for routes added with
129/// `Routes::domain("{account}.example.com", …)`:
130///
131/// ```
132/// # use renox::prelude::*;
133/// use renox::DomainParams;
134///
135/// async fn home(domain: DomainParams) -> String {
136///     format!("Welcome to {}", domain.get("account").unwrap_or("?"))
137/// }
138///
139/// # let _: Routes =
140/// Routes::new().domain("{account}.example.com", Routes::new().get("/", home))
141/// # ;
142/// ```
143///
144/// Empty for other routes.
145#[derive(Debug, Clone, Default)]
146pub struct DomainParams(Arc<HashMap<String, String>>);
147
148impl DomainParams {
149    /// A parameter of the domain pattern, e.g. `account`.
150    pub fn get(&self, name: &str) -> Option<&str> {
151        self.0.get(name).map(String::as_str)
152    }
153}
154
155impl<S: Send + Sync> FromRequestParts<S> for DomainParams {
156    type Rejection = Infallible;
157
158    async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self, Infallible> {
159        Ok(parts
160            .extensions
161            .get::<DomainParams>()
162            .cloned()
163            .unwrap_or_default())
164    }
165}
166
167#[cfg(test)]
168mod tests {
169    use super::*;
170
171    #[test]
172    fn patterns_match_hosts() {
173        let admin = DomainPattern::parse("Admin.Example.com").unwrap();
174        assert!(admin.matches("admin.example.com").is_some());
175        assert!(admin.matches("ADMIN.example.com").is_some());
176        assert!(admin.matches("example.com").is_none());
177        assert!(admin.matches("x.admin.example.com").is_none());
178
179        let tenant = DomainParams(Arc::new(
180            DomainPattern::parse("{account}.example.com")
181                .unwrap()
182                .matches("acme.example.com")
183                .unwrap(),
184        ));
185        assert_eq!(tenant.get("account"), Some("acme"));
186        assert!(
187            DomainPattern::parse("{account}.example.com")
188                .unwrap()
189                .matches("example.com")
190                .is_none()
191        );
192
193        for bad in [
194            "",
195            "https://x.com",
196            "x.com:80",
197            "x.com/a",
198            "{}.x.com",
199            "a..b",
200            "{a b}.x.com",
201        ] {
202            assert!(DomainPattern::parse(bad).is_err(), "{bad}");
203        }
204        // An empty label never fills a parameter.
205        let tenant = DomainPattern::parse("{account}.example.com").unwrap();
206        assert!(tenant.matches(".example.com").is_none());
207    }
208
209    #[test]
210    fn hosts_drop_their_port_ipv6_included() {
211        let host_of = |value: &str| {
212            let req = Request::builder()
213                .header(axum::http::header::HOST, value)
214                .body(axum::body::Body::empty())
215                .unwrap();
216            host(&req)
217        };
218        assert_eq!(host_of("[::1]:3000").as_deref(), Some("[::1]"));
219        assert_eq!(host_of("shop.test:8080").as_deref(), Some("shop.test"));
220    }
221}