use std::{error::Error, sync::Arc};
use bytes::{Buf, Bytes};
use futures::{Sink, Stream, StreamExt, TryFutureExt};
use http::{
HeaderMap, HeaderValue, Method, Uri,
header::{AsHeaderName, IntoHeaderName},
uri::{Authority, PathAndQuery, Scheme},
};
use snafu::{Report, ResultExt, Snafu};
use tracing::Instrument;
use crate::{
client::Client,
dhttp::protocol::InitialRawMessageStreamError,
error::Code,
message::{
stream::{InitialMessageStreamError, MessageStreamError, ReadStream, WriteStream},
unify::{MalformedMessageError, Message, MessageStage, ReadToStringError},
},
pool::ConnectError,
qpack::field::{MalformedHeaderSection, Protocol},
quic::{self, agent},
};
#[derive(Clone)]
pub struct PendingRequest<'c, C: quic::Connect> {
client: &'c Client<C>,
request: Message,
auto_close: bool,
}
impl<'c, C: quic::Connect> std::fmt::Debug for PendingRequest<'c, C>
where
Client<C>: std::fmt::Debug,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PendingRequest")
.field("client", &self.client)
.field("request", &self.request)
.field("auto_close", &self.auto_close)
.finish()
}
}
#[derive(Debug, Snafu)]
#[snafu(module)]
pub enum RequestError<E: Error + 'static> {
#[snafu(transparent)]
Connect { source: ConnectError<E> },
#[snafu(transparent)]
Connection { source: quic::ConnectionError },
#[snafu(display("request stream error"))]
RequestStream { source: quic::StreamError },
#[snafu(display("response stream error"))]
ResponseStream { source: quic::StreamError },
#[snafu(display("request cannot be sent due to malformed header"))]
MalformedRequestHeader { source: MalformedHeaderSection },
#[snafu(display(
"header section too large to fit into a single frame, maybe too many header fields"
))]
HeaderTooLarge,
#[snafu(display(
"trailer section too large to fit into a single frame, maybe too many header fields"
))]
TrailerTooLarge,
#[snafu(display("data frame payload too large, try smaller chunk size"))]
DataFrameTooLarge,
#[snafu(display("response from peer is malformed"))]
MalformedResponse,
}
impl<C: quic::Connect> Client<C> {
pub fn new_request(&self) -> PendingRequest<'_, C> {
PendingRequest {
client: self,
request: Message::unresolved_request(),
auto_close: true,
}
}
}
impl<C: quic::Connect> PendingRequest<'_, C> {
pub fn with_method(mut self, method: Method) -> Self {
self.request.header_mut().set_method(method);
self
}
pub fn with_scheme(mut self, scheme: Scheme) -> Self {
self.request.header_mut().set_scheme(scheme);
self
}
pub fn with_authority(mut self, authority: Authority) -> Self {
self.request.header_mut().set_authority(authority);
self
}
pub fn with_path(mut self, path: PathAndQuery) -> Self {
self.request.header_mut().set_path(path);
self
}
pub fn with_protocol(mut self, protocol: Protocol) -> Self {
self.request.header_mut().set_protocol(protocol);
self
}
pub fn with_uri(mut self, uri: Uri) -> Self {
self.request.header_mut().set_uri(uri);
self
}
pub fn headers(&self) -> &HeaderMap {
&self.request.header().header_map
}
pub fn headers_mut(&mut self) -> &mut HeaderMap {
&mut self.request.header_mut().header_map
}
pub fn with_header(mut self, name: impl IntoHeaderName, value: HeaderValue) -> Self {
self.headers_mut().insert(name, value);
self
}
pub fn with_headers(mut self, headers: HeaderMap) -> Self {
*self.headers_mut() = headers;
self
}
pub fn with_body(mut self, body: impl Buf) -> Self {
self.request.set_body(body);
self
}
pub fn auto_close(mut self, auto_close: bool) -> Self {
self.auto_close = auto_close;
self
}
pub fn trailers(&self) -> &HeaderMap {
self.request.trailers()
}
pub fn trailers_mut(&mut self) -> &mut HeaderMap {
self.request.trailers_mut()
}
pub fn with_trailer(mut self, name: impl IntoHeaderName, value: HeaderValue) -> Self {
self.trailers_mut().insert(name, value);
self
}
pub fn with_trailers(mut self, trailers: HeaderMap) -> Self {
*self.trailers_mut() = trailers;
self
}
}
impl<'a, C: quic::Connect + Sync> PendingRequest<'a, C>
where
C::Connection: Send + 'static,
<C::Connection as quic::ManageStream>::StreamReader: Send,
<C::Connection as quic::ManageStream>::StreamWriter: Send,
{
#[tracing::instrument(
level = "debug",
target = "h3x::client",
name = "execute_request",
skip_all,
fields(
method = %self.request.header().method(),
uri = %self.request.header().uri(),
)
)]
pub async fn execute(mut self) -> Result<(Request, Response), RequestError<C::Error>> {
self.request
.header()
.check_pseudo()
.context(request_error::MalformedRequestHeaderSnafu)?;
if tracing::enabled!(tracing::Level::DEBUG) {
let span = tracing::Span::current();
if !span.has_field("method") {
span.record("method", self.request.header().method().as_str());
}
if !span.has_field("uri") {
span.record("uri", self.request.header().uri().to_string());
}
}
let authority = self.request.header().authority().expect("checked");
loop {
let connection = self.client.connect(authority.clone()).await?;
tracing::trace!(target: "h3x::client", %authority, "connected");
let (mut read_stream, mut write_stream) = match connection
.initial_message_stream()
.await
{
Ok(pair) => pair,
Err(InitialMessageStreamError::InitialRawStream { source }) => match source {
InitialRawMessageStreamError::Connection { source } => {
tracing::debug!(
target: "h3x::client",
?source,
"connection error on reused connection, retrying..."
);
continue;
}
InitialRawMessageStreamError::ResponseStream { source } => {
return Err(RequestError::ResponseStream { source });
}
InitialRawMessageStreamError::Goaway { .. } => {
tracing::debug!(target: "h3x::client", "connection goaway, retrying...");
continue;
}
},
Err(InitialMessageStreamError::QPackProtocolDisabled { .. }) => {
unreachable!("Client always initializes the QPack protocol")
}
};
let Ok(local_agent) = connection.local_agent().await else {
continue;
};
let Ok(remote_agent) = connection.remote_agent().await else {
continue;
};
let remote_agent = remote_agent.expect("checked by Client::connect");
let send_request = async {
if self.auto_close && self.request.is_chunked() {
write_stream.close_message(&mut self.request).await
} else {
write_stream.send_message(&mut self.request).await
}
};
let mut response = Message::unresolved_response();
#[derive(Debug, PartialEq)]
enum Stream {
Request,
Response,
}
return match tokio::try_join!(
send_request.map_err(|e| (Stream::Request, e)),
read_stream
.read_message_header(&mut response)
.map_err(|e| (Stream::Response, e)),
) {
Ok(..) => {
let request = Request {
message: self.request,
stream: write_stream,
agent: local_agent,
};
let response = Response {
message: response,
stream: read_stream,
agent: remote_agent,
};
return Ok((request, response));
}
Err((stream, MessageStreamError::HeaderTooLarge)) => {
debug_assert_eq!(stream, Stream::Response);
Err(RequestError::HeaderTooLarge)
}
Err((stream, MessageStreamError::TrailerTooLarge)) => {
debug_assert_eq!(stream, Stream::Response);
Err(RequestError::TrailerTooLarge)
}
Err((stream, MessageStreamError::DataFrameTooLarge)) => {
debug_assert_eq!(stream, Stream::Request);
Err(RequestError::DataFrameTooLarge)
}
Err((stream, MessageStreamError::MalformedIncomingMessage)) => {
debug_assert_eq!(stream, Stream::Response);
Err(RequestError::MalformedResponse)
}
Err((stream, MessageStreamError::Quic { source })) => match stream {
Stream::Request => Err(RequestError::RequestStream { source }),
Stream::Response => Err(RequestError::ResponseStream { source }),
},
Err((.., MessageStreamError::Goaway { .. })) => {
self.request = self.request.to_unsend();
tracing::debug!(target: "h3x::client", "connection goaway, retrying...");
continue;
}
};
}
}
pub async fn get(self, uri: Uri) -> Result<(Request, Response), RequestError<C::Error>> {
self.with_method(Method::GET).with_uri(uri).execute().await
}
pub async fn post(self, uri: Uri) -> Result<(Request, Response), RequestError<C::Error>> {
self.with_method(Method::POST).with_uri(uri).execute().await
}
pub async fn put(self, uri: Uri) -> Result<(Request, Response), RequestError<C::Error>> {
self.with_method(Method::PUT).with_uri(uri).execute().await
}
pub async fn delete(self, uri: Uri) -> Result<(Request, Response), RequestError<C::Error>> {
self.with_method(Method::DELETE)
.with_uri(uri)
.execute()
.await
}
pub async fn head(self, uri: Uri) -> Result<(Request, Response), RequestError<C::Error>> {
self.with_method(Method::HEAD).with_uri(uri).execute().await
}
pub async fn options(self, uri: Uri) -> Result<(Request, Response), RequestError<C::Error>> {
self.with_method(Method::OPTIONS)
.with_uri(uri)
.execute()
.await
}
pub async fn connect(self, uri: Uri) -> Result<(Request, Response), RequestError<C::Error>> {
self.with_method(Method::CONNECT)
.with_uri(uri)
.execute()
.await
}
pub async fn patch(self, uri: Uri) -> Result<(Request, Response), RequestError<C::Error>> {
self.with_method(Method::PATCH)
.with_uri(uri)
.execute()
.await
}
pub async fn trace(self, uri: Uri) -> Result<(Request, Response), RequestError<C::Error>> {
self.with_method(Method::TRACE)
.with_uri(uri)
.execute()
.await
}
}
pub struct Request {
message: Message,
stream: WriteStream,
agent: Option<Arc<dyn agent::LocalAgent>>,
}
impl Request {
pub fn method(&self) -> Method {
self.message.header().method()
}
pub fn scheme(&self) -> Option<Scheme> {
self.message.header().scheme()
}
pub fn authority(&self) -> Option<Authority> {
self.message.header().authority()
}
pub fn path(&self) -> Option<PathAndQuery> {
self.message.header().path()
}
pub fn uri(&self) -> Uri {
self.message.header().uri()
}
pub fn headers(&self) -> &HeaderMap {
&self.message.header().header_map
}
fn check_message_operation(
&mut self,
operation: &str,
check: impl FnOnce(&mut Self) -> Result<(), MalformedMessageError>,
) {
if self.message.is_malformed() {
tracing::warn!(
target: "h3x::client", operation,
"Request is malformed, operation will not affect the request stream",
);
}
if let Err(error) = check(self) {
tracing::warn!(
target: "h3x::client", operation, error = %Report::from_error(error),
"Operation malformed the request message, request stream will be cancelled with H3_REQUEST_CANCELLED",
);
self.message.set_malformed();
}
}
pub async fn write(
&mut self,
content: impl Buf + Send,
) -> Result<&mut Self, MessageStreamError> {
self.check_message_operation("write_streaming_body", |this| {
this.message.streaming_body()?;
Ok(())
});
self.stream
.send_message_streaming_body(&mut self.message, content)
.await?;
Ok(self)
}
pub async fn flush(&mut self) -> Result<&mut Self, MessageStreamError> {
self.stream.flush_message(&mut self.message).await?;
Ok(self)
}
pub fn as_sink<B: Buf + Send>(&mut self) -> impl Sink<B, Error = MessageStreamError> {
crate::message::stream::unfold::write::unfold(
self,
async |request: &mut Self, buf: B| {
request.write(buf).await?;
Ok(request)
},
async |request: &mut Self| {
request.flush().await?;
Ok(request)
},
async |request: &mut Self| {
request.close().await?;
Ok(request)
},
)
}
pub fn into_sink<B: Buf + Send>(self) -> impl Sink<B, Error = MessageStreamError> {
crate::message::stream::unfold::write::unfold(
self,
async |request: Self, buf: B| {
let mut request = request;
request.write(buf).await?;
Ok(request)
},
async |request: Self| {
let mut request = request;
request.flush().await?;
Ok(request)
},
async |request: Self| {
let mut request = request;
request.close().await?;
Ok(request)
},
)
}
pub fn trailers(&self) -> &HeaderMap {
self.message.trailers()
}
pub fn trailers_mut(&mut self) -> &mut HeaderMap {
self.check_message_operation("modify_trailers", |this| {
if this.message.stage() >= MessageStage::Trailer {
return Err(MalformedMessageError::TrailerAlreadySent);
}
Ok(())
});
self.message.trailers_mut()
}
pub fn set_trailer(&mut self, name: impl IntoHeaderName, value: HeaderValue) -> &mut Self {
self.trailers_mut().insert(name, value);
self
}
pub fn set_trailers(&mut self, map: HeaderMap) -> &mut Self {
*self.trailers_mut() = map;
self
}
pub async fn close(&mut self) -> Result<(), MessageStreamError> {
self.stream.close_message(&mut self.message).await
}
pub async fn cancel(&mut self, code: Code) -> Result<(), MessageStreamError> {
self.stream.cancel(code).await
}
pub fn write_stream(&mut self) -> &mut WriteStream {
&mut self.stream
}
pub fn agent(&self) -> Option<&Arc<dyn agent::LocalAgent>> {
self.agent.as_ref()
}
pub(crate) fn drop(&mut self) -> Option<impl Future<Output = ()> + Send + use<>> {
if self.message.is_complete() || self.message.is_dropped() {
return None;
}
let mut stream = self.stream.take();
let mut message = self.message.take();
Some(async move { _ = stream.close_message(&mut message).await })
}
}
impl Drop for Request {
fn drop(&mut self) {
if let Some(future) = self.drop() {
tokio::spawn(future.in_current_span());
}
}
}
pub struct Response {
message: Message,
stream: ReadStream,
agent: Arc<dyn agent::RemoteAgent>,
}
impl Response {
pub async fn next_response(&mut self) -> Result<&mut Self, MessageStreamError> {
self.stream.read_message_header(&mut self.message).await?;
Ok(self)
}
pub fn status(&self) -> http::StatusCode {
self.message.header().status()
}
pub fn headers(&mut self) -> &HeaderMap {
&self.message.header().header_map
}
pub fn header(&mut self, name: impl AsHeaderName) -> Option<&HeaderValue> {
self.headers().get(name)
}
pub async fn read(&mut self) -> Option<Result<Bytes, MessageStreamError>> {
self.stream.read_message(&mut self.message).await
}
pub async fn read_all(&mut self) -> Result<impl Buf, MessageStreamError> {
self.stream.read_message_full_body(&mut self.message).await
}
pub async fn read_to_bytes(&mut self) -> Result<Bytes, MessageStreamError> {
self.stream
.read_message_body_to_bytes(&mut self.message)
.await
}
pub async fn read_to_string(&mut self) -> Result<String, ReadToStringError> {
self.stream
.read_message_body_to_string(&mut self.message)
.await
}
pub async fn as_stream(&mut self) -> impl Stream<Item = Result<Bytes, MessageStreamError>> {
futures::stream::unfold(self, async |this| {
this.read().await.map(|item| (item, this))
})
.fuse()
}
pub async fn into_stream(self) -> impl Stream<Item = Result<Bytes, MessageStreamError>> {
futures::stream::unfold(self, async |mut this| {
this.read().await.map(|item| (item, this))
})
.fuse()
}
pub async fn trailers(&mut self) -> Result<&HeaderMap, MessageStreamError> {
self.stream.read_message_trailer(&mut self.message).await
}
pub async fn stop(&mut self, code: Code) -> Result<(), MessageStreamError> {
self.stream.stop(code).await
}
pub fn read_stream(&mut self) -> &mut ReadStream {
&mut self.stream
}
pub fn agent(&self) -> &Arc<dyn agent::RemoteAgent> {
&self.agent
}
}