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