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}