1use 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 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
149pub 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 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 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 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 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 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}