1use 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#[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 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
84pub(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 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
102pub(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#[derive(Debug, Clone)]
126pub(crate) struct MatchedDomain(pub Arc<str>);
127
128#[derive(Debug, Clone, Default)]
146pub struct DomainParams(Arc<HashMap<String, String>>);
147
148impl DomainParams {
149 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 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}