Skip to main content

ddns/resolvers/
weak.rs

1use std::{
2    fmt, io,
3    sync::{Arc, Weak},
4};
5
6use dquic::qresolve::{Publish, PublishFuture, RecordStream, Resolve, ResolveFuture};
7use futures::FutureExt;
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 async fn publish_typed(&self, name: &str, packet: &[u8]) -> Result<(), WeakPublishError> {
98        let Some(resolver) = self.inner.upgrade() else {
99            return weak_publish_error::DroppedSnafu.fail();
100        };
101        resolver
102            .publish(name, packet)
103            .await
104            .context(weak_publish_error::PublishSnafu)
105    }
106}
107
108impl<R: ?Sized> Publish for WeakResolver<R>
109where
110    R: Publish + 'static,
111{
112    fn publish<'a>(&'a self, name: &'a str, packet: &'a [u8]) -> PublishFuture<'a> {
113        async move {
114            self.publish_typed(name, packet)
115                .await
116                .map_err(io::Error::other)
117        }
118        .boxed()
119    }
120}
121
122#[cfg(test)]
123mod tests {
124    use std::{fmt, sync::Arc};
125
126    use dquic::{
127        qbase::net::addr::EndpointAddr,
128        qresolve::{Publish, Resolve, Source},
129    };
130    use futures::{FutureExt, StreamExt};
131
132    use super::*;
133
134    #[derive(Debug)]
135    struct TestResolver;
136
137    impl fmt::Display for TestResolver {
138        fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
139            f.write_str("test resolver")
140        }
141    }
142
143    impl Resolve for TestResolver {
144        fn lookup<'a>(&'a self, _name: &'a str) -> dquic::qresolve::ResolveFuture<'a> {
145            async move {
146                let endpoint = EndpointAddr::direct("127.0.0.1:4433".parse().unwrap());
147                Ok(futures::stream::iter([(Source::System, endpoint)]).boxed())
148            }
149            .boxed()
150        }
151    }
152
153    impl Publish for TestResolver {
154        fn publish<'a>(
155            &'a self,
156            _name: &'a str,
157            _packet: &'a [u8],
158        ) -> dquic::qresolve::PublishFuture<'a> {
159            async move { Ok(()) }.boxed()
160        }
161    }
162
163    #[tokio::test]
164    async fn lookup_after_target_drop_returns_typed_error() {
165        let strong = Arc::new(TestResolver);
166        let resolver = WeakResolver::new(Arc::downgrade(&strong));
167        drop(strong);
168
169        let error = match resolver.lookup_typed("example.test").await {
170            Ok(_) => panic!("dropped weak resolver must not resolve"),
171            Err(error) => error,
172        };
173
174        assert!(matches!(error, WeakLookupError::Dropped));
175    }
176
177    #[tokio::test]
178    async fn lookup_forwards_while_target_is_alive() {
179        let strong = Arc::new(TestResolver);
180        let resolver = WeakResolver::new(Arc::downgrade(&strong));
181
182        let mut stream = resolver.lookup_typed("example.test").await.unwrap();
183        let (_source, endpoint) = stream.next().await.expect("forwarded endpoint");
184
185        assert_eq!(
186            endpoint,
187            EndpointAddr::direct("127.0.0.1:4433".parse().unwrap())
188        );
189    }
190
191    #[tokio::test]
192    async fn publish_forwards_while_target_is_alive() {
193        let strong = Arc::new(TestResolver);
194        let resolver = WeakResolver::new(Arc::downgrade(&strong));
195
196        resolver
197            .publish_typed("example.test", b"packet")
198            .await
199            .unwrap();
200    }
201}