saddle-framework 0.2.0

The single business-facing facade for Saddle applications
Documentation
use std::{
    pin::Pin,
    task::{Context, Poll},
};

use saddle_admission::ManagedResponse;
use saddle_runtime::{
    compiled_route::ResponseOutcomeClass,
    response_framing::{FixedResponseWrite, ResponseFramingAdapter},
};
use tokio::{io::AsyncWrite, net::TcpStream};

const HEADER_BYTES: usize = 160;

#[derive(Clone, Copy)]
pub(crate) struct Http1ResponsePlan {
    _private: (),
}

impl Http1ResponsePlan {
    pub(crate) const fn connection_close() -> Self {
        Self { _private: () }
    }
}

pub(crate) struct Http1Framing;

pub(crate) struct Http1WriteState {
    header: [u8; HEADER_BYTES],
    header_len: usize,
    header_written: usize,
    body_written: usize,
    payload_capacity: usize,
    encoded: bool,
    closed: bool,
}

#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum Http1WriteError {
    PayloadCapacity,
    HeaderCapacity,
    Socket,
}

impl ResponseFramingAdapter for Http1Framing {
    type Plan = Http1ResponsePlan;
    type Error = Http1WriteError;
    type WriteState = Http1WriteState;

    fn begin(
        &self,
        _plan: Self::Plan,
        payload_capacity: usize,
    ) -> Result<Self::WriteState, Self::Error> {
        Ok(Http1WriteState {
            header: [0; HEADER_BYTES],
            header_len: 0,
            header_written: 0,
            body_written: 0,
            payload_capacity,
            encoded: false,
            closed: false,
        })
    }
}

impl FixedResponseWrite for Http1WriteState {
    type Error = Http1WriteError;

    fn poll_write(
        &mut self,
        socket: &mut TcpStream,
        class: ResponseOutcomeClass,
        payload: &ManagedResponse,
        context: &mut Context<'_>,
    ) -> Poll<Result<(), Self::Error>> {
        if !self.encoded {
            if payload.as_slice().len() > self.payload_capacity {
                return Poll::Ready(Err(Http1WriteError::PayloadCapacity));
            }
            self.header_len = encode_header(&mut self.header, class, payload.as_slice().len())?;
            self.encoded = true;
        }

        while self.header_written < self.header_len {
            match Pin::new(&mut *socket)
                .poll_write(context, &self.header[self.header_written..self.header_len])
            {
                Poll::Pending => return Poll::Pending,
                Poll::Ready(Err(_)) | Poll::Ready(Ok(0)) => {
                    return Poll::Ready(Err(Http1WriteError::Socket));
                }
                Poll::Ready(Ok(written)) => self.header_written += written,
            }
        }
        while self.body_written < payload.as_slice().len() {
            match Pin::new(&mut *socket)
                .poll_write(context, &payload.as_slice()[self.body_written..])
            {
                Poll::Pending => return Poll::Pending,
                Poll::Ready(Err(_)) | Poll::Ready(Ok(0)) => {
                    return Poll::Ready(Err(Http1WriteError::Socket));
                }
                Poll::Ready(Ok(written)) => self.body_written += written,
            }
        }
        if !self.closed {
            match Pin::new(socket).poll_shutdown(context) {
                Poll::Pending => return Poll::Pending,
                Poll::Ready(Err(_)) => return Poll::Ready(Err(Http1WriteError::Socket)),
                Poll::Ready(Ok(())) => self.closed = true,
            }
        }
        Poll::Ready(Ok(()))
    }
}

fn encode_header(
    output: &mut [u8; HEADER_BYTES],
    class: ResponseOutcomeClass,
    content_length: usize,
) -> Result<usize, Http1WriteError> {
    let status = match class {
        ResponseOutcomeClass::Success => b"200 OK".as_slice(),
        ResponseOutcomeClass::InvalidRequest => b"400 Bad Request".as_slice(),
        ResponseOutcomeClass::BusinessRejected => b"422 Unprocessable Entity".as_slice(),
        ResponseOutcomeClass::Unavailable => b"503 Service Unavailable".as_slice(),
        ResponseOutcomeClass::Internal => b"500 Internal Server Error".as_slice(),
    };
    const VERSION: &[u8] = b"HTTP/1.1 ";
    const LENGTH: &[u8] = b"\r\nContent-Length: ";
    const SUFFIX: &[u8] =
        b"\r\nContent-Type: application/octet-stream\r\nConnection: close\r\n\r\n";

    let mut cursor = 0;
    append(output, &mut cursor, VERSION)?;
    append(output, &mut cursor, status)?;
    append(output, &mut cursor, LENGTH)?;
    append_decimal(output, &mut cursor, content_length)?;
    append(output, &mut cursor, SUFFIX)?;
    Ok(cursor)
}

fn append(
    output: &mut [u8; HEADER_BYTES],
    cursor: &mut usize,
    value: &[u8],
) -> Result<(), Http1WriteError> {
    let end = cursor
        .checked_add(value.len())
        .ok_or(Http1WriteError::HeaderCapacity)?;
    let target = output
        .get_mut(*cursor..end)
        .ok_or(Http1WriteError::HeaderCapacity)?;
    target.copy_from_slice(value);
    *cursor = end;
    Ok(())
}

fn append_decimal(
    output: &mut [u8; HEADER_BYTES],
    cursor: &mut usize,
    mut value: usize,
) -> Result<(), Http1WriteError> {
    let mut reversed = [0_u8; 20];
    let mut count = 0;
    loop {
        reversed[count] =
            b'0' + u8::try_from(value % 10).map_err(|_| Http1WriteError::HeaderCapacity)?;
        count += 1;
        value /= 10;
        if value == 0 {
            break;
        }
    }
    for digit in reversed[..count].iter().rev() {
        append(output, cursor, std::slice::from_ref(digit))?;
    }
    Ok(())
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn five_classes_have_exhaustive_canonical_headers() {
        let cases = [
            (ResponseOutcomeClass::Success, b"HTTP/1.1 200 OK".as_slice()),
            (
                ResponseOutcomeClass::InvalidRequest,
                b"HTTP/1.1 400 Bad Request".as_slice(),
            ),
            (
                ResponseOutcomeClass::BusinessRejected,
                b"HTTP/1.1 422 Unprocessable Entity".as_slice(),
            ),
            (
                ResponseOutcomeClass::Unavailable,
                b"HTTP/1.1 503 Service Unavailable".as_slice(),
            ),
            (
                ResponseOutcomeClass::Internal,
                b"HTTP/1.1 500 Internal Server Error".as_slice(),
            ),
        ];
        for (class, prefix) in cases {
            let mut bytes = [0; HEADER_BYTES];
            let length = encode_header(&mut bytes, class, 0).unwrap();
            let header = &bytes[..length];
            assert!(header.starts_with(prefix));
            assert!(
                header
                    .windows(b"\r\nContent-Length: 0\r\n".len())
                    .any(|value| value == b"\r\nContent-Length: 0\r\n")
            );
            assert!(header.ends_with(b"Connection: close\r\n\r\n"));
        }
    }

    #[test]
    fn decimal_length_is_canonical() {
        let mut bytes = [0; HEADER_BYTES];
        let length = encode_header(&mut bytes, ResponseOutcomeClass::Success, 10_042).unwrap();
        assert!(
            bytes[..length]
                .windows(b"Content-Length: 10042\r\n".len())
                .any(|value| value == b"Content-Length: 10042\r\n")
        );
    }
}