prns-runtime-tokio 0.3.6

Tokio host runtime for Personal Reticulum
Documentation
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
        ));
    }
}