1use std::{fmt, io};
2
3use dquic::qresolve::{
4 EndpointAddr, Family, Publish, PublishFuture, RecordStream, Resolve, ResolveFuture,
5};
6use futures::{FutureExt, future::BoxFuture};
7use snafu::{ResultExt, Snafu};
8use tokio::sync::{Notify, OnceCell};
9
10use crate::resolvers::endpoint_candidates::{
11 EndpointCandidateFuture, EndpointLookup, ResolveEndpointCandidates,
12};
13
14#[derive(Debug, Snafu)]
15#[snafu(module, visibility(pub))]
16pub enum DeferredLookupError {
17 #[snafu(display("deferred resolver has not been initialized"))]
18 Uninitialized,
19 #[snafu(display("deferred resolver lookup failed"))]
20 Lookup { source: io::Error },
21}
22
23#[derive(Debug, Snafu)]
24#[snafu(module, visibility(pub))]
25pub enum DeferredPublishError {
26 #[snafu(display("deferred resolver has not been initialized"))]
27 Uninitialized,
28 #[snafu(display("deferred resolver publish failed"))]
29 Publish { source: io::Error },
30}
31
32#[derive(Debug, Snafu)]
33#[snafu(module, visibility(pub))]
34pub enum SetDeferredResolverError {
35 #[snafu(display("deferred resolver has already been initialized"))]
36 AlreadyInitialized,
37}
38
39pub struct DeferredResolver<R> {
40 inner: OnceCell<R>,
41 initialized: Notify,
42}
43
44impl<R> fmt::Debug for DeferredResolver<R> {
45 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
46 f.debug_struct("DeferredResolver")
47 .field("initialized", &self.inner.get().is_some())
48 .finish()
49 }
50}
51
52impl<R> Default for DeferredResolver<R> {
53 fn default() -> Self {
54 Self::new()
55 }
56}
57
58impl<R> DeferredResolver<R> {
59 #[must_use]
60 pub fn new() -> Self {
61 Self {
62 inner: OnceCell::new(),
63 initialized: Notify::new(),
64 }
65 }
66
67 pub fn set(&self, resolver: R) -> Result<(), SetDeferredResolverError> {
68 if self.inner.set(resolver).is_err() {
69 return set_deferred_resolver_error::AlreadyInitializedSnafu.fail();
70 }
71 self.initialized.notify_waiters();
72 Ok(())
73 }
74
75 #[must_use]
76 pub fn get(&self) -> Option<&R> {
77 self.inner.get()
78 }
79
80 async fn wait(&self) -> &R {
81 loop {
82 let initialized = self.initialized.notified();
83 if let Some(resolver) = self.get() {
84 return resolver;
85 }
86 initialized.await;
87 }
88 }
89}
90
91impl<R> fmt::Display for DeferredResolver<R>
92where
93 R: fmt::Display,
94{
95 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
96 match self.inner.get() {
97 Some(resolver) => write!(f, "DeferredResolver({resolver})"),
98 None => f.write_str("DeferredResolver(uninitialized)"),
99 }
100 }
101}
102
103impl<R> DeferredResolver<R>
104where
105 R: Resolve + 'static,
106{
107 pub async fn lookup_typed(
108 &self,
109 hostname: &str,
110 servname: &str,
111 family: Option<Family>,
112 ) -> Result<RecordStream, DeferredLookupError> {
113 let Some(resolver) = self.get() else {
114 return deferred_lookup_error::UninitializedSnafu.fail();
115 };
116 resolver
117 .lookup(hostname, servname, family)
118 .await
119 .context(deferred_lookup_error::LookupSnafu)
120 }
121}
122
123impl<R> Resolve for DeferredResolver<R>
124where
125 R: Resolve + 'static,
126{
127 fn lookup<'a>(
128 &'a self,
129 hostname: &'a str,
130 servname: &'a str,
131 family: Option<Family>,
132 ) -> ResolveFuture<'a> {
133 async move { self.wait().await.lookup(hostname, servname, family).await }.boxed()
134 }
135}
136
137impl<R> ResolveEndpointCandidates for DeferredResolver<R>
138where
139 R: ResolveEndpointCandidates + 'static,
140{
141 fn lookup_endpoint_candidates<'a>(
142 &'a self,
143 name: &'a str,
144 lookup: EndpointLookup,
145 ) -> EndpointCandidateFuture<'a> {
146 async move {
147 self.wait()
148 .await
149 .lookup_endpoint_candidates(name, lookup)
150 .await
151 }
152 .boxed()
153 }
154}
155
156impl<R> DeferredResolver<R>
157where
158 R: Publish + 'static,
159{
160 pub fn publish_typed<'a>(
161 &'a self,
162 name: &'a str,
163 endpoints: &mut dyn Iterator<Item = EndpointAddr>,
164 ) -> BoxFuture<'a, Result<(), DeferredPublishError>> {
165 let endpoints: Vec<_> = endpoints.collect();
166 async move {
167 let Some(resolver) = self.get() else {
168 return deferred_publish_error::UninitializedSnafu.fail();
169 };
170 let mut endpoints = endpoints.into_iter();
171 resolver
172 .publish(name, &mut endpoints)
173 .await
174 .context(deferred_publish_error::PublishSnafu)
175 }
176 .boxed()
177 }
178}
179
180impl<R> Publish for DeferredResolver<R>
181where
182 R: Publish + 'static,
183{
184 fn publish<'a>(
185 &'a self,
186 name: &'a str,
187 endpoints: &mut dyn Iterator<Item = EndpointAddr>,
188 ) -> PublishFuture<'a> {
189 let endpoints: Vec<_> = endpoints.collect();
190 async move {
191 let resolver = self.wait().await;
192 let mut endpoints = endpoints.into_iter();
193 resolver.publish(name, &mut endpoints).await
194 }
195 .boxed()
196 }
197}
198
199#[cfg(test)]
200mod tests {
201 use std::{fmt, time::Duration};
202
203 use dquic::{
204 qbase::net::addr::EndpointAddr,
205 qresolve::{Publish, Resolve, Source},
206 };
207 use futures::{FutureExt, StreamExt};
208
209 use super::*;
210
211 #[derive(Debug)]
212 struct TestResolver;
213
214 impl fmt::Display for TestResolver {
215 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
216 f.write_str("test resolver")
217 }
218 }
219
220 impl Resolve for TestResolver {
221 fn lookup<'a>(
222 &'a self,
223 _hostname: &'a str,
224 _servname: &'a str,
225 _family: Option<Family>,
226 ) -> dquic::qresolve::ResolveFuture<'a> {
227 async move {
228 let endpoint = EndpointAddr::direct("127.0.0.1:4433".parse().unwrap());
229 Ok(futures::stream::iter([(Source::System, endpoint)]).boxed())
230 }
231 .boxed()
232 }
233 }
234
235 impl Publish for TestResolver {
236 fn publish<'a>(
237 &'a self,
238 _name: &'a str,
239 endpoints: &mut dyn Iterator<Item = dquic::qresolve::EndpointAddr>,
240 ) -> dquic::qresolve::PublishFuture<'a> {
241 let _endpoints: Vec<_> = endpoints.collect();
242 async move { Ok(()) }.boxed()
243 }
244 }
245
246 #[tokio::test]
247 async fn lookup_before_set_returns_typed_uninitialized_error() {
248 let resolver: DeferredResolver<TestResolver> = DeferredResolver::new();
249
250 let error = match resolver.lookup_typed("example.test", "", None).await {
251 Ok(_) => panic!("uninitialized resolver must not resolve"),
252 Err(error) => error,
253 };
254
255 assert!(matches!(error, DeferredLookupError::Uninitialized));
256 }
257
258 #[tokio::test]
259 async fn resolve_trait_lookup_waits_until_set() {
260 let resolver = DeferredResolver::new();
261 let mut lookup = resolver.lookup("example.test", "", None);
262
263 assert!(
264 tokio::time::timeout(Duration::from_millis(10), &mut lookup)
265 .await
266 .is_err(),
267 "trait lookup must not fail fast before set"
268 );
269
270 resolver.set(TestResolver).expect("first set succeeds");
271
272 let mut stream = lookup.await.expect("lookup completes after set");
273 let (_source, endpoint) = stream.next().await.expect("forwarded endpoint");
274 assert_eq!(
275 endpoint,
276 EndpointAddr::direct("127.0.0.1:4433".parse().unwrap())
277 );
278 }
279
280 #[tokio::test]
281 async fn lookup_after_set_forwards_to_inner_resolver() {
282 let resolver = DeferredResolver::new();
283 resolver.set(TestResolver).expect("first set succeeds");
284
285 let mut stream = resolver
286 .lookup_typed("example.test", "", None)
287 .await
288 .unwrap();
289 let (_source, endpoint) = stream.next().await.expect("forwarded endpoint");
290
291 assert_eq!(
292 endpoint,
293 EndpointAddr::direct("127.0.0.1:4433".parse().unwrap())
294 );
295 }
296
297 #[tokio::test]
298 async fn publish_after_set_forwards_to_inner_resolver() {
299 let resolver = DeferredResolver::new();
300 resolver.set(TestResolver).expect("first set succeeds");
301
302 let mut endpoints = std::iter::empty();
303 resolver
304 .publish_typed("example.test", &mut endpoints)
305 .await
306 .unwrap();
307 }
308}