Skip to main content

ddns/resolvers/
weak.rs

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