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