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