Skip to main content

rama_http/layer/dns/dns_resolve/
layer.rs

1use super::DnsResolveModeService;
2use crate::HeaderName;
3use rama_core::Layer;
4
5/// Layer which can extend `Dns` (see `rama_core`) overwrites with mappings.
6///
7/// See [the module level documentation](crate::layer::dns) for more information.
8#[derive(Debug, Clone)]
9pub struct DnsResolveModeLayer {
10    header_name: HeaderName,
11}
12
13impl DnsResolveModeLayer {
14    /// Creates a new [`DnsResolveModeLayer`].
15    pub const fn new(name: HeaderName) -> Self {
16        Self { header_name: name }
17    }
18}
19
20impl<S> Layer<S> for DnsResolveModeLayer {
21    type Service = DnsResolveModeService<S>;
22
23    fn layer(&self, inner: S) -> Self::Service {
24        DnsResolveModeService::new(inner, self.header_name.clone())
25    }
26
27    fn into_layer(self, inner: S) -> Self::Service {
28        DnsResolveModeService::new(inner, self.header_name)
29    }
30}
31
32#[cfg(test)]
33mod tests {
34    use super::*;
35    use crate::{Request, layer::dns::DnsResolveMode};
36    use rama_core::{Service, extensions::ExtensionsRef, service::service_fn};
37    use std::convert::Infallible;
38
39    #[tokio::test]
40    async fn test_dns_resolve_mode_layer() {
41        let svc = DnsResolveModeLayer::new(HeaderName::from_static("x-dns-resolve")).into_layer(
42            service_fn(async |req: Request<()>| {
43                assert_eq!(
44                    req.extensions().get_ref::<DnsResolveMode>().unwrap(),
45                    &DnsResolveMode::eager()
46                );
47                Ok::<_, Infallible>(())
48            }),
49        );
50
51        let req = Request::builder()
52            .header("x-dns-resolve", "eager")
53            .uri("http://example.com")
54            .body(())
55            .unwrap();
56
57        svc.serve(req).await.unwrap();
58    }
59}