mod_dns_resolver/
lib.rs

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        /// An insecure (non-DNSSEC) zone.
198        Insecure(String),
199        /// A zone with an explicit DNSSEC secure flag.
200        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        /// A bare list of zones (each an insecure string or `{zone, secure}`).
211        Zones(Vec<ZoneSpec>),
212        /// The full form, which additionally supports forcing SERVFAIL for a
213        /// set of owner names.
214        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}