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")
);
}
}