Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
107 changes: 75 additions & 32 deletions crates/starknet_transaction_prover/src/server/tls.rs
Original file line number Diff line number Diff line change
@@ -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;
Expand All @@ -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;
Expand All @@ -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);

Expand All @@ -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
Expand All @@ -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<PrepareStream, PrepareStreamFuture, ServedStream>(
listener: TcpListener,
stop_handle: StopHandle,
methods: Methods,
server_config: ServerConfig,
cors_layer: Option<CorsLayer>,
ohttp_layer: Option<OhttpJsonrpseeLayer>,
prepare_stream: PrepareStream,
) where
PrepareStream: Fn(TcpStream, SocketAddr) -> PrepareStreamFuture + Clone + Send + 'static,
PrepareStreamFuture: Future<Output = Option<ServedStream>> + 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! {
Expand All @@ -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.
Expand Down
81 changes: 81 additions & 0 deletions crates/starknet_transaction_prover/src/server/tls_test.rs
Original file line number Diff line number Diff line change
@@ -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");
}