Skip to main content

iota_network_stack/
server.rs

1// Copyright (c) Mysten Labs, Inc.
2// Modifications Copyright (c) 2024 IOTA Stiftung
3// SPDX-License-Identifier: Apache-2.0
4
5use std::{
6    convert::Infallible,
7    num::NonZeroUsize,
8    task::{Context, Poll},
9};
10
11use eyre::{Result, eyre};
12use tokio_rustls::rustls::ServerConfig;
13use tonic::{
14    body::Body,
15    codegen::http::{HeaderValue, Request, Response},
16    server::NamedService,
17};
18use tower::{Layer, Service, ServiceBuilder};
19use tower_http::{
20    propagate_header::PropagateHeaderLayer, set_header::SetRequestHeaderLayer, trace::TraceLayer,
21};
22
23use crate::{
24    concurrency::ServiceConcurrencyLimit,
25    config::Config,
26    metrics::{
27        DefaultMetricsCallbackProvider, GRPC_ENDPOINT_PATH_HEADER, MetricsCallbackProvider,
28        MetricsHandler,
29    },
30    multiaddr::{Multiaddr, Protocol},
31};
32
33pub struct ServerBuilder<M: MetricsCallbackProvider = DefaultMetricsCallbackProvider> {
34    config: Config,
35    metrics_provider: M,
36    router: tonic::service::Routes,
37    health_reporter: tonic_health::server::HealthReporter,
38}
39
40impl<M: MetricsCallbackProvider> ServerBuilder<M> {
41    pub fn from_config(config: &Config, metrics_provider: M) -> Self {
42        let (health_reporter, health_service) = tonic_health::server::health_reporter();
43        let router = tonic::service::Routes::new(health_service);
44
45        Self {
46            config: config.to_owned(),
47            metrics_provider,
48            router,
49            health_reporter,
50        }
51    }
52
53    pub fn health_reporter(&self) -> tonic_health::server::HealthReporter {
54        self.health_reporter.clone()
55    }
56
57    /// Add a new service to this Server.
58    pub fn add_service<S>(mut self, svc: S) -> Self
59    where
60        S: Service<Request<Body>, Response = Response<Body>, Error = Infallible>
61            + NamedService
62            + Clone
63            + Send
64            + Sync
65            + 'static,
66        S::Future: Send + 'static,
67    {
68        self.router = self.router.add_service(svc);
69        self
70    }
71
72    /// Add a new service to this Server with its own concurrency limit,
73    /// enforced independently of every other service on this server.
74    ///
75    /// With `load_shed` enabled, requests over the limit are rejected
76    /// immediately with gRPC `RESOURCE_EXHAUSTED`; otherwise they wait for a
77    /// slot to free up.
78    pub fn add_service_with_concurrency_limit<S>(
79        mut self,
80        svc: S,
81        limit: NonZeroUsize,
82        load_shed: bool,
83    ) -> Self
84    where
85        S: Service<Request<Body>, Response = Response<Body>, Error = Infallible>
86            + NamedService
87            + Clone
88            + Send
89            + Sync
90            + 'static,
91        S::Future: Send + 'static,
92    {
93        self.router = self
94            .router
95            .add_service(ServiceConcurrencyLimit::new(svc, limit, load_shed));
96        self
97    }
98
99    pub async fn bind(self, addr: &Multiaddr, tls_config: Option<ServerConfig>) -> Result<Server> {
100        let http_config = self
101            .config
102            .http_config()
103            // Temporarily continue allowing clients to connection without TLS even when the server
104            // is configured with a tls_config
105            .allow_insecure(true);
106
107        let request_timeout = self.config.request_timeout;
108        let metrics_provider = self.metrics_provider;
109        let metrics = MetricsHandler::new(metrics_provider.clone());
110        let request_metrics = TraceLayer::new_for_grpc()
111            .on_request(metrics.clone())
112            .on_response(metrics.clone())
113            .on_failure(metrics);
114
115        fn add_path_to_request_header<T>(request: &Request<T>) -> Option<HeaderValue> {
116            let path = request.uri().path();
117            Some(HeaderValue::from_str(path).unwrap())
118        }
119
120        let limiting_layers = ServiceBuilder::new()
121            .option_layer(
122                self.config
123                    .load_shed
124                    .unwrap_or_default()
125                    .then_some(tower::load_shed::LoadShedLayer::new()),
126            )
127            .option_layer(
128                self.config
129                    .global_concurrency_limit
130                    .map(tower::limit::GlobalConcurrencyLimitLayer::new),
131            );
132
133        let route_layers = ServiceBuilder::new()
134            .map_request(|mut request: http::Request<_>| {
135                if let Some(connect_info) = request.extensions().get::<iota_http::ConnectInfo>() {
136                    let tonic_connect_info = tonic::transport::server::TcpConnectInfo {
137                        local_addr: Some(connect_info.local_addr),
138                        remote_addr: Some(connect_info.remote_addr),
139                    };
140                    request.extensions_mut().insert(tonic_connect_info);
141                }
142                request
143            })
144            .layer(RequestLifetimeLayer { metrics_provider })
145            .layer(SetRequestHeaderLayer::overriding(
146                GRPC_ENDPOINT_PATH_HEADER.clone(),
147                add_path_to_request_header,
148            ))
149            .layer(request_metrics)
150            .layer(PropagateHeaderLayer::new(GRPC_ENDPOINT_PATH_HEADER.clone()))
151            .layer_fn(move |service| {
152                crate::grpc_timeout::GrpcTimeout::new(service, request_timeout)
153            });
154
155        let mut builder = iota_http::Builder::new().config(http_config);
156
157        let has_tls = tls_config.is_some();
158        if let Some(tls_config) = tls_config {
159            builder = builder.tls_config(tls_config);
160        }
161
162        let server_handle = builder
163            .serve(
164                addr,
165                limiting_layers.service(self.router.into_axum_router().layer(route_layers)),
166            )
167            .map_err(|e| eyre!(e))?;
168
169        let mut local_addr = update_tcp_port_in_multiaddr(addr, server_handle.local_addr().port());
170        if has_tls {
171            local_addr = local_addr.rewrite_http_to_https();
172        }
173        Ok(Server {
174            server_handle,
175            local_addr,
176            health_reporter: self.health_reporter,
177        })
178    }
179}
180
181/// TLS server name to use for the public IOTA validator interface.
182pub const IOTA_TLS_SERVER_NAME: &str = "iota";
183
184pub struct Server {
185    server_handle: iota_http::ServerHandle,
186    local_addr: Multiaddr,
187    health_reporter: tonic_health::server::HealthReporter,
188}
189
190impl Server {
191    pub async fn serve(self) -> Result<(), tonic::transport::Error> {
192        self.server_handle.wait_for_shutdown().await;
193        Ok(())
194    }
195
196    pub fn trigger_shutdown(&self) {
197        self.server_handle.trigger_shutdown();
198    }
199
200    pub fn local_addr(&self) -> &Multiaddr {
201        &self.local_addr
202    }
203
204    pub fn health_reporter(&self) -> tonic_health::server::HealthReporter {
205        self.health_reporter.clone()
206    }
207
208    pub fn handle(&self) -> &iota_http::ServerHandle {
209        &self.server_handle
210    }
211}
212
213fn update_tcp_port_in_multiaddr(addr: &Multiaddr, port: u16) -> Multiaddr {
214    addr.replace(1, |protocol| {
215        if let Protocol::Tcp(_) = protocol {
216            Some(Protocol::Tcp(port))
217        } else {
218            panic!("expected tcp protocol at index 1");
219        }
220    })
221    .expect("tcp protocol at index 1")
222}
223
224#[cfg(test)]
225mod test {
226    use std::{
227        ops::Deref,
228        sync::{Arc, Mutex},
229        time::Duration,
230    };
231
232    use fastcrypto::{ed25519::Ed25519KeyPair, traits::KeyPair};
233    use tonic::Code;
234    use tonic_health::pb::{HealthCheckRequest, health_client::HealthClient};
235
236    use crate::{Multiaddr, config::Config, metrics::MetricsCallbackProvider};
237
238    #[tokio::test]
239    async fn test_metrics_layer_successful() {
240        #[derive(Clone)]
241        struct Metrics {
242            /// a flag to figure out whether the
243            /// on_request method has been called.
244            metrics_called: Arc<Mutex<bool>>,
245        }
246
247        impl MetricsCallbackProvider for Metrics {
248            fn on_request(&self, path: String) {
249                assert_eq!(path, "/grpc.health.v1.Health/Check");
250            }
251
252            fn on_response(
253                &self,
254                path: String,
255                _latency: Duration,
256                status: u16,
257                grpc_status_code: Code,
258            ) {
259                assert_eq!(path, "/grpc.health.v1.Health/Check");
260                assert_eq!(status, 200);
261                assert_eq!(grpc_status_code, Code::Ok);
262                let mut m = self.metrics_called.lock().unwrap();
263                *m = true
264            }
265        }
266
267        let metrics = Metrics {
268            metrics_called: Arc::new(Mutex::new(false)),
269        };
270
271        let address: Multiaddr = "/ip4/127.0.0.1/tcp/0/http".parse().unwrap();
272        let config = Config::new();
273        let keypair = Ed25519KeyPair::generate(&mut rand::thread_rng());
274
275        let server = config
276            .server_builder_with_metrics(metrics.clone())
277            .bind(
278                &address,
279                Some(iota_tls::create_rustls_server_config(
280                    keypair.copy().private(),
281                    "test".to_string(),
282                )),
283            )
284            .await
285            .unwrap();
286
287        let address = server.local_addr().to_owned();
288        let channel = config
289            .connect(
290                &address,
291                iota_tls::create_rustls_client_config(
292                    keypair.public().to_owned(),
293                    "test".to_string(),
294                    None,
295                ),
296            )
297            .await
298            .unwrap();
299        let mut client = HealthClient::new(channel);
300
301        client
302            .check(HealthCheckRequest {
303                service: "".to_owned(),
304            })
305            .await
306            .unwrap();
307
308        server.server_handle.shutdown().await;
309
310        assert!(metrics.metrics_called.lock().unwrap().deref());
311    }
312
313    #[tokio::test]
314    async fn test_metrics_layer_error() {
315        #[derive(Clone)]
316        struct Metrics {
317            /// a flag to figure out whether the
318            /// on_request method has been called.
319            metrics_called: Arc<Mutex<bool>>,
320        }
321
322        impl MetricsCallbackProvider for Metrics {
323            fn on_request(&self, path: String) {
324                assert_eq!(path, "/grpc.health.v1.Health/Check");
325            }
326
327            fn on_response(
328                &self,
329                path: String,
330                _latency: Duration,
331                status: u16,
332                grpc_status_code: Code,
333            ) {
334                assert_eq!(path, "/grpc.health.v1.Health/Check");
335                assert_eq!(status, 200);
336                // According to https://github.com/grpc/grpc/blob/master/doc/statuscodes.md#status-codes-and-their-use-in-grpc
337                // code 5 is not_found , which is what we expect to get in this case
338                assert_eq!(grpc_status_code, Code::NotFound);
339                let mut m = self.metrics_called.lock().unwrap();
340                *m = true
341            }
342        }
343
344        let metrics = Metrics {
345            metrics_called: Arc::new(Mutex::new(false)),
346        };
347
348        let address: Multiaddr = "/ip4/127.0.0.1/tcp/0/http".parse().unwrap();
349        let config = Config::new();
350        let keypair = Ed25519KeyPair::generate(&mut rand::thread_rng());
351
352        let server = config
353            .server_builder_with_metrics(metrics.clone())
354            .bind(
355                &address,
356                Some(iota_tls::create_rustls_server_config(
357                    keypair.copy().private(),
358                    "test".to_string(),
359                )),
360            )
361            .await
362            .unwrap();
363        let address = server.local_addr().to_owned();
364        let channel = config
365            .connect(
366                &address,
367                iota_tls::create_rustls_client_config(
368                    keypair.public().to_owned(),
369                    "test".to_string(),
370                    None,
371                ),
372            )
373            .await
374            .unwrap();
375        let mut client = HealthClient::new(channel);
376
377        // Call the healthcheck for a service that doesn't exist
378        // that should give us back an error with code 5 (not_found)
379        // https://github.com/grpc/grpc/blob/master/doc/statuscodes.md#status-codes-and-their-use-in-grpc
380        let _ = client
381            .check(HealthCheckRequest {
382                service: "non-existing-service".to_owned(),
383            })
384            .await;
385
386        server.server_handle.shutdown().await;
387
388        assert!(metrics.metrics_called.lock().unwrap().deref());
389    }
390
391    async fn test_multiaddr(address: Multiaddr) {
392        let config = Config::new();
393        let keypair = Ed25519KeyPair::generate(&mut rand::thread_rng());
394
395        let server_handle = config
396            .server_builder()
397            .bind(
398                &address,
399                Some(iota_tls::create_rustls_server_config(
400                    keypair.copy().private(),
401                    "test".to_string(),
402                )),
403            )
404            .await
405            .unwrap();
406        let address = server_handle.local_addr().to_owned();
407        let channel = config
408            .connect(
409                &address,
410                iota_tls::create_rustls_client_config(
411                    keypair.public().to_owned(),
412                    "test".to_string(),
413                    None,
414                ),
415            )
416            .await
417            .unwrap();
418        let mut client = HealthClient::new(channel);
419
420        client
421            .check(HealthCheckRequest {
422                service: "".to_owned(),
423            })
424            .await
425            .unwrap();
426
427        server_handle.server_handle.shutdown().await;
428    }
429
430    #[tokio::test]
431    async fn dns() {
432        let address: Multiaddr = "/dns/localhost/tcp/0/http".parse().unwrap();
433        test_multiaddr(address).await;
434    }
435
436    #[tokio::test]
437    async fn ip4() {
438        let address: Multiaddr = "/ip4/127.0.0.1/tcp/0/http".parse().unwrap();
439        test_multiaddr(address).await;
440    }
441
442    #[tokio::test]
443    async fn ip6() {
444        let address: Multiaddr = "/ip6/::1/tcp/0/http".parse().unwrap();
445        test_multiaddr(address).await;
446    }
447}
448
449#[derive(Clone)]
450struct RequestLifetimeLayer<M: MetricsCallbackProvider> {
451    metrics_provider: M,
452}
453
454impl<M: MetricsCallbackProvider, S> Layer<S> for RequestLifetimeLayer<M> {
455    type Service = RequestLifetime<M, S>;
456
457    fn layer(&self, inner: S) -> Self::Service {
458        RequestLifetime {
459            inner,
460            metrics_provider: self.metrics_provider.clone(),
461            path: None,
462        }
463    }
464}
465
466#[derive(Clone)]
467struct RequestLifetime<M: MetricsCallbackProvider, S> {
468    inner: S,
469    metrics_provider: M,
470    path: Option<String>,
471}
472
473impl<M: MetricsCallbackProvider, S, RequestBody> Service<Request<RequestBody>>
474    for RequestLifetime<M, S>
475where
476    S: Service<Request<RequestBody>>,
477{
478    type Response = S::Response;
479    type Error = S::Error;
480    type Future = S::Future;
481
482    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
483        self.inner.poll_ready(cx)
484    }
485
486    fn call(&mut self, request: Request<RequestBody>) -> Self::Future {
487        if self.path.is_none() {
488            let path = request.uri().path().to_string();
489            self.metrics_provider.on_start(&path);
490            self.path = Some(path);
491        }
492        self.inner.call(request)
493    }
494}
495
496impl<M: MetricsCallbackProvider, S> Drop for RequestLifetime<M, S> {
497    fn drop(&mut self) {
498        if let Some(path) = &self.path {
499            self.metrics_provider.on_drop(path)
500        }
501    }
502}