Skip to main content

sui_network/validator/
server.rs

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