Skip to main content

ddns/resolvers/
deferred.rs

1use std::{fmt, io};
2
3use dquic::qresolve::{
4    EndpointAddr, Family, Publish, PublishFuture, RecordStream, Resolve, ResolveFuture,
5};
6use futures::{FutureExt, future::BoxFuture};
7use snafu::{ResultExt, Snafu};
8use tokio::sync::{Notify, OnceCell};
9
10use crate::resolvers::endpoint_candidates::{
11    EndpointCandidateFuture, EndpointLookup, ResolveEndpointCandidates,
12};
13
14#[derive(Debug, Snafu)]
15#[snafu(module, visibility(pub))]
16pub enum DeferredLookupError {
17    #[snafu(display("deferred resolver has not been initialized"))]
18    Uninitialized,
19    #[snafu(display("deferred resolver lookup failed"))]
20    Lookup { source: io::Error },
21}
22
23#[derive(Debug, Snafu)]
24#[snafu(module, visibility(pub))]
25pub enum DeferredPublishError {
26    #[snafu(display("deferred resolver has not been initialized"))]
27    Uninitialized,
28    #[snafu(display("deferred resolver publish failed"))]
29    Publish { source: io::Error },
30}
31
32#[derive(Debug, Snafu)]
33#[snafu(module, visibility(pub))]
34pub enum SetDeferredResolverError {
35    #[snafu(display("deferred resolver has already been initialized"))]
36    AlreadyInitialized,
37}
38
39pub struct DeferredResolver<R> {
40    inner: OnceCell<R>,
41    initialized: Notify,
42}
43
44impl<R> fmt::Debug for DeferredResolver<R> {
45    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
46        f.debug_struct("DeferredResolver")
47            .field("initialized", &self.inner.get().is_some())
48            .finish()
49    }
50}
51
52impl<R> Default for DeferredResolver<R> {
53    fn default() -> Self {
54        Self::new()
55    }
56}
57
58impl<R> DeferredResolver<R> {
59    #[must_use]
60    pub fn new() -> Self {
61        Self {
62            inner: OnceCell::new(),
63            initialized: Notify::new(),
64        }
65    }
66
67    pub fn set(&self, resolver: R) -> Result<(), SetDeferredResolverError> {
68        if self.inner.set(resolver).is_err() {
69            return set_deferred_resolver_error::AlreadyInitializedSnafu.fail();
70        }
71        self.initialized.notify_waiters();
72        Ok(())
73    }
74
75    #[must_use]
76    pub fn get(&self) -> Option<&R> {
77        self.inner.get()
78    }
79
80    async fn wait(&self) -> &R {
81        loop {
82            let initialized = self.initialized.notified();
83            if let Some(resolver) = self.get() {
84                return resolver;
85            }
86            initialized.await;
87        }
88    }
89}
90
91impl<R> fmt::Display for DeferredResolver<R>
92where
93    R: fmt::Display,
94{
95    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
96        match self.inner.get() {
97            Some(resolver) => write!(f, "DeferredResolver({resolver})"),
98            None => f.write_str("DeferredResolver(uninitialized)"),
99        }
100    }
101}
102
103impl<R> DeferredResolver<R>
104where
105    R: Resolve + 'static,
106{
107    pub async fn lookup_typed(
108        &self,
109        hostname: &str,
110        servname: &str,
111        family: Option<Family>,
112    ) -> Result<RecordStream, DeferredLookupError> {
113        let Some(resolver) = self.get() else {
114            return deferred_lookup_error::UninitializedSnafu.fail();
115        };
116        resolver
117            .lookup(hostname, servname, family)
118            .await
119            .context(deferred_lookup_error::LookupSnafu)
120    }
121}
122
123impl<R> Resolve for DeferredResolver<R>
124where
125    R: Resolve + 'static,
126{
127    fn lookup<'a>(
128        &'a self,
129        hostname: &'a str,
130        servname: &'a str,
131        family: Option<Family>,
132    ) -> ResolveFuture<'a> {
133        async move { self.wait().await.lookup(hostname, servname, family).await }.boxed()
134    }
135}
136
137impl<R> ResolveEndpointCandidates for DeferredResolver<R>
138where
139    R: ResolveEndpointCandidates + 'static,
140{
141    fn lookup_endpoint_candidates<'a>(
142        &'a self,
143        name: &'a str,
144        lookup: EndpointLookup,
145    ) -> EndpointCandidateFuture<'a> {
146        async move {
147            self.wait()
148                .await
149                .lookup_endpoint_candidates(name, lookup)
150                .await
151        }
152        .boxed()
153    }
154}
155
156impl<R> DeferredResolver<R>
157where
158    R: Publish + 'static,
159{
160    pub fn publish_typed<'a>(
161        &'a self,
162        name: &'a str,
163        endpoints: &mut dyn Iterator<Item = EndpointAddr>,
164    ) -> BoxFuture<'a, Result<(), DeferredPublishError>> {
165        let endpoints: Vec<_> = endpoints.collect();
166        async move {
167            let Some(resolver) = self.get() else {
168                return deferred_publish_error::UninitializedSnafu.fail();
169            };
170            let mut endpoints = endpoints.into_iter();
171            resolver
172                .publish(name, &mut endpoints)
173                .await
174                .context(deferred_publish_error::PublishSnafu)
175        }
176        .boxed()
177    }
178}
179
180impl<R> Publish for DeferredResolver<R>
181where
182    R: Publish + 'static,
183{
184    fn publish<'a>(
185        &'a self,
186        name: &'a str,
187        endpoints: &mut dyn Iterator<Item = EndpointAddr>,
188    ) -> PublishFuture<'a> {
189        let endpoints: Vec<_> = endpoints.collect();
190        async move {
191            let resolver = self.wait().await;
192            let mut endpoints = endpoints.into_iter();
193            resolver.publish(name, &mut endpoints).await
194        }
195        .boxed()
196    }
197}
198
199#[cfg(test)]
200mod tests {
201    use std::{fmt, time::Duration};
202
203    use dquic::{
204        qbase::net::addr::EndpointAddr,
205        qresolve::{Publish, Resolve, Source},
206    };
207    use futures::{FutureExt, StreamExt};
208
209    use super::*;
210
211    #[derive(Debug)]
212    struct TestResolver;
213
214    impl fmt::Display for TestResolver {
215        fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
216            f.write_str("test resolver")
217        }
218    }
219
220    impl Resolve for TestResolver {
221        fn lookup<'a>(
222            &'a self,
223            _hostname: &'a str,
224            _servname: &'a str,
225            _family: Option<Family>,
226        ) -> dquic::qresolve::ResolveFuture<'a> {
227            async move {
228                let endpoint = EndpointAddr::direct("127.0.0.1:4433".parse().unwrap());
229                Ok(futures::stream::iter([(Source::System, endpoint)]).boxed())
230            }
231            .boxed()
232        }
233    }
234
235    impl Publish for TestResolver {
236        fn publish<'a>(
237            &'a self,
238            _name: &'a str,
239            endpoints: &mut dyn Iterator<Item = dquic::qresolve::EndpointAddr>,
240        ) -> dquic::qresolve::PublishFuture<'a> {
241            let _endpoints: Vec<_> = endpoints.collect();
242            async move { Ok(()) }.boxed()
243        }
244    }
245
246    #[tokio::test]
247    async fn lookup_before_set_returns_typed_uninitialized_error() {
248        let resolver: DeferredResolver<TestResolver> = DeferredResolver::new();
249
250        let error = match resolver.lookup_typed("example.test", "", None).await {
251            Ok(_) => panic!("uninitialized resolver must not resolve"),
252            Err(error) => error,
253        };
254
255        assert!(matches!(error, DeferredLookupError::Uninitialized));
256    }
257
258    #[tokio::test]
259    async fn resolve_trait_lookup_waits_until_set() {
260        let resolver = DeferredResolver::new();
261        let mut lookup = resolver.lookup("example.test", "", None);
262
263        assert!(
264            tokio::time::timeout(Duration::from_millis(10), &mut lookup)
265                .await
266                .is_err(),
267            "trait lookup must not fail fast before set"
268        );
269
270        resolver.set(TestResolver).expect("first set succeeds");
271
272        let mut stream = lookup.await.expect("lookup completes after set");
273        let (_source, endpoint) = stream.next().await.expect("forwarded endpoint");
274        assert_eq!(
275            endpoint,
276            EndpointAddr::direct("127.0.0.1:4433".parse().unwrap())
277        );
278    }
279
280    #[tokio::test]
281    async fn lookup_after_set_forwards_to_inner_resolver() {
282        let resolver = DeferredResolver::new();
283        resolver.set(TestResolver).expect("first set succeeds");
284
285        let mut stream = resolver
286            .lookup_typed("example.test", "", None)
287            .await
288            .unwrap();
289        let (_source, endpoint) = stream.next().await.expect("forwarded endpoint");
290
291        assert_eq!(
292            endpoint,
293            EndpointAddr::direct("127.0.0.1:4433".parse().unwrap())
294        );
295    }
296
297    #[tokio::test]
298    async fn publish_after_set_forwards_to_inner_resolver() {
299        let resolver = DeferredResolver::new();
300        resolver.set(TestResolver).expect("first set succeeds");
301
302        let mut endpoints = std::iter::empty();
303        resolver
304            .publish_typed("example.test", &mut endpoints)
305            .await
306            .unwrap();
307    }
308}