diff --git a/crates/starknet_transaction_prover/src/server/tls.rs b/crates/starknet_transaction_prover/src/server/tls.rs index c28f597c390..aa1b1f1f808 100644 --- a/crates/starknet_transaction_prover/src/server/tls.rs +++ b/crates/starknet_transaction_prover/src/server/tls.rs @@ -1,5 +1,6 @@ //! TLS helpers for serving JSON-RPC over HTTPS. +use std::future::Future; use std::net::SocketAddr; use std::path::Path; use std::sync::Arc; @@ -14,8 +15,10 @@ use jsonrpsee::server::{ ServerBuilder, ServerConfig, ServerHandle, + StopHandle, }; -use tokio::net::TcpListener; +use tokio::io::{AsyncRead, AsyncWrite}; +use tokio::net::{TcpListener, TcpStream}; use tokio_rustls::rustls::pki_types::pem::PemObject; use tokio_rustls::rustls::pki_types::{CertificateDer, PrivateKeyDer}; use tokio_rustls::rustls::ServerConfig as RustlsServerConfig; @@ -29,6 +32,10 @@ use tracing::warn; use crate::server::{HealthLayer, OhttpJsonrpseeLayer, RequestLogLayer, RequestSpanLayer}; +#[cfg(test)] +#[path = "tls_test.rs"] +mod tls_test; + /// Maximum time allowed for a TLS handshake before the connection is dropped. const TLS_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10); @@ -52,11 +59,6 @@ pub async fn start_tls_server( .max_connections(max_connections) .max_request_body_size(max_request_body_size) .build(); - // See `prover_http_middleware!` for the full layer-order rationale. - let svc_builder = ServerBuilder::default() - .set_config(server_config) - .set_http_middleware(prover_http_middleware!(cors_layer, ohttp_layer)) - .to_service_builder(); let listener = TcpListener::bind(addr) .await @@ -67,6 +69,67 @@ pub async fn start_tls_server( let methods: Methods = methods.into(); let (stop_handle, server_handle) = stop_channel(); + let prepare_stream = move |socket, remote_addr| { + let tls_acceptor = tls_acceptor.clone(); + async move { + match tokio::time::timeout(TLS_HANDSHAKE_TIMEOUT, tls_acceptor.accept(socket)).await { + Ok(Ok(stream)) => Some(stream), + Ok(Err(err)) => { + warn!( + remote_address = %remote_addr, + error = %err, + "TLS handshake failed" + ); + None + } + Err(_) => { + warn!( + remote_address = %remote_addr, + "TLS handshake timed out" + ); + None + } + } + } + }; + + spawn_accept_loop( + listener, + stop_handle, + methods, + server_config, + cors_layer, + ohttp_layer, + prepare_stream, + ); + + Ok((local_addr, server_handle)) +} + +/// Spawns the loop that accepts connections until `stop_handle` signals shutdown. +/// +/// `prepare_stream` turns an accepted socket into the stream to serve, or `None` to drop it. +/// Each connection's service holds a `StopHandle` clone, so `ServerHandle::stopped()` stays pending +/// until the in-flight requests finish. +fn spawn_accept_loop( + listener: TcpListener, + stop_handle: StopHandle, + methods: Methods, + server_config: ServerConfig, + cors_layer: Option, + ohttp_layer: Option, + prepare_stream: PrepareStream, +) where + PrepareStream: Fn(TcpStream, SocketAddr) -> PrepareStreamFuture + Clone + Send + 'static, + PrepareStreamFuture: Future> + Send, + ServedStream: AsyncRead + AsyncWrite + Unpin + Send + 'static, +{ + // See `prover_http_middleware!` for the full layer-order rationale. + let svc_builder = ServerBuilder::default() + .set_config(server_config) + .set_http_middleware(prover_http_middleware!(cors_layer, ohttp_layer)) + .to_service_builder(); + tokio::spawn(async move { loop { let accept_result = tokio::select! { @@ -82,49 +145,29 @@ pub async fn start_tls_server( } }; - let tls_acceptor = tls_acceptor.clone(); let stop_handle = stop_handle.clone(); let methods = methods.clone(); let svc_builder = svc_builder.clone(); + let prepare_stream = prepare_stream.clone(); tokio::spawn(async move { - let tls_stream = - match tokio::time::timeout(TLS_HANDSHAKE_TIMEOUT, tls_acceptor.accept(socket)) - .await - { - Ok(Ok(stream)) => stream, - Ok(Err(err)) => { - warn!( - remote_address = %remote_addr, - error = %err, - "TLS handshake failed" - ); - return; - } - Err(_) => { - warn!( - remote_address = %remote_addr, - "TLS handshake timed out" - ); - return; - } - }; + let Some(stream) = prepare_stream(socket, remote_addr).await else { + return; + }; let svc = svc_builder.build(methods, stop_handle.clone()); if let Err(err) = - serve_with_graceful_shutdown(tls_stream, svc, stop_handle.shutdown()).await + serve_with_graceful_shutdown(stream, svc, stop_handle.shutdown()).await { warn!( remote_address = %remote_addr, error = %err, - "HTTPS connection terminated with error" + "Connection terminated with error" ); } }); } }); - - Ok((local_addr, server_handle)) } /// Loads a certificate chain and private key from PEM files and builds a TLS acceptor. diff --git a/crates/starknet_transaction_prover/src/server/tls_test.rs b/crates/starknet_transaction_prover/src/server/tls_test.rs new file mode 100644 index 00000000000..475ed3bca31 --- /dev/null +++ b/crates/starknet_transaction_prover/src/server/tls_test.rs @@ -0,0 +1,81 @@ +//! Drives [`super::spawn_accept_loop`] over plain TCP, with an identity `prepare_stream` in place +//! of the TLS handshake so no certificate is needed. + +use std::sync::Arc; +use std::time::Duration; + +use jsonrpsee::core::RpcResult; +use jsonrpsee::server::{stop_channel, Methods, ServerConfig}; +use jsonrpsee::RpcModule; +use tokio::net::TcpListener; +use tokio::sync::Notify; + +use super::spawn_accept_loop; + +/// Lets the test hold a request inside the handler for as long as it needs. +struct ParkedHandler { + entered: Notify, + release: Notify, +} + +/// Handing `build` a different `StopHandle` in `spawn_accept_loop` would let `stopped()` resolve +/// mid-request, dropping the runtime on top of an in-flight proof, and would fail this test. +#[tokio::test] +async fn stopped_waits_for_an_in_flight_request_to_finish() { + let handler = Arc::new(ParkedHandler { entered: Notify::new(), release: Notify::new() }); + + let mut module = RpcModule::new(handler.clone()); + module + .register_async_method("test_park", |_params, handler, _extensions| async move { + handler.entered.notify_one(); + handler.release.notified().await; + RpcResult::Ok("released".to_string()) + }) + .expect("Failed to register the parking method"); + + let listener = TcpListener::bind("127.0.0.1:0").await.expect("Failed to bind test listener"); + let local_addr = listener.local_addr().expect("Failed to read the listener address"); + let (stop_handle, server_handle) = stop_channel(); + + spawn_accept_loop( + listener, + stop_handle, + Methods::from(module), + ServerConfig::builder().build(), + None, + None, + |socket, _remote_addr| async move { Some(socket) }, + ); + + let request = tokio::spawn(async move { + reqwest::Client::new() + .post(format!("http://{local_addr}")) + .header("content-type", "application/json") + .body(r#"{"jsonrpc":"2.0","id":1,"method":"test_park","params":[]}"#) + .send() + .await + .expect("Request failed") + .text() + .await + .expect("Failed to read the response body") + }); + + handler.entered.notified().await; + server_handle.stop().expect("Failed to stop the server"); + + // The accept loop has dropped its handle by now, so only the draining connection holds one. + let resolved_early = + tokio::time::timeout(Duration::from_millis(300), server_handle.clone().stopped()).await; + assert!(resolved_early.is_err(), "stopped() resolved while a request was still in flight"); + + handler.release.notify_one(); + let body = request.await.expect("Request task panicked"); + assert!( + body.contains("released"), + "the in-flight request did not complete across the stop: {body}" + ); + + tokio::time::timeout(Duration::from_secs(5), server_handle.stopped()) + .await + .expect("stopped() never resolved after the in-flight request finished"); +}