use std::convert::Infallible;
use std::pin::Pin;
use std::task::{Context, Poll};
use bytes::Bytes;
use http::{Request, Response};
use http_body_util::{BodyExt, Empty};
use hyper::body::{Body, Frame as HttpBodyFrame, Incoming, SizeHint};
use hyper_util::rt::{TokioExecutor, TokioIo};
use tokio::net::TcpStream;
use tonic::body::Body as TonicBody;
use tonic::Status;
use tower::ServiceExt as _;
use crate::{Backend, BoxError, ResBody, Runtime};
pub(crate) fn serve(rt: &Runtime, req: &mut Request<Incoming>) -> Response<ResBody> {
let (response, ws_fut) = match h2ts_server::accept(req) {
Ok(pair) => pair,
Err(e) => return e.rejection_response().map(box_empty),
};
let rt = rt.clone();
tokio::spawn(async move {
let ws = match ws_fut.await {
Ok(ws) => ws,
Err(e) => {
tracing::debug!("h2ts upgrade failed: {e}");
return;
}
};
match &rt.backend {
Backend::InProcess(routes) => {
let routes = routes.clone();
let max = rt.cfg.max_message_bytes;
let service = hyper::service::service_fn(move |req: Request<Incoming>| {
let routes = routes.clone();
async move {
let deadline = crate::metadata::parse_grpc_timeout(req.headers());
let req = req.map(|body| TonicBody::new(GrpcSizeLimit::new(body, max)));
match deadline {
Some(d) => match tokio::time::timeout(d, routes.oneshot(req)).await {
Ok(result) => result,
Err(_) => Ok(Status::deadline_exceeded("deadline exceeded")
.into_http()),
},
None => routes.oneshot(req).await,
}
}
});
let io = TokioIo::new(h2ts_server::WsByteStream::new(ws));
let builder = hyper::server::conn::http2::Builder::new(TokioExecutor::new());
let conn = builder.serve_connection(io, service);
if let Err(e) = crate::drain::serve_graceful!(&rt.drain, conn) {
tracing::debug!("h2ts connection ended: {e}");
}
}
Backend::Upstream(_) => {
let Some(authority) = rt.cfg.upstream_authority.clone() else {
tracing::debug!("h2ts proxy: no upstream authority configured");
return;
};
match TcpStream::connect(&authority).await {
Ok(tcp) => {
let _ = h2ts_server::bridge(ws, tcp).await;
}
Err(e) => tracing::debug!("h2ts proxy connect {authority} failed: {e}"),
}
}
}
});
response.map(box_empty)
}
fn box_empty(body: Empty<Bytes>) -> ResBody {
body.map_err(|e: Infallible| -> BoxError { match e {} }).boxed_unsync()
}
struct GrpcSizeLimit<B> {
inner: B,
max: usize,
header_seen: usize,
len_buf: [u8; 4],
body_remaining: usize,
}
impl<B> GrpcSizeLimit<B> {
fn new(inner: B, max: usize) -> Self {
Self { inner, max, header_seen: 0, len_buf: [0; 4], body_remaining: 0 }
}
fn inspect(&mut self, data: &[u8]) -> Result<(), BoxError> {
let mut i = 0;
while i < data.len() {
if self.header_seen < 5 {
if self.header_seen >= 1 {
self.len_buf[self.header_seen - 1] = data[i];
}
self.header_seen += 1;
i += 1;
if self.header_seen == 5 {
let len = u32::from_be_bytes(self.len_buf) as usize;
if len > self.max {
return Err(Box::new(Status::resource_exhausted(format!(
"request message exceeds size limit ({} bytes)",
self.max
))) as BoxError);
}
self.body_remaining = len;
}
} else {
let take = self.body_remaining.min(data.len() - i);
self.body_remaining -= take;
i += take;
if self.body_remaining == 0 {
self.header_seen = 0;
}
}
}
Ok(())
}
}
impl<B> Body for GrpcSizeLimit<B>
where
B: Body<Data = Bytes> + Unpin,
B::Error: Into<BoxError>,
{
type Data = Bytes;
type Error = BoxError;
fn poll_frame(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Result<HttpBodyFrame<Bytes>, BoxError>>> {
let this = self.as_mut().get_mut();
match Pin::new(&mut this.inner).poll_frame(cx) {
Poll::Ready(Some(Ok(frame))) => {
if let Some(data) = frame.data_ref() {
if let Err(e) = this.inspect(data) {
return Poll::Ready(Some(Err(e)));
}
}
Poll::Ready(Some(Ok(frame)))
}
Poll::Ready(Some(Err(e))) => Poll::Ready(Some(Err(e.into()))),
Poll::Ready(None) => Poll::Ready(None),
Poll::Pending => Poll::Pending,
}
}
fn is_end_stream(&self) -> bool {
self.inner.is_end_stream()
}
fn size_hint(&self) -> SizeHint {
self.inner.size_hint()
}
}