use std::{
future::Future,
pin::Pin,
sync::Arc,
task::{Context, Poll},
};
use bytes::{Buf, Bytes};
use http::{Request, Response};
use http_body::Body;
use tokio::sync::mpsc;
use crate::{
client::core::{
self, Error,
body::{DecodedLength, Incoming as IncomingBody},
dispatch::{self, TrySendError},
error::BoxError,
proto::{Dispatched, headers},
},
error::H3Error,
};
#[allow(dead_code)] pub(crate) type ClientRx<B> = dispatch::Receiver<Request<B>, Response<IncomingBody>>;
#[allow(dead_code)] pub(crate) struct ConnTask<B>
where
B: Body,
B: 'static,
B: Unpin,
B: Send,
B::Data: Send,
{
send_request: hpx_h3::client::SendRequest<hpx_h3::quinn::OpenStreams, Bytes>,
req_rx: ClientRx<B>,
close_rx: mpsc::Receiver<hpx_h3::error::ConnectionError>,
}
impl<B> ConnTask<B>
where
B: Body + 'static + Unpin + Send,
B::Data: Send + 'static + Into<Bytes>,
B::Error: Into<BoxError> + Send + Sync + 'static,
{
#[allow(dead_code)] pub(crate) fn new(
send_request: hpx_h3::client::SendRequest<hpx_h3::quinn::OpenStreams, Bytes>,
req_rx: ClientRx<B>,
close_rx: mpsc::Receiver<hpx_h3::error::ConnectionError>,
) -> Self {
Self {
send_request,
req_rx,
close_rx,
}
}
}
impl<B> Future for ConnTask<B>
where
B: Body + 'static + Unpin + Send,
B::Data: Send + 'static + Into<Bytes>,
B::Error: Into<BoxError> + Send + Sync + 'static,
{
type Output = core::Result<Dispatched>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
match this.close_rx.poll_recv(cx) {
Poll::Ready(Some(_err)) => {
trace!("h3 connection closed by driver task");
return Poll::Ready(Ok(Dispatched::Shutdown));
}
Poll::Ready(None) => {
trace!("h3 driver task dropped close_tx");
return Poll::Ready(Ok(Dispatched::Shutdown));
}
Poll::Pending => {}
}
match this.req_rx.poll_recv(cx) {
Poll::Ready(Some((req, cb))) => {
if cb.is_canceled() {
trace!("h3 request callback is canceled");
return Poll::Pending;
}
let mut send_request = this.send_request.clone();
let join_handle = tokio::spawn(async move {
let result = drive_request(&mut send_request, req).await;
let cb_result =
result.map_err(|(error, message)| TrySendError { error, message });
cb.send(cb_result);
});
drop(join_handle);
Poll::Pending
}
Poll::Ready(None) => {
trace!("h3 dispatch::Sender dropped");
Poll::Ready(Ok(Dispatched::Shutdown))
}
Poll::Pending => Poll::Pending,
}
}
}
#[allow(clippy::result_large_err)]
pub(crate) async fn drive_request<B>(
send_request: &mut hpx_h3::client::SendRequest<hpx_h3::quinn::OpenStreams, Bytes>,
req: Request<B>,
) -> Result<Response<IncomingBody>, (Error, Option<Request<B>>)>
where
B: Body + Send + Unpin,
B::Data: Into<Bytes>,
B::Error: Into<BoxError>,
{
let (head, mut body) = req.into_parts();
let mut req = http::Request::from_parts(head, ());
let is_connect = req.method() == http::Method::CONNECT;
headers::strip_connection_headers(req.headers_mut(), true);
if !is_connect
&& let Some(len) = body.size_hint().exact()
&& (len != 0 || headers::method_has_defined_payload_semantics(req.method()))
{
headers::set_content_length_if_missing(req.headers_mut(), len);
}
let mut stream = send_request
.send_request(req)
.await
.map_err(|e| (stream_error_to_error(e), None::<Request<B>>))?;
if !is_connect {
super::dispatch::send_request_body(&mut stream, &mut body).await?;
}
let response = match super::dispatch::handle_response(&mut stream).await? {
super::dispatch::RecvResponseResult::EarlyReturn(response) => return Ok(response),
super::dispatch::RecvResponseResult::Success(response) => response,
};
if is_connect {
let (body_tx, body_rx) =
IncomingBody::new_channel(DecodedLength::CHUNKED, false);
drop(body_tx); let mut response = response.map(|_| body_rx);
response.extensions_mut().insert(H3WebSocket {
inner: Arc::new(tokio::sync::Mutex::new(H3WebSocketInner {
stream,
recv_buf: Vec::new(),
})),
read_timeout: None,
write_timeout: None,
max_frame_size: DEFAULT_MAX_FRAME_SIZE,
});
return Ok(response);
}
let response = super::dispatch::collect_response_body(&mut stream, response).await?;
Ok(response)
}
pub(crate) fn stream_error_to_error(err: hpx_h3::error::StreamError) -> Error {
Error::new_body(stream_error_to_h3_error(err))
}
pub(crate) fn stream_error_to_h3_error(err: hpx_h3::error::StreamError) -> H3Error {
match err {
hpx_h3::error::StreamError::StreamError { code, .. } => H3Error::StreamReset {
code: code.value(),
stream_id: 0,
},
hpx_h3::error::StreamError::RemoteTerminate { code, .. } => H3Error::StreamReset {
code: code.value(),
stream_id: 0,
},
hpx_h3::error::StreamError::HeaderTooBig { .. } => H3Error::MaxConcurrentStreamsExceeded,
other => H3Error::Other(Box::new(other)),
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum WsMessage {
Text(String),
Binary(Vec<u8>),
Close {
code: Option<u16>,
reason: Option<String>,
},
Ping(Vec<u8>),
Pong(Vec<u8>),
}
#[derive(Debug)]
pub enum WsError {
Stream(hpx_h3::error::StreamError),
Timeout,
FrameTooLarge,
FrameMalformed(&'static str),
}
impl std::fmt::Display for WsError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
WsError::Stream(e) => write!(f, "h3 stream error: {e}"),
WsError::Timeout => write!(f, "WebSocket operation timed out"),
WsError::FrameTooLarge => write!(f, "WebSocket frame payload exceeds max size"),
WsError::FrameMalformed(reason) => write!(f, "WebSocket frame malformed: {reason}"),
}
}
}
impl std::error::Error for WsError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
WsError::Stream(e) => Some(e),
_ => None,
}
}
}
impl From<hpx_h3::error::StreamError> for WsError {
fn from(e: hpx_h3::error::StreamError) -> Self {
WsError::Stream(e)
}
}
const DEFAULT_MAX_FRAME_SIZE: usize = 1024 * 1024;
struct H3WebSocketInner {
stream: hpx_h3::client::RequestStream<hpx_h3::quinn::BidiStream<Bytes>, Bytes>,
recv_buf: Vec<u8>,
}
#[derive(Clone)]
pub struct H3WebSocket {
inner: Arc<tokio::sync::Mutex<H3WebSocketInner>>,
read_timeout: Option<std::time::Duration>,
write_timeout: Option<std::time::Duration>,
max_frame_size: usize,
}
impl std::fmt::Debug for H3WebSocket {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("H3WebSocket")
.field("read_timeout", &self.read_timeout)
.field("write_timeout", &self.write_timeout)
.field("max_frame_size", &self.max_frame_size)
.finish()
}
}
impl H3WebSocket {
#[doc(hidden)]
#[allow(dead_code)]
pub fn new(
stream: hpx_h3::client::RequestStream<hpx_h3::quinn::BidiStream<Bytes>, Bytes>,
) -> Self {
H3WebSocket {
inner: Arc::new(tokio::sync::Mutex::new(H3WebSocketInner {
stream,
recv_buf: Vec::new(),
})),
read_timeout: None,
write_timeout: None,
max_frame_size: DEFAULT_MAX_FRAME_SIZE,
}
}
#[inline]
pub fn from_response<B>(response: &mut http::Response<B>) -> Option<Self> {
response.extensions_mut().remove::<H3WebSocket>()
}
#[allow(dead_code)]
pub fn set_read_timeout(&mut self, timeout: Option<std::time::Duration>) {
self.read_timeout = timeout;
}
#[allow(dead_code)]
pub fn set_write_timeout(&mut self, timeout: Option<std::time::Duration>) {
self.write_timeout = timeout;
}
#[allow(dead_code)]
pub fn set_max_frame_size(&mut self, max_size: usize) {
self.max_frame_size = max_size;
}
#[allow(dead_code)]
pub async fn send_text(&mut self, data: &str) -> Result<(), WsError> {
self.send_frame(0x1, data.as_bytes()).await
}
#[allow(dead_code)]
pub async fn send_binary(&mut self, data: &[u8]) -> Result<(), WsError> {
self.send_frame(0x2, data).await
}
#[allow(dead_code)]
pub async fn send_close(
&mut self,
code: Option<u16>,
reason: Option<&str>,
) -> Result<(), WsError> {
let payload = build_close_payload(code, reason);
self.send_frame(0x8, &payload).await
}
#[allow(dead_code)]
pub async fn send_ping(&mut self, data: &[u8]) -> Result<(), WsError> {
self.send_frame(0x9, data).await
}
#[allow(dead_code)]
pub async fn send_pong(&mut self, data: &[u8]) -> Result<(), WsError> {
self.send_frame(0xA, data).await
}
#[allow(dead_code)]
pub async fn recv(&mut self) -> Result<Option<WsMessage>, WsError> {
let mut inner = self.inner.lock().await;
loop {
if let Some(msg) = Self::try_parse_frame(&mut inner.recv_buf, self.max_frame_size)? {
if let WsMessage::Ping(ref data) = msg {
drop(inner);
self.send_pong(data).await?;
inner = self.inner.lock().await;
continue;
}
return Ok(Some(msg));
}
let chunk = if let Some(timeout) = self.read_timeout {
match tokio::time::timeout(timeout, inner.stream.recv_data()).await {
Ok(Ok(Some(data))) => {
let mut data = data;
let len = data.remaining();
data.copy_to_bytes(len)
}
Ok(Ok(None)) => {
return Ok(None);
}
Ok(Err(e)) => return Err(WsError::Stream(e)),
Err(_elapsed) => return Err(WsError::Timeout),
}
} else {
match inner.stream.recv_data().await {
Ok(Some(data)) => {
let mut data = data;
let len = data.remaining();
data.copy_to_bytes(len)
}
Ok(None) => return Ok(None),
Err(e) => return Err(WsError::Stream(e)),
}
};
inner.recv_buf.extend_from_slice(&chunk);
}
}
#[allow(dead_code)]
pub async fn finish(&mut self) -> Result<(), WsError> {
Ok(self.inner.lock().await.stream.finish().await?)
}
async fn send_frame(&mut self, opcode: u8, payload: &[u8]) -> Result<(), WsError> {
let frame = build_ws_frame(opcode, payload);
let frame_bytes = Bytes::from(frame);
let mut stream = self.inner.lock().await;
let send_result = if let Some(timeout) = self.write_timeout {
tokio::time::timeout(timeout, stream.stream.send_data(frame_bytes))
.await
.map_err(|_elapsed| WsError::Timeout)?
} else {
stream.stream.send_data(frame_bytes).await
};
Ok(send_result?)
}
fn try_parse_frame(
buf: &mut Vec<u8>,
max_frame_size: usize,
) -> Result<Option<WsMessage>, WsError> {
if buf.len() < 2 {
return Ok(None);
}
let byte0 = buf[0];
let byte1 = buf[1];
let _fin = byte0 & 0x80 != 0;
let opcode = byte0 & 0x0F;
let masked = byte1 & 0x80 != 0;
let mut payload_len = (byte1 & 0x7F) as u64;
let mut header_len = 2usize;
if payload_len == 126 {
if buf.len() < 4 {
return Ok(None);
}
payload_len = u16::from_be_bytes([buf[2], buf[3]]) as u64;
header_len = 4;
} else if payload_len == 127 {
if buf.len() < 10 {
return Ok(None);
}
payload_len = u64::from_be_bytes([
buf[2], buf[3], buf[4], buf[5], buf[6], buf[7], buf[8], buf[9],
]);
header_len = 10;
}
if masked {
header_len += 4;
}
let payload_len_usize = match usize::try_from(payload_len) {
Ok(n) => n,
Err(_) => return Err(WsError::FrameMalformed("payload length overflow")),
};
if payload_len_usize > max_frame_size {
return Err(WsError::FrameTooLarge);
}
let total_len = header_len + payload_len_usize;
if buf.len() < total_len {
return Ok(None);
}
let mut payload: Vec<u8> = buf[header_len..total_len].to_vec();
if masked {
let mask_key = [
buf[header_len - 4],
buf[header_len - 3],
buf[header_len - 2],
buf[header_len - 1],
];
for (i, byte) in payload.iter_mut().enumerate() {
*byte ^= mask_key[i % 4];
}
}
buf.drain(..total_len);
let msg = match opcode {
0x1 => {
let text = String::from_utf8(payload)
.map_err(|_| WsError::FrameMalformed("text frame contains invalid UTF-8"))?;
WsMessage::Text(text)
}
0x2 => WsMessage::Binary(payload),
0x8 => {
let (code, reason) = parse_close_payload(&payload);
WsMessage::Close { code, reason }
}
0x9 => WsMessage::Ping(payload),
0xA => WsMessage::Pong(payload),
_ => return Err(WsError::FrameMalformed("unknown opcode")),
};
Ok(Some(msg))
}
}
fn build_ws_frame(opcode: u8, payload: &[u8]) -> Vec<u8> {
let fin_and_opcode = 0x80 | (opcode & 0x0F); let mask_key = generate_mask_key();
let len = payload.len();
let mut frame = Vec::with_capacity(2 + 4 + len + 8); frame.push(fin_and_opcode);
if len < 126 {
frame.push(0x80 | len as u8);
} else if len <= 0xFFFF {
frame.push(0x80 | 126);
frame.extend_from_slice(&(len as u16).to_be_bytes());
} else {
frame.push(0x80 | 127);
frame.extend_from_slice(&(len as u64).to_be_bytes());
}
frame.extend_from_slice(&mask_key);
for (i, byte) in payload.iter().enumerate() {
frame.push(byte ^ mask_key[i % 4]);
}
frame
}
fn generate_mask_key() -> [u8; 4] {
use std::time::{SystemTime, UNIX_EPOCH};
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0);
let mut x = nanos as u64;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
x.to_ne_bytes()[..4]
.try_into()
.unwrap_or([0x37, 0xfa, 0x21, 0x3d])
}
fn build_close_payload(code: Option<u16>, reason: Option<&str>) -> Vec<u8> {
let mut payload = Vec::new();
if let Some(c) = code {
payload.extend_from_slice(&c.to_be_bytes());
}
if let Some(r) = reason {
payload.extend_from_slice(r.as_bytes());
}
payload
}
fn parse_close_payload(payload: &[u8]) -> (Option<u16>, Option<String>) {
let code = if payload.len() >= 2 {
Some(u16::from_be_bytes([payload[0], payload[1]]))
} else {
None
};
let reason = if payload.len() > 2 {
String::from_utf8(payload[2..].to_vec()).ok()
} else {
None
};
(code, reason)
}
#[cfg(test)]
mod tests {
use std::{net::SocketAddr, sync::Arc, time::Duration};
use bytes::Bytes;
use quinn::{Endpoint, ServerConfig, crypto::rustls::QuicServerConfig};
use rustls::{
RootCertStore, ServerConfig as RustlsServerConfig,
pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer},
};
use tower::Service;
use super::*;
use crate::{
Body,
client::{
conn::quic::{H3Connection, QuicConnector},
core::{body::Incoming as IncomingBody, dispatch, http3::Http3Options},
http::ConnectRequest,
},
};
type TestResult<T> = Result<T, Box<dyn std::error::Error + Send + Sync>>;
#[tokio::test]
async fn conn_task_drives_request_through_real_h3_server() -> TestResult<()> {
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.key_pair.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::builder_with_provider(provider)
.with_safe_default_protocol_versions()
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("bad protocol versions: {e}").into()
})?
.with_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let server_task = tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
match incoming.await {
Ok(quinn_conn) => {
tokio::spawn(async move {
let h3_quinn_conn = hpx_h3::quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<
hpx_h3::quinn::Connection,
Bytes,
> = match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => return,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (_req, mut stream) =
match resolver.resolve_request().await {
Ok(parts) => parts,
Err(_) => continue,
};
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => continue,
};
if stream.send_response(resp).await.is_err() {
continue;
}
while matches!(stream.recv_data().await, Ok(Some(_))) {}
let _ = stream.finish().await;
}
Ok(None) => break,
Err(_) => break,
}
}
});
}
Err(_) => break,
}
}
});
let _ = server_task;
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_provider)
.with_safe_default_protocol_versions()
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("bad protocol versions: {e}").into()
})?
.with_root_certificates(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_addr: SocketAddr = "127.0.0.1:0".parse()?;
let client_endpoint = Endpoint::client(client_addr)?;
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config,
Http3Options::default(),
);
let uri: http::Uri = format!("https://127.0.0.1:{}/", server_addr.port()).parse()?;
let connect_req = ConnectRequest::new(uri, None);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => return Err(format!("poll_ready should be Ok, got {other:?}").into()),
}
let h3_conn: H3Connection =
tokio::time::timeout(Duration::from_secs(5), connector.call(connect_req))
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"QuicConnector::call should resolve within 5s".into()
})??;
let (mut tx, rx) = dispatch::channel::<Request<Body>, Response<IncomingBody>>();
let conn_task: ConnTask<Body> = ConnTask::new(h3_conn.send_request, rx, h3_conn.close_rx);
let task_handle = tokio::spawn(async move { conn_task.await });
let req = Request::get(format!("https://127.0.0.1:{}/", server_addr.port()))
.body(Body::empty())?;
let promise =
tx.try_send(req)
.map_err(|_req| -> Box<dyn std::error::Error + Send + Sync> {
"dispatch::Sender::try_send failed".into()
})?;
let response_result = tokio::time::timeout(Duration::from_secs(5), promise)
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"response promise should resolve within 5s".into()
})?
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"oneshot receiver error".into()
})?;
let response =
response_result.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("dispatch returned error: {e:?}").into()
})?;
assert_eq!(
response.status(),
http::StatusCode::OK,
"expected 200 OK from ConnTask-driven request"
);
assert_eq!(
response.version(),
http::Version::HTTP_3,
"expected HTTP/3 version from ConnTask-driven request"
);
drop(tx);
let _ = tokio::time::timeout(Duration::from_secs(2), task_handle).await;
Ok(())
}
}