use crate::courierust_body::Body;
use crate::courierust_bytes::Bytes;
use crate::courierust_client::ClientConfig;
use crate::courierust_error::{Error, Result};
use crate::courierust_h2::connection::{Config as H2Config, Connection, Event};
use crate::courierust_h2::error::ErrorCode;
use crate::courierust_h2::priority::Priority;
use crate::courierust_hpack::HeaderField;
use crate::courierust_http::header::HeaderMap;
use crate::courierust_http::response::ResponseHead;
use crate::courierust_http::status::StatusCode;
use crate::courierust_http::version::Version;
use crate::courierust_net::stats::{ActiveH2Streams, Counting, Stats};
use crate::courierust_net::ConnStream;
use std::collections::{HashMap, VecDeque};
use std::net::SocketAddr;
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::sync::mpsc::{channel, Receiver, Sender};
use std::sync::Arc;
use std::thread;
use std::time::Duration;
type DriverConn<'a> = Connection<Counting<&'a ConnStream>, Counting<&'a ConnStream>>;
pub struct H2Response {
pub head: ResponseHead,
pub body: Body,
pub trailers: Arc<std::sync::Mutex<Option<HeaderMap>>>,
}
pub enum H2Cmd {
Request {
fields: Vec<HeaderField>,
body: Option<Bytes>,
end_stream: bool,
priority: Priority,
timeout: Option<std::time::Duration>,
reply: Sender<Result<H2Response>>,
},
RequestStream {
fields: Vec<HeaderField>,
body: Receiver<Result<Bytes>>,
priority: Priority,
timeout: Option<std::time::Duration>,
reply: Sender<Result<H2Response>>,
},
Shutdown,
}
#[derive(Clone)]
pub struct H2Conn {
pub tx: Sender<H2Cmd>,
pub peer: SocketAddr,
pub accepting: Arc<std::sync::atomic::AtomicBool>,
reservations: Arc<AtomicUsize>,
body_load: Arc<AtomicUsize>,
ewma_service_us: Arc<AtomicU64>,
}
const H2_BODY_UNIT: usize = 64 * 1024;
const H2_BODY_WEIGHT_CAP: usize = 256;
const H2_EWMA_DIVISOR: u64 = 1000;
const H2_EWMA_CAP_US: u64 = 10_000;
fn body_weight(bytes: usize) -> usize {
bytes.div_ceil(H2_BODY_UNIT).min(H2_BODY_WEIGHT_CAP)
}
impl H2Conn {
pub(crate) fn reserve(&self, body_bytes: usize) {
self.reservations.fetch_add(1, Ordering::AcqRel);
self.body_load
.fetch_add(body_weight(body_bytes), Ordering::AcqRel);
}
pub(crate) fn release(&self, body_bytes: usize) {
let _ = self
.reservations
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |value| {
Some(value.saturating_sub(1))
});
let _ = self
.body_load
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |value| {
Some(value.saturating_sub(body_weight(body_bytes)))
});
}
pub(crate) fn is_idle(&self) -> bool {
self.reservations.load(Ordering::Acquire) == 0
}
pub(crate) fn load(&self) -> usize {
let streams = self.reservations.load(Ordering::Acquire);
let body = self.body_load.load(Ordering::Acquire);
let ewma = self
.ewma_service_us
.load(Ordering::Acquire)
.min(H2_EWMA_CAP_US);
streams
.saturating_add(body)
.saturating_add((ewma / H2_EWMA_DIVISOR) as usize)
}
pub(crate) fn note_service_us(&self, sample_us: u64) {
let mut current = self.ewma_service_us.load(Ordering::Relaxed);
loop {
let next = if current == 0 {
sample_us.min(H2_EWMA_CAP_US)
} else {
((current * 7 + sample_us) / 8).min(H2_EWMA_CAP_US)
};
match self.ewma_service_us.compare_exchange_weak(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => break,
Err(observed) => current = observed,
}
}
}
}
struct Pending {
reply: Option<Sender<Result<H2Response>>>,
body_tx: Option<Sender<Result<Bytes>>>,
body_rx: Option<Receiver<Result<Bytes>>>,
trailers: Arc<std::sync::Mutex<Option<HeaderMap>>>,
body_len: usize,
deadline: Option<std::time::Instant>,
}
fn deadline_from(timeout: Option<Duration>) -> Option<std::time::Instant> {
timeout.map(|t| std::time::Instant::now() + t)
}
struct StreamBody {
rx: Receiver<Result<Bytes>>,
pending_chunk: Option<Bytes>,
}
const MAX_DEFERRED: usize = 1024;
fn stream_limit_reached(conn: &DriverConn<'_>, pending: &HashMap<u32, Pending>) -> bool {
let limit = conn.peer_settings().max_concurrent_streams as usize;
limit != 0 && pending.len() >= limit
}
pub(crate) fn start(stream: ConnStream, cfg: &ClientConfig) -> Result<H2Conn> {
start_inner(stream, cfg, Vec::new(), None)
}
pub(crate) fn start_upgraded(
stream: ConnStream,
cfg: &ClientConfig,
seed: Vec<u8>,
reply: Sender<Result<H2Response>>,
) -> Result<H2Conn> {
start_inner(stream, cfg, seed, Some(reply))
}
fn start_inner(
stream: ConnStream,
cfg: &ClientConfig,
seed: Vec<u8>,
upgrade_reply: Option<Sender<Result<H2Response>>>,
) -> Result<H2Conn> {
let peer = stream.peer_addr();
let (tx, rx) = channel::<H2Cmd>();
let cfg = cfg.clone();
let accepting = Arc::new(std::sync::atomic::AtomicBool::new(true));
let accepting2 = accepting.clone();
let reservations = Arc::new(AtomicUsize::new(0));
let body_load = Arc::new(AtomicUsize::new(0));
let ewma_service_us = Arc::new(AtomicU64::new(0));
let (reads, writes) = match cfg.stats.as_deref() {
Some(s) => (s.h2_read_syscalls.clone(), s.h2_write_syscalls.clone()),
None => (Arc::new(AtomicUsize::new(0)), Arc::new(AtomicUsize::new(0))),
};
if let Some(s) = cfg.stats.as_deref() {
s.h2_connections.fetch_add(1, Ordering::Relaxed);
s.h2_connections_active.fetch_add(1, Ordering::Relaxed);
}
let stats = cfg.stats.clone();
thread::Builder::new()
.name("courierust-h2-driver".into())
.spawn(move || {
driver(
stream,
rx,
cfg,
accepting2,
seed,
upgrade_reply,
reads,
writes,
stats,
);
})?;
Ok(H2Conn {
tx,
peer,
accepting,
reservations,
body_load,
ewma_service_us,
})
}
const DRIVER_READ_TIMEOUT: Duration = Duration::from_millis(5);
struct ActiveGuard(Arc<AtomicUsize>);
impl Drop for ActiveGuard {
fn drop(&mut self) {
Stats::decrement(&self.0, 1);
}
}
#[allow(clippy::too_many_arguments)]
fn driver(
stream: ConnStream,
rx: Receiver<H2Cmd>,
cfg: ClientConfig,
accepting: Arc<AtomicBool>,
seed: Vec<u8>,
upgrade_reply: Option<Sender<Result<H2Response>>>,
reads: Arc<AtomicUsize>,
writes: Arc<AtomicUsize>,
stats: Option<Arc<Stats>>,
) {
let _ = stream.configure(Some(DRIVER_READ_TIMEOUT));
let stats = stats.as_deref();
let _active_guard = stats.map(|s| ActiveGuard(s.h2_connections_active.clone()));
let mut conn = if seed.is_empty() {
Connection::new(
Counting::new(&stream, reads.clone(), writes.clone()),
Counting::new(&stream, reads, writes),
h2_config(&cfg),
)
} else {
Connection::new_with_seed(
Counting::new(&stream, reads.clone(), writes.clone()),
Counting::new(&stream, reads, writes),
h2_config(&cfg),
&seed,
)
};
let mut pending: HashMap<u32, Pending> = HashMap::new();
let mut stream_bodies: HashMap<u32, StreamBody> = HashMap::new();
let mut goaway = false;
let mut deferred: VecDeque<H2Cmd> = VecDeque::new();
let mut stream_stats = ActiveH2Streams::new(stats);
if let Some(reply) = upgrade_reply {
if conn.register_upgrade_stream().is_ok() {
let (body_tx, body_rx) = channel::<Result<Bytes>>();
let trailers = Arc::new(std::sync::Mutex::new(None));
pending.insert(
1,
Pending {
reply: Some(reply),
body_tx: Some(body_tx),
body_rx: Some(body_rx),
trailers,
body_len: 0,
deadline: deadline_from(cfg.read_timeout),
},
);
if let Some(s) = stats {
s.h2_streams_total.fetch_add(1, Ordering::Relaxed);
}
}
}
let started = std::time::Instant::now();
let mut last_rx = started;
let mut last_ping: Option<std::time::Instant> = None;
let _ = conn.poll();
loop {
if !retry_deferred(
&mut conn,
&mut pending,
&mut stream_bodies,
&mut goaway,
&mut deferred,
cfg.read_timeout,
stats,
) {
cleanup(&mut conn, &mut pending, &mut stream_bodies);
return;
}
let mut got_cmd = false;
while let Ok(cmd) = rx.try_recv() {
got_cmd = true;
if !handle_cmd(
&mut conn,
&mut pending,
&mut stream_bodies,
&mut goaway,
&mut deferred,
cmd,
cfg.read_timeout,
stats,
) {
cleanup(&mut conn, &mut pending, &mut stream_bodies);
return;
}
}
stream_stats.set(conn.open_stream_count());
let has_work =
got_cmd || !deferred.is_empty() || !pending.is_empty() || !stream_bodies.is_empty();
if has_work {
match conn.poll_available(64) {
Ok(true) => last_rx = std::time::Instant::now(),
Ok(false) => {}
Err(e) => {
accepting.store(false, std::sync::atomic::Ordering::Release);
fail_all(&mut pending, e);
break;
}
};
drain_events(
&mut conn,
&mut pending,
&mut goaway,
&accepting,
cfg.max_body,
cfg.read_timeout,
);
drain_stream_bodies(&mut conn, &mut stream_bodies);
check_timeouts(&mut conn, &mut pending, stats);
if !retry_deferred(
&mut conn,
&mut pending,
&mut stream_bodies,
&mut goaway,
&mut deferred,
cfg.read_timeout,
stats,
) {
cleanup(&mut conn, &mut pending, &mut stream_bodies);
return;
}
if conn.is_closed() {
accepting.store(false, std::sync::atomic::Ordering::Release);
fail_all(&mut pending, Error::eof());
break;
}
if !apply_liveness(
&mut conn,
&mut pending,
&accepting,
&cfg,
started,
&mut last_rx,
&mut last_ping,
true,
) {
break;
}
} else {
match rx.recv_timeout(Duration::from_millis(200)) {
Ok(cmd) => {
if !handle_cmd(
&mut conn,
&mut pending,
&mut stream_bodies,
&mut goaway,
&mut deferred,
cmd,
cfg.read_timeout,
stats,
) {
break;
}
}
Err(std::sync::mpsc::RecvTimeoutError::Timeout) => {
match conn.poll_available(64) {
Ok(true) => last_rx = std::time::Instant::now(),
Ok(false) => {}
Err(e) => {
accepting.store(false, std::sync::atomic::Ordering::Release);
fail_all(&mut pending, e);
break;
}
};
drain_events(
&mut conn,
&mut pending,
&mut goaway,
&accepting,
cfg.max_body,
cfg.read_timeout,
);
drain_stream_bodies(&mut conn, &mut stream_bodies);
check_timeouts(&mut conn, &mut pending, stats);
if !retry_deferred(
&mut conn,
&mut pending,
&mut stream_bodies,
&mut goaway,
&mut deferred,
cfg.read_timeout,
stats,
) {
break;
}
if conn.is_closed() {
accepting.store(false, std::sync::atomic::Ordering::Release);
fail_all(&mut pending, Error::eof());
break;
}
if !apply_liveness(
&mut conn,
&mut pending,
&accepting,
&cfg,
started,
&mut last_rx,
&mut last_ping,
false,
) {
break;
}
}
Err(std::sync::mpsc::RecvTimeoutError::Disconnected) => break,
}
}
}
cleanup(&mut conn, &mut pending, &mut stream_bodies);
}
#[allow(clippy::too_many_arguments)]
fn apply_liveness(
conn: &mut DriverConn<'_>,
pending: &mut HashMap<u32, Pending>,
accepting: &Arc<AtomicBool>,
cfg: &ClientConfig,
started: std::time::Instant,
last_rx: &mut std::time::Instant,
last_ping: &mut Option<std::time::Instant>,
has_work: bool,
) -> bool {
use std::sync::atomic::Ordering;
let now = std::time::Instant::now();
if let Some(t) = cfg.h2_settings_timeout {
if conn.settings_ack_pending() && now.duration_since(started) >= t {
conn.send_goaway(ErrorCode::SettingsTimeout, b"peer did not ACK SETTINGS");
accepting.store(false, Ordering::Release);
fail_all(
pending,
Error::h2(
ErrorCode::SettingsTimeout.as_u32(),
"peer did not ACK our SETTINGS",
),
);
return false;
}
}
if !has_work {
if let Some(t) = cfg.h2_idle_timeout {
if now.duration_since(*last_rx) >= t {
conn.send_goaway(ErrorCode::NoError, b"idle timeout");
accepting.store(false, Ordering::Release);
return false;
}
}
}
if let Some(interval) = cfg.h2_ping_interval {
if now.duration_since(*last_rx) >= interval {
match *last_ping {
None => {
let nanos = now.duration_since(started).as_nanos() as u64;
conn.send_ping(nanos.to_be_bytes());
*last_ping = Some(now);
}
Some(sent) => {
if *last_rx < sent {
if let Some(pt) = cfg.h2_ping_timeout {
if now.duration_since(sent) >= pt {
accepting.store(false, Ordering::Release);
fail_all(pending, Error::eof());
return false;
}
}
} else {
*last_ping = None;
}
}
}
} else {
*last_ping = None;
}
}
true
}
#[allow(clippy::too_many_arguments)]
fn handle_cmd(
conn: &mut DriverConn<'_>,
pending: &mut HashMap<u32, Pending>,
stream_bodies: &mut HashMap<u32, StreamBody>,
goaway: &mut bool,
deferred: &mut VecDeque<H2Cmd>,
cmd: H2Cmd,
timeout: Option<Duration>,
stats: Option<&Stats>,
) -> bool {
match cmd {
H2Cmd::Shutdown => {
conn.send_goaway(ErrorCode::NoError, b"client shutdown");
false
}
H2Cmd::Request {
fields,
body,
end_stream,
priority,
timeout: request_timeout,
reply,
} => {
if *goaway {
let _ = reply.send(Err(Error::canceled("connection received GOAWAY")));
return true;
}
if stream_limit_reached(conn, pending) {
if deferred.len() < MAX_DEFERRED {
deferred.push_back(H2Cmd::Request {
fields,
body,
end_stream,
priority,
timeout: request_timeout,
reply,
});
} else {
let _ = reply.send(Err(Error::h2(
ErrorCode::RefusedStream.as_u32(),
"peer SETTINGS_MAX_CONCURRENT_STREAMS exhausted",
)));
}
return true;
}
match conn.open_request(priority) {
Ok(sid) => {
if let Some(s) = stats {
s.h2_streams_total.fetch_add(1, Ordering::Relaxed);
}
let (body_tx, body_rx) = channel::<Result<Bytes>>();
let trailers = Arc::new(std::sync::Mutex::new(None));
let body_empty = body.is_none() && end_stream;
if let Err(e) = conn.send_headers(sid, &fields, body_empty) {
let _ = reply.send(Err(e));
return true;
}
if let Some(b) = body {
if let Err(e) = conn.send_data(sid, b, end_stream) {
let _ = reply.send(Err(e));
return true;
}
}
pending.insert(
sid,
Pending {
reply: Some(reply),
body_tx: Some(body_tx),
body_rx: Some(body_rx),
trailers,
body_len: 0,
deadline: deadline_from(request_timeout.or(timeout)),
},
);
}
Err(e) => {
let _ = reply.send(Err(e));
}
}
true
}
H2Cmd::RequestStream {
fields,
body,
priority,
timeout: request_timeout,
reply,
} => {
if *goaway {
let _ = reply.send(Err(Error::canceled("connection received GOAWAY")));
return true;
}
if stream_limit_reached(conn, pending) {
if deferred.len() < MAX_DEFERRED {
deferred.push_back(H2Cmd::RequestStream {
fields,
body,
priority,
timeout: request_timeout,
reply,
});
} else {
let _ = reply.send(Err(Error::h2(
ErrorCode::RefusedStream.as_u32(),
"peer SETTINGS_MAX_CONCURRENT_STREAMS exhausted",
)));
}
return true;
}
match conn.open_request(priority) {
Ok(sid) => {
if let Some(s) = stats {
s.h2_streams_total.fetch_add(1, Ordering::Relaxed);
}
let (body_tx, body_rx) = channel::<Result<Bytes>>();
let trailers = Arc::new(std::sync::Mutex::new(None));
if let Err(e) = conn.send_headers(sid, &fields, false) {
let _ = reply.send(Err(e));
return true;
}
stream_bodies.insert(
sid,
StreamBody {
rx: body,
pending_chunk: None,
},
);
pending.insert(
sid,
Pending {
reply: Some(reply),
body_tx: Some(body_tx),
body_rx: Some(body_rx),
trailers,
body_len: 0,
deadline: deadline_from(request_timeout.or(timeout)),
},
);
}
Err(e) => {
let _ = reply.send(Err(e));
}
}
true
}
}
}
fn check_timeouts(
conn: &mut DriverConn<'_>,
pending: &mut HashMap<u32, Pending>,
stats: Option<&Stats>,
) {
use std::sync::atomic::Ordering;
let now = std::time::Instant::now();
let timed_out: Vec<u32> = pending
.iter()
.filter(|(_, p)| p.deadline.is_some_and(|d| now >= d))
.map(|(sid, _)| *sid)
.collect();
if timed_out.is_empty() {
return;
}
if let Some(s) = stats {
s.h2_streams_timed_out
.fetch_add(timed_out.len(), Ordering::Relaxed);
}
for sid in timed_out {
conn.send_rst(sid, ErrorCode::Cancel);
if let Some(p) = pending.remove(&sid) {
let err = Error::timeout("h2 response did not complete in time");
let _ = p.body_tx.map(|tx| tx.send(Err(err.clone())));
let _ = p.reply.map(|r| r.send(Err(err)));
}
}
}
fn retry_deferred(
conn: &mut DriverConn<'_>,
pending: &mut HashMap<u32, Pending>,
stream_bodies: &mut HashMap<u32, StreamBody>,
goaway: &mut bool,
deferred: &mut VecDeque<H2Cmd>,
timeout: Option<Duration>,
stats: Option<&Stats>,
) -> bool {
while let Some(cmd) = deferred.pop_front() {
if stream_limit_reached(conn, pending) {
deferred.push_front(cmd);
return true;
}
if !handle_cmd(
conn,
pending,
stream_bodies,
goaway,
deferred,
cmd,
timeout,
stats,
) {
return false;
}
}
true
}
fn drain_stream_bodies(conn: &mut DriverConn<'_>, bodies: &mut HashMap<u32, StreamBody>) {
let mut done = Vec::new();
for (&sid, b) in bodies.iter_mut() {
if let Some(chunk) = b.pending_chunk.take() {
match conn.send_data(sid, chunk.clone(), false) {
Ok(_) => {}
Err(e) if e.kind == crate::courierust_error::ErrorKind::Overflow => {
b.pending_chunk = Some(chunk);
continue;
}
Err(_) => {
done.push(sid);
continue;
}
}
}
loop {
match b.rx.try_recv() {
Ok(Ok(chunk)) => {
if chunk.is_empty() {
continue;
}
match conn.send_data(sid, chunk.clone(), false) {
Ok(_) => {}
Err(e) if e.kind == crate::courierust_error::ErrorKind::Overflow => {
b.pending_chunk = Some(chunk);
break;
}
Err(_) => {
done.push(sid);
break;
}
}
}
Ok(Err(_)) => {
let _ = conn.send_data(sid, Bytes::new(), true);
done.push(sid);
break;
}
Err(std::sync::mpsc::TryRecvError::Disconnected) => {
let _ = conn.send_data(sid, Bytes::new(), true);
done.push(sid);
break;
}
Err(std::sync::mpsc::TryRecvError::Empty) => break,
}
}
}
for sid in done {
bodies.remove(&sid);
}
}
fn drain_events(
conn: &mut DriverConn<'_>,
pending: &mut HashMap<u32, Pending>,
goaway: &mut bool,
accepting: &Arc<AtomicBool>,
max_body: usize,
timeout: Option<Duration>,
) {
while let Some(ev) = conn.next_event() {
match ev {
Event::Headers {
stream_id,
headers,
end_stream,
..
} => {
if let Some(p) = pending.get_mut(&stream_id) {
match build_response(&headers) {
Ok(head) => {
if end_stream {
if let Some(reply) = p.reply.take() {
let _ = reply.send(Ok(H2Response {
head,
body: Body::Empty,
trailers: p.trailers.clone(),
}));
}
pending.remove(&stream_id);
} else {
let body_rx = p.body_rx.take();
if let Some(reply) = p.reply.take() {
let body = match body_rx {
Some(rx) => Body::Channel(rx),
None => Body::Empty,
};
let _ = reply.send(Ok(H2Response {
head,
body,
trailers: p.trailers.clone(),
}));
}
if let Some(t) = timeout {
p.deadline = Some(std::time::Instant::now() + t);
}
}
}
Err(e) => {
let _ = p.reply.take().map(|r| r.send(Err(e)));
pending.remove(&stream_id);
}
}
}
}
Event::Data {
stream_id,
data,
end_stream,
} => {
if let Some(p) = pending.get_mut(&stream_id) {
p.body_len = p.body_len.saturating_add(data.len());
if p.body_len > max_body {
conn.send_rst(stream_id, ErrorCode::EnhanceYourCalm);
let err = Error::overflow("response body exceeds limit");
let _ = p.body_tx.take().map(|tx| tx.send(Err(err.clone())));
let _ = p.reply.take().map(|r| r.send(Err(err)));
pending.remove(&stream_id);
continue;
}
let _ = p.body_tx.as_ref().map(|tx| tx.send(Ok(data)));
if let Some(t) = timeout {
p.deadline = Some(std::time::Instant::now() + t);
}
if end_stream {
pending.remove(&stream_id);
}
}
}
Event::Trailers { stream_id, headers } => {
if let Some(p) = pending.get_mut(&stream_id) {
let mut map = HeaderMap::new();
for f in &headers {
map.append(f.name.clone(), f.value.clone());
}
*p.trailers.lock().unwrap() = Some(map);
pending.remove(&stream_id);
}
}
Event::Rst {
stream_id,
error_code,
} => {
if let Some(p) = pending.remove(&stream_id) {
let err = Error::h2(error_code.as_u32(), "stream reset by peer");
let _ = p.body_tx.map(|tx| tx.send(Err(err.clone())));
let _ = p.reply.map(|r| r.send(Err(err)));
}
}
Event::StreamError {
stream_id,
error_code,
message,
} => {
if let Some(p) = pending.remove(&stream_id) {
let err = Error::h2(error_code.as_u32(), message.as_str());
let _ = p.body_tx.map(|tx| tx.send(Err(err.clone())));
let _ = p.reply.map(|r| r.send(Err(err)));
}
}
Event::GoAway {
error_code,
last_stream_id,
..
} => {
accepting.store(false, std::sync::atomic::Ordering::Release);
*goaway = true;
let dead: Vec<u32> = pending
.keys()
.copied()
.filter(|&s| s > last_stream_id)
.collect();
for sid in dead {
if let Some(p) = pending.remove(&sid) {
let err = Error::h2(error_code.as_u32(), "peer sent GOAWAY");
let _ = p.body_tx.map(|tx| tx.send(Err(err.clone())));
let _ = p.reply.map(|r| r.send(Err(err)));
}
}
}
_ => {}
}
}
}
fn build_response(headers: &[HeaderField]) -> Result<ResponseHead> {
let mut status = StatusCode::OK;
let mut map = HeaderMap::new();
for f in headers {
if f.name.is_pseudo() {
if f.name.as_str() == ":status" {
let code: u16 = std::str::from_utf8(f.value.as_bytes())
.ok()
.and_then(|s| s.parse().ok())
.ok_or_else(|| Error::protocol("missing/invalid :status"))?;
status = StatusCode::from_u16(code);
}
} else {
map.append(f.name.clone(), f.value.clone());
}
}
Ok(ResponseHead {
status,
version: Version::HTTP_2,
headers: map,
})
}
fn h2_config(cfg: &ClientConfig) -> H2Config {
let mut c = H2Config {
client: true,
max_send_buffer: cfg.max_body,
auto_release_credit: true,
..Default::default()
};
c.local_settings.enable_push = 0;
if c.local_settings.initial_window_size < 256 * 1024 {
c.local_settings.initial_window_size = 256 * 1024;
}
if let Ok(hl) = cfg.max_header_list.try_into() {
c.local_settings.max_header_list_size = hl;
}
c
}
pub(crate) fn upgrade_settings_b64(cfg: &ClientConfig) -> String {
let c = h2_config(cfg);
crate::courierust_crypto::base64::encode_url_no_pad(&c.local_settings.to_wire())
}
pub(crate) enum UpgradeOutcome {
Upgraded(Vec<u8>),
Declined(ResponseHead, Vec<u8>),
}
pub(crate) fn h2c_upgrade_handshake(
mut stream: &std::net::TcpStream,
request_wire: &[u8],
) -> Result<UpgradeOutcome> {
use std::io::{Read as _, Write as _};
stream
.write_all(request_wire)
.map_err(|e| Error::io(e.to_string()))?;
stream.flush().map_err(|e| Error::io(e.to_string()))?;
let mut buf = [0u8; 8192];
let mut filled = 0usize;
let mut head_end = None;
while filled < buf.len() {
let n = stream
.read(&mut buf[filled..])
.map_err(|e| Error::io(e.to_string()))?;
if n == 0 {
return Err(Error::eof());
}
filled += n;
if let Some(i) = find_subslice(&buf[..filled], b"\r\n\r\n") {
head_end = Some(i + 4);
break;
}
}
let end = head_end.ok_or_else(|| Error::protocol("101 response head too large"))?;
let (status, version, headers) = parse_head_headers(&buf[..end])?;
let leftover = buf[end..filled].to_vec();
if status == StatusCode::SWITCHING_PROTOCOLS {
Ok(UpgradeOutcome::Upgraded(leftover))
} else {
Ok(UpgradeOutcome::Declined(
ResponseHead {
status,
version,
headers,
},
leftover,
))
}
}
pub(crate) fn build_upgrade_request(
req: &crate::courierust_http::request::Request<Body>,
authority: &str,
settings_b64: &str,
user_agent: Option<&str>,
) -> Result<Vec<u8>> {
let body = match &req.body {
Body::Empty => None,
Body::Bytes(b) => Some(b),
Body::Channel(_) | Body::Stream(_) => {
return Err(Error::protocol(
"streaming request bodies cannot use the h2c Upgrade",
));
}
};
let mut out = Vec::with_capacity(256);
out.extend_from_slice(req.method.as_str().as_bytes());
out.push(b' ');
out.extend_from_slice(req.uri.as_bytes());
out.extend_from_slice(b" HTTP/1.1\r\n");
out.extend_from_slice(b"Host: ");
out.extend_from_slice(authority.as_bytes());
out.extend_from_slice(b"\r\n");
out.extend_from_slice(b"Connection: Upgrade, HTTP2-Settings\r\n");
out.extend_from_slice(b"Upgrade: h2c\r\n");
out.extend_from_slice(b"HTTP2-Settings: ");
out.extend_from_slice(settings_b64.as_bytes());
out.extend_from_slice(b"\r\n");
if let Some(b) = body {
if !b.is_empty() {
out.extend_from_slice(b"Content-Length: ");
let cl = crate::courierust_h1::IToA::new(b.len());
out.extend_from_slice(cl.as_slice());
out.extend_from_slice(b"\r\n");
}
}
for (n, v) in req.headers.iter() {
let name = n.as_str();
if crate::courierust_h1::is_hop_by_hop(name) || name == "host" {
continue;
}
out.extend_from_slice(name.as_bytes());
out.extend_from_slice(b": ");
out.extend_from_slice(v.as_bytes());
out.extend_from_slice(b"\r\n");
}
if !req.headers.contains_key("user-agent") {
if let Some(ua) = user_agent {
out.extend_from_slice(b"User-Agent: ");
out.extend_from_slice(ua.as_bytes());
out.extend_from_slice(b"\r\n");
}
}
out.extend_from_slice(b"\r\n");
if let Some(b) = body {
out.extend_from_slice(b.as_slice());
}
Ok(out)
}
fn parse_head_headers(head: &[u8]) -> Result<(StatusCode, Version, HeaderMap)> {
let end =
find_subslice(head, b"\r\n").ok_or_else(|| Error::protocol("malformed response head"))?;
let (status, version) = crate::courierust_h1::parse_status_line(&head[..end])?;
let mut map = HeaderMap::new();
let mut pos = end + 2;
while pos < head.len() {
let line_end = match find_subslice(&head[pos..], b"\r\n") {
Some(i) => pos + i,
None => head.len(),
};
let line = &head[pos..line_end];
if line.is_empty() {
break;
}
let colon = line
.iter()
.position(|&b| b == b':')
.ok_or_else(|| Error::protocol("malformed header line"))?;
let name = crate::courierust_http::header::HeaderName::from_bytes(&line[..colon])?;
let mut val = &line[colon + 1..];
while val.first() == Some(&b' ') || val.first() == Some(&b'\t') {
val = &val[1..];
}
map.append(
name,
crate::courierust_http::header::HeaderValue::from_bytes(val)?,
);
pos = line_end + 2;
}
Ok((status, version, map))
}
fn find_subslice(haystack: &[u8], needle: &[u8]) -> Option<usize> {
haystack.windows(needle.len()).position(|w| w == needle)
}
fn fail_all(pending: &mut HashMap<u32, Pending>, err: Error) {
for (_, p) in pending.drain() {
let _ = p.body_tx.map(|tx| tx.send(Err(err.clone())));
let _ = p.reply.map(|r| r.send(Err(err.clone())));
}
}
fn cleanup(
_conn: &mut DriverConn<'_>,
pending: &mut HashMap<u32, Pending>,
_bodies: &mut HashMap<u32, StreamBody>,
) {
fail_all(pending, Error::canceled("connection closed"));
}
pub fn request_fields(
req: &crate::courierust_http::request::Request<Body>,
scheme: &str,
authority: &str,
) -> Vec<HeaderField> {
let head = crate::courierust_http::request::RequestHead {
method: req.method.clone(),
uri: req.uri.clone(),
version: req.version,
headers: req.headers.clone(),
};
head.to_h2_fields(scheme, Some(authority))
}
#[cfg(test)]
mod tests {
use super::*;
fn test_conn() -> H2Conn {
let (tx, _rx) = channel::<H2Cmd>();
H2Conn {
tx,
peer: "127.0.0.1:1".parse().unwrap(),
accepting: Arc::new(AtomicBool::new(true)),
reservations: Arc::new(AtomicUsize::new(0)),
body_load: Arc::new(AtomicUsize::new(0)),
ewma_service_us: Arc::new(AtomicU64::new(0)),
}
}
#[test]
fn reserve_release_balances_to_idle() {
let conn = test_conn();
assert!(conn.is_idle(), "fresh connection must be idle");
conn.reserve(0);
conn.reserve(1 << 20);
assert!(!conn.is_idle());
conn.release(0);
conn.release(1 << 20);
assert!(
conn.is_idle(),
"balanced reserve/release must return to idle"
);
assert_eq!(conn.load(), 0);
}
#[test]
fn load_weights_bodies_over_stream_count() {
let conn = test_conn();
conn.reserve(1 << 20);
let big_upload_load = conn.load();
conn.release(1 << 20);
for _ in 0..4 {
conn.reserve(0);
}
let four_small_load = conn.load();
assert!(
big_upload_load > four_small_load,
"one 1 MiB upload must outweigh four header-only RPCs: {big_upload_load} vs {four_small_load}"
);
conn.release(0);
conn.release(0);
conn.release(0);
conn.release(0);
conn.reserve(64 * 1024);
assert_eq!(conn.load(), 2, "one stream + one 64 KiB body unit");
}
#[test]
fn ewma_updates_and_caps() {
let conn = test_conn();
assert_eq!(conn.load(), 0, "no samples yet -> no latency term");
conn.note_service_us(1_000); conn.note_service_us(1_000);
conn.note_service_us(1_000);
assert_eq!(conn.load(), 1, "1 ms of service time adds one unit");
conn.note_service_us(60_000_000);
assert!(conn.load() <= 10 + 1, "EWMA term must stay within its cap");
}
#[test]
fn idle_wins_over_stale_ewma() {
let conn = test_conn();
conn.note_service_us(60_000_000);
assert!(conn.is_idle(), "released connection must be idle");
assert!(
conn.load() > 0,
"the EWMA term still shows in load, which is why idle-first matters"
);
conn.reserve(0);
assert!(!conn.is_idle());
conn.release(0);
assert!(conn.is_idle());
}
}