Skip to main content

ddns/resolvers/
deferred.rs

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