Skip to main content

ddns/resolvers/
weak.rs

1use std::{
2    fmt, io,
3    sync::{Arc, Weak},
4};
5
6use dquic::qresolve::{
7    EndpointAddr, Family, Publish, PublishFuture, RecordStream, Resolve, ResolveFuture,
8};
9use futures::{FutureExt, future::BoxFuture};
10use snafu::{ResultExt, Snafu};
11
12#[derive(Debug, Snafu)]
13#[snafu(module, visibility(pub))]
14pub enum WeakLookupError {
15    #[snafu(display("weak resolver target has been dropped"))]
16    Dropped,
17    #[snafu(display("weak resolver lookup failed"))]
18    Lookup { source: io::Error },
19}
20
21#[derive(Debug, Snafu)]
22#[snafu(module, visibility(pub))]
23pub enum WeakPublishError {
24    #[snafu(display("weak resolver target has been dropped"))]
25    Dropped,
26    #[snafu(display("weak resolver publish failed"))]
27    Publish { source: io::Error },
28}
29
30pub struct WeakResolver<R: ?Sized> {
31    inner: Weak<R>,
32}
33
34impl<R: ?Sized> fmt::Debug for WeakResolver<R> {
35    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
36        f.debug_struct("WeakResolver")
37            .field("alive", &self.inner.strong_count().gt(&0))
38            .finish()
39    }
40}
41
42impl<R: ?Sized> Clone for WeakResolver<R> {
43    fn clone(&self) -> Self {
44        Self {
45            inner: self.inner.clone(),
46        }
47    }
48}
49
50impl<R: ?Sized> WeakResolver<R> {
51    #[must_use]
52    pub fn new(inner: Weak<R>) -> Self {
53        Self { inner }
54    }
55
56    pub fn upgrade(&self) -> Result<Arc<R>, WeakLookupError> {
57        self.inner.upgrade().ok_or(WeakLookupError::Dropped)
58    }
59}
60
61impl<R: ?Sized> fmt::Display for WeakResolver<R>
62where
63    R: fmt::Display,
64{
65    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
66        match self.inner.upgrade() {
67            Some(resolver) => write!(f, "WeakResolver({resolver})"),
68            None => f.write_str("WeakResolver(dropped)"),
69        }
70    }
71}
72
73impl<R: ?Sized> WeakResolver<R>
74where
75    R: Resolve + 'static,
76{
77    pub async fn lookup_typed(
78        &self,
79        hostname: &str,
80        servname: &str,
81        family: Option<Family>,
82    ) -> Result<RecordStream, WeakLookupError> {
83        let resolver = self.upgrade()?;
84        resolver
85            .lookup(hostname, servname, family)
86            .await
87            .context(weak_lookup_error::LookupSnafu)
88    }
89}
90
91impl<R: ?Sized> Resolve for WeakResolver<R>
92where
93    R: Resolve + 'static,
94{
95    fn lookup<'a>(
96        &'a self,
97        hostname: &'a str,
98        servname: &'a str,
99        family: Option<Family>,
100    ) -> ResolveFuture<'a> {
101        async move {
102            self.lookup_typed(hostname, servname, family)
103                .await
104                .map_err(io::Error::other)
105        }
106        .boxed()
107    }
108}
109
110impl<R: ?Sized> WeakResolver<R>
111where
112    R: Publish + 'static,
113{
114    pub fn publish_typed<'a>(
115        &'a self,
116        name: &'a str,
117        endpoints: &mut dyn Iterator<Item = EndpointAddr>,
118    ) -> BoxFuture<'a, Result<(), WeakPublishError>> {
119        let endpoints: Vec<_> = endpoints.collect();
120        async move {
121            let Some(resolver) = self.inner.upgrade() else {
122                return weak_publish_error::DroppedSnafu.fail();
123            };
124            let mut endpoints = endpoints.into_iter();
125            resolver
126                .publish(name, &mut endpoints)
127                .await
128                .context(weak_publish_error::PublishSnafu)
129        }
130        .boxed()
131    }
132}
133
134impl<R: ?Sized> Publish for WeakResolver<R>
135where
136    R: Publish + 'static,
137{
138    fn publish<'a>(
139        &'a self,
140        name: &'a str,
141        endpoints: &mut dyn Iterator<Item = EndpointAddr>,
142    ) -> PublishFuture<'a> {
143        let endpoints: Vec<_> = endpoints.collect();
144        async move {
145            let mut endpoints = endpoints.into_iter();
146            self.publish_typed(name, &mut endpoints)
147                .await
148                .map_err(io::Error::other)
149        }
150        .boxed()
151    }
152}
153
154#[cfg(test)]
155mod tests {
156    use std::{fmt, sync::Arc};
157
158    use dquic::{
159        qbase::net::addr::EndpointAddr,
160        qresolve::{Publish, Resolve, Source},
161    };
162    use futures::{FutureExt, StreamExt};
163
164    use super::*;
165
166    #[derive(Debug)]
167    struct TestResolver;
168
169    impl fmt::Display for TestResolver {
170        fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
171            f.write_str("test resolver")
172        }
173    }
174
175    impl Resolve for TestResolver {
176        fn lookup<'a>(
177            &'a self,
178            _hostname: &'a str,
179            _servname: &'a str,
180            _family: Option<Family>,
181        ) -> dquic::qresolve::ResolveFuture<'a> {
182            async move {
183                let endpoint = EndpointAddr::direct("127.0.0.1:4433".parse().unwrap());
184                Ok(futures::stream::iter([(Source::System, endpoint)]).boxed())
185            }
186            .boxed()
187        }
188    }
189
190    impl Publish for TestResolver {
191        fn publish<'a>(
192            &'a self,
193            _name: &'a str,
194            endpoints: &mut dyn Iterator<Item = dquic::qresolve::EndpointAddr>,
195        ) -> dquic::qresolve::PublishFuture<'a> {
196            let _endpoints: Vec<_> = endpoints.collect();
197            async move { Ok(()) }.boxed()
198        }
199    }
200
201    #[tokio::test]
202    async fn lookup_after_target_drop_returns_typed_error() {
203        let strong = Arc::new(TestResolver);
204        let resolver = WeakResolver::new(Arc::downgrade(&strong));
205        drop(strong);
206
207        let error = match resolver.lookup_typed("example.test", "", None).await {
208            Ok(_) => panic!("dropped weak resolver must not resolve"),
209            Err(error) => error,
210        };
211
212        assert!(matches!(error, WeakLookupError::Dropped));
213    }
214
215    #[tokio::test]
216    async fn lookup_forwards_while_target_is_alive() {
217        let strong = Arc::new(TestResolver);
218        let resolver = WeakResolver::new(Arc::downgrade(&strong));
219
220        let mut stream = resolver
221            .lookup_typed("example.test", "", None)
222            .await
223            .unwrap();
224        let (_source, endpoint) = stream.next().await.expect("forwarded endpoint");
225
226        assert_eq!(
227            endpoint,
228            EndpointAddr::direct("127.0.0.1:4433".parse().unwrap())
229        );
230    }
231
232    #[tokio::test]
233    async fn publish_forwards_while_target_is_alive() {
234        let strong = Arc::new(TestResolver);
235        let resolver = WeakResolver::new(Arc::downgrade(&strong));
236
237        let mut endpoints = std::iter::empty();
238        resolver
239            .publish_typed("example.test", &mut endpoints)
240            .await
241            .unwrap();
242    }
243}