use crate::{client::H3Connector, executor::SharedExec};
use futures::{FutureExt, future::BoxFuture};
use hyper::{
Request, Response, Uri,
body::{Body, Bytes},
http::uri::{Authority, PathAndQuery, Scheme},
rt::Executor,
};
use std::future::Future;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use crate::client_body::H3IncomingClient;
fn is_connection_closing(err: &h3::error::StreamError) -> bool {
let msg = err.to_string();
msg.starts_with("Remote is closing the connection") || msg.starts_with("Connection error:")
}
pub async fn send_request_inner<CONN, B>(
req: hyper::Request<B>,
mut send_request: h3::client::SendRequest<CONN::OS, Bytes>,
executor: &SharedExec,
closing: Arc<AtomicBool>,
) -> Result<Response<H3IncomingClient<CONN::RS, Bytes>>, crate::Error>
where
CONN: H3Connector,
B: Body + Send + 'static + Unpin,
B::Data: Send,
B::Error: Into<crate::Error> + Send,
{
let (parts, body) = req.into_parts();
let head_req = hyper::Request::from_parts(parts, ());
tracing::trace!("sending h3 req header: {:?}", head_req);
let stream = match send_request.send_request(head_req).await {
Ok(stream) => stream,
Err(e) => {
if is_connection_closing(&e) {
closing.store(true, Ordering::SeqCst);
}
return Err(e.into());
}
};
let (w, mut r) = stream.split();
let (cancel_tx, cancel_rx) = tokio::sync::oneshot::channel();
let mut body_fut = Box::pin(crate::client_body::send_h3_client_body::<CONN::BS, _>(
w, body, cancel_rx,
));
match futures::future::poll_fn(|cx| match body_fut.as_mut().poll(cx) {
std::task::Poll::Ready(res) => std::task::Poll::Ready(Some(res)),
std::task::Poll::Pending => std::task::Poll::Ready(None),
})
.await
{
Some(res) => {
res?;
}
None => {
executor.execute(async move {
if let Err(e) = body_fut.await {
tracing::warn!("h3 client body send failed: {e}");
}
});
}
};
tracing::trace!("recv header");
let resp = match r.recv_response().await {
Ok(resp) => resp,
Err(e) => {
tracing::error!("recv header error: {e}");
if is_connection_closing(&e) {
closing.store(true, Ordering::SeqCst);
}
return Err(e.into());
}
};
let (resp, _) = resp.into_parts();
let resp_body = H3IncomingClient::new(r, Some(cancel_tx));
tracing::trace!("return resp");
Ok(hyper::Response::from_parts(resp, resp_body))
}
#[allow(clippy::type_complexity)]
pub struct RequestSender<CONN: H3Connector> {
conn: CONN,
send_request: Option<h3::client::SendRequest<CONN::OS, Bytes>>,
driver_rx: Option<tokio::sync::oneshot::Receiver<()>>,
make_send_request_fut: Option<
BoxFuture<
'static,
Result<
(
h3::client::SendRequest<CONN::OS, Bytes>,
tokio::sync::oneshot::Receiver<()>,
),
crate::Error,
>,
>,
>,
executor: SharedExec,
base_scheme: Option<Scheme>,
base_authority: Option<Authority>,
connect_error: Option<crate::Error>,
closing: Arc<AtomicBool>,
}
impl<CONN> RequestSender<CONN>
where
CONN: H3Connector,
{
pub fn new(conn: CONN, uri: Uri, executor: SharedExec) -> Self {
let base_scheme = uri.scheme().cloned();
let base_authority = uri.authority().cloned();
Self {
conn,
send_request: None,
driver_rx: None,
make_send_request_fut: None,
executor,
base_scheme,
base_authority,
connect_error: None,
closing: Arc::new(AtomicBool::new(false)),
}
}
fn retire_connection(&mut self) {
self.send_request = None;
self.driver_rx = None;
self.closing = Arc::new(AtomicBool::new(false));
}
}
impl<CONN, B> tower::Service<Request<B>> for RequestSender<CONN>
where
CONN: H3Connector,
B: Body + Send + 'static + Unpin,
B::Data: Send,
B::Error: Into<crate::Error> + Send,
{
type Response = Response<H3IncomingClient<CONN::RS, Bytes>>;
type Error = crate::Error;
type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;
fn poll_ready(
&mut self,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
if let Some(rx) = &mut self.driver_rx {
match rx.try_recv() {
Ok(()) => {
tracing::trace!("driver is closed, reconnecting.");
self.retire_connection();
}
Err(tokio::sync::oneshot::error::TryRecvError::Empty) => {
}
Err(tokio::sync::oneshot::error::TryRecvError::Closed) => {
tracing::trace!("driver is closed, reconnecting.");
self.retire_connection();
}
}
}
if self.closing.load(Ordering::SeqCst) {
tracing::trace!("connection is closing (goaway), reconnecting.");
self.retire_connection();
}
if self.connect_error.is_some() {
return std::task::Poll::Ready(Ok(()));
}
if self.send_request.is_some() {
tracing::trace!("exp poll_ready cache hit.");
assert!(self.make_send_request_fut.is_none());
assert!(self.driver_rx.is_some());
return std::task::Poll::Ready(Ok(()));
}
if self.make_send_request_fut.is_none() {
let conn = self.conn.clone();
let executor = self.executor.clone();
self.make_send_request_fut = Some(Box::pin(async move {
let conn = conn.connect().await?;
let (mut driver, send_request) = h3::client::new(conn).await?;
let (tx, rx) = tokio::sync::oneshot::channel();
executor.execute(async move {
let res = std::future::poll_fn(|cx| driver.poll_close(cx)).await;
tracing::trace!("h3 driver ended: {res:?}");
let _ = tx.send(());
});
Ok((send_request, rx))
}));
}
self.make_send_request_fut
.as_mut()
.unwrap()
.poll_unpin(cx)
.map(|res| match res {
Ok((send_request, rx)) => {
self.send_request = Some(send_request);
self.driver_rx = Some(rx);
self.make_send_request_fut = None;
Ok(())
}
Err(e) => {
self.make_send_request_fut = None;
self.connect_error = Some(e);
Ok(())
}
})
}
fn call(&mut self, mut req: Request<B>) -> Self::Future {
if let Some(e) = self.connect_error.take() {
return Box::pin(async move { Err(e) });
}
let (scheme, authority) = match (self.base_scheme.clone(), self.base_authority.clone()) {
(Some(scheme), Some(authority)) => (scheme, authority),
(None, _) => {
return Box::pin(async move {
Err(crate::Error::from("h3 client base URI is missing a scheme"))
});
}
(_, None) => {
return Box::pin(async move {
Err(crate::Error::from(
"h3 client base URI is missing an authority",
))
});
}
};
let path_and_query = req
.uri()
.path_and_query()
.cloned()
.unwrap_or_else(|| PathAndQuery::from_static("/"));
let uri2 = match Uri::builder()
.scheme(scheme)
.authority(authority)
.path_and_query(path_and_query)
.build()
{
Ok(uri2) => uri2,
Err(e) => {
return Box::pin(async move { Err(crate::Error::from(e)) });
}
};
let send_request = match &self.send_request {
Some(sr) => sr.clone(),
None => {
return Box::pin(async move {
Err(crate::Error::from(
"h3 request sender is not ready: poll_ready must return Ready(Ok) before call",
))
});
}
};
*req.uri_mut() = uri2;
let executor = self.executor.clone();
let closing = self.closing.clone();
Box::pin(async move {
crate::client_conn::send_request_inner::<CONN, B>(req, send_request, &executor, closing)
.await
})
}
}