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