1use dns_resolver::Resolver;
2use futures::future::BoxFuture;
3use hickory_resolver::proto::rr::Name;
4use policy::MtaStsPolicy;
5use std::collections::BTreeMap;
6use std::sync::{Arc, Mutex};
7use std::time::{Duration, Instant};
8
9lruttl::declare_cache! {
10static CACHE: LruCacheWithTtl<Name, CachedPolicy>::new("mta_sts_policy", 64 * 1024);
12}
13
14pub mod dns;
15pub mod policy;
16
17#[derive(Clone, Debug)]
18struct CachedPolicy {
19 pub id: String,
20 pub policy: Arc<MtaStsPolicy>,
21}
22
23struct Getter {}
24
25impl policy::Get for Getter {
26 fn http_get<'a>(&'a self, url: &'a str) -> BoxFuture<'a, anyhow::Result<String>> {
27 Box::pin(async move {
28 let response = reqwest::Client::builder()
29 .redirect(reqwest::redirect::Policy::none())
32 .timeout(std::time::Duration::from_secs(20))
33 .build()?
34 .request(reqwest::Method::GET, url)
35 .send()
36 .await?;
37
38 let status = response.status();
42 if status != reqwest::StatusCode::OK {
43 anyhow::bail!("failed to GET {url}: {status}");
44 }
45
46 let content_type = response
56 .headers()
57 .get(reqwest::header::CONTENT_TYPE)
58 .ok_or_else(|| anyhow::anyhow!("missing required Content-Type header"))?;
59
60 let content_type = content_type.to_str()?;
61
62 let ct = if let Some((ct, _)) = content_type.split_once(';') {
63 ct.trim()
64 } else {
65 content_type.trim()
66 };
67 if ct != "text/plain" {
68 anyhow::bail!("Content-Type must be text/plain, got {content_type}");
69 }
70
71 Ok(response.text().await?)
72 })
73 }
74}
75
76static TEST_POLICIES: Mutex<Option<Arc<BTreeMap<String, Arc<MtaStsPolicy>>>>> = Mutex::new(None);
81
82pub fn set_test_policies(policies: BTreeMap<String, MtaStsPolicy>) {
86 let map = policies
87 .into_iter()
88 .map(|(domain, policy)| (domain, Arc::new(policy)))
89 .collect();
90 *TEST_POLICIES.lock().unwrap() = Some(Arc::new(map));
91}
92
93pub async fn get_policy_for_domain(
94 policy_domain: &str,
95 resolver: Option<&dyn Resolver>,
96) -> anyhow::Result<Arc<MtaStsPolicy>> {
97 if let Some(policies) = TEST_POLICIES.lock().unwrap().clone() {
98 let domain = policy_domain.trim_end_matches('.');
99 return policies
100 .get(domain)
101 .cloned()
102 .ok_or_else(|| anyhow::anyhow!("no MTA-STS policy for {domain}"));
103 }
104
105 match resolver {
106 Some(resolver) => get_policy_for_domain_impl(policy_domain, resolver, &Getter {}).await,
107 None => {
108 let resolver = dns_resolver::get_resolver();
109 get_policy_for_domain_impl(policy_domain, &**resolver, &Getter {}).await
110 }
111 }
112}
113
114fn cache_lookup(name: &Name) -> Option<CachedPolicy> {
115 CACHE.get(name)
116}
117
118async fn get_policy_for_domain_impl(
119 policy_domain: &str,
120 resolver: &dyn Resolver,
121 getter: &dyn policy::Get,
122) -> anyhow::Result<Arc<MtaStsPolicy>> {
123 let name = Name::from_str_relaxed(policy_domain)?.to_lowercase();
124
125 if let Some(cached) = cache_lookup(&name) {
126 let still_valid = dns::resolve_dns_record(policy_domain, resolver)
129 .await
130 .map(|r| cached.id == r.id)
131 .unwrap_or(true);
132
133 if still_valid {
134 return Ok(Arc::clone(&cached.policy));
135 }
136 }
137
138 let record = dns::resolve_dns_record(policy_domain, resolver).await?;
139
140 let policy = Arc::new(policy::load_policy_for_domain(policy_domain, getter).await?);
141
142 let expires = Instant::now() + Duration::from_secs(policy.max_age);
143
144 CACHE
145 .insert(
146 name,
147 CachedPolicy {
148 id: record.id,
149 policy: Arc::clone(&policy),
150 },
151 expires.into(),
152 )
153 .await;
154
155 Ok(policy)
156}
157
158