use std::{future::poll_fn, task::Poll};
use bytes::{Buf, Bytes};
use http_body::Body;
use squiche::h3::Header;
use crate::{
h3::common::{H3App, H3Error, headers::header_map_to_h3},
quic::connection::{QuicScionConn, WeakConnectionHandle},
};
pub(crate) async fn send_headers<A: H3App>(
handle: &WeakConnectionHandle<A>,
stream_id: u64,
headers: Vec<Header>,
) -> Result<(), H3Error> {
poll_fn(|cx| {
let Some(handle) = handle.upgrade() else {
return Poll::Ready(Err(H3Error::ConnectionClosed));
};
let mut guard = handle.lock();
let QuicScionConn { inner, app, .. } = &mut *guard;
let (h3, streams) = app.h3_streams();
let Some(h3) = h3 else {
return Poll::Ready(Err(H3Error::ConnectionClosed));
};
match h3.send_response(inner, stream_id, &headers, false) {
Ok(()) => {
drop(guard);
handle.notify();
Poll::Ready(Ok(()))
}
Err(squiche::h3::Error::StreamBlocked) => {
let Some(st) = streams.get_mut(&stream_id) else {
return Poll::Ready(Err(H3Error::ConnectionClosed));
};
st.write_waker = Some(cx.waker().clone());
drop(guard);
handle.notify();
Poll::Pending
}
Err(err) => Poll::Ready(Err(H3Error::H3(err))),
}
})
.await
}
pub(crate) async fn send_data<A: H3App>(
handle: &WeakConnectionHandle<A>,
stream_id: u64,
mut bytes: Bytes,
fin: bool,
) -> Result<(), H3Error> {
poll_fn(move |cx| {
let Some(handle) = handle.upgrade() else {
return Poll::Ready(Err(H3Error::ConnectionClosed));
};
let mut guard = handle.lock();
let QuicScionConn { inner, app, .. } = &mut *guard;
let (h3, streams) = app.h3_streams();
let Some(h3) = h3 else {
return Poll::Ready(Err(H3Error::ConnectionClosed));
};
loop {
if bytes.is_empty() && !fin {
drop(guard);
handle.notify();
return Poll::Ready(Ok(()));
}
match h3.send_body(inner, stream_id, bytes.as_ref(), fin) {
Ok(written) => {
bytes.advance(written);
if bytes.is_empty() {
drop(guard);
handle.notify();
return Poll::Ready(Ok(()));
}
}
Err(squiche::h3::Error::Done) | Err(squiche::h3::Error::StreamBlocked) => {
let Some(st) = streams.get_mut(&stream_id) else {
return Poll::Ready(Err(H3Error::ConnectionClosed));
};
st.write_waker = Some(cx.waker().clone());
drop(guard);
handle.notify();
return Poll::Pending;
}
Err(err) => return Poll::Ready(Err(H3Error::H3(err))),
}
}
})
.await
}
pub(crate) async fn send_trailers<A: H3App>(
handle: &WeakConnectionHandle<A>,
stream_id: u64,
trailers: Vec<Header>,
) -> Result<(), H3Error> {
poll_fn(|cx| {
let Some(handle) = handle.upgrade() else {
return Poll::Ready(Err(H3Error::ConnectionClosed));
};
let mut guard = handle.lock();
let QuicScionConn { inner, app, .. } = &mut *guard;
let (h3, streams) = app.h3_streams();
let Some(h3) = h3 else {
return Poll::Ready(Err(H3Error::ConnectionClosed));
};
match h3.send_additional_headers(inner, stream_id, &trailers, true, true) {
Ok(()) => {
drop(guard);
handle.notify();
Poll::Ready(Ok(()))
}
Err(squiche::h3::Error::StreamBlocked) => {
let Some(st) = streams.get_mut(&stream_id) else {
return Poll::Ready(Err(H3Error::ConnectionClosed));
};
st.write_waker = Some(cx.waker().clone());
drop(guard);
handle.notify();
Poll::Pending
}
Err(err) => Poll::Ready(Err(H3Error::H3(err))),
}
})
.await
}
pub(crate) async fn pump_body<A, B>(
handle: &WeakConnectionHandle<A>,
stream_id: u64,
body: B,
) -> bool
where
A: H3App + 'static,
B: Body + Send + 'static,
B::Data: Send,
B::Error: Send,
{
let mut body = std::pin::pin!(body);
let mut trailers_sent = false;
loop {
match poll_fn(|cx| body.as_mut().poll_frame(cx)).await {
Some(Ok(frame)) => {
match frame.into_data() {
Ok(mut data) => {
let bytes = data.copy_to_bytes(data.remaining());
if let Err(err) = send_data(handle, stream_id, bytes, false).await {
tracing::debug!(?err, stream_id, "failed to send body data");
return false;
}
}
Err(non_data) => {
if let Ok(map) = non_data.into_trailers() {
let trailers = header_map_to_h3(&map);
if !trailers.is_empty() {
if let Err(err) = send_trailers(handle, stream_id, trailers).await {
tracing::debug!(?err, stream_id, "failed to send trailers");
return false;
}
trailers_sent = true;
break;
}
}
}
}
}
Some(Err(_err)) => {
tracing::debug!(stream_id, "body errored before completion");
return false;
}
None => break,
}
}
if trailers_sent {
true
} else {
send_data(handle, stream_id, Bytes::new(), true)
.await
.is_ok()
}
}