use crate::config::Config;
use crate::dns::DnsResolver;
use crate::error::{Error, SocketError};
use crate::headers::{Headers, WellKnownHeader, well_known_header_bytes};
use crate::parser::chunked::{ChunkedDecoder, FeedResult};
use crate::parser::has_complete_headers;
use crate::parser::uri::{Host, Uri};
use crate::parser::version::Version;
use crate::parser::{BodyReadStrategy, Response};
use crate::socket::BlockingSocket;
use crate::transport::pool::PooledBuffers;
use crate::util::IpAddr;
use alloc::string::String;
use alloc::vec::Vec;
use bytes::{Bytes, BytesMut};
use core::net::SocketAddr;
use core::time::Duration;
#[derive(Debug, Clone)]
pub struct RawResponse {
pub status_code: u16,
pub reason: String,
pub headers: Headers,
pub version: Version,
pub body_bytes: Bytes,
pub decoded_chunked_trailers: Option<Headers>,
}
pub struct Connection<'a, S> {
socket: &'a mut S,
max_header_size: usize,
max_body_size: usize,
reusable: bool,
buf: BytesMut,
}
impl<'a, S: BlockingSocket> Connection<'a, S> {
#[cfg_attr(not(test), allow(dead_code))]
pub fn new(
socket: &'a mut S,
max_header_size: usize,
max_body_size: usize,
) -> Self {
Self::with_buffers(socket, max_header_size, max_body_size, PooledBuffers::default())
}
pub fn with_buffers(
socket: &'a mut S,
max_header_size: usize,
max_body_size: usize,
buffers: PooledBuffers,
) -> Self {
Self {
socket,
max_header_size,
max_body_size,
reusable: true,
buf: buffers.buf,
}
}
#[must_use]
pub fn take_buffers(&mut self) -> PooledBuffers {
self.buf.clear();
PooledBuffers {
buf: core::mem::replace(&mut self.buf, BytesMut::new()),
}
}
pub fn send_request(
&mut self,
head: &[u8],
body: &[u8],
) -> Result<(), Error> {
self.write_all_vectored(&[head, body])?;
if request_has_connection_close(head) {
self.reusable = false;
}
Ok(())
}
fn write_all_vectored(
&mut self,
bufs: &[&[u8]],
) -> Result<(), Error> {
let mut idx = 0usize;
let mut off = 0usize;
loop {
while idx < bufs.len() {
let Some(cur) = bufs.get(idx).copied() else {
return Ok(());
};
if off < cur.len() {
break;
}
idx = idx.saturating_add(1);
off = 0;
}
if idx >= bufs.len() {
return Ok(());
}
let Some(cur) = bufs.get(idx).copied() else {
return Ok(());
};
let first = cur.get(off..).unwrap_or(&[]);
let second = bufs.get(idx.saturating_add(1)).copied().unwrap_or(&[]);
let n = if second.is_empty() {
self.socket.write(first).map_err(Error::Socket)?
} else {
self
.socket
.write_vectored(&[first, second])
.map_err(Error::Socket)?
};
if n == 0 {
return Err(Error::Socket(SocketError::NotConnected));
}
let mut remaining = n;
while remaining > 0 {
let Some(slice) = bufs.get(idx).copied() else {
break;
};
let avail = slice.len().saturating_sub(off);
if remaining < avail {
off = off.saturating_add(remaining);
remaining = 0;
} else {
remaining = remaining.saturating_sub(avail);
idx = idx.saturating_add(1);
off = 0;
}
}
}
}
pub fn read_raw_response(
&mut self,
expect_body: bool,
) -> Result<RawResponse, Error> {
let max_header_size = self.max_header_size;
let chunk_cap = max_header_size.min(8192);
self.buf.clear();
if self.buf.capacity() < chunk_cap {
self
.buf
.reserve(chunk_cap.saturating_sub(self.buf.capacity()));
}
loop {
while !has_complete_headers(&self.buf) {
if self.buf.len() > max_header_size {
return Err(Error::ResponseHeaderTooLarge);
}
let n = self.read_socket_into_buf(chunk_cap)?;
if n == 0 {
return Err(Error::Socket(SocketError::NotConnected));
}
if headers_section_len(&self.buf).is_some_and(|hdr_len| hdr_len > max_header_size)
|| (!has_complete_headers(&self.buf) && self.buf.len() > max_header_size)
{
return Err(Error::ResponseHeaderTooLarge);
}
}
let (status_code, reason_bytes, header_refs, version, remaining_after_headers) =
Response::scan_headers_only(&self.buf).map_err(Error::Parse)?;
let consumed = self.buf.len().saturating_sub(remaining_after_headers.len());
if (100..200).contains(&status_code) {
let _ = self.buf.split_to(consumed);
continue;
}
let body_strategy = if expect_body {
match Response::body_read_strategy_refs(&header_refs, status_code, version) {
Ok(s) => Some(s),
Err(e) => {
self.reusable = false;
return Err(Error::Parse(e));
},
}
} else {
None
};
let reason = Response::reason_owned(reason_bytes);
let wire_spans = Response::try_wire_header_spans(self.buf.get(..consumed).unwrap_or(&[]), &header_refs);
let headers = if let Some(spans) = wire_spans {
Headers::from_spans(self.buf.split_to(consumed).freeze(), spans)
} else {
let headers = Response::headers_from_refs(&header_refs);
let _ = self.buf.split_to(consumed);
headers
};
let (body_bytes, decoded_chunked_trailers) = if let Some(strategy) = body_strategy {
if matches!(strategy, BodyReadStrategy::UntilClose) {
self.reusable = false;
}
self.read_body(strategy)?
} else {
if !self.buf.is_empty() {
self.reusable = false;
self.buf.clear();
}
(Bytes::new(), None)
};
if connection_option_present(&headers, "close") {
self.reusable = false;
} else if !version.defaults_to_persistent() {
if !connection_option_present(&headers, "keep-alive") {
self.reusable = false;
}
}
return Ok(RawResponse {
status_code,
reason,
headers,
version,
body_bytes,
decoded_chunked_trailers,
});
}
}
fn read_socket_into_buf(
&mut self,
max: usize,
) -> Result<usize, Error> {
if max == 0 {
return Ok(0);
}
let existing_spare = self.buf.capacity().saturating_sub(self.buf.len());
if existing_spare < max {
self.buf.reserve(max.saturating_sub(existing_spare));
}
let Self { socket, buf, .. } = self;
let uninit = buf.spare_capacity_mut();
let to_read = uninit.len().min(max);
if to_read == 0 {
return Ok(0);
}
let dst = unsafe { core::slice::from_raw_parts_mut(uninit.as_mut_ptr().cast::<u8>(), to_read) };
match socket.read(dst) {
Ok(n) => {
unsafe {
buf.set_len(buf.len().saturating_add(n));
}
Ok(n)
},
Err(e) => {
if e == SocketError::TimedOut {
let _ = socket.shutdown();
}
Err(Error::Socket(e))
},
}
}
fn read_body(
&mut self,
strategy: BodyReadStrategy,
) -> Result<(Bytes, Option<Headers>), Error> {
let max_body = self.max_body_size;
match strategy {
BodyReadStrategy::NoBody => {
if !self.buf.is_empty() {
self.reusable = false;
self.buf.clear();
}
Ok((Bytes::new(), None))
},
BodyReadStrategy::ContentLength(len) => {
if len > max_body {
self.reusable = false;
return Err(Error::BodyExceedsLimit(max_body));
}
if self.buf.len() > len {
self.reusable = false;
}
let bytes_needed = len.saturating_sub(self.buf.len().min(len));
if self.buf.len() > len {
self.buf.truncate(len);
}
if bytes_needed > 0 {
self.buf.reserve(bytes_needed);
let mut bytes_read = 0usize;
while bytes_read < bytes_needed {
let to_read = bytes_needed.saturating_sub(bytes_read);
let n = self.read_socket_into_buf(to_read)?;
if n == 0 {
return Err(Error::Socket(SocketError::NotConnected));
}
bytes_read = bytes_read.saturating_add(n);
}
}
Ok((self.buf.split_to(len).freeze(), None))
},
BodyReadStrategy::Chunked => {
if self.buf.len() > max_body {
self.reusable = false;
return Err(Error::BodyExceedsLimit(max_body));
}
let mut decoder = ChunkedDecoder::new();
let mut decoded = Vec::new();
loop {
match decoder.feed(self.buf.as_ref(), Some(&mut decoded)) {
Ok(FeedResult::Done { rest }) => {
let rest_len = rest.len();
let framed = self.buf.len().saturating_sub(rest_len);
if rest_len > 0 {
self.reusable = false;
let _ = self.buf.split_to(framed);
} else {
self.buf.clear();
}
if decoded.len() > max_body {
self.reusable = false;
return Err(Error::BodyExceedsLimit(max_body));
}
return Ok((Bytes::from(decoded), Some(decoder.take_trailers())));
},
Ok(FeedResult::NeedMore { consumed }) => {
if consumed > 0 {
let _ = self.buf.split_to(consumed);
}
if decoded.len() > max_body {
self.reusable = false;
return Err(Error::BodyExceedsLimit(max_body));
}
},
Err(e) => {
self.reusable = false;
return Err(Error::Parse(e));
},
}
let n = self.read_socket_into_buf(8192)?;
if n == 0 {
return Err(Error::Socket(SocketError::NotConnected));
}
if self.buf.len() > max_body {
self.reusable = false;
return Err(Error::BodyExceedsLimit(max_body));
}
}
},
BodyReadStrategy::UntilClose => {
if self.buf.len() > max_body {
self.reusable = false;
return Err(Error::BodyExceedsLimit(max_body));
}
loop {
let n = self.read_socket_into_buf(8192)?;
if n == 0 {
break;
}
if self.buf.len() > max_body {
self.reusable = false;
return Err(Error::BodyExceedsLimit(max_body));
}
}
Ok((self.buf.split().freeze(), None))
},
}
}
pub const fn is_reusable(&self) -> bool {
self.reusable
}
}
#[cfg_attr(not(test), allow(dead_code))]
pub fn connect<'a, S, D>(
socket: &'a mut S,
dns: &D,
uri: &Uri,
config: &Config,
reused: bool,
) -> Result<Connection<'a, S>, Error>
where
S: BlockingSocket,
D: DnsResolver,
{
connect_with_buffers(socket, dns, uri, config, reused, PooledBuffers::default())
}
pub fn connect_with_buffers<'a, S, D>(
socket: &'a mut S,
dns: &D,
uri: &Uri,
config: &Config,
reused: bool,
buffers: PooledBuffers,
) -> Result<Connection<'a, S>, Error>
where
S: BlockingSocket,
D: DnsResolver,
{
if !reused {
let authority = uri.authority().ok_or(Error::InvalidUrl)?;
let port = uri.port_or_default();
let host_for_sni = match authority.host() {
Host::RegName(name) => String::from(*name),
Host::IpAddr(addr) => crate::util::format_ip_for_host(*addr),
};
let addresses: Vec<IpAddr> = match authority.host() {
Host::RegName(name) => dns.resolve(name).map_err(Error::Dns)?,
Host::IpAddr(addr) => alloc::vec![*addr],
};
if addresses.is_empty() {
return Err(Error::Dns(crate::error::DnsError::NoAddressesFound));
}
let connect_ms = duration_ms_u32(config.timeout_connect()).unwrap_or(0);
socket
.set_connect_timeout(connect_ms)
.map_err(Error::Socket)?;
let mut last_error = None;
for addr in &addresses {
match socket.connect(&SocketAddr::new(*addr, port), host_for_sni.as_str()) {
Ok(()) => {
last_error = None;
break;
},
Err(e) => last_error = Some(e),
}
}
if let Some(e) = last_error {
return Err(Error::Socket(e));
}
}
apply_io_timeouts(socket, config)?;
Ok(Connection::with_buffers(
socket,
config.max_response_header_size(),
config.max_response_body_size(),
buffers,
))
}
fn apply_io_timeouts<S: BlockingSocket>(
socket: &mut S,
config: &Config,
) -> Result<(), Error> {
let read_ms = duration_ms_u32(config.timeout_read()).unwrap_or(0);
socket.set_read_timeout(read_ms).map_err(Error::Socket)?;
let write_ms = duration_ms_u32(config.timeout_write()).unwrap_or(0);
socket.set_write_timeout(write_ms).map_err(Error::Socket)?;
Ok(())
}
fn duration_ms_u32(d: Option<Duration>) -> Option<u32> {
Some(u32::try_from(d?.as_millis()).unwrap_or(u32::MAX))
}
fn headers_section_len(data: &[u8]) -> Option<usize> {
data
.windows(4)
.position(|w| w == b"\r\n\r\n")
.map(|i| i + 4)
.or_else(|| data.windows(2).position(|w| w == b"\n\n").map(|i| i + 2))
}
fn connection_option_present(
headers: &Headers,
option: &str,
) -> bool {
headers
.get_all(Headers::CONNECTION)
.iter()
.any(|v| v.split(',').any(|t| t.trim().eq_ignore_ascii_case(option)))
}
fn request_has_connection_close(request_bytes: &[u8]) -> bool {
let headers = headers_section_len(request_bytes).map_or(request_bytes, |n| request_bytes.get(..n).unwrap_or(&[]));
let mut lines = headers.split(|&b| b == b'\n');
let _ = lines.next();
for raw_line in lines {
let line = raw_line.strip_suffix(b"\r").unwrap_or(raw_line);
if line.is_empty() {
break;
}
let Some(colon) = line.iter().position(|&b| b == b':') else {
continue;
};
let name = line.get(..colon).unwrap_or(&[]);
if well_known_header_bytes(name) != Some(WellKnownHeader::Connection) {
continue;
}
let value_bytes = line.get(colon + 1..).unwrap_or(&[]);
let Ok(value) = core::str::from_utf8(value_bytes) else {
continue;
};
if value
.split(',')
.any(|t| t.trim().eq_ignore_ascii_case("close"))
{
return true;
}
}
false
}