Skip to main content

ddns/resolvers/
deferred.rs

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