Skip to main content

miden_node_utils/grpc/
layers.rs

1use std::net::{IpAddr, SocketAddr};
2use std::task::{Context as TaskContext, Poll};
3
4use tower::{Layer, Service};
5use tower_governor::GovernorError;
6use tower_governor::key_extractor::{KeyExtractor, SmartIpKeyExtractor};
7
8/// The originating client IP, resolved by [`ResolveClientIpLayer`] and stored in a request's
9/// extensions.
10///
11/// gRPC handlers can read this via `ClientIp::from_extensions(request.extensions())` to obtain the
12/// load-balancer-aware client address without re-implementing IP extraction.
13#[derive(Debug, Clone, Copy, PartialEq, Eq)]
14pub struct ClientIp(pub IpAddr);
15
16impl ClientIp {
17    /// Returns the client IP resolved into `extensions` by [`ResolveClientIpLayer`], or `None` if
18    /// it could not be determined.
19    pub fn from_extensions(extensions: &http::Extensions) -> Option<IpAddr> {
20        extensions.get::<Self>().map(|ip| ip.0)
21    }
22}
23
24/// A [`tower::Layer`] that resolves the originating client IP and stores it in the request's
25/// extensions as [`ClientIp`].
26///
27/// IP resolution reuses [`GrpcIpExtractor`], so clients behind a load balancer or reverse proxy are
28/// identified by their forwarded IP (via `X-Forwarded-For` / `X-Real-Ip` / `Forwarded` headers),
29/// falling back to the peer address. Resolving once at the transport layer lets handlers read the
30/// result instead of re-deriving it.
31#[derive(Debug, Clone, Copy, Default)]
32pub struct ResolveClientIpLayer;
33
34impl<S> Layer<S> for ResolveClientIpLayer {
35    type Service = ResolveClientIp<S>;
36
37    fn layer(&self, inner: S) -> Self::Service {
38        ResolveClientIp { inner }
39    }
40}
41
42/// The service produced by [`ResolveClientIpLayer`].
43#[derive(Debug, Clone, Copy)]
44pub struct ResolveClientIp<S> {
45    inner: S,
46}
47
48impl<S, B> Service<http::Request<B>> for ResolveClientIp<S>
49where
50    S: Service<http::Request<B>>,
51{
52    type Response = S::Response;
53    type Error = S::Error;
54    type Future = S::Future;
55
56    fn poll_ready(&mut self, cx: &mut TaskContext<'_>) -> Poll<Result<(), Self::Error>> {
57        self.inner.poll_ready(cx)
58    }
59
60    fn call(&mut self, mut request: http::Request<B>) -> Self::Future {
61        if let Ok(ip) = GrpcIpExtractor::default().extract(&request) {
62            request.extensions_mut().insert(ClientIp(ip));
63        }
64        self.inner.call(request)
65    }
66}
67
68/// Wraps [`SmartIpKeyExtractor`] by providing a fallback to the client IP address provided by the
69/// gRPC transport.
70///
71/// [`SmartIpKeyExtractor`]'s own fallback of checking the peer IP directly fails because we are in
72/// a gRPC transport and not the typical `SocketAddr` as it expects.
73#[derive(Debug, Clone, Copy, PartialEq, Eq)]
74pub struct GrpcIpExtractor(SmartIpKeyExtractor);
75
76impl Default for GrpcIpExtractor {
77    fn default() -> Self {
78        Self(SmartIpKeyExtractor)
79    }
80}
81
82impl GrpcIpExtractor {
83    #[expect(clippy::result_large_err, reason = "this is a third party error type")]
84    fn extract_tonic_address<T>(
85        request: &http::Request<T>,
86    ) -> Result<<Self as KeyExtractor>::Key, GovernorError> {
87        request
88            .extensions()
89            .get::<tonic::transport::server::TcpConnectInfo>()
90            .and_then(tonic::transport::server::TcpConnectInfo::remote_addr)
91            .as_ref()
92            .map(SocketAddr::ip)
93            .ok_or(GovernorError::UnableToExtractKey)
94    }
95}
96
97impl KeyExtractor for GrpcIpExtractor {
98    type Key = IpAddr;
99
100    #[expect(clippy::result_large_err, reason = "error type is dictated by tower-governor")]
101    fn extract<T>(
102        &self,
103        request: &http::Request<T>,
104    ) -> Result<Self::Key, tower_governor::GovernorError> {
105        self.0.extract(request).or_else(|_| Self::extract_tonic_address(request))
106    }
107}