1use ::config::{any_err, get_or_create_sub_module, serialize_options, SerdeWrappedValue};
2use dns_resolver::{
3 get_resolver, ptr_host, resolve_a_or_aaaa, reverse_ip, AggregateResolver, HickoryResolver,
4 IpLookupStrategy, Resolver, TestResolver,
5};
6use kumo_address::host_or_socket::HostOrSocketAddress;
7use mailexchanger::{
8 set_mx_concurrency_limit, set_mx_negative_cache_ttl, set_mx_timeout, MailExchanger,
9};
10use mlua::{Lua, LuaSerdeExt, Value};
11use parking_lot::Mutex;
12use std::collections::HashMap;
13use std::net::IpAddr;
14use std::str::FromStr;
15use std::sync::{Arc, LazyLock};
16use std::time::Duration;
17
18mod config;
19mod hickory_backend;
20mod resolv_conf_loader;
21mod unbound_backend;
22
23use crate::config::DnsResolverConfig;
24
25static RESOLVERS: LazyLock<Mutex<HashMap<String, Arc<Box<dyn Resolver>>>>> =
26 LazyLock::new(|| Mutex::new(HashMap::new()));
27
28pub fn get_resolver_instance(
29 opt_resolver_name: &Option<String>,
30) -> anyhow::Result<Arc<Box<dyn Resolver>>> {
31 if let Some(name) = opt_resolver_name {
32 return RESOLVERS
33 .lock()
34 .get(name)
35 .cloned()
36 .ok_or_else(|| anyhow::anyhow!("resolver {name} is not defined"));
37 }
38
39 Ok(get_resolver())
40}
41
42pub fn get_opt_resolver(
43 opt_resolver_name: &Option<String>,
44) -> anyhow::Result<Option<Arc<Box<dyn Resolver>>>> {
45 if let Some(name) = opt_resolver_name {
46 let r = RESOLVERS
47 .lock()
48 .get(name)
49 .cloned()
50 .ok_or_else(|| anyhow::anyhow!("resolver {name} is not defined"))?;
51 Ok(Some(r))
52 } else {
53 Ok(None)
54 }
55}
56
57pub fn register(lua: &Lua) -> anyhow::Result<()> {
58 let dns_mod = get_or_create_sub_module(lua, "dns")?;
59
60 dns_mod.set(
61 "lookup_mx",
62 lua.create_async_function(
63 |lua, (domain, opt_resolver_name): (String, Option<String>)| async move {
64 let opt_resolver = get_opt_resolver(&opt_resolver_name).map_err(any_err)?;
65 let mx = MailExchanger::resolve_via(&domain, opt_resolver.as_ref().map(|r| &***r))
66 .await
67 .map_err(any_err)?;
68 Ok(lua.to_value_with(&*mx, serialize_options()))
69 },
70 )?,
71 )?;
72
73 dns_mod.set(
74 "set_mx_concurrency_limit",
75 lua.create_function(move |_lua, limit: usize| {
76 set_mx_concurrency_limit(limit);
77 Ok(())
78 })?,
79 )?;
80
81 dns_mod.set(
82 "set_mx_timeout",
83 lua.create_function(move |lua, duration: Value| {
84 let duration: duration_serde::Wrap<Duration> = lua.from_value(duration)?;
85 set_mx_timeout(duration.into_inner()).map_err(any_err)
86 })?,
87 )?;
88
89 dns_mod.set(
90 "set_mx_negative_cache_ttl",
91 lua.create_function(move |lua, duration: Value| {
92 let duration: duration_serde::Wrap<Duration> = lua.from_value(duration)?;
93 set_mx_negative_cache_ttl(duration.into_inner()).map_err(any_err)
94 })?,
95 )?;
96
97 dns_mod.set(
98 "set_mta_sts_enabled",
99 lua.create_function(move |_lua, enabled: bool| {
100 mailexchanger::set_mta_sts_enabled(enabled);
101 Ok(())
102 })?,
103 )?;
104
105 dns_mod.set(
106 "ptr_host",
107 lua.create_function(move |_lua, ip: String| {
108 let ip: IpAddr = ip.parse().map_err(any_err)?;
109 Ok(ptr_host(ip))
110 })?,
111 )?;
112
113 dns_mod.set(
114 "reverse_ip",
115 lua.create_function(move |_lua, ip: String| {
116 let ip: IpAddr = ip.parse().map_err(any_err)?;
117 Ok(reverse_ip(ip))
118 })?,
119 )?;
120
121 dns_mod.set(
122 "rbl_lookup",
123 lua.create_async_function(
124 |_lua, (ip_str, bl_domain, opt_resolver_name): (String, String, Option<String>)| async move {
125 let resolver = get_resolver_instance(&opt_resolver_name).map_err(any_err)?;
126
127 let address: HostOrSocketAddress = ip_str.parse().map_err(any_err)?;
128 let reversed_ip = reverse_ip(address.ip().ok_or_else(||mlua::Error::external(format!("{ip_str} is not a valid IpAddr or SocketAddr")))?);
129 let name = format!("{reversed_ip}.{bl_domain}.");
130
131 let answers = resolver.resolve_ip(&name).await.map_err(any_err)?;
132 match answers.first() {
133 Some(ip) => {
134 let txt = resolver.resolve_txt(&name).await.map(|a| a.as_txt().join("")).ok();
135 Ok((Some(ip.to_string()), txt))
136 }
137 None => {
138 Ok((None, None))
139 }
140 }
141 },
142 )?,
143 )?;
144
145 dns_mod.set(
146 "lookup_ptr",
147 lua.create_async_function(
148 |lua, (ip_str, opt_resolver_name): (String, Option<String>)| async move {
149 let resolver = get_resolver_instance(&opt_resolver_name).map_err(any_err)?;
150 let addr = std::net::IpAddr::from_str(&ip_str).map_err(any_err)?;
151 let answer = resolver.resolve_ptr(addr).await.map_err(any_err)?;
152 Ok(lua.to_value_with(&*answer, serialize_options()))
153 },
154 )?,
155 )?;
156
157 dns_mod.set(
158 "lookup_txt",
159 lua.create_async_function(
160 |_lua, (domain, opt_resolver_name): (String, Option<String>)| async move {
161 let resolver = get_resolver_instance(&opt_resolver_name).map_err(any_err)?;
162 let answer = resolver.resolve_txt(&domain).await.map_err(any_err)?;
163 Ok(answer.as_txt())
164 },
165 )?,
166 )?;
167
168 dns_mod.set(
169 "lookup_addr",
170 lua.create_async_function(
171 |_lua,
172 (domain, opt_resolver_name, strategy): (
173 String,
174 Option<String>,
175 Option<SerdeWrappedValue<IpLookupStrategy>>,
176 )| async move {
177 let opt_resolver = get_opt_resolver(&opt_resolver_name).map_err(any_err)?;
178 let result = resolve_a_or_aaaa(
179 &domain,
180 opt_resolver.as_ref().map(|r| &***r),
181 strategy.map(|v| v.0).unwrap_or_default(),
182 )
183 .await
184 .map_err(any_err)?;
185 let result: Vec<String> = result
186 .into_iter()
187 .map(|item| item.addr.to_string())
188 .collect();
189 Ok(result)
190 },
191 )?,
192 )?;
193
194 #[derive(serde::Deserialize, Debug)]
195 #[serde(untagged)]
196 enum ZoneSpec {
197 Insecure(String),
199 Detailed {
201 zone: String,
202 #[serde(default)]
203 secure: bool,
204 },
205 }
206
207 #[derive(serde::Deserialize, Debug)]
208 #[serde(untagged)]
209 enum TestResolverConfig {
210 Zones(Vec<ZoneSpec>),
212 Detailed {
215 zones: Vec<ZoneSpec>,
216 #[serde(default)]
217 servfail: Vec<String>,
218 },
219 }
220
221 impl TestResolverConfig {
222 fn make_resolver(&self) -> anyhow::Result<TestResolver> {
223 let mut resolver = TestResolver::default();
224
225 let (zones, servfail): (&[ZoneSpec], &[String]) = match self {
226 Self::Zones(zones) => (zones, &[]),
227 Self::Detailed { zones, servfail } => (zones, servfail),
228 };
229
230 for zone in zones {
231 resolver = match zone {
232 ZoneSpec::Insecure(zone) => resolver.with_zone(zone),
233 ZoneSpec::Detailed { zone, secure: true } => resolver.with_secure_zone(zone),
234 ZoneSpec::Detailed {
235 zone,
236 secure: false,
237 } => resolver.with_zone(zone),
238 }
239 .map_err(|err| anyhow::anyhow!("{err}"))?;
240 }
241
242 for name in servfail {
243 resolver = resolver.with_servfail(name);
244 }
245
246 Ok(resolver)
247 }
248 }
249
250 #[derive(serde::Deserialize, Debug)]
251 enum KumoResolverConfig {
252 Hickory(DnsResolverConfig),
253 HickorySystemConfig,
254 Unbound(DnsResolverConfig),
255 Test(TestResolverConfig),
256 Aggregate(Vec<KumoResolverConfig>),
257 }
258
259 impl KumoResolverConfig {
260 fn make_resolver(&self, path: &str) -> anyhow::Result<Box<dyn Resolver>> {
261 match self {
262 Self::Hickory(config) => Ok(Box::new(
263 hickory_backend::build_hickory_resolver(config)
264 .map_err(|e| anyhow::anyhow!("{path}: {e}"))?,
265 )),
266 Self::HickorySystemConfig => Ok(Box::new(HickoryResolver::new()?)),
267 Self::Unbound(config) => Ok(Box::new(
268 unbound_backend::build_unbound_resolver(config)
269 .map_err(|e| anyhow::anyhow!("{path}: {e}"))?,
270 )),
271 Self::Test(config) => Ok(Box::new(config.make_resolver()?)),
272 Self::Aggregate(children) => {
273 let mut resolver = AggregateResolver::new();
274 for (idx, child) in children.iter().enumerate() {
275 let child_path = format!("{path}.Aggregate[{idx}]");
276 resolver.push_resolver(child.make_resolver(&child_path)?);
277 }
278 Ok(Box::new(resolver))
279 }
280 }
281 }
282 }
283
284 dns_mod.set(
285 "configure_resolver",
286 lua.create_function(move |lua, config: mlua::Value| {
287 match lua.from_value::<KumoResolverConfig>(config.clone()) {
288 Ok(config) => {
289 let resolver = config
290 .make_resolver("configure_resolver")
291 .map_err(any_err)?;
292 dns_resolver::reconfigure_resolver(resolver);
293 Ok(())
294 }
295 Err(err1) => match lua.from_value::<DnsResolverConfig>(config) {
296 Ok(config) => {
297 let resolver = hickory_backend::build_hickory_resolver(&config)
298 .map_err(any_err)?;
299 dns_resolver::reconfigure_resolver(resolver);
300 Ok(())
301 }
302 Err(err2) => Err(mlua::Error::external(format!(
303 "failed to parse config as either KumoResolverConfig ({err1:#}) or DnsResolverConfig ({err2:#})"
304 ))),
305 },
306 }
307 })?,
308 )?;
309
310 dns_mod.set(
311 "define_resolver",
312 lua.create_function(move |lua, (name, config): (String, mlua::Value)| {
313 let config = lua
314 .from_value::<KumoResolverConfig>(config.clone())
315 .map_err(any_err)?;
316 let path = format!("define_resolver({name:?})");
317 let resolver = config.make_resolver(&path).map_err(any_err)?;
318
319 RESOLVERS.lock().insert(name, resolver.into());
320
321 Ok(())
322 })?,
323 )?;
324
325 dns_mod.set(
326 "configure_unbound_resolver",
327 lua.create_function(move |lua, config: mlua::Value| {
328 let config: DnsResolverConfig = lua.from_value(config)?;
329 let resolver = unbound_backend::build_unbound_resolver(&config).map_err(any_err)?;
330 dns_resolver::reconfigure_resolver(resolver);
331 Ok(())
332 })?,
333 )?;
334
335 dns_mod.set(
336 "configure_test_resolver",
337 lua.create_function(move |lua, config: mlua::Value| {
338 let config = lua
339 .from_value::<TestResolverConfig>(config)
340 .map_err(any_err)?;
341 let resolver = config.make_resolver().map_err(any_err)?;
342 dns_resolver::reconfigure_resolver(resolver);
343 Ok(())
344 })?,
345 )?;
346
347 dns_mod.set(
348 "configure_test_mta_sts",
349 lua.create_function(
350 move |_lua, policies: std::collections::BTreeMap<String, String>| {
351 let parsed = policies
352 .into_iter()
353 .map(|(domain, text)| {
354 let policy =
355 mta_sts::policy::MtaStsPolicy::parse(&text).map_err(any_err)?;
356 Ok((domain, policy))
357 })
358 .collect::<mlua::Result<std::collections::BTreeMap<_, _>>>()?;
359 mta_sts::set_test_policies(parsed);
360 Ok(())
361 },
362 )?,
363 )?;
364
365 dns_mod.set(
366 "load_resolv_conf",
367 lua.create_function(move |lua, path: Option<String>| {
368 let config = resolv_conf_loader::load_resolv_conf(path.as_deref()).map_err(any_err)?;
369 lua.to_value_with(&config, serialize_options())
370 })?,
371 )?;
372
373 Ok(())
374}