use std::panic::AssertUnwindSafe;
use std::sync::{Arc, Weak};
use futures_util::stream::{FuturesUnordered, StreamExt};
use futures_util::FutureExt;
use tokio::sync::mpsc;
use tokio::sync::Mutex;
use crate::engine::InstantMillis;
use crate::identity::IdentityHash;
use crate::routing::links::request::RequestId;
use crate::routing::links::LinkId;
use crate::routing::request_handlers::RequestPathHash;
use crate::units::RttMillis;
use crate::wire::DestinationHash;
use super::node_facade::PrnsNodeHandle;
use super::request_endpoints::{dispatch_request, Decline, InboundRequest, RequestEndpointSet};
use super::request_endpoints::{ResponseCapacityExceeded, ResponseSink};
pub(super) const REQUEST_QUEUE_DEPTH: usize = 1024;
const MAX_IN_FLIGHT: usize = 256;
pub(super) struct RunnerRequest {
pub destination: DestinationHash,
pub link_id: LinkId,
pub request_id: RequestId,
pub requester: Option<IdentityHash>,
pub path_hash: RequestPathHash,
pub requested_at: InstantMillis,
pub rtt: RttMillis,
pub data: std::vec::Vec<u8>,
}
enum RunnerResponse {
Buffered(std::vec::Vec<u8>),
StaticFile {
name: &'static str,
bytes: &'static [u8],
},
OpenBytes {
file: std::fs::File,
byte_len: u64,
},
OpenFile {
name: std::string::String,
file: std::fs::File,
byte_len: u64,
},
}
impl ResponseSink for RunnerResponse {
fn put_packed(&mut self, bytes: &[u8]) -> Result<(), ResponseCapacityExceeded> {
match self {
Self::Buffered(body) => ResponseSink::put_packed(body, bytes),
Self::StaticFile { .. } | Self::OpenBytes { .. } | Self::OpenFile { .. } => {
Err(ResponseCapacityExceeded)
}
}
}
fn put_bytes(&mut self, bytes: &[u8]) -> Result<(), ResponseCapacityExceeded> {
match self {
Self::Buffered(body) => ResponseSink::put_bytes(body, bytes),
Self::StaticFile { .. } | Self::OpenBytes { .. } | Self::OpenFile { .. } => {
Err(ResponseCapacityExceeded)
}
}
}
fn put_static_file(
&mut self,
name: &'static str,
bytes: &'static [u8],
) -> Result<(), ResponseCapacityExceeded> {
match self {
Self::Buffered(body) if body.is_empty() => {
*self = Self::StaticFile { name, bytes };
Ok(())
}
_ => Err(ResponseCapacityExceeded),
}
}
fn put_open_bytes(
&mut self,
file: std::fs::File,
byte_len: u64,
) -> Result<(), ResponseCapacityExceeded> {
match self {
Self::Buffered(body) if body.is_empty() => {
*self = Self::OpenBytes { file, byte_len };
Ok(())
}
_ => Err(ResponseCapacityExceeded),
}
}
fn put_open_file(
&mut self,
name: &str,
file: std::fs::File,
byte_len: u64,
) -> Result<(), ResponseCapacityExceeded> {
match self {
Self::Buffered(body) if body.is_empty() => {
*self = Self::OpenFile {
name: name.to_owned(),
file,
byte_len,
};
Ok(())
}
_ => Err(ResponseCapacityExceeded),
}
}
}
pub(super) async fn run_router<St, R: RequestEndpointSet<St>>(
state: &St,
mut requests: mpsc::Receiver<RunnerRequest>,
commands: PrnsNodeHandle,
) {
let mut in_flight = FuturesUnordered::new();
let mut response_lanes: std::collections::HashMap<LinkId, Weak<Mutex<()>>> =
std::collections::HashMap::new();
loop {
let accepting = in_flight.len() < MAX_IN_FLIGHT;
tokio::select! {
biased;
Some(()) = in_flight.next(), if !in_flight.is_empty() => {}
request = requests.recv(), if accepting => match request {
Some(request) => {
response_lanes.retain(|_, lane| lane.strong_count() > 0);
let response_lane = response_lanes
.get(&request.link_id)
.and_then(Weak::upgrade)
.unwrap_or_else(|| {
let lane = Arc::new(Mutex::new(()));
response_lanes.insert(request.link_id, Arc::downgrade(&lane));
lane
});
in_flight.push(dispatch_guarded::<St, R>(
state,
&commands,
request,
response_lane,
));
}
None => break,
},
}
}
}
async fn dispatch_guarded<St, R: RequestEndpointSet<St>>(
state: &St,
commands: &PrnsNodeHandle,
request: RunnerRequest,
response_lane: Arc<Mutex<()>>,
) {
let link_id = request.link_id;
if AssertUnwindSafe(dispatch::<St, R>(state, commands, request, response_lane))
.catch_unwind()
.await
.is_err()
{
commands.close_link(link_id);
}
}
async fn dispatch<St, R: RequestEndpointSet<St>>(
state: &St,
commands: &PrnsNodeHandle,
request: RunnerRequest,
response_lane: Arc<Mutex<()>>,
) {
let link_id = request.link_id;
let inbound = InboundRequest::new(
request.destination,
request.link_id,
request.request_id,
request.requester,
request.requested_at,
request.rtt,
&request.data,
);
let responder = inbound.respond_token();
let mut body = RunnerResponse::Buffered(std::vec::Vec::new());
match dispatch_request::<St, R>(state, request.path_hash, inbound, &mut body).await {
Ok(()) => {
let _response_guard = response_lane.lock().await;
let result = match body {
RunnerResponse::Buffered(body) => {
commands.respond_owned_packed_settled(responder, body).await
}
RunnerResponse::StaticFile { name, bytes } => {
commands
.respond_static_file_settled(responder, name, bytes)
.await
}
RunnerResponse::OpenBytes { file, byte_len } => {
commands
.respond_bytes_streaming(
responder,
byte_len,
tokio::fs::File::from_std(file),
)
.await
}
RunnerResponse::OpenFile {
name,
file,
byte_len,
} => {
commands
.respond_open_file_settled(responder, &name, file, byte_len)
.await
}
};
if let Err(error) = result {
eprintln!(
"REQUEST_RESPONSE_FAILURE link_id={:?} error={error}",
link_id.as_bytes()
);
commands.close_link(link_id);
}
}
Err(Decline::Ignore) => {}
Err(Decline::CloseLink) => {
commands.close_link(responder.link_id);
}
Err(Decline::ResponseTooLarge) => {}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::engine::{IssuedCommand, PrnsCommand};
use crate::manifold::driver::HostCommand;
use crate::routing::request_handlers::RequestPathHash;
use crate::runtime::request_endpoints::{RequestContext, RequestEndpointPolicy};
#[test]
fn static_file_sink_preserves_filename_and_borrowed_bytes() {
static FILE: [u8; 32] = [0x42; 32];
let mut response = RunnerResponse::Buffered(std::vec::Vec::new());
ResponseSink::put_static_file(&mut response, "source.zip", &FILE).unwrap();
let RunnerResponse::StaticFile { name, bytes } = response else {
panic!("static file response");
};
assert_eq!(name, "source.zip");
assert_eq!(bytes.as_ptr(), FILE.as_ptr());
}
#[test]
fn open_file_sink_retains_the_handle_without_reading_it() {
let source = std::fs::File::open("Cargo.toml").unwrap();
let mut response = RunnerResponse::Buffered(std::vec::Vec::new());
ResponseSink::put_open_file(&mut response, "source.zip", source, 42).unwrap();
let RunnerResponse::OpenFile { name, byte_len, .. } = response else {
panic!("open file response");
};
assert_eq!(name, "source.zip");
assert_eq!(byte_len, 42);
}
struct PanickingRequestEndpointSet;
impl RequestEndpointSet<()> for PanickingRequestEndpointSet {
const REGISTRATIONS: &'static [(&'static str, RequestEndpointPolicy)] = &[];
async fn dispatch(
_context: RequestContext<'_, ()>,
_path_hash: RequestPathHash,
) -> Result<(), Decline> {
std::panic::panic_any("request handler")
}
}
#[tokio::test]
async fn a_panicking_request_handler_closes_its_link() {
let (commands, mut command_rx) = mpsc::unbounded_channel();
let handle = PrnsNodeHandle::over(commands);
let link_id = LinkId::new([0x44; 16]);
dispatch_guarded::<(), PanickingRequestEndpointSet>(
&(),
&handle,
RunnerRequest {
destination: DestinationHash::new([0x33; 16]),
link_id,
request_id: RequestId([0x55; 16]),
requester: None,
path_hash: RequestPathHash::new([0x66; 16]),
requested_at: InstantMillis(700),
rtt: RttMillis::new(80),
data: std::vec::Vec::new(),
},
Arc::new(Mutex::new(())),
)
.await;
assert!(matches!(
command_rx.recv().await,
Some(HostCommand::Engine(IssuedCommand {
command: PrnsCommand::CloseLink(close),
..
})) if close.link_id == link_id
));
}
}