diff --git a/ldk-server-client/src/client.rs b/ldk-server-client/src/client.rs index 6d0d50d0..75dac350 100644 --- a/ldk-server-client/src/client.rs +++ b/ldk-server-client/src/client.rs @@ -527,6 +527,7 @@ impl LdkServerClient { body, buf: Vec::new(), trailers_checked: false, + terminated: false, _marker: std::marker::PhantomData, }) } @@ -615,6 +616,7 @@ pub struct GrpcStream { body: hyper::Body, buf: Vec, trailers_checked: bool, + terminated: bool, _marker: std::marker::PhantomData, } @@ -626,46 +628,53 @@ impl GrpcStream { /// /// Returns `None` if the stream has ended. pub async fn next_message(&mut self) -> Option> { + if self.terminated { + return None; + } + loop { // Try to decode a complete gRPC frame from the buffer if self.buf.len() >= GRPC_FRAME_HEADER_LEN { if self.buf[0] != 0 { - return Some(Err(LdkServerError::new( + return self.terminate_with_error(LdkServerError::new( InternalError, "gRPC stream compression is not supported", - ))); + )); } let msg_len = u32::from_be_bytes([self.buf[1], self.buf[2], self.buf[3], self.buf[4]]) as usize; if msg_len > MAX_GRPC_STREAM_MESSAGE_LEN { - return Some(Err(LdkServerError::new( + return self.terminate_with_error(LdkServerError::new( InternalError, format!( "gRPC stream message exceeds maximum size of {} bytes", MAX_GRPC_STREAM_MESSAGE_LEN ), - ))); + )); } let frame_len = match GRPC_FRAME_HEADER_LEN.checked_add(msg_len) { Some(frame_len) => frame_len, None => { - return Some(Err(LdkServerError::new( + return self.terminate_with_error(LdkServerError::new( InternalError, "gRPC stream frame length overflow", - ))); + )); }, }; if self.buf.len() >= frame_len { let proto_bytes = &self.buf[GRPC_FRAME_HEADER_LEN..frame_len]; - let result = M::decode(proto_bytes).map_err(|e| { - LdkServerError::new( - InternalError, - format!("Failed to decode gRPC stream message: {}", e), - ) - }); + let message = match M::decode(proto_bytes) { + Ok(message) => message, + Err(e) => { + return self.terminate_with_error(LdkServerError::new( + InternalError, + format!("Failed to decode gRPC stream message: {}", e), + )); + }, + }; self.buf.drain(..frame_len); - return Some(result); + return Some(Ok(message)); } } @@ -673,10 +682,10 @@ impl GrpcStream { match self.body.data().await { Some(Ok(chunk)) => self.buf.extend_from_slice(&chunk), Some(Err(e)) => { - return Some(Err(LdkServerError::new( + return self.terminate_with_error(LdkServerError::new( InternalError, format!("Failed to read gRPC stream: {}", e), - ))); + )); }, None => { if self.trailers_checked { @@ -689,6 +698,12 @@ impl GrpcStream { } } + fn terminate_with_error(&mut self, error: LdkServerError) -> Option> { + self.terminated = true; + self.buf.clear(); + Some(Err(error)) + } + async fn finish_stream(&mut self) -> Option> { match self.body.trailers().await { Ok(Some(trailers)) => { @@ -698,10 +713,10 @@ impl GrpcStream { }, Ok(None) => {}, Err(e) => { - return Some(Err(LdkServerError::new( + return self.terminate_with_error(LdkServerError::new( InternalError, format!("Failed to read gRPC stream trailers: {}", e), - ))); + )); }, } @@ -807,6 +822,7 @@ mod tests { body, buf: Vec::new(), trailers_checked: false, + terminated: false, _marker: std::marker::PhantomData, }; @@ -826,6 +842,7 @@ mod tests { body, buf: Vec::new(), trailers_checked: false, + terminated: false, _marker: std::marker::PhantomData, }; @@ -838,6 +855,7 @@ mod tests { MAX_GRPC_STREAM_MESSAGE_LEN ) ); + assert!(stream.next_message().await.is_none()); } #[tokio::test] @@ -850,12 +868,34 @@ mod tests { body, buf: Vec::new(), trailers_checked: false, + terminated: false, _marker: std::marker::PhantomData, }; let result = stream.next_message().await.unwrap().unwrap_err(); assert_eq!(result.error_code, InternalError); assert_eq!(result.message, "gRPC stream compression is not supported"); + assert!(stream.next_message().await.is_none()); + } + + #[tokio::test] + async fn test_event_stream_terminates_after_decode_error() { + let (mut sender, body) = Body::channel(); + sender.send_data(vec![0u8, 0, 0, 0, 1, 0xff].into()).await.unwrap(); + drop(sender); + + let mut stream: EventStream = GrpcStream { + body, + buf: Vec::new(), + trailers_checked: false, + terminated: false, + _marker: std::marker::PhantomData, + }; + + let result = stream.next_message().await.unwrap().unwrap_err(); + assert_eq!(result.error_code, InternalError); + assert!(result.message.starts_with("Failed to decode gRPC stream message:")); + assert!(stream.next_message().await.is_none()); } #[test] diff --git a/ldk-server-grpc/src/grpc.rs b/ldk-server-grpc/src/grpc.rs index 59d15764..e25e1f10 100644 --- a/ldk-server-grpc/src/grpc.rs +++ b/ldk-server-grpc/src/grpc.rs @@ -24,6 +24,8 @@ pub const GRPC_STATUS_INTERNAL: u32 = 13; pub const GRPC_STATUS_UNAVAILABLE: u32 = 14; pub const GRPC_STATUS_UNAUTHENTICATED: u32 = 16; +const MAX_GRPC_MESSAGE_HEADER_LEN: usize = 4 * 1024; + /// A gRPC status with code and human-readable message. #[derive(Debug)] pub struct GrpcStatus { @@ -166,16 +168,25 @@ fn ok_trailers() -> http::HeaderMap { trailers } +fn grpc_message_header_value(message: &str) -> Option { + if message.is_empty() || message.len() > MAX_GRPC_MESSAGE_HEADER_LEN { + return None; + } + + let encoded = percent_encode(message); + if encoded.len() > MAX_GRPC_MESSAGE_HEADER_LEN { + return None; + } + + http::HeaderValue::from_str(&encoded).ok() +} + /// Build trailers for a gRPC error response. fn error_trailers(status: &GrpcStatus) -> http::HeaderMap { let mut trailers = http::HeaderMap::with_capacity(2); trailers.insert("grpc-status", http::HeaderValue::from_str(&status.code.to_string()).unwrap()); - if !status.message.is_empty() { - // Percent-encode the message per gRPC spec. - let encoded = percent_encode(&status.message); - if let Ok(val) = http::HeaderValue::from_str(&encoded) { - trailers.insert("grpc-message", val); - } + if let Some(value) = grpc_message_header_value(&status.message) { + trailers.insert("grpc-message", value); } trailers } @@ -193,11 +204,8 @@ pub fn grpc_error_response(status: GrpcStatus) -> http::Response { .header("grpc-accept-encoding", "identity") .header("content-length", "0") .header("grpc-status", status.code.to_string()); - if !status.message.is_empty() { - let encoded = percent_encode(&status.message); - if let Ok(val) = http::HeaderValue::from_str(&encoded) { - builder = builder.header("grpc-message", val); - } + if let Some(value) = grpc_message_header_value(&status.message) { + builder = builder.header("grpc-message", value); } builder.body(GrpcBody::Empty).unwrap() } @@ -350,6 +358,26 @@ mod tests { assert_eq!(response.headers().get("content-length").unwrap(), "0"); } + #[test] + fn test_grpc_error_response_omits_oversized_message() { + let response = grpc_error_response(GrpcStatus::new( + GRPC_STATUS_INVALID_ARGUMENT, + "%".repeat(MAX_GRPC_MESSAGE_HEADER_LEN), + )); + + assert!(response.headers().get("grpc-message").is_none()); + } + + #[test] + fn test_error_trailers_omit_oversized_message() { + let trailers = error_trailers(&GrpcStatus::new( + GRPC_STATUS_INVALID_ARGUMENT, + "a".repeat(MAX_GRPC_MESSAGE_HEADER_LEN + 1), + )); + + assert!(trailers.get("grpc-message").is_none()); + } + #[test] fn test_decode_too_short() { assert!(decode_grpc_body(&[0, 0, 0]).is_err()); diff --git a/ldk-server/src/main.rs b/ldk-server/src/main.rs index 28ea10de..efdd5e0d 100644 --- a/ldk-server/src/main.rs +++ b/ldk-server/src/main.rs @@ -38,7 +38,7 @@ use prost::Message; use tokio::net::TcpListener; use tokio::select; use tokio::signal::unix::SignalKind; -use tokio::sync::broadcast; +use tokio::sync::{broadcast, Semaphore}; use crate::api::node_to_proto_custom_tlv; use crate::io::persist::paginated_kv_store::PaginatedKVStore; @@ -58,6 +58,9 @@ use crate::util::{systemd, write_new}; const API_KEY_FILE: &str = "api_key"; const FULL_VERSION: &str = concat!(env!("CARGO_PKG_VERSION"), " (", env!("GIT_HASH"), ")"); +const MAX_CONCURRENT_HTTP2_STREAMS: u32 = 32; +const MAX_PENDING_TLS_HANDSHAKES: usize = 64; +const TLS_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10); pub fn get_default_data_dir() -> Option { #[cfg(target_os = "macos")] @@ -369,6 +372,7 @@ fn main() { } }; let tls_acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(server_config)); + let tls_handshake_semaphore = Arc::new(Semaphore::new(MAX_PENDING_TLS_HANDSHAKES)); info!("gRPC service listening on {}", config_file.grpc_service_addr); systemd::notify_ready(); @@ -642,6 +646,14 @@ fn main() { res = grpc_listener.accept() => { match res { Ok((stream, _)) => { + let handshake_permit = + match Arc::clone(&tls_handshake_semaphore).try_acquire_owned() { + Ok(permit) => permit, + Err(_) => { + debug!("TLS handshake limit reached, rejecting connection"); + continue; + }, + }; let node_service = NodeService::new( Arc::clone(&node), Arc::clone(&paginated_store), @@ -653,14 +665,26 @@ fn main() { ); let acceptor = tls_acceptor.clone(); runtime.spawn(async move { - match acceptor.accept(stream).await { - Ok(tls_stream) => { + match tokio::time::timeout( + TLS_HANDSHAKE_TIMEOUT, + acceptor.accept(stream), + ) + .await + { + Ok(Ok(tls_stream)) => { + // Only the handshake holds a slot. Holding it for the whole + // connection would let an unauthenticated peer block new + // connections by keeping established ones idle. + drop(handshake_permit); let io_stream = TokioIo::new(tls_stream); - if let Err(err) = http2::Builder::new(TokioExecutor::new()).serve_connection(io_stream, node_service).await { + let mut builder = http2::Builder::new(TokioExecutor::new()); + builder.max_concurrent_streams(MAX_CONCURRENT_HTTP2_STREAMS); + if let Err(err) = builder.serve_connection(io_stream, node_service).await { error!("Failed to serve TLS connection: {err}"); } }, - Err(e) => error!("TLS handshake failed: {e}"), + Ok(Err(e)) => error!("TLS handshake failed: {e}"), + Err(_) => debug!("TLS handshake timed out"), } }); }, diff --git a/ldk-server/src/service.rs b/ldk-server/src/service.rs index 83c6fb7e..cdd0f201 100644 --- a/ldk-server/src/service.rs +++ b/ldk-server/src/service.rs @@ -10,6 +10,7 @@ use std::future::Future; use std::pin::Pin; use std::sync::Arc; +use std::time::Duration; use http_body_util::{BodyExt, Limited}; use hyper::body::Incoming; @@ -40,7 +41,7 @@ use ldk_server_grpc::grpc::{ GRPC_STATUS_UNAUTHENTICATED, GRPC_STATUS_UNAVAILABLE, GRPC_STATUS_UNIMPLEMENTED, }; use prost::Message; -use tokio::sync::{broadcast, mpsc}; +use tokio::sync::{broadcast, mpsc, Semaphore}; use crate::api::bolt11_claim_for_hash::handle_bolt11_claim_for_hash_request; use crate::api::bolt11_fail_for_hash::handle_bolt11_fail_for_hash_request; @@ -88,6 +89,11 @@ const GRPC_SERVICE_PREFIX: &str = "/api.LightningNode/"; // Maximum request body size: 10 MB const MAX_BODY_SIZE: usize = 10 * 1024 * 1024; +const MAX_CONCURRENT_BODY_READS: usize = 8; +// A client that stalls mid-body would otherwise hold one of the few body-read slots +// indefinitely. +const REQUEST_BODY_TIMEOUT: Duration = Duration::from_secs(30); +static REQUEST_BODY_SEMAPHORE: Semaphore = Semaphore::const_new(MAX_CONCURRENT_BODY_READS); #[derive(Clone)] pub(crate) struct NodeService { @@ -217,6 +223,10 @@ impl Service> for NodeService { if let Err(status) = validate_grpc_request(&req) { return Box::pin(async move { Ok(grpc_error_response(status)) }); } + if req.headers().get("x-auth").is_none() { + let status = GrpcStatus::new(GRPC_STATUS_UNAUTHENTICATED, "Missing x-auth metadata"); + return Box::pin(async move { Ok(grpc_error_response(status)) }); + } let context = Arc::clone(&self.context); let path = req.uri().path().to_string(); @@ -257,14 +267,35 @@ impl Service> for NodeService { let shutdown_rx = self.shutdown_rx.clone(); let (request_parts, request_body) = req.into_parts(); let future: Self::Future = Box::pin(async move { + let body_permit = match REQUEST_BODY_SEMAPHORE.try_acquire() { + Ok(permit) => permit, + Err(_) => { + return Ok(grpc_error_response(GrpcStatus::new( + GRPC_STATUS_UNAVAILABLE, + "Too many concurrent requests", + ))); + }, + }; let content_length = match request_content_length(&request_parts.headers) { Ok(content_length) => content_length, Err(status) => return Ok(grpc_error_response(status)), }; - let body_bytes = match read_request_body(request_body, content_length).await { - Ok(bytes) => bytes, - Err(status) => return Ok(grpc_error_response(status)), + let body_bytes = match tokio::time::timeout( + REQUEST_BODY_TIMEOUT, + read_request_body(request_body, content_length), + ) + .await + { + Ok(Ok(bytes)) => bytes, + Ok(Err(status)) => return Ok(grpc_error_response(status)), + Err(_) => { + return Ok(grpc_error_response(GrpcStatus::new( + GRPC_STATUS_UNAVAILABLE, + "Timed out reading request body", + ))); + }, }; + drop(body_permit); let auth_req = Request::from_parts(request_parts, ()); if let Err(e) = validate_auth(&auth_req, &api_key, &body_bytes) {