use crate::application::ServeState;
use armature_h1::{H2Fallback, Transport};
use bytes::Bytes;
use hyper::body::Incoming as IncomingBody;
use hyper::service::service_fn;
use hyper_util::rt::{TokioExecutor, TokioIo};
use std::future::Future;
use std::io::IoSlice;
use std::net::SocketAddr;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
pub(crate) struct Replay {
inner: Box<dyn Transport>,
head: Bytes,
}
impl Replay {
pub(crate) fn new(inner: Box<dyn Transport>, head: Bytes) -> Self {
Self { inner, head }
}
pub(crate) fn wrap(inner: Box<dyn Transport>, head: Bytes) -> Box<dyn Transport> {
if head.is_empty() {
inner
} else {
Box::new(Self::new(inner, head))
}
}
}
impl AsyncRead for Replay {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
if !self.head.is_empty() {
let n = self.head.len().min(buf.remaining());
let chunk = self.head.split_to(n);
buf.put_slice(&chunk);
return Poll::Ready(Ok(()));
}
Pin::new(&mut self.inner).poll_read(cx, buf)
}
}
impl AsyncWrite for Replay {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Pin::new(&mut self.inner).poll_write(cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.inner).poll_flush(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.inner).poll_shutdown(cx)
}
fn is_write_vectored(&self) -> bool {
self.inner.is_write_vectored()
}
fn poll_write_vectored(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
bufs: &[IoSlice<'_>],
) -> Poll<std::io::Result<usize>> {
Pin::new(&mut self.inner).poll_write_vectored(cx, bufs)
}
}
pub(crate) struct HyperH2 {
state: ServeState,
builder: hyper::server::conn::http2::Builder<TokioExecutor>,
}
impl HyperH2 {
pub(crate) fn new(
state: ServeState,
builder: hyper::server::conn::http2::Builder<TokioExecutor>,
) -> Self {
Self { state, builder }
}
}
impl H2Fallback for HyperH2 {
fn handle(
&self,
io: Box<dyn Transport>,
buffered: Bytes,
peer: Option<SocketAddr>,
) -> Pin<Box<dyn Future<Output = ()>>> {
let state = match peer {
Some(peer) => self.state.for_peer(peer),
None => self.state.clone(),
};
let builder = self.builder.clone();
Box::pin(async move {
let io = TokioIo::new(Replay::wrap(io, buffered));
let service = service_fn(move |req: hyper::Request<IncomingBody>| {
let state = state.clone();
async move { crate::application::handle_request(req, state).await }
});
if let Err(err) = builder.serve_connection(io, service).await {
crate::logging::error!(error = %err, "Error serving HTTP/2 connection");
}
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
#[tokio::test]
async fn replayed_bytes_are_read_before_the_transport() {
let (mut client, server) = tokio::io::duplex(4096);
let mut replay = Replay::new(Box::new(server), Bytes::from_static(b"PREFACE"));
client.write_all(b"THEREST").await.expect("write");
let mut out = [0u8; 14];
let mut filled = 0;
while filled < out.len() {
let n = replay.read(&mut out[filled..]).await.expect("read");
assert_ne!(n, 0, "unexpected EOF at {filled}");
filled += n;
}
assert_eq!(
&out[..],
b"PREFACETHEREST",
"the bytes read to classify the connection must arrive first, or \
hyper sees a stream missing its opening frames"
);
}
#[tokio::test]
async fn writes_bypass_the_replay_buffer() {
let (mut client, server) = tokio::io::duplex(4096);
let mut replay = Replay::new(Box::new(server), Bytes::from_static(b"PREFACE"));
replay.write_all(b"SETTINGS").await.expect("write");
replay.flush().await.expect("flush");
let mut out = vec![0u8; 8];
client.read_exact(&mut out).await.expect("read");
assert_eq!(&out[..], b"SETTINGS");
}
#[tokio::test]
async fn wrap_hands_back_the_transport_when_there_is_nothing_to_replay() {
let (mut client, server) = tokio::io::duplex(4096);
let server: Box<dyn Transport> = Box::new(server);
let addr = std::ptr::addr_of!(*server) as *const ();
let mut io = Replay::wrap(server, Bytes::new());
assert_eq!(
std::ptr::addr_of!(*io) as *const (),
addr,
"an empty head must not be wrapped"
);
client.write_all(b"THEREST").await.expect("write");
let mut out = [0u8; 7];
io.read_exact(&mut out).await.expect("read");
assert_eq!(&out[..], b"THEREST");
}
#[tokio::test]
async fn wrap_splices_a_non_empty_head() {
let (mut client, server) = tokio::io::duplex(4096);
let mut io = Replay::wrap(Box::new(server), Bytes::from_static(b"PREFACE"));
client.write_all(b"THEREST").await.expect("write");
let mut out = [0u8; 14];
io.read_exact(&mut out).await.expect("read");
assert_eq!(&out[..], b"PREFACETHEREST");
}
}