use crate::courierust_body::Body;
use crate::courierust_bytes::Bytes;
use crate::courierust_client::ClientConfig;
use crate::courierust_error::{Error, Result};
use crate::courierust_h1;
use crate::courierust_http::header::{HeaderMap, HeaderName, HeaderValue};
use crate::courierust_http::request::Request;
use crate::courierust_http::response::{Response, ResponseHead};
use crate::courierust_http::status::StatusCode;
use crate::courierust_http::version::Version;
use crate::courierust_io::{BufReader, BufWriter, Scratch};
use crate::courierust_net::{self, ConnStream};
use std::net::SocketAddr;
use std::sync::Arc;
pub struct H1Connection {
stream: Arc<ConnStream>,
reader: BufReader<Arc<ConnStream>>,
writer: BufWriter<Arc<ConnStream>>,
scratch: Scratch,
version: Version,
reusable: bool,
}
impl H1Connection {
pub fn connect(
addr: SocketAddr,
tls: Option<&crate::courierust_tls::TlsConnector>,
hostname: &str,
cfg: &ClientConfig,
) -> Result<Self> {
let stream = courierust_net::connect(&addr, cfg.connect_timeout)?;
let conn = match tls {
Some(c) => {
courierust_net::configure(&stream, cfg.handshake_timeout)?;
let conn = ConnStream::tls_client(stream, c, hostname)?;
if let Some(alpn) = conn.alpn() {
if alpn.as_slice() == b"h2" {
return Err(Error::protocol(
"server negotiated h2 via ALPN, but the client is configured for HTTP/1.1",
));
}
}
conn
}
None => {
courierust_net::configure(&stream, cfg.read_timeout)?;
ConnStream::plain(stream)
}
};
let _ = conn.configure(cfg.read_timeout);
let conn = Arc::new(conn);
Ok(Self {
reader: BufReader::new(conn.clone(), 16 * 1024),
writer: BufWriter::new(conn.clone(), 16 * 1024),
stream: conn,
scratch: Scratch::new(),
version: Version::HTTP_11,
reusable: true,
})
}
pub(crate) fn from_stream_seeded(
stream: ConnStream,
cfg: &ClientConfig,
seed: &[u8],
) -> Result<Self> {
let _ = stream.configure(cfg.read_timeout);
let conn = Arc::new(stream);
let mut reader = BufReader::new(conn.clone(), 16 * 1024);
if !seed.is_empty() {
reader.seed(seed);
}
Ok(Self {
reader,
writer: BufWriter::new(conn.clone(), 16 * 1024),
stream: conn,
scratch: Scratch::new(),
version: Version::HTTP_11,
reusable: true,
})
}
pub fn is_reusable(&self) -> bool {
self.reusable
}
pub fn peer_addr(&self) -> SocketAddr {
self.stream.peer_addr()
}
pub fn send(
&mut self,
req: &Request<Body>,
cfg: &ClientConfig,
host_header: &str,
) -> Result<Response<Body>> {
let mut headers = HeaderMap::with_capacity(req.headers.len() + 4);
for (n, v) in req.headers.iter() {
if courierust_h1::is_hop_by_hop(n.as_str()) {
continue;
}
headers.append(n.clone(), v.clone());
}
headers.insert(
HeaderName::from_lowercase("host"),
HeaderValue::from_bytes(host_header.as_bytes())?,
);
if !headers.contains_key("user-agent") {
if let Some(ua) = &cfg.user_agent {
headers.insert(
HeaderName::from_lowercase("user-agent"),
HeaderValue::from_bytes(ua.as_bytes())?,
);
}
}
let body = match &req.body {
Body::Empty => None,
Body::Bytes(b) => Some(b),
Body::Channel(_) => {
return Err(Error::protocol("streaming request bodies require h2"));
}
};
if let Some(b) = body {
let cl = courierust_h1::IToA::new(b.len());
headers.insert(
HeaderName::from_lowercase("content-length"),
HeaderValue::from_bytes(cl.as_slice())?,
);
}
let head = self.scratch.body();
courierust_h1::write_request_head(head, &req.method, &req.uri, Version::HTTP_11, &headers)?;
self.writer.write_all(head)?;
if let Some(b) = body {
self.writer.write_all(b)?;
}
self.writer.flush()?;
self.read_response(cfg)
}
fn read_response(&mut self, cfg: &ClientConfig) -> Result<Response<Body>> {
let (reader, scratch) = (&mut self.reader, &mut self.scratch);
let status_line = scratch.line();
reader.read_until_into(b'\n', 16 * 1024, status_line)?;
let (status, version) = courierust_h1::parse_status_line(status_line)?;
self.version = version;
let mut status = status;
let mut headers = courierust_h1::read_headers_scratch(reader, scratch)?;
while status.is_informational() {
let line = scratch.line();
reader.read_until_into(b'\n', 16 * 1024, line)?;
let (s, _) = courierust_h1::parse_status_line(line)?;
status = s;
headers = courierust_h1::read_headers_scratch(reader, scratch)?;
}
let head = ResponseHead {
status,
version,
headers,
};
self.finish_response(cfg, head)
}
pub fn finish_response(
&mut self,
cfg: &ClientConfig,
head: ResponseHead,
) -> Result<Response<Body>> {
let status = head.status;
let version = head.version;
let (reader, scratch) = (&mut self.reader, &mut self.scratch);
let mut close_delimited = false;
let body = match courierust_h1::body_length(&head.headers, None, Some(status))? {
courierust_h1::BodyLen::None => {
if status == StatusCode::NO_CONTENT
|| status == StatusCode::NOT_MODIFIED
|| status.is_informational()
|| status.is_redirection()
{
Body::Empty
} else {
close_delimited = true;
Body::Bytes(read_until_eof_scratch(reader, cfg.max_body, scratch)?)
}
}
courierust_h1::BodyLen::Length(0) => Body::Empty,
courierust_h1::BodyLen::Length(n) => Body::Bytes(
courierust_h1::read_body_fixed_scratch(reader, n, cfg.max_body, scratch)?,
),
courierust_h1::BodyLen::Chunked => Body::Bytes(
courierust_h1::read_body_chunked_scratch(reader, cfg.max_body, scratch)?,
),
};
self.reusable = !close_delimited
&& !courierust_h1::wants_close(&head.headers)
&& version == Version::HTTP_11;
Ok(head.with_body(body))
}
}
fn read_until_eof_scratch(
reader: &mut BufReader<Arc<ConnStream>>,
max: usize,
scratch: &mut Scratch,
) -> Result<Bytes> {
let out = scratch.body();
loop {
let b = match reader.fill_buf() {
Ok([]) => break,
Ok(b) => b,
Err(_) => break,
};
let n = b.len();
if out.len() + n > max {
return Err(Error::overflow("body exceeds limit"));
}
out.extend_from_slice(b);
reader.consume(n);
}
Ok(Bytes::from(core::mem::take(out)))
}