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}