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