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::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}
36
37impl<R> fmt::Debug for DeferredResolver<R> {
38    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
39        f.debug_struct("DeferredResolver")
40            .field("initialized", &self.inner.get().is_some())
41            .finish()
42    }
43}
44
45impl<R> Default for DeferredResolver<R> {
46    fn default() -> Self {
47        Self::new()
48    }
49}
50
51impl<R> DeferredResolver<R> {
52    #[must_use]
53    pub fn new() -> Self {
54        Self {
55            inner: OnceCell::new(),
56        }
57    }
58
59    pub fn set(&self, resolver: R) -> Result<(), SetDeferredResolverError> {
60        if self.inner.set(resolver).is_err() {
61            return set_deferred_resolver_error::AlreadyInitializedSnafu.fail();
62        }
63        Ok(())
64    }
65
66    #[must_use]
67    pub fn get(&self) -> Option<&R> {
68        self.inner.get()
69    }
70}
71
72impl<R> fmt::Display for DeferredResolver<R>
73where
74    R: fmt::Display,
75{
76    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
77        match self.inner.get() {
78            Some(resolver) => write!(f, "DeferredResolver({resolver})"),
79            None => f.write_str("DeferredResolver(uninitialized)"),
80        }
81    }
82}
83
84impl<R> DeferredResolver<R>
85where
86    R: Resolve + 'static,
87{
88    pub async fn lookup_typed(&self, name: &str) -> Result<RecordStream, DeferredLookupError> {
89        let Some(resolver) = self.get() else {
90            return deferred_lookup_error::UninitializedSnafu.fail();
91        };
92        resolver
93            .lookup(name)
94            .await
95            .context(deferred_lookup_error::LookupSnafu)
96    }
97}
98
99impl<R> Resolve for DeferredResolver<R>
100where
101    R: Resolve + 'static,
102{
103    fn lookup<'a>(&'a self, name: &'a str) -> ResolveFuture<'a> {
104        async move { self.lookup_typed(name).await.map_err(io::Error::other) }.boxed()
105    }
106}
107
108impl<R> DeferredResolver<R>
109where
110    R: Publish + 'static,
111{
112    pub async fn publish_typed(
113        &self,
114        name: &str,
115        packet: &[u8],
116    ) -> Result<(), DeferredPublishError> {
117        let Some(resolver) = self.get() else {
118            return deferred_publish_error::UninitializedSnafu.fail();
119        };
120        resolver
121            .publish(name, packet)
122            .await
123            .context(deferred_publish_error::PublishSnafu)
124    }
125}
126
127impl<R> Publish for DeferredResolver<R>
128where
129    R: Publish + 'static,
130{
131    fn publish<'a>(&'a self, name: &'a str, packet: &'a [u8]) -> PublishFuture<'a> {
132        async move {
133            self.publish_typed(name, packet)
134                .await
135                .map_err(io::Error::other)
136        }
137        .boxed()
138    }
139}
140
141#[cfg(test)]
142mod tests {
143    use std::fmt;
144
145    use dquic::{
146        qbase::net::addr::EndpointAddr,
147        qresolve::{Publish, Resolve, Source},
148    };
149    use futures::{FutureExt, StreamExt};
150
151    use super::*;
152
153    #[derive(Debug)]
154    struct TestResolver;
155
156    impl fmt::Display for TestResolver {
157        fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
158            f.write_str("test resolver")
159        }
160    }
161
162    impl Resolve for TestResolver {
163        fn lookup<'a>(&'a self, _name: &'a str) -> dquic::qresolve::ResolveFuture<'a> {
164            async move {
165                let endpoint = EndpointAddr::direct("127.0.0.1:4433".parse().unwrap());
166                Ok(futures::stream::iter([(Source::System, endpoint)]).boxed())
167            }
168            .boxed()
169        }
170    }
171
172    impl Publish for TestResolver {
173        fn publish<'a>(
174            &'a self,
175            _name: &'a str,
176            _packet: &'a [u8],
177        ) -> dquic::qresolve::PublishFuture<'a> {
178            async move { Ok(()) }.boxed()
179        }
180    }
181
182    #[tokio::test]
183    async fn lookup_before_set_returns_typed_uninitialized_error() {
184        let resolver: DeferredResolver<TestResolver> = DeferredResolver::new();
185
186        let error = match resolver.lookup_typed("example.test").await {
187            Ok(_) => panic!("uninitialized resolver must not resolve"),
188            Err(error) => error,
189        };
190
191        assert!(matches!(error, DeferredLookupError::Uninitialized));
192    }
193
194    #[tokio::test]
195    async fn lookup_after_set_forwards_to_inner_resolver() {
196        let resolver = DeferredResolver::new();
197        resolver.set(TestResolver).expect("first set succeeds");
198
199        let mut stream = resolver.lookup_typed("example.test").await.unwrap();
200        let (_source, endpoint) = stream.next().await.expect("forwarded endpoint");
201
202        assert_eq!(
203            endpoint,
204            EndpointAddr::direct("127.0.0.1:4433".parse().unwrap())
205        );
206    }
207
208    #[tokio::test]
209    async fn publish_after_set_forwards_to_inner_resolver() {
210        let resolver = DeferredResolver::new();
211        resolver.set(TestResolver).expect("first set succeeds");
212
213        resolver
214            .publish_typed("example.test", b"packet")
215            .await
216            .unwrap();
217    }
218}