use std::cell::RefCell;
use std::collections::{HashMap, VecDeque};
use std::future::Future;
use std::pin::Pin;
use std::rc::{Rc, Weak};
use std::task::{Context, Poll, Waker};
use futures::channel::{mpsc, oneshot};
use futures::future::{poll_fn, LocalBoxFuture};
use futures::stream::{self, FuturesUnordered, LocalBoxStream};
use futures::{FutureExt, SinkExt, Stream, StreamExt};
use crate::errors::{ErrorCode, H2Error};
use crate::flow::SendWindow;
use crate::frames::{serialize_frame, Frame, FrameDecoder, Settings, DEFAULT_MAX_FRAME_SIZE};
use crate::hpack::{Header, HpackDecoder, HpackEncoder};
use crate::transport::Transport;
const CONNECTION_PREFACE: &[u8] = b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n";
const SPEC_INITIAL_WINDOW: i64 = 65535;
const MAX_HEADER_BLOCK_SIZE: usize = 1 << 20; const FORBIDDEN_HEADERS: [&str; 6] = [
"connection",
"host",
"keep-alive",
"proxy-connection",
"transfer-encoding",
"upgrade",
];
#[derive(Default)]
pub struct RequestInit {
pub method: Option<String>,
pub path: Option<String>,
pub authority: Option<String>,
pub scheme: Option<String>,
pub headers: Vec<(String, String)>,
pub body: RequestBody,
}
#[derive(Default)]
pub enum RequestBody {
#[default]
Empty,
Bytes(Vec<u8>),
Stream(LocalBoxStream<'static, Vec<u8>>),
}
impl RequestBody {
pub fn stream<S>(chunks: S) -> Self
where
S: Stream<Item = Vec<u8>> + 'static,
{
RequestBody::Stream(chunks.boxed_local())
}
fn is_empty(&self) -> bool {
match self {
RequestBody::Empty => true,
RequestBody::Bytes(b) => b.is_empty(),
RequestBody::Stream(_) => false,
}
}
}
impl From<Vec<u8>> for RequestBody {
fn from(v: Vec<u8>) -> Self {
RequestBody::Bytes(v)
}
}
impl From<&[u8]> for RequestBody {
fn from(v: &[u8]) -> Self {
RequestBody::Bytes(v.to_vec())
}
}
impl From<String> for RequestBody {
fn from(v: String) -> Self {
RequestBody::Bytes(v.into_bytes())
}
}
impl From<&str> for RequestBody {
fn from(v: &str) -> Self {
RequestBody::Bytes(v.as_bytes().to_vec())
}
}
#[derive(Default)]
struct RecvState {
queue: VecDeque<Vec<u8>>,
buffered: usize,
ended: bool,
error: Option<H2Error>,
waker: Option<Waker>,
}
pub struct ResponseBody {
recv: Rc<RefCell<RecvState>>,
conn: Weak<RefCell<ConnState>>,
stream_id: u32,
}
impl ResponseBody {
fn replenish(&self, n: usize) {
if n == 0 {
return;
}
if let Some(conn) = self.conn.upgrade() {
conn.borrow().replenish_recv_window(self.stream_id, n);
}
}
}
impl Stream for ResponseBody {
type Item = Result<Vec<u8>, H2Error>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
let mut recv = this.recv.borrow_mut();
if let Some(chunk) = recv.queue.pop_front() {
let n = chunk.len();
recv.buffered -= n;
drop(recv); this.replenish(n); Poll::Ready(Some(Ok(chunk)))
} else if let Some(e) = recv.error.take() {
Poll::Ready(Some(Err(e)))
} else if recv.ended {
Poll::Ready(None)
} else {
recv.waker = Some(cx.waker().clone());
Poll::Pending
}
}
}
impl Drop for ResponseBody {
fn drop(&mut self) {
let remaining = self.recv.borrow().buffered;
if remaining > 0 {
self.replenish(remaining);
}
}
}
pub struct Response {
pub status: u16,
pub headers: HashMap<String, String>,
pub raw_headers: Vec<Header>,
body: ResponseBody,
trailers: Rc<RefCell<Option<HashMap<String, String>>>>,
}
impl Response {
pub fn into_body(self) -> ResponseBody {
self.body
}
pub async fn bytes(&mut self) -> Result<Vec<u8>, H2Error> {
let mut out = Vec::new();
while let Some(chunk) = self.body.next().await {
out.extend_from_slice(&chunk?);
}
Ok(out)
}
pub async fn text(&mut self) -> Result<String, H2Error> {
Ok(String::from_utf8_lossy(&self.bytes().await?).into_owned())
}
pub fn trailers(&self) -> Option<HashMap<String, String>> {
self.trailers.borrow().clone()
}
}
#[derive(Default, Clone)]
pub struct ConnectOptions {
pub header_table_size: Option<usize>,
pub enable_push: Option<bool>,
pub initial_window_size: Option<u32>,
pub max_frame_size: Option<usize>,
pub connection_window_size: Option<u32>,
}
struct Head {
status: u16,
headers: HashMap<String, String>,
raw: Vec<Header>,
}
fn collect_headers(raw: Vec<Header>) -> Head {
let mut headers: HashMap<String, String> = HashMap::new();
let mut status = 0u16;
for h in &raw {
if h.name == ":status" {
status = h.value.parse().unwrap_or(0);
continue;
}
if h.name.starts_with(':') {
continue;
}
match headers.get(&h.name) {
Some(existing) => {
let sep = if h.name == "cookie" { "; " } else { ", " };
let joined = format!("{existing}{sep}{}", h.value);
headers.insert(h.name.clone(), joined);
}
None => {
headers.insert(h.name.clone(), h.value.clone());
}
}
}
Head {
status,
headers,
raw,
}
}
struct StreamState {
id: u32,
send_window: SendWindow,
head_tx: Option<oneshot::Sender<Result<Head, H2Error>>>,
recv: Rc<RefCell<RecvState>>,
trailers: Rc<RefCell<Option<HashMap<String, String>>>>,
got_head: bool,
local_closed: bool,
remote_closed: bool,
}
impl StreamState {
fn new(id: u32, initial_send_window: i64) -> Self {
Self {
id,
send_window: SendWindow::new(initial_send_window),
head_tx: None,
recv: Rc::new(RefCell::new(RecvState::default())),
trailers: Rc::new(RefCell::new(None)),
got_head: false,
local_closed: false,
remote_closed: false,
}
}
fn receive_headers(&mut self, raw: Vec<Header>, end_stream: bool) {
if !self.got_head {
let head = collect_headers(raw);
if (100..200).contains(&head.status) {
return;
}
self.got_head = true;
if let Some(tx) = self.head_tx.take() {
let _ = tx.send(Ok(head));
}
} else {
*self.trailers.borrow_mut() = Some(collect_headers(raw).headers);
}
if end_stream {
self.end_body();
}
}
fn receive_data(&mut self, data: &[u8], end_stream: bool) {
let mut recv = self.recv.borrow_mut();
if !data.is_empty() && recv.error.is_none() && !recv.ended {
recv.queue.push_back(data.to_vec());
recv.buffered += data.len();
}
if end_stream {
recv.ended = true;
}
if let Some(w) = recv.waker.take() {
w.wake();
}
}
fn receive_reset(&mut self, error_code: u32) {
let code = ErrorCode::from_value(error_code).unwrap_or(ErrorCode::ProtocolError);
self.fail(H2Error::stream(
code,
format!("stream {} reset by peer", self.id),
self.id,
));
}
fn fail(&mut self, err: H2Error) {
self.send_window.close();
if !self.got_head {
self.got_head = true;
if let Some(tx) = self.head_tx.take() {
let _ = tx.send(Err(err.clone()));
}
}
let mut recv = self.recv.borrow_mut();
if recv.error.is_none() {
recv.error = Some(err);
}
if let Some(w) = recv.waker.take() {
w.wake();
}
}
fn end_body(&mut self) {
let mut recv = self.recv.borrow_mut();
if !recv.ended {
recv.ended = true;
if let Some(w) = recv.waker.take() {
w.wake();
}
}
}
}
struct RemoteSettings {
initial_window_size: i64,
max_frame_size: usize,
#[allow(dead_code)]
header_table_size: usize,
#[allow(dead_code)]
enable_push: bool,
max_concurrent_streams: u32,
}
impl Default for RemoteSettings {
fn default() -> Self {
Self {
initial_window_size: SPEC_INITIAL_WINDOW,
max_frame_size: DEFAULT_MAX_FRAME_SIZE,
header_table_size: 4096,
enable_push: true,
max_concurrent_streams: u32::MAX,
}
}
}
enum HeaderBlockKind {
Response,
Push,
}
struct PendingHeaderBlock {
stream_id: u32,
kind: HeaderBlockKind,
end_stream: bool,
promised_stream_id: Option<u32>,
fragments: Vec<Vec<u8>>,
size: usize,
}
struct PingWaiter {
resolve: oneshot::Sender<Result<f64, H2Error>>,
sent_at: f64,
}
struct ConnState {
out_tx: Option<mpsc::UnboundedSender<Vec<u8>>>,
task_tx: mpsc::UnboundedSender<LocalBoxFuture<'static, ()>>,
encoder: HpackEncoder,
decoder: HpackDecoder,
frame_decoder: FrameDecoder,
streams: HashMap<u32, StreamState>,
next_stream_id: u32,
conn_send_window: SendWindow,
remote: RemoteSettings,
pending_header_block: Option<PendingHeaderBlock>,
pings: HashMap<[u8; 8], PingWaiter>,
ping_counter: u32,
slot_waiters: Vec<oneshot::Sender<()>>,
closed: bool,
close_error: Option<H2Error>,
goaway_received: bool,
highest_promised: u32,
}
impl ConnState {
fn write_raw(&self, bytes: Vec<u8>) {
if !self.closed {
if let Some(tx) = &self.out_tx {
let _ = tx.unbounded_send(bytes);
}
}
}
fn send_frame(&self, frame: Frame) {
self.write_raw(serialize_frame(&frame));
}
fn on_bytes(&mut self, chunk: &[u8]) {
let frames = match self.frame_decoder.push(chunk) {
Ok(f) => f,
Err(e) => {
self.connection_error(e);
return;
}
};
for frame in frames {
if let Err(e) = self.dispatch(frame) {
self.connection_error(e);
return;
}
}
}
fn dispatch(&mut self, frame: Frame) -> Result<(), H2Error> {
if self.pending_header_block.is_some() && !matches!(frame, Frame::Continuation { .. }) {
return Err(H2Error::new(
ErrorCode::ProtocolError,
"expected CONTINUATION frame",
));
}
match frame {
Frame::Settings { ack, settings } => {
if ack {
return Ok(());
}
self.apply_remote_settings(&settings)?;
self.send_frame(Frame::Settings {
ack: true,
settings: Settings::default(),
});
}
Frame::Headers {
stream_id,
header_block_fragment,
end_stream,
end_headers,
..
} => {
let size = header_block_fragment.len();
self.pending_header_block = Some(PendingHeaderBlock {
stream_id,
kind: HeaderBlockKind::Response,
end_stream,
promised_stream_id: None,
fragments: vec![header_block_fragment],
size,
});
self.guard_header_block_size()?;
if end_headers {
self.complete_header_block()?;
}
}
Frame::Continuation {
stream_id,
header_block_fragment,
end_headers,
} => {
match &mut self.pending_header_block {
Some(pb) if pb.stream_id == stream_id => {
pb.size += header_block_fragment.len();
pb.fragments.push(header_block_fragment);
}
_ => {
return Err(H2Error::new(
ErrorCode::ProtocolError,
"unexpected CONTINUATION",
))
}
}
self.guard_header_block_size()?;
if end_headers {
self.complete_header_block()?;
}
}
Frame::PushPromise {
stream_id,
promised_stream_id,
header_block_fragment,
end_headers,
} => {
let size = header_block_fragment.len();
self.pending_header_block = Some(PendingHeaderBlock {
stream_id,
kind: HeaderBlockKind::Push,
end_stream: false,
promised_stream_id: Some(promised_stream_id),
fragments: vec![header_block_fragment],
size,
});
self.guard_header_block_size()?;
if end_headers {
self.complete_header_block()?;
}
}
Frame::Data {
stream_id,
data,
end_stream,
} => {
if let Some(s) = self.streams.get_mut(&stream_id) {
s.receive_data(&data, end_stream);
if end_stream {
s.remote_closed = true;
}
} else if !data.is_empty() {
self.send_frame(Frame::WindowUpdate {
stream_id: 0,
window_size_increment: data.len() as u32,
});
}
if end_stream {
self.retire_if_fully_closed(stream_id);
}
}
Frame::RstStream {
stream_id,
error_code,
} => {
if let Some(mut s) = self.streams.remove(&stream_id) {
s.receive_reset(error_code);
}
}
Frame::WindowUpdate {
stream_id,
window_size_increment,
} => {
if window_size_increment == 0 {
if stream_id == 0 {
return Err(H2Error::new(ErrorCode::ProtocolError, "zero WINDOW_UPDATE"));
}
self.reset_stream(stream_id, ErrorCode::ProtocolError);
return Ok(());
}
if stream_id == 0 {
self.conn_send_window.update(window_size_increment as i64);
} else if let Some(s) = self.streams.get_mut(&stream_id) {
s.send_window.update(window_size_increment as i64);
}
}
Frame::Ping { ack, opaque_data } => {
if ack {
if let Some(w) = self.pings.remove(&opaque_data) {
let _ = w.resolve.send(Ok(now_millis() - w.sent_at));
}
} else {
self.send_frame(Frame::Ping {
ack: true,
opaque_data,
});
}
}
Frame::Goaway {
last_stream_id,
error_code,
..
} => {
self.goaway_received = true;
let code = ErrorCode::from_value(error_code).unwrap_or(ErrorCode::NoError);
let err = H2Error::new(code, "peer sent GOAWAY");
let doomed: Vec<u32> = self
.streams
.keys()
.copied()
.filter(|&id| id > last_stream_id)
.collect();
for id in doomed {
if let Some(mut s) = self.streams.remove(&id) {
s.fail(err.clone());
}
}
self.wake_slot_waiters(); if error_code != 0 {
self.destroy(err);
}
}
Frame::Priority { .. } => {} }
Ok(())
}
fn guard_header_block_size(&self) -> Result<(), H2Error> {
if let Some(pb) = &self.pending_header_block {
if pb.size > MAX_HEADER_BLOCK_SIZE {
return Err(H2Error::new(
ErrorCode::EnhanceYourCalm,
"header block exceeds the maximum size",
));
}
}
Ok(())
}
fn complete_header_block(&mut self) -> Result<(), H2Error> {
let pb = self
.pending_header_block
.take()
.expect("header block present");
let block: Vec<u8> = if pb.fragments.len() == 1 {
pb.fragments.into_iter().next().unwrap()
} else {
pb.fragments.concat()
};
let headers = self.decoder.decode(&block)?;
match pb.kind {
HeaderBlockKind::Response => {
if let Some(s) = self.streams.get_mut(&pb.stream_id) {
s.receive_headers(headers, pb.end_stream);
if pb.end_stream {
s.remote_closed = true;
}
}
if pb.end_stream {
self.retire_if_fully_closed(pb.stream_id);
}
}
HeaderBlockKind::Push => {
let promised = pb.promised_stream_id.unwrap_or(0);
if promised > self.highest_promised {
self.highest_promised = promised;
}
self.send_frame(Frame::RstStream {
stream_id: promised,
error_code: ErrorCode::RefusedStream.value(),
});
}
}
Ok(())
}
fn apply_remote_settings(&mut self, s: &Settings) -> Result<(), H2Error> {
if let Some(iw) = s.initial_window_size {
if iw > 0x7fff_ffff {
return Err(H2Error::new(
ErrorCode::FlowControlError,
"SETTINGS_INITIAL_WINDOW_SIZE exceeds 2^31-1",
));
}
let delta = iw as i64 - self.remote.initial_window_size;
self.remote.initial_window_size = iw as i64;
for stream in self.streams.values_mut() {
stream.send_window.adjust(delta);
}
}
if let Some(mfs) = s.max_frame_size {
if !(16384..=16_777_215).contains(&mfs) {
return Err(H2Error::new(
ErrorCode::ProtocolError,
"SETTINGS_MAX_FRAME_SIZE out of range",
));
}
self.remote.max_frame_size = mfs as usize;
}
if let Some(hts) = s.header_table_size {
self.remote.header_table_size = hts as usize;
}
if let Some(ep) = s.enable_push {
self.remote.enable_push = ep;
}
if let Some(mcs) = s.max_concurrent_streams {
self.remote.max_concurrent_streams = mcs;
self.wake_slot_waiters(); }
Ok(())
}
fn send_headers(&self, id: u32, block: Vec<u8>, end_stream: bool) {
let max = self.remote.max_frame_size;
if block.len() <= max {
self.send_frame(Frame::Headers {
stream_id: id,
header_block_fragment: block,
end_stream,
end_headers: true,
priority: None,
});
return;
}
self.send_frame(Frame::Headers {
stream_id: id,
header_block_fragment: block[..max].to_vec(),
end_stream,
end_headers: false,
priority: None,
});
let mut offset = max;
while offset < block.len() {
let next = (offset + max).min(block.len());
self.send_frame(Frame::Continuation {
stream_id: id,
header_block_fragment: block[offset..next].to_vec(),
end_headers: next >= block.len(),
});
offset = next;
}
}
fn reset_stream(&mut self, id: u32, code: ErrorCode) {
self.send_frame(Frame::RstStream {
stream_id: id,
error_code: code.value(),
});
if let Some(mut s) = self.streams.remove(&id) {
s.fail(H2Error::stream(code, format!("stream {id} reset"), id));
}
self.wake_slot_waiters();
}
fn retire_if_fully_closed(&mut self, id: u32) {
let fully_closed = self
.streams
.get(&id)
.is_some_and(|s| s.local_closed && s.remote_closed);
if fully_closed {
self.streams.remove(&id);
self.wake_slot_waiters(); }
}
fn replenish_recv_window(&self, stream_id: u32, n: usize) {
if self.closed || n == 0 {
return;
}
let inc = n as u32;
if self.streams.contains_key(&stream_id) {
self.send_frame(Frame::WindowUpdate {
stream_id,
window_size_increment: inc,
});
}
self.send_frame(Frame::WindowUpdate {
stream_id: 0,
window_size_increment: inc,
});
}
fn active_streams(&self) -> usize {
self.streams.keys().filter(|id| *id % 2 == 1).count()
}
fn can_open_stream(&self) -> bool {
self.active_streams() < self.remote.max_concurrent_streams as usize
}
fn wake_slot_waiters(&mut self) {
for tx in self.slot_waiters.drain(..) {
let _ = tx.send(());
}
}
fn connection_error(&mut self, err: H2Error) {
self.send_frame(Frame::Goaway {
last_stream_id: self.highest_promised,
error_code: err.code.value(),
debug_data: Vec::new(),
});
self.destroy(err);
}
fn destroy(&mut self, err: H2Error) {
if self.closed {
return;
}
self.closed = true;
self.close_error = Some(err.clone());
self.conn_send_window.close();
let ids: Vec<u32> = self.streams.keys().copied().collect();
for id in ids {
if let Some(mut s) = self.streams.remove(&id) {
s.fail(err.clone());
}
}
for (_, w) in self.pings.drain() {
let _ = w.resolve.send(Err(err.clone()));
}
self.wake_slot_waiters(); self.out_tx = None;
}
}
#[derive(Clone)]
pub struct H2Connection {
shared: Rc<RefCell<ConnState>>,
}
impl H2Connection {
pub fn is_closed(&self) -> bool {
self.shared.borrow().closed
}
pub fn active_streams(&self) -> usize {
self.shared.borrow().active_streams()
}
pub fn can_open_stream(&self) -> bool {
self.shared.borrow().can_open_stream()
}
pub async fn request(&self, mut init: RequestInit) -> Result<Response, H2Error> {
let body = std::mem::take(&mut init.body);
let has_body = !body.is_empty();
loop {
let rx = {
let mut st = self.shared.borrow_mut();
if st.closed {
return Err(st.close_error.clone().unwrap_or_else(|| {
H2Error::new(ErrorCode::InternalError, "connection closed")
}));
}
if st.goaway_received {
return Err(H2Error::new(
ErrorCode::RefusedStream,
"connection is going away",
));
}
if st.can_open_stream() {
None
} else {
let (tx, rx) = oneshot::channel();
st.slot_waiters.push(tx);
Some(rx)
}
};
match rx {
None => break,
Some(rx) => {
let _ = rx.await;
}
}
}
let id;
let head_rx;
let recv;
let task_tx;
let trailers;
{
let mut st = self.shared.borrow_mut();
if st.closed {
return Err(st.close_error.clone().unwrap_or_else(|| {
H2Error::new(ErrorCode::InternalError, "connection closed")
}));
}
if st.goaway_received {
return Err(H2Error::new(
ErrorCode::RefusedStream,
"connection is going away",
));
}
id = st.next_stream_id;
st.next_stream_id += 2;
let (htx, hrx) = oneshot::channel();
let initial = st.remote.initial_window_size;
let mut stream = StreamState::new(id, initial);
stream.head_tx = Some(htx);
stream.local_closed = !has_body;
trailers = stream.trailers.clone();
recv = stream.recv.clone(); st.streams.insert(id, stream);
head_rx = hrx;
let headers = build_request_headers(&init);
let block = st.encoder.encode(&headers);
st.send_headers(id, block, !has_body);
task_tx = st.task_tx.clone();
}
if has_body {
let pump = pump_body(self.shared.clone(), id, body);
let _ = task_tx.unbounded_send(pump.boxed_local());
}
match head_rx.await {
Ok(Ok(head)) => Ok(Response {
status: head.status,
headers: head.headers,
raw_headers: head.raw,
body: ResponseBody {
recv,
conn: Rc::downgrade(&self.shared),
stream_id: id,
},
trailers,
}),
Ok(Err(e)) => Err(e),
Err(_canceled) => Err(self
.shared
.borrow()
.close_error
.clone()
.unwrap_or_else(|| H2Error::new(ErrorCode::InternalError, "connection closed"))),
}
}
pub async fn ping(&self) -> Result<f64, H2Error> {
let rx = {
let mut st = self.shared.borrow_mut();
if st.closed {
return Err(st.close_error.clone().unwrap_or_else(|| {
H2Error::new(ErrorCode::InternalError, "connection closed")
}));
}
st.ping_counter = st.ping_counter.wrapping_add(1);
let mut opaque = [0u8; 8];
opaque[4..8].copy_from_slice(&st.ping_counter.to_be_bytes());
let (tx, rx) = oneshot::channel();
st.pings.insert(
opaque,
PingWaiter {
resolve: tx,
sent_at: now_millis(),
},
);
st.send_frame(Frame::Ping {
ack: false,
opaque_data: opaque,
});
rx
};
match rx.await {
Ok(res) => res,
Err(_canceled) => Err(H2Error::new(ErrorCode::InternalError, "connection closed")),
}
}
pub fn close(&self) {
let mut st = self.shared.borrow_mut();
if st.closed {
return;
}
st.send_frame(Frame::Goaway {
last_stream_id: st.highest_promised,
error_code: 0,
debug_data: Vec::new(),
});
st.destroy(H2Error::new(
ErrorCode::NoError,
"connection closed by client",
));
}
}
async fn pump_body(shared: Rc<RefCell<ConnState>>, id: u32, body: RequestBody) {
let mut chunks: LocalBoxStream<'static, Vec<u8>> = match body {
RequestBody::Empty => return,
RequestBody::Bytes(bytes) => stream::once(async move { bytes }).boxed_local(),
RequestBody::Stream(s) => s,
};
while let Some(chunk) = chunks.next().await {
if chunk.is_empty() {
continue;
}
if !pump_chunk(&shared, id, &chunk).await {
return; }
}
let mut st = shared.borrow_mut();
if st.streams.contains_key(&id) {
st.send_frame(Frame::Data {
stream_id: id,
data: Vec::new(),
end_stream: true,
});
if let Some(s) = st.streams.get_mut(&id) {
s.local_closed = true;
}
st.retire_if_fully_closed(id);
}
}
async fn pump_chunk(shared: &Rc<RefCell<ConnState>>, id: u32, chunk: &[u8]) -> bool {
let mut offset = 0;
while offset < chunk.len() {
let alive = poll_fn(|cx| {
let mut st = shared.borrow_mut();
if st.closed || !st.streams.contains_key(&id) {
return Poll::Ready(false);
}
let conn_ready = st.conn_send_window.is_ready();
let stream_ready = st
.streams
.get(&id)
.map(|s| s.send_window.is_ready())
.unwrap_or(false);
if conn_ready && stream_ready {
Poll::Ready(true)
} else {
if !conn_ready {
st.conn_send_window.register_waker(cx.waker());
}
if !stream_ready {
if let Some(s) = st.streams.get_mut(&id) {
s.send_window.register_waker(cx.waker());
}
}
Poll::Pending
}
})
.await;
if !alive {
return false;
}
let mut st = shared.borrow_mut();
if st.closed
|| st
.streams
.get(&id)
.map(|s| s.send_window.is_closed())
.unwrap_or(true)
{
return false;
}
let conn_w = st.conn_send_window.value();
let stream_w = st.streams.get(&id).unwrap().send_window.value();
let max = st.remote.max_frame_size as i64;
let remaining = (chunk.len() - offset) as i64;
let grant = remaining.min(conn_w).min(stream_w).min(max);
if grant <= 0 {
continue; }
st.conn_send_window.consume(grant);
st.streams.get_mut(&id).unwrap().send_window.consume(grant);
let slice = chunk[offset..offset + grant as usize].to_vec();
st.send_frame(Frame::Data {
stream_id: id,
data: slice,
end_stream: false,
});
offset += grant as usize;
}
true
}
#[cfg(target_arch = "wasm32")]
fn now_millis() -> f64 {
js_sys::Date::now()
}
#[cfg(not(target_arch = "wasm32"))]
fn now_millis() -> f64 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs_f64() * 1000.0)
.unwrap_or(0.0)
}
fn build_request_headers(init: &RequestInit) -> Vec<Header> {
let method = init
.method
.clone()
.unwrap_or_else(|| "GET".into())
.to_uppercase();
let scheme = init.scheme.clone().unwrap_or_else(|| "http".into());
let path = init.path.clone().unwrap_or_else(|| "/".into());
let mut headers = vec![
Header::new(":method", method),
Header::new(":scheme", scheme),
];
if let Some(auth) = &init.authority {
headers.push(Header::new(":authority", auth.clone()));
}
headers.push(Header::new(":path", path));
for (raw_name, value) in &init.headers {
let name = raw_name.to_ascii_lowercase();
if name.starts_with(':') || FORBIDDEN_HEADERS.contains(&name.as_str()) {
continue;
}
if name == "authorization" || name == "cookie" {
headers.push(Header::never_indexed(name, value.clone()));
} else {
headers.push(Header::new(name, value.clone()));
}
}
headers
}
pub fn connect(
transport: Transport,
options: ConnectOptions,
) -> (H2Connection, impl Future<Output = ()>) {
let (out_tx, out_rx) = mpsc::unbounded();
let (task_tx, task_rx) = mpsc::unbounded();
let local_max_frame_size = options.max_frame_size.unwrap_or(DEFAULT_MAX_FRAME_SIZE);
let local_initial_window = options.initial_window_size.unwrap_or(1024 * 1024);
let conn_recv_window = options.connection_window_size.unwrap_or(64 * 1024 * 1024);
let header_table_size = options.header_table_size.unwrap_or(4096);
let enable_push = options.enable_push.unwrap_or(true);
let state = ConnState {
out_tx: Some(out_tx),
task_tx,
encoder: HpackEncoder::new(),
decoder: HpackDecoder::new(header_table_size),
frame_decoder: FrameDecoder::new(local_max_frame_size),
streams: HashMap::new(),
next_stream_id: 1,
conn_send_window: SendWindow::new(SPEC_INITIAL_WINDOW),
remote: RemoteSettings::default(),
pending_header_block: None,
pings: HashMap::new(),
ping_counter: 0,
slot_waiters: Vec::new(),
closed: false,
close_error: None,
goaway_received: false,
highest_promised: 0,
};
let shared = Rc::new(RefCell::new(state));
{
let st = shared.borrow();
st.write_raw(CONNECTION_PREFACE.to_vec());
st.send_frame(Frame::Settings {
ack: false,
settings: Settings {
header_table_size: Some(header_table_size as u32),
enable_push: Some(enable_push),
initial_window_size: Some(local_initial_window),
max_frame_size: Some(local_max_frame_size as u32),
..Default::default()
},
});
let grow = conn_recv_window as i64 - SPEC_INITIAL_WINDOW;
if grow > 0 {
st.send_frame(Frame::WindowUpdate {
stream_id: 0,
window_size_increment: grow as u32,
});
}
}
let driver = drive(
shared.clone(),
transport.reader,
transport.writer,
out_rx,
task_rx,
);
(H2Connection { shared }, driver)
}
async fn drive(
shared: Rc<RefCell<ConnState>>,
mut reader: crate::transport::ByteStream,
mut writer: crate::transport::ByteSink,
mut out_rx: mpsc::UnboundedReceiver<Vec<u8>>,
task_rx: mpsc::UnboundedReceiver<LocalBoxFuture<'static, ()>>,
) {
let read = {
let shared = shared.clone();
async move {
while let Some(chunk) = reader.next().await {
if !chunk.is_empty() {
shared.borrow_mut().on_bytes(&chunk);
}
if shared.borrow().closed {
break;
}
}
shared
.borrow_mut()
.destroy(H2Error::new(ErrorCode::NoError, "transport closed by peer"));
}
};
let write = async move {
while let Some(bytes) = out_rx.next().await {
if writer.send(bytes).await.is_err() {
shared.borrow_mut().destroy(H2Error::new(
ErrorCode::InternalError,
"transport write failed",
));
break;
}
}
};
let tasks = run_tasks(task_rx);
let read = read.fuse();
let tasks = tasks.fuse();
let write = write.fuse();
futures::pin_mut!(read, write, tasks);
futures::future::poll_fn(|cx| {
let _ = read.as_mut().poll(cx);
let _ = tasks.as_mut().poll(cx);
write.as_mut().poll(cx)
})
.await;
}
async fn run_tasks(mut task_rx: mpsc::UnboundedReceiver<LocalBoxFuture<'static, ()>>) {
let mut pending: FuturesUnordered<LocalBoxFuture<'static, ()>> = FuturesUnordered::new();
poll_fn(move |cx| {
while let Poll::Ready(Some(task)) = task_rx.poll_next_unpin(cx) {
pending.push(task);
}
while let Poll::Ready(Some(())) = pending.poll_next_unpin(cx) {}
Poll::Pending
})
.await
}