diff --git a/Cargo.lock b/Cargo.lock index 602ddce..04025ef 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2000,9 +2000,15 @@ name = "summit-checkpointer" version = "0.1.0" dependencies = [ "anyhow", + "bytes", "chrono", "clap", "config", + "futures-util", + "http", + "http-body", + "http-body-util", + "hyper", "jsonrpsee", "reqwest", "serde", @@ -2012,6 +2018,7 @@ dependencies = [ "thiserror 1.0.69", "tokio", "tokio-util", + "tower", "tracing", "tracing-subscriber", ] diff --git a/Cargo.toml b/Cargo.toml index e412275..bee98fb 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -39,4 +39,13 @@ serde_cbor = "0.11" chrono = { version = "0.4", features = ["serde"] } # Graceful shutdown -tokio-util = { version = "0.7", features = ["rt"] } +tokio-util = { version = "0.7", features = ["rt", "io"] } + +# HTTP streaming / Tower service +hyper = { version = "1" } +http-body = "1" +http-body-util = "0.1" +http = "1" +bytes = "1" +tower = { version = "0.5", features = ["util"] } +futures-util = "0.3" diff --git a/src/main.rs b/src/main.rs index 93d3e57..99bd625 100644 --- a/src/main.rs +++ b/src/main.rs @@ -68,11 +68,13 @@ async fn main() -> Result<()> { tracing::info!("Block monitor initialized, starting main loop"); let addr: SocketAddr = format!("0.0.0.0:{}", cli.port).parse().unwrap(); - let rpc_handle = tokio::spawn(RpcServer::new().start_server(addr)); + let server_handle = RpcServer::new().start_server(addr).await; + // Run until cancelled monitor.run_until_cancelled(shutdown_token).await?; - rpc_handle.abort(); + server_handle.stop().expect("Failed to stop RPC server"); + server_handle.stopped().await; // Cleanup: save state tracing::info!("Saving final state..."); diff --git a/src/server/server.rs b/src/server/server.rs index 8da1bda..4968a4a 100644 --- a/src/server/server.rs +++ b/src/server/server.rs @@ -1,17 +1,23 @@ -use std::{fs, net::SocketAddr, path::Path}; +use std::{convert::Infallible, fs, net::SocketAddr, path::Path}; +use futures_util::StreamExt; +use http_body::Frame; +use http_body_util::StreamBody; +use hyper::body::Incoming; use jsonrpsee::{ core::{async_trait, RpcResult}, - server::ServerBuilder, + server::{serve_with_graceful_shutdown, stop_channel, HttpBody, ServerBuilder, ServerHandle}, types::{ErrorCode, ErrorObjectOwned}, }; -use tokio::io::AsyncWriteExt as _; +use tokio::{io::AsyncWriteExt as _, net::TcpListener}; +use tokio_util::io::ReaderStream; +use tower::Service; use tracing::info; use crate::server::api::CheckpointerRpcServer; pub const SNAPSHOT_FILE_PREFIX: &str = "epoch_"; -pub const DATA_DISK_DIR: &str = "/home/ubuntu/checkpoints"; +pub const DATA_DISK_DIR: &str = "/persistent/checkpoints"; pub struct RpcServer; @@ -21,17 +27,126 @@ impl RpcServer { Self } - pub async fn start_server(self, addr: SocketAddr) { - let server = ServerBuilder::default().build(addr).await.expect("Failed to start rpc"); + pub async fn start_server(self, addr: SocketAddr) -> ServerHandle { + let listener = TcpListener::bind(addr).await.expect("Failed to bind RPC server"); + let (stop_handle, server_handle) = stop_channel(); - let handle = server.start(self.into_rpc()); + let rpc_module = self.into_rpc(); + let svc_builder = ServerBuilder::default().to_service_builder(); - info!("JSON-RPC Server started at {}", addr); + tokio::spawn(async move { + loop { + let sock = tokio::select! { + res = listener.accept() => { + match res { + Ok((stream, _)) => stream, + Err(e) => { + tracing::error!("TCP accept error: {e}"); + continue; + } + } + } + _ = stop_handle.clone().shutdown() => break, + }; + + let rpc_module = rpc_module.clone(); + let svc_builder = svc_builder.clone(); + let conn_stop = stop_handle.clone(); + let shutdown_stop = stop_handle.clone(); + + let svc = tower::service_fn(move |req: http::Request| { + let rpc_module = rpc_module.clone(); + let stop_handle = conn_stop.clone(); + let svc_builder = svc_builder.clone(); + + async move { + if req.method() == http::Method::GET { + if let Some(epoch) = parse_snapshot_path(req.uri().path()) { + return Ok::<_, Infallible>(handle_snapshot_stream(epoch).await); + } + } - handle.stopped().await; + let mut jsonrpc_svc = svc_builder.build(rpc_module, stop_handle); + Ok(match jsonrpc_svc.call(req).await { + Ok(resp) => resp, + Err(e) => { + tracing::error!("JSON-RPC service error: {e}"); + http::Response::builder() + .status(http::StatusCode::INTERNAL_SERVER_ERROR) + .body(HttpBody::from(format!("Internal error: {e}"))) + .expect("response build") + } + }) + } + }); + + tokio::spawn(async move { + if let Err(e) = + serve_with_graceful_shutdown(sock, svc, shutdown_stop.shutdown()).await + { + tracing::error!("Connection error: {e}"); + } + }); + } + }); + + info!("JSON-RPC Server started at {}", addr); + server_handle } } +fn parse_snapshot_path(path: &str) -> Option { + path.strip_prefix("/snapshots/").and_then(|rest| rest.trim_end_matches('/').parse::().ok()) +} + +async fn handle_snapshot_stream(epoch: u64) -> http::Response { + let snapshot_path = format!( + "{DATA_DISK_DIR}/{SNAPSHOT_FILE_PREFIX}{epoch}/{SNAPSHOT_FILE_PREFIX}{epoch}.tar.gz", + ); + + let file = match tokio::fs::File::open(&snapshot_path).await { + Ok(f) => f, + Err(e) if e.kind() == std::io::ErrorKind::NotFound => { + return http::Response::builder() + .status(http::StatusCode::NOT_FOUND) + .body(HttpBody::from(format!("No snapshot for epoch {epoch}"))) + .expect("response build"); + } + Err(e) => { + return http::Response::builder() + .status(http::StatusCode::INTERNAL_SERVER_ERROR) + .body(HttpBody::from(format!("Failed to open snapshot: {e}"))) + .expect("response build"); + } + }; + + let metadata = match file.metadata().await { + Ok(m) => m, + Err(e) => { + return http::Response::builder() + .status(http::StatusCode::INTERNAL_SERVER_ERROR) + .body(HttpBody::from(format!("Failed to read file metadata: {e}"))) + .expect("response build"); + } + }; + let file_size = metadata.len(); + + let reader = ReaderStream::with_capacity(file, 1024 * 1024); + let stream = reader.map(|result| result.map(Frame::data)); + let body = HttpBody::new(StreamBody::new(stream)); + + http::Response::builder() + .status(http::StatusCode::OK) + .header(http::header::CONTENT_TYPE, "application/gzip") + .header(http::header::CONTENT_LENGTH, file_size) + .header( + http::header::CONTENT_DISPOSITION, + format!("attachment; filename=\"epoch_{epoch}.tar.gz\""), + ) + .body(body) + .expect("response build") +} + #[async_trait] impl CheckpointerRpcServer for RpcServer { /// Health check endpoint that returns "OK" if service is running @@ -78,8 +193,13 @@ impl CheckpointerRpcServer for RpcServer { /// Get an encrypted snapshot from this servers database async fn get_encrypted_snapshot(&self, epoch: u64) -> RpcResult> { + tracing::warn!( + epoch, + "get_encrypted_snapshot is deprecated; use GET /snapshots/{{epoch}} for streaming" + ); + let snapshot_path = format!( - "{DATA_DISK_DIR}/{SNAPSHOT_FILE_PREFIX}{epoch}/{SNAPSHOT_FILE_PREFIX}{epoch}.tar.lz4", + "{DATA_DISK_DIR}/{SNAPSHOT_FILE_PREFIX}{epoch}/{SNAPSHOT_FILE_PREFIX}{epoch}.tar.gz", ); if !fs::exists(&snapshot_path).unwrap_or_default() {