use std::{
collections::VecDeque,
mem::{ManuallyDrop, MaybeUninit},
os::fd::IntoRawFd,
};
use color_eyre::{
Result,
eyre::{Context, ContextCompat, bail, eyre},
};
use heapless::index_map::FnvIndexMap;
use io_uring::{IoUring, cqueue};
use smallvec::SmallVec;
use crate::request::Request;
const BUF_SIZE: usize = 16 * 1024;
#[derive(Default, Debug, Clone)]
struct RecvTypeSized {
header_boundary: usize,
expect_body_size: usize,
}
#[derive(Default, Debug, Clone)]
enum RecvType {
#[default]
NotYetKnown,
Sized(RecvTypeSized),
Chunked,
}
impl RecvType {
fn from_headers(headers: &[httparse::Header], header_boundary: usize) -> Result<Self> {
for h in headers.iter() {
if h.name.eq_ignore_ascii_case("content-length") {
let size: usize = str::from_utf8(h.value)?.parse()?;
return Ok(Self::Sized(RecvTypeSized {
header_boundary,
expect_body_size: size,
}));
} else if h.name.eq_ignore_ascii_case("transfer-encoding")
&& h.value.eq_ignore_ascii_case("chunked".as_bytes())
{
return Ok(Self::Chunked);
}
}
Ok(Self::NotYetKnown)
}
}
#[derive(Default, Debug)]
struct RecvState {
typ: RecvType,
resp: Option<Vec<u8>>,
half_packet: Option<Vec<u8>>,
}
impl RecvState {
#[tracing::instrument(skip(packet), fields(packet_len = %packet.len()))]
fn try_from_first_packet(packet: &mut &[u8]) -> Result<Self> {
let mut headers = [httparse::EMPTY_HEADER; 128];
let mut response = httparse::Response::new(&mut headers);
let header_boundary = match response.parse(packet) {
Ok(httparse::Status::Complete(header_boundary)) => header_boundary,
Ok(httparse::Status::Partial) => {
return Ok(Self {
typ: RecvType::NotYetKnown,
resp: None,
half_packet: Some(packet.to_vec()),
});
}
Err(e) => return Err(e.into()),
};
let recv_type = RecvType::from_headers(response.headers, header_boundary)
.wrap_err("todo recv type parsing failed")?;
tracing::debug!(header_boundary, recv_type = ?recv_type, "parsed headers from first packet");
let capacity = match (recv_type.clone(), packet.len() == BUF_SIZE) {
(RecvType::Sized(size), _) => size.expect_body_size,
(RecvType::NotYetKnown, false) => packet.len(),
_ => BUF_SIZE * 2,
};
let mut v = Vec::with_capacity(capacity);
v.extend_from_slice(&packet[..header_boundary]);
*packet = &packet[header_boundary..];
Ok(Self {
typ: recv_type,
resp: Some(v),
half_packet: None,
})
}
fn try_from_next_packet(&mut self, packet: &mut &[u8]) -> Result<()> {
let half_packet = self
.half_packet
.as_mut()
.expect("should be called on half filled one");
half_packet.extend_from_slice(packet);
let mut headers = [httparse::EMPTY_HEADER; 128];
let mut response = httparse::Response::new(&mut headers);
let header_boundary = match response.parse(half_packet) {
Ok(httparse::Status::Complete(header_boundary)) => header_boundary,
Ok(httparse::Status::Partial) => {
return Ok(());
}
Err(e) => return Err(e.into()),
};
let recv_type = RecvType::from_headers(response.headers, header_boundary)
.wrap_err("todo recv type parsing failed")?;
let capacity = match (recv_type.clone(), half_packet.len() == BUF_SIZE) {
(RecvType::Sized(size), _) => size.expect_body_size,
(RecvType::NotYetKnown, false) => half_packet.len(),
_ => BUF_SIZE * 2,
};
let mut v = Vec::with_capacity(capacity);
v.extend_from_slice(&half_packet[..header_boundary]);
self.resp.replace(v);
let _ = half_packet.drain(..header_boundary);
Ok(())
}
}
#[derive(Clone, Debug, PartialEq)]
enum FlowState {
Connecting,
Sending,
Receiving,
Closing,
}
struct FlowMetadataConnecting {
socket_address: socket2::SockAddr,
request: Request,
}
#[repr(C)]
union FlowMetadata {
connecting: ManuallyDrop<FlowMetadataConnecting>,
}
#[derive(Debug, Clone)]
struct Flow {
pool_id: u16,
state: FlowState,
fd: io_uring::types::Fd,
}
pub struct Concuring<const C: usize>
where
bitmaps::BitsImpl<C>: bitmaps::Bits,
{
ring: IoUring,
occupancy: bitmaps::Bitmap<C>,
flows: FnvIndexMap<u32, Flow, C>,
metadata: [MaybeUninit<FlowMetadata>; C],
buffers: [MaybeUninit<[u8; BUF_SIZE]>; C],
recv_states: [RecvState; C],
pub finished_responses: VecDeque<(u32, Result<http::Response<Vec<u8>>>)>,
}
impl<const C: usize> Concuring<C>
where
bitmaps::BitsImpl<C>: bitmaps::Bits,
{
pub fn new() -> Result<Self> {
assert!(
C.is_power_of_two() && C <= 1024 && C >= 2,
"C must be a power of two between 2 and 1024, but was {}",
C
);
Ok(Self {
ring: io_uring::IoUring::new((C * 4) as u32)?,
occupancy: Default::default(),
flows: Default::default(),
buffers: unsafe { MaybeUninit::uninit().assume_init() },
metadata: unsafe { MaybeUninit::uninit().assume_init() },
recv_states: array_init::array_init(|_| Default::default()),
finished_responses: Default::default(),
})
}
#[tracing::instrument(skip(self, request), fields(flow_id, occupancy = %self.occupancy.len()))]
pub fn try_submit(&mut self, flow_id: u32, request: Request) -> Result<()> {
let pool_id = self
.occupancy
.first_false_index()
.wrap_err("concurrency limit")? as u16;
tracing::debug!(slot = %pool_id, "found available slot");
let (socket, socket_address) = request
.create_socket()
.wrap_err_with(|| format!("create_socket for flow {}", flow_id))?;
let fd = io_uring::types::Fd(socket.into_raw_fd());
let flow = Flow {
pool_id,
state: FlowState::Connecting,
fd,
};
match self.flows.entry(flow_id) {
heapless::index_map::Entry::Occupied(_) => bail!("duplicate flow_id: {}", flow_id),
heapless::index_map::Entry::Vacant(vacant) => vacant
.insert(flow)
.expect("map should not be full after occupancy check"),
};
let metadata_slot = &mut self.metadata[pool_id as usize];
metadata_slot.write(FlowMetadata {
connecting: ManuallyDrop::new(FlowMetadataConnecting {
socket_address,
request,
}),
});
let addr = unsafe { &metadata_slot.assume_init_ref().connecting.socket_address };
self.occupancy.set(pool_id as usize, true);
let addr_ptr = addr.as_ptr() as *const _;
let addr_len = addr.len();
let sqe_connect = io_uring::opcode::Connect::new(fd, addr_ptr, addr_len)
.build()
.user_data(flow_id as u64);
unsafe {
self.ring
.submission()
.push(&sqe_connect)
.expect("submission queue is full");
}
self.ring.submit().wrap_err("submitting to io_uring")?;
Ok(())
}
#[tracing::instrument(skip(self, cqe), fields(flow_id = %cqe.user_data(), result = %cqe.result()))]
fn handle_negative_cqe(&mut self, cqe: &cqueue::Entry) {
let flow_id = cqe.user_data() as u32;
let mut flow_entry = match self.flows.entry(flow_id) {
heapless::index_map::Entry::Occupied(occupied_entry) => occupied_entry,
heapless::index_map::Entry::Vacant(_) => panic!("should always exist"),
};
let flow = flow_entry.get();
let pool_id = flow.pool_id as usize;
let mut sub_q = self.ring.submission();
let close_sqe = io_uring::opcode::Close::new(flow.fd)
.build()
.user_data(cqe.user_data());
match flow.state {
FlowState::Connecting => {
tracing::error!(fd = %flow.fd.0, "connect failed");
flow_entry.get_mut().state = FlowState::Closing;
unsafe {
sub_q.push(&close_sqe).expect("submission for queue failed");
}
self.finished_responses
.push_back((flow_id, Err(eyre!("conn error {}", cqe.result()))));
}
FlowState::Sending => {
tracing::error!(fd = %flow.fd.0, "send failed");
self.recv_states[pool_id] = Default::default();
flow_entry.get_mut().state = FlowState::Closing;
unsafe {
sub_q.push(&close_sqe).expect("submission for queue failed");
}
self.finished_responses
.push_back((flow_id, Err(eyre!("send error {}", cqe.result()))));
}
FlowState::Receiving => {
todo!();
}
FlowState::Closing => {
let _ = flow_entry.remove();
self.occupancy.set(pool_id, false);
} };
}
#[tracing::instrument(skip(self, cqe, flow), fields(flow_id = %cqe.user_data(), fd = %flow.fd.0, slot = %flow.pool_id))]
fn handle_connecting(&mut self, cqe: &cqueue::Entry, flow: Flow) -> FlowState {
tracing::debug!("connection established, transitioning to Sending");
let metadata_slot = &mut self.metadata[flow.pool_id as usize];
let metadata_connecting = unsafe { &metadata_slot.assume_init_ref().connecting };
let serialized = &metadata_connecting.request.data;
let sqe_send = io_uring::opcode::Send::new(
flow.fd,
serialized.as_ptr() as *const _,
serialized.len() as _,
)
.build()
.user_data(cqe.user_data());
let mut sub_q = self.ring.submission();
unsafe {
sub_q.push(&sqe_send).expect("submission queue is full");
}
FlowState::Sending
}
#[tracing::instrument(skip(self, cqe, flow), fields(flow_id = %cqe.user_data(), fd = %flow.fd.0, slot = %flow.pool_id, bytes_sent = %cqe.result()))]
fn handle_sending(&mut self, cqe: &cqueue::Entry, flow: Flow) -> FlowState {
tracing::debug!("send completed, transitioning to Receiving");
let pool_id = flow.pool_id as usize;
let metadata_slot = &mut self.metadata[pool_id];
unsafe {
let metadata_connecting = &mut metadata_slot.assume_init_mut().connecting;
ManuallyDrop::drop(metadata_connecting);
}
let sqe_recv = io_uring::opcode::Recv::new(
flow.fd,
self.buffers[pool_id].as_mut_ptr() as *mut u8,
BUF_SIZE as _,
)
.build()
.user_data(cqe.user_data());
let mut sub_q = self.ring.submission();
unsafe {
sub_q.push(&sqe_recv).expect("submission queue is full");
}
FlowState::Receiving
}
fn read_chunks(
prev_packet: &mut Option<Vec<u8>>,
packet: &[u8],
into: &mut Vec<u8>,
) -> Result<bool, httparse::InvalidChunkSize> {
let mut packet = if let Some(prev_packet) = prev_packet
&& !prev_packet.is_empty()
{
prev_packet.extend_from_slice(packet);
prev_packet.as_slice()
} else {
packet
};
while !packet.is_empty() {
match httparse::parse_chunk_size(packet)? {
httparse::Status::Complete((offset, length)) => {
let length = length as usize;
let next_start = offset + length + 2;
if next_start > packet.len() {
break; }
if length == 0 {
if let Some(prev_packet) = prev_packet {
prev_packet.clear();
}
return Ok(false);
}
into.extend_from_slice(&packet[offset..offset + length]);
packet = &packet[next_start..];
}
httparse::Status::Partial => {
break; }
}
}
let v = packet.to_vec();
let _ = prev_packet.replace(v);
Ok(true)
}
#[tracing::instrument(skip(self, cqe, flow), fields(flow_id = %cqe.user_data(), fd = %flow.fd.0, slot = %flow.pool_id, bytes_received = %cqe.result()))]
fn handle_receiving(&mut self, cqe: &cqueue::Entry, flow: Flow) -> FlowState {
let pool_id = flow.pool_id as usize;
let flow_id = cqe.user_data() as u32;
macro_rules! finalize {
() => {{
tracing::debug!(slot = %pool_id, "finalizing response");
let partial = self.recv_states[pool_id].resp.take();
let res_resp = match partial {
Some(partial) => parse_response(&partial),
None => Err(eyre!("response parsing {}", flow_id)),
};
self.finished_responses.push_back((flow_id, res_resp));
let sqe_close = io_uring::opcode::Close::new(flow.fd)
.build()
.user_data(cqe.user_data());
unsafe {
self.ring
.submission()
.push(&sqe_close)
.expect("submission queue is full");
}
FlowState::Closing
}};
}
let bytes_read = cqe.result();
if bytes_read > 0 {
let buf = &mut self.buffers[pool_id];
let buf = unsafe { buf.assume_init_mut() };
macro_rules! recv_more {
() => {{
let buf = &mut self.buffers[pool_id];
let buf = unsafe { buf.assume_init_mut() };
let sqe_recv =
io_uring::opcode::Recv::new(flow.fd, buf.as_mut_ptr(), buf.len() as _)
.build()
.user_data(cqe.user_data());
unsafe {
self.ring
.submission()
.push(&sqe_recv)
.expect("submission queue is full");
}
FlowState::Receiving
}};
}
let mut packet = &buf[..(bytes_read as _)];
let recv_state = &mut self.recv_states[pool_id];
tracing::trace!(
slot = %pool_id,
recv_type = ?recv_state.typ,
has_resp = %recv_state.resp.is_some(),
has_half_packet = %recv_state.half_packet.is_some(),
"processing received data"
);
if recv_state.half_packet.is_none() && recv_state.resp.is_none() {
*recv_state = RecvState::try_from_first_packet(&mut packet).expect("todo");
if recv_state.resp.is_none() {
return recv_more!();
}
} else if recv_state.resp.is_none() {
recv_state.try_from_next_packet(&mut packet).expect("todo");
if recv_state.resp.is_none() {
return recv_more!();
}
}
let recv_buf = recv_state.resp.as_mut().expect("its not none we checked");
match &recv_state.typ {
RecvType::Sized(size) => {
recv_buf.extend_from_slice(packet);
if recv_buf.len() >= size.expect_body_size + size.header_boundary {
finalize!()
} else {
recv_more!()
}
}
RecvType::Chunked => {
match Self::read_chunks(&mut recv_state.half_packet, packet, recv_buf) {
Ok(true) => recv_more!(),
Ok(false) => finalize!(),
Err(e) => panic!("{}", e),
}
}
RecvType::NotYetKnown => {
recv_buf.extend_from_slice(packet);
recv_more!()
}
}
} else if bytes_read == 0 {
finalize!()
} else {
unreachable!("we checked the cqe.result() before");
}
}
#[tracing::instrument(skip(self), fields(active_flows = %self.flows.len(), occupancy = %self.occupancy.len()))]
pub fn wait(&mut self) -> Option<()> {
let finished_resp_len_prev = self.finished_responses.len();
loop {
if self.flows.is_empty() {
return None;
}
self.ring.submit_and_wait(1).expect("waiting ring");
let cqes: SmallVec<[cqueue::Entry; 32]> = self.ring.completion().collect();
tracing::trace!(cqe_count = %cqes.len(), "processing completion queue entries");
for cqe in cqes {
let flow_id = cqe.user_data() as u32;
if cqe.result() < 0 {
self.handle_negative_cqe(&cqe);
continue;
}
let flow = self
.flows
.get(&flow_id)
.expect("should always exist")
.clone();
let pool_id = flow.pool_id as usize;
if flow.state == FlowState::Closing {
self.occupancy.set(pool_id, false);
let _ = self.flows.remove(&flow_id);
continue;
}
let old_state = flow.state.clone();
let new_state = match flow.state.clone() {
FlowState::Connecting => {
self.handle_connecting(&cqe, flow)
}
FlowState::Sending => {
self.handle_sending(&cqe, flow)
}
FlowState::Receiving => {
self.handle_receiving(&cqe, flow)
}
FlowState::Closing => unreachable!("we checked this above"),
};
tracing::debug!(
flow_id,
slot = %pool_id,
from_state = ?old_state,
to_state = ?new_state,
"state transition"
);
self.flows
.get_mut(&flow_id)
.expect("should always exist")
.state = new_state;
}
if self.finished_responses.len() != finished_resp_len_prev {
return Some(());
}
}
}
}
fn parse_response(buf: &[u8]) -> Result<http::Response<Vec<u8>>> {
let mut headers = [httparse::EMPTY_HEADER; 64];
let mut parsed_response = httparse::Response::new(&mut headers);
let header_size = match parsed_response.parse(buf)? {
httparse::Status::Complete(size) => size,
httparse::Status::Partial => bail!("incomplete HTTP response"),
};
let body_bytes = &buf[header_size..];
let mut response_builder = http::Response::builder()
.status(parsed_response.code.unwrap_or(0))
.version(if parsed_response.version.unwrap_or(0) == 1 {
http::Version::HTTP_11
} else {
http::Version::HTTP_10
});
for header in parsed_response.headers {
response_builder = response_builder.header(header.name, header.value);
}
response_builder
.body(Vec::from(body_bytes))
.wrap_err("parsing resp body")
}