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}