diff --git a/ldk-server-client/src/client.rs b/ldk-server-client/src/client.rs index 87680c50..2ac06e35 100644 --- a/ldk-server-client/src/client.rs +++ b/ldk-server-client/src/client.rs @@ -612,6 +612,7 @@ impl LdkServerClient { body, buf: Vec::new(), trailers_checked: false, + terminated: false, _marker: std::marker::PhantomData, }) } @@ -701,6 +702,7 @@ pub struct GrpcStream { body: hyper::Body, buf: Vec, trailers_checked: bool, + terminated: bool, _marker: std::marker::PhantomData, } @@ -712,46 +714,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)); } } @@ -759,10 +768,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 { @@ -775,6 +784,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)) => { @@ -784,10 +799,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), - ))); + )); }, } @@ -893,6 +908,7 @@ mod tests { body, buf: Vec::new(), trailers_checked: false, + terminated: false, _marker: std::marker::PhantomData, }; @@ -912,6 +928,7 @@ mod tests { body, buf: Vec::new(), trailers_checked: false, + terminated: false, _marker: std::marker::PhantomData, }; @@ -924,6 +941,7 @@ mod tests { MAX_GRPC_STREAM_MESSAGE_LEN ) ); + assert!(stream.next_message().await.is_none()); } #[tokio::test] @@ -936,12 +954,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 06deb959..95e8579f 100644 --- a/ldk-server-grpc/src/grpc.rs +++ b/ldk-server-grpc/src/grpc.rs @@ -25,6 +25,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 { @@ -167,16 +169,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 } @@ -194,11 +205,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() } @@ -351,6 +359,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/Cargo.toml b/ldk-server/Cargo.toml index 554a7ba3..a1eaa76b 100644 --- a/ldk-server/Cargo.toml +++ b/ldk-server/Cargo.toml @@ -40,4 +40,5 @@ default = [] experimental-lsps2-support = [] [dev-dependencies] +tokio = { version = "1.38.0", default-features = false, features = ["test-util"] } futures-util = "0.3.31" diff --git a/ldk-server/src/main.rs b/ldk-server/src/main.rs index 18afa931..ccaa6a63 100644 --- a/ldk-server/src/main.rs +++ b/ldk-server/src/main.rs @@ -39,7 +39,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::macaroons::MacaroonStore; @@ -54,6 +54,9 @@ use crate::util::{systemd, write_new}; const LDK_NODE_POSTGRES_LOCK_FILE: &str = "ldk_node_postgres.lock"; pub(crate) 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")] @@ -397,6 +400,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(); @@ -712,7 +716,15 @@ fn main() { }, res = grpc_listener.accept() => { match res { - Ok((stream, _)) => { + Ok((stream, peer_addr)) => { + 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(&macaroon_store), @@ -720,17 +732,30 @@ fn main() { metrics_auth_header.clone(), event_sender.clone(), shutdown_rx.clone(), + peer_addr.ip(), ); 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 ec2af133..4a452b24 100644 --- a/ldk-server/src/service.rs +++ b/ldk-server/src/service.rs @@ -7,9 +7,12 @@ // You may not use this file except in accordance with one or both of these // licenses. +use std::collections::BTreeMap; use std::future::Future; +use std::net::IpAddr; use std::pin::Pin; -use std::sync::Arc; +use std::sync::{Arc, Mutex}; +use std::time::Duration; use http_body_util::{BodyExt, Limited}; use hyper::body::Incoming; @@ -104,9 +107,58 @@ 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; +// Share this limit across all connections from the same source IP, after +// authenticating the macaroon and checking method permissions. +const MAX_BODY_READS_PER_PEER: usize = 2; +// 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_LIMITER: RequestBodyLimiter = RequestBodyLimiter::new(); + +// Track only active reads; the map has at most MAX_CONCURRENT_BODY_READS entries. +struct RequestBodyLimiter { + active: Mutex>, +} + +impl RequestBodyLimiter { + const fn new() -> Self { + Self { active: Mutex::new(BTreeMap::new()) } + } + + fn try_acquire(&self, peer_ip: IpAddr) -> Result, GrpcStatus> { + // Treat IPv4 and its IPv4-mapped IPv6 representation as the same peer. + let peer_ip = peer_ip.to_canonical(); + let mut active = self.active.lock().unwrap(); + if active.get(&peer_ip).copied().unwrap_or(0) >= MAX_BODY_READS_PER_PEER + || active.values().sum::() >= MAX_CONCURRENT_BODY_READS + { + return Err(GrpcStatus::new(GRPC_STATUS_UNAVAILABLE, "Too many concurrent requests")); + } + *active.entry(peer_ip).or_default() += 1; + Ok(RequestBodyPermit { limiter: self, peer_ip }) + } +} + +struct RequestBodyPermit<'a> { + limiter: &'a RequestBodyLimiter, + peer_ip: IpAddr, +} + +impl Drop for RequestBodyPermit<'_> { + fn drop(&mut self) { + let mut active = self.limiter.active.lock().unwrap(); + let count = active.get_mut(&self.peer_ip).unwrap(); + *count -= 1; + if *count == 0 { + active.remove(&self.peer_ip); + } + } +} #[derive(Clone)] pub(crate) struct NodeService { + peer_ip: IpAddr, context: Arc, macaroon_store: Arc, metrics: Option>, @@ -119,10 +171,18 @@ impl NodeService { pub(crate) fn new( node: Arc, macaroon_store: Arc, metrics: Option>, metrics_auth_header: Option, event_sender: broadcast::Sender, - shutdown_rx: tokio::sync::watch::Receiver, + shutdown_rx: tokio::sync::watch::Receiver, peer_ip: IpAddr, ) -> Self { let context = Arc::new(Context { node }); - Self { context, macaroon_store, metrics, metrics_auth_header, event_sender, shutdown_rx } + Self { + context, + macaroon_store, + metrics, + metrics_auth_header, + event_sender, + shutdown_rx, + peer_ip, + } } } @@ -219,12 +279,15 @@ impl Service> for NodeService { let event_sender = self.event_sender.clone(); let shutdown_rx = self.shutdown_rx.clone(); let (request_parts, request_body) = req.into_parts(); + let peer_ip = self.peer_ip; let future: Self::Future = Box::pin(async move { let (issuer, body_bytes) = match read_authorized_request( &macaroon_store, &method, &request_parts.headers, request_body, + peer_ip, + &REQUEST_BODY_LIMITER, ) .await { @@ -576,7 +639,8 @@ fn validate_request_body_len( } async fn read_authorized_request( - store: &MacaroonStore, method: &str, headers: &HeaderMap, body: B, + store: &MacaroonStore, method: &str, headers: &HeaderMap, body: B, peer_ip: IpAddr, + limiter: &RequestBodyLimiter, ) -> Result<(Arc, bytes::Bytes), GrpcStatus> where B: hyper::body::Body, @@ -601,21 +665,39 @@ where _ => {}, } let content_length = request_content_length(headers)?; - let limited_body = Limited::new(body, MAX_BODY_SIZE); - let bytes = match limited_body.collect().await { - Ok(collected) => collected.to_bytes(), - Err(_) => { - return Err(GrpcStatus::new( - GRPC_STATUS_INVALID_ARGUMENT, - "Request body too large or failed to read", - )); - }, - }; - validate_request_body_len(content_length, bytes.len())?; + let bytes = read_request_body(body, content_length, peer_ip, limiter).await?; let info = store.finish_request(request, method, &bytes).map_err(ldk_error_to_grpc_status)?; Ok((info, bytes)) } +async fn read_request_body( + body: B, content_length: Option, peer_ip: IpAddr, limiter: &RequestBodyLimiter, +) -> Result +where + B: hyper::body::Body, + B::Error: Into>, +{ + let _permit = limiter.try_acquire(peer_ip)?; + tokio::time::timeout(REQUEST_BODY_TIMEOUT, async move { + let limited_body = Limited::new(body, MAX_BODY_SIZE); + let bytes = match limited_body.collect().await { + Ok(collected) => collected.to_bytes(), + Err(_) => { + return Err(GrpcStatus::new( + GRPC_STATUS_INVALID_ARGUMENT, + "Request body too large or failed to read", + )); + }, + }; + validate_request_body_len(content_length, bytes.len())?; + Ok(bytes) + }) + .await + .unwrap_or_else(|_| { + Err(GrpcStatus::new(GRPC_STATUS_UNAVAILABLE, "Timed out reading request body")) + }) +} + /// Map an `LdkServerError` to a `GrpcStatus`. pub(crate) fn ldk_error_to_grpc_status(e: LdkServerError) -> GrpcStatus { let code = match e.error_code { @@ -633,6 +715,135 @@ mod tests { use super::*; use crate::macaroons::test_util::{admin_token, bind_request, test_store}; + fn stalled_body( + ) -> impl hyper::body::Body { + http_body_util::StreamBody::new(futures_util::stream::pending::< + Result, std::convert::Infallible>, + >()) + } + + #[test] + fn test_request_body_sustained_peer_saturation() { + tokio::runtime::Builder::new_current_thread().enable_all().build().unwrap().block_on( + async { + tokio::time::pause(); + let limiter = RequestBodyLimiter::new(); + let (_directory, store) = test_store("body-saturation"); + let token = admin_token(&store); + let attacker = "192.0.2.1".parse().unwrap(); + let client = "192.0.2.2".parse().unwrap(); + + // Refill the attacker's slots after each deadline, as replacement + // streams on either existing or new connections would do. + for _ in 0..10 { + let mut stalled = Vec::new(); + for _ in 0..MAX_BODY_READS_PER_PEER { + let mut read = + Box::pin(read_request_body(stalled_body(), None, attacker, &limiter)); + assert!(futures_util::poll!(&mut read).is_pending()); + stalled.push(read); + } + for _ in 0..MAX_CONCURRENT_BODY_READS { + let err = read_request_body(stalled_body(), None, attacker, &limiter) + .await + .unwrap_err(); + assert_eq!(err.code, GRPC_STATUS_UNAVAILABLE); + assert_eq!(err.message, "Too many concurrent requests"); + } + + // A correctly signed request can still collect its body and + // authenticate while the attacker's bodies remain unfinished. + let body = encode_grpc_frame(&[]); + let timestamp = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_secs(); + let bound = bind_request(&token, GET_NODE_INFO_PATH, &body, timestamp); + let mut headers = HeaderMap::new(); + headers.insert("macaroon", bound.parse().unwrap()); + let (_, received) = read_authorized_request( + &store, + GET_NODE_INFO_PATH, + &headers, + http_body_util::Full::new(body.clone()), + client, + &limiter, + ) + .await + .unwrap(); + assert_eq!(received, body); + + tokio::time::advance(REQUEST_BODY_TIMEOUT).await; + for read in stalled { + let err = read.await.unwrap_err(); + assert_eq!(err.code, GRPC_STATUS_UNAVAILABLE); + assert_eq!(err.message, "Timed out reading request body"); + } + assert!(limiter.active.lock().unwrap().is_empty()); + } + }, + ); + } + + #[test] + fn test_request_body_global_limit_and_peer_cleanup() { + let limiter = RequestBodyLimiter::new(); + let mut permits = Vec::new(); + for i in 1..=MAX_CONCURRENT_BODY_READS { + permits.push(limiter.try_acquire(IpAddr::from([192, 0, 2, i as u8])).unwrap()); + } + let next_peer = "192.0.2.100".parse().unwrap(); + assert!(limiter.try_acquire(next_peer).is_err()); + assert_eq!(limiter.active.lock().unwrap().len(), MAX_CONCURRENT_BODY_READS); + permits.pop(); + permits.push(limiter.try_acquire(next_peer).unwrap()); + drop(permits); + assert!(limiter.active.lock().unwrap().is_empty()); + + let peer = "192.0.2.1".parse().unwrap(); + let mapped_peer = "::ffff:192.0.2.1".parse().unwrap(); + let _first = limiter.try_acquire(peer).unwrap(); + let _second = limiter.try_acquire(mapped_peer).unwrap(); + assert!(limiter.try_acquire(peer).is_err()); + assert!(limiter.try_acquire(mapped_peer).is_err()); + } + + #[test] + fn test_request_body_releases_slots_on_cancel_and_error() { + tokio::runtime::Builder::new_current_thread().enable_all().build().unwrap().block_on( + async { + let limiter = RequestBodyLimiter::new(); + let peer = "192.0.2.1".parse().unwrap(); + let mut read = Box::pin(read_request_body(stalled_body(), None, peer, &limiter)); + assert!(futures_util::poll!(&mut read).is_pending()); + drop(read); + assert!(limiter.active.lock().unwrap().is_empty()); + + let err = read_request_body( + http_body_util::Full::new(bytes::Bytes::from_static(b"short")), + Some(10), + peer, + &limiter, + ) + .await + .unwrap_err(); + assert_eq!(err.code, GRPC_STATUS_INVALID_ARGUMENT); + assert!(limiter.active.lock().unwrap().is_empty()); + + let err = read_request_body( + http_body_util::Full::new(bytes::Bytes::from(vec![0; MAX_BODY_SIZE + 1])), + None, + peer, + &limiter, + ) + .await + .unwrap_err(); + assert_eq!(err.code, GRPC_STATUS_INVALID_ARGUMENT); + assert!(limiter.active.lock().unwrap().is_empty()); + }, + ); + } + struct UnreadBody; impl hyper::body::Body for UnreadBody { type Data = bytes::Bytes; @@ -646,6 +857,8 @@ mod tests { #[tokio::test] async fn macaroon_request_clock_skew() { + let limiter = RequestBodyLimiter::new(); + let peer_ip = "192.0.2.200".parse().unwrap(); use ldk_server_grpc::grpc::GRPC_STATUS_OK; let (_directory, store) = test_store("http-clock-skew"); @@ -673,20 +886,30 @@ mod tests { let mut headers = HeaderMap::new(); headers.insert("macaroon", bound.parse().unwrap()); let body = http_body_util::Full::new(bytes::Bytes::copy_from_slice(&bytes)); - let status = - match read_authorized_request(&store, GET_NODE_INFO_PATH, &headers, body).await { - Ok((_, received)) => { - assert_eq!(received.as_ref(), &bytes); - GRPC_STATUS_OK - }, - Err(error) => error.code, - }; + let status = match read_authorized_request( + &store, + GET_NODE_INFO_PATH, + &headers, + body, + peer_ip, + &limiter, + ) + .await + { + Ok((_, received)) => { + assert_eq!(received.as_ref(), &bytes); + GRPC_STATUS_OK + }, + Err(error) => error.code, + }; assert_eq!(status, expected, "timestamp offset: {offset}"); } } #[tokio::test] async fn rejected_requests_do_not_poll_the_body() { + let limiter = RequestBodyLimiter::new(); + let peer_ip = "192.0.2.200".parse().unwrap(); let (_directory, store) = test_store("http-admission"); let token = admin_token(&store); let admin = store.authenticate(CREATE_MACAROON_PATH, Some(&token)).unwrap(); @@ -697,6 +920,10 @@ mod tests { let unknown = bind_request(&token, "UnmappedMethod", b"", timestamp); let stale = bind_request(&token, GET_NODE_INFO_PATH, b"", timestamp - 61); let wrong_method = bind_request(&token, GET_BALANCES_PATH, b"", timestamp); + // Authentication errors must take precedence even when all body slots are held. + let permits: Vec<_> = (1..=MAX_CONCURRENT_BODY_READS) + .map(|i| limiter.try_acquire(IpAddr::from([192, 0, 2, i as u8])).unwrap()) + .collect(); for (credential, method, expected) in [ (None, GET_NODE_INFO_PATH, GRPC_STATUS_UNAUTHENTICATED), (None, GET_PERMISSIONS_PATH, GRPC_STATUS_UNAUTHENTICATED), @@ -712,46 +939,74 @@ mod tests { headers.insert("macaroon", token.parse().unwrap()); } let error = - read_authorized_request(&store, method, &headers, UnreadBody).await.unwrap_err(); + read_authorized_request(&store, method, &headers, UnreadBody, peer_ip, &limiter) + .await + .unwrap_err(); assert_eq!(error.code, expected); } + drop(permits); let mut headers = HeaderMap::new(); let bound = bind_request(&reader.token, GET_NODE_INFO_PATH, b"request", timestamp); headers.insert("macaroon", bound.parse().unwrap()); let body = http_body_util::Full::new(bytes::Bytes::from_static(b"request")); let (_, bytes) = - read_authorized_request(&store, GET_NODE_INFO_PATH, &headers, body).await.unwrap(); + read_authorized_request(&store, GET_NODE_INFO_PATH, &headers, body, peer_ip, &limiter) + .await + .unwrap(); assert_eq!(bytes.as_ref(), b"request"); let changed_body = http_body_util::Full::new(bytes::Bytes::from_static(b"changed")); assert_eq!( - read_authorized_request(&store, GET_NODE_INFO_PATH, &headers, changed_body) - .await - .unwrap_err() - .code, + read_authorized_request( + &store, + GET_NODE_INFO_PATH, + &headers, + changed_body, + peer_ip, + &limiter + ) + .await + .unwrap_err() + .code, GRPC_STATUS_UNAUTHENTICATED ); // Authorized requests still have both declared and actual body-size limits. headers.insert("content-length", (MAX_BODY_SIZE + 1).to_string().parse().unwrap()); assert_eq!( - read_authorized_request(&store, GET_NODE_INFO_PATH, &headers, UnreadBody) - .await - .unwrap_err() - .code, + read_authorized_request( + &store, + GET_NODE_INFO_PATH, + &headers, + UnreadBody, + peer_ip, + &limiter + ) + .await + .unwrap_err() + .code, GRPC_STATUS_INVALID_ARGUMENT ); headers.remove("content-length"); let oversized = http_body_util::Full::new(bytes::Bytes::from(vec![0; MAX_BODY_SIZE + 1])); assert_eq!( - read_authorized_request(&store, GET_NODE_INFO_PATH, &headers, oversized) - .await - .unwrap_err() - .code, + read_authorized_request( + &store, + GET_NODE_INFO_PATH, + &headers, + oversized, + peer_ip, + &limiter + ) + .await + .unwrap_err() + .code, GRPC_STATUS_INVALID_ARGUMENT ); } #[tokio::test] async fn malformed_macaroon_headers_are_rejected_before_reading_the_body() { + let limiter = RequestBodyLimiter::new(); + let peer_ip = "192.0.2.200".parse().unwrap(); use hyper::header::HeaderValue; use ldk_server_macaroons::MAX_MACAROON_BYTES; @@ -768,6 +1023,8 @@ mod tests { GET_NODE_INFO_PATH, &headers, http_body_util::Full::new(bytes::Bytes::from_static(body)), + peer_ip, + &limiter, ) .await .unwrap(); @@ -789,15 +1046,24 @@ mod tests { if let Some(header) = header { headers.insert("macaroon", HeaderValue::from_bytes(&header).unwrap()); } - let error = read_authorized_request(&store, GET_NODE_INFO_PATH, &headers, UnreadBody) - .await - .unwrap_err(); + let error = read_authorized_request( + &store, + GET_NODE_INFO_PATH, + &headers, + UnreadBody, + peer_ip, + &limiter, + ) + .await + .unwrap_err(); assert_eq!(error.code, GRPC_STATUS_UNAUTHENTICATED, "header case: {case}"); } } #[tokio::test] async fn policy_expiry_during_body_read_is_rejected() { + let limiter = RequestBodyLimiter::new(); + let peer_ip = "192.0.2.200".parse().unwrap(); use std::time::{Duration, SystemTime, UNIX_EPOCH}; let (_directory, store) = test_store("http-expiry"); @@ -817,6 +1083,8 @@ mod tests { GET_NODE_INFO_PATH, &headers, http_body_util::Full::new(bytes.clone()), + peer_ip, + &limiter, ) .await .unwrap(); @@ -831,7 +1099,7 @@ mod tests { })); let error = tokio::time::timeout( Duration::from_secs(10), - read_authorized_request(&store, GET_NODE_INFO_PATH, &headers, body), + read_authorized_request(&store, GET_NODE_INFO_PATH, &headers, body, peer_ip, &limiter), ) .await .unwrap()