use std::collections::VecDeque;
use bytes::{Bytes, BytesMut};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct H2Limits {
pub max_message_size: u64,
pub max_message_body_size: u64,
pub max_headers_size: u64,
pub max_header_count: u16,
pub max_concurrent_streams: u32,
pub max_connection_buffer_size: u64,
pub max_premature_resets: u32,
pub max_idle_frames: u32,
pub output_high_water: u64,
pub max_encoder_table_size: u64,
pub read_chunk_size: u64,
pub idle_capacity: u64,
pub read_timeout: f64,
pub write_timeout: f64,
pub receive_timeout: f64,
pub send_timeout: f64,
}
impl Default for H2Limits {
fn default() -> Self {
Limits::default().into()
}
}
impl From<Limits> for H2Limits {
fn from(limits: Limits) -> Self {
Self {
max_message_size: limits.max_message_size,
max_message_body_size: limits.max_message_body_size,
max_headers_size: limits.max_headers_size,
max_header_count: limits.max_header_count,
max_concurrent_streams: limits.max_concurrent_streams,
max_connection_buffer_size: limits.max_connection_buffer_size,
max_premature_resets: limits.max_premature_resets,
max_idle_frames: limits.max_idle_frames,
output_high_water: limits.output_high_water,
max_encoder_table_size: limits.max_encoder_table_size,
read_chunk_size: limits.read_chunk_size,
idle_capacity: limits.idle_capacity,
read_timeout: limits.read_timeout,
write_timeout: limits.write_timeout,
receive_timeout: limits.receive_timeout,
send_timeout: limits.send_timeout,
}
}
}
pub mod frames;
pub use frames::{Code, Flag, Frame, FrameHeader, FrameType, Settings, PREFACE};
use crate::helpers::fields::HeaderField;
use crate::helpers::hpack::{Decoder as HPACKDecoder, Encoder as HPACKEncoder};
use crate::models::{Body, ConnectionID, Headers, Limits, Message, Method, Role, StreamID, Version};
use crate::tls::Security;
use crate::protocol::base::{Connection, Stream};
use crate::protocol::common::{self, Buffer, Error};
use crate::helpers::sync;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StreamState {
Idle,
ReservedLocal,
ReservedRemote,
Open,
HalfClosedLocal,
HalfClosedRemote,
Closed,
}
impl StreamState {
pub fn receivable(&self) -> bool {
matches!(self, Self::Idle | Self::Open | Self::HalfClosedLocal)
}
pub fn sendable(&self) -> bool {
matches!(self, Self::Idle | Self::Open | Self::HalfClosedRemote)
}
pub fn close_local(&self) -> Self {
match self {
Self::Open | Self::Idle => Self::HalfClosedLocal,
_ => Self::Closed,
}
}
pub fn close_remote(&self) -> Self {
match self {
Self::Open | Self::Idle => Self::HalfClosedRemote,
_ => Self::Closed,
}
}
}
pub struct H2Stream {
id: StreamID,
state: StreamState,
window_local: i64,
window_remote: i64,
block: Vec<u8>,
head: u64,
body: BytesMut,
headers: Option<Message>,
method: Option<Method>,
pending_reset: Option<u64>,
}
impl H2Stream {
pub fn new(id: StreamID, window_local: i64, window_remote: i64) -> Self {
Self {
id,
state: StreamState::Idle,
window_local,
window_remote,
block: Vec::new(),
head: 0,
body: BytesMut::new(),
headers: None,
method: None,
pending_reset: None,
}
}
pub fn state(&self) -> StreamState {
self.state
}
pub fn received(&self) -> u64 {
self.head + self.body.len() as u64
}
pub fn window_local(&self) -> i64 {
self.window_local
}
pub fn window_remote(&self) -> i64 {
self.window_remote
}
}
impl Stream for H2Stream {
fn id(&self) -> StreamID {
self.id
}
async fn reset(&mut self, code: u64) {
self.state = StreamState::Closed;
self.pending_reset = Some(code);
}
}
pub struct H2Connection<T> {
transport: T,
role: Role,
id: ConnectionID,
limits: H2Limits,
buffer: Buffer,
streams: common::StreamMap<StreamID, H2Stream>,
hpack_encoder: HPACKEncoder,
hpack_decoder: HPACKDecoder,
settings_local: Settings,
settings_remote: Settings,
window_local: i64,
window_remote: i64,
next_stream_id: u64,
highest_peer_stream_id: u64,
started: bool,
goaway: Option<u32>,
ready: VecDeque<Message>,
out: BytesMut,
block: Vec<u8>,
fields: Vec<HeaderField>,
buffered_bound: u64,
premature_resets: u32,
idle_frames: u32,
request_finalizer: crate::finalizer::RequestFinalizer,
response_finalizer: crate::finalizer::ResponseFinalizer,
security: Security,
}
impl<T> H2Connection<T>
where
T: AsyncRead + AsyncWrite + Unpin,
{
pub fn new(transport: T, role: Role, id: ConnectionID, limits: impl Into<H2Limits>) -> Self {
let limits = limits.into();
Self::resume(transport, role, id, limits, Buffer::new())
}
pub fn resume(transport: T, role: Role, id: ConnectionID, limits: impl Into<H2Limits>, buffer: Buffer) -> Self {
let limits = limits.into();
let settings_local = Settings { max_concurrent_streams: Some(limits.max_concurrent_streams), ..Settings::default() };
let mut hpack_encoder = HPACKEncoder::new();
hpack_encoder.set_capacity_limit(limits.max_encoder_table_size as usize);
let mut buffer = buffer;
buffer.set_chunk_size(limits.read_chunk_size as usize);
Self {
transport,
role,
id,
limits,
buffer,
streams: common::StreamMap::default(),
hpack_encoder,
hpack_decoder: HPACKDecoder::new(),
settings_local,
settings_remote: Settings::peer(),
window_local: Settings::DEFAULT_INITIAL_WINDOW_SIZE as i64,
window_remote: Settings::DEFAULT_INITIAL_WINDOW_SIZE as i64,
next_stream_id: if role.is_client() { 1 } else { 2 },
highest_peer_stream_id: 0,
started: false,
goaway: None,
ready: VecDeque::new(),
out: BytesMut::new(),
block: Vec::new(),
fields: Vec::new(),
buffered_bound: 0,
premature_resets: 0,
idle_frames: 0,
request_finalizer: crate::finalizer::RequestFinalizer::default(),
response_finalizer: crate::finalizer::ResponseFinalizer::new(None),
security: Security::default(),
}
}
pub fn limits(&self) -> H2Limits {
self.limits
}
pub fn with_request_finalizer(mut self, finalizer: crate::finalizer::RequestFinalizer) -> Self {
self.request_finalizer = finalizer;
self
}
pub fn with_response_finalizer(mut self, finalizer: crate::finalizer::ResponseFinalizer) -> Self {
self.response_finalizer = finalizer;
self
}
pub fn with_security(mut self, security: Security) -> Self {
self.security = security;
self
}
pub fn settings_local(&self) -> &Settings {
&self.settings_local
}
pub fn settings_remote(&self) -> &Settings {
&self.settings_remote
}
pub fn hpack_encoder(&self) -> &HPACKEncoder {
&self.hpack_encoder
}
pub fn hpack_decoder(&self) -> &HPACKDecoder {
&self.hpack_decoder
}
pub fn local_stream_ceiling(&self) -> usize {
let advertised = self.settings_remote.max_concurrent_streams.unwrap_or(self.limits.max_concurrent_streams);
(advertised as usize).max(1)
}
pub fn stream(&self, stream_id: StreamID) -> Option<&H2Stream> {
self.streams.get(&stream_id)
}
pub fn open_stream(&mut self, stream_id: StreamID) -> Result<&mut H2Stream, Error> {
self.streams
.get_mut(&stream_id)
.ok_or_else(|| Error::Protocol(format!("stream {} is no longer open", stream_id.0)))
}
pub async fn start(&mut self) -> Result<(), Error> {
if self.started {
return Ok(());
}
self.started = true;
if self.role.is_client() {
self.out.extend_from_slice(PREFACE);
} else {
let preface = self.buffer.require(&mut self.transport, PREFACE.len(), self.limits.read_timeout).await?;
if preface != PREFACE {
return Err(Error::Protocol("connection preface is not the HTTP/2 preface".into()));
}
self.buffer.consume(PREFACE.len());
}
self.hpack_decoder.set_max_capacity(self.settings_local.header_table_size as usize);
self.hpack_decoder.set_max_decoded_size(self.limits.max_headers_size as usize);
let settings = Frame::Settings { ack: false, params: self.settings_local.parameters() };
self.queue(&settings);
self.flush_out().await
}
pub fn queue(&mut self, frame: &Frame) {
frame.encode_into(&mut self.out);
}
pub async fn flush_out(&mut self) -> Result<(), Error> {
if self.out.is_empty() {
return Ok(());
}
let out = std::mem::take(&mut self.out);
let transport = &mut self.transport;
let result = sync::Timeout::within(self.limits.write_timeout, async move {
transport.write_all(&out).await?;
transport.flush().await.map(|()| out)
})
.await;
match result? {
Ok(out) => {
self.out = out;
self.out.clear();
common::Buffer::reclaim_bytes(&mut self.out, self.limits.idle_capacity as usize);
Ok(())
}
Err(error) => Err(error.into()),
}
}
pub async fn write(&mut self, frame: &Frame) -> Result<(), Error> {
self.queue(frame);
self.flush_out().await
}
pub async fn receive_message(&mut self) -> Result<Message, Error> {
loop {
if let Some(message) = self.ready.pop_front() {
return Ok(message);
}
if let Some(message) = self.pump().await? {
return Ok(message);
}
if self.goaway.is_some() && self.streams.is_empty() {
return Err(Error::Closed);
}
}
}
pub async fn send_message(&mut self, message: Message) -> Result<(), Error> {
let mut message = message;
self.request_finalizer.finalize(self.role, &mut message);
self.response_finalizer.finalize(self.role, self.security.secure, &mut message);
self.start().await?;
self.flush_resets().await?;
let stream_id = match message.stream_id {
Some(stream_id) => stream_id,
None => {
let ceiling = self.local_stream_ceiling();
if self.streams.len() >= ceiling {
return Err(Error::Limit(format!("more than {ceiling} streams are open at once")));
}
let stream_id = StreamID(self.next_stream_id);
self.next_stream_id += 2;
stream_id
}
};
let window_local = self.settings_local.initial_window_size as i64;
let window_remote = self.settings_remote.initial_window_size as i64;
self.streams.entry(stream_id).or_insert_with(|| H2Stream::new(stream_id, window_local, window_remote));
self.fields.clear();
common::Fields::write(&message, &mut self.fields)?;
let mut block = std::mem::take(&mut self.block);
block.clear();
self.hpack_encoder.encode_into(&mut block, &self.fields);
let body = match message.body.take() {
Some(body) => Some(body.into_bytes().await?),
None => None,
};
let body = body.filter(|body| !body.is_empty());
let trailers = message.trailers.as_ref().filter(|trailers| !trailers.is_empty());
let tunneling = message.method == Some(Method::CONNECT)
|| (matches!(message.status_code, Some(200..=299))
&& self.streams.get(&stream_id).is_some_and(|stream| stream.method == Some(Method::CONNECT)));
let open = tunneling || message.is_informational();
let end_stream = !open && body.is_none() && trailers.is_none();
let written = self.write_block(stream_id, &block, end_stream).await;
self.block = block;
common::Buffer::reclaim_octets(&mut self.block, self.limits.idle_capacity as usize);
written?;
if let Some(body) = body {
self.write_data(stream_id, &body, trailers.is_none()).await?;
}
if let Some(trailers) = trailers {
self.fields.clear();
self.fields.extend_from_slice(trailers.fields());
let mut block = std::mem::take(&mut self.block);
block.clear();
self.hpack_encoder.encode_into(&mut block, &self.fields);
let written = self.write_block(stream_id, &block, true).await;
self.block = block;
common::Buffer::reclaim_octets(&mut self.block, self.limits.idle_capacity as usize);
written?;
}
if let Some(stream) = self.streams.get_mut(&stream_id) {
if !open {
stream.state = stream.state.close_local();
}
if message.is_request() {
stream.method = message.method;
}
}
self.retire(stream_id);
if self.buffer.is_empty() {
self.flush_out().await?;
}
Ok(())
}
pub async fn reset(&mut self, stream_id: StreamID, error_code: u32) -> Result<(), Error> {
self.streams.remove(&stream_id);
self.write(&Frame::RstStream { stream_id, error_code }).await
}
pub async fn overloaded(&mut self, reason: impl Into<String>) -> Error {
let goaway = Frame::GoAway {
last_stream_id: StreamID(self.highest_peer_stream_id),
error_code: Code::ENHANCE_YOUR_CALM,
debug_data: Vec::new(),
};
let _ = self.write(&goaway).await;
Error::Limit(reason.into())
}
pub async fn idle(&mut self) -> Result<(), Error> {
self.idle_frames = self.idle_frames.saturating_add(1);
if self.idle_frames > self.limits.max_idle_frames {
let reason = format!("more than {} frames arrived without advancing a stream", self.limits.max_idle_frames);
return Err(self.overloaded(reason).await);
}
Ok(())
}
pub fn buffered(&self) -> u64 {
self.streams.values().map(|stream| stream.body.len() as u64).sum()
}
pub fn overbuffered(&mut self) -> bool {
let limit = self.limits.max_connection_buffer_size;
if self.buffered_bound <= limit {
return false;
}
self.buffered_bound = self.buffered();
self.buffered_bound > limit
}
pub fn retire(&mut self, stream_id: StreamID) {
if self.streams.get(&stream_id).is_some_and(|stream| stream.state == StreamState::Closed) {
self.streams.remove(&stream_id);
}
}
pub async fn read_frame(&mut self) -> Result<Frame, Error> {
let max_frame_size = self.settings_local.max_frame_size;
loop {
if let Some(frame) = Frame::parse(self.buffer.as_bytes_mut(), max_frame_size)? {
return Ok(frame);
}
self.flush_out().await?;
if !self.buffer.fill(&mut self.transport, self.limits.read_timeout).await? {
return Err(Error::Closed);
}
}
}
pub async fn pump(&mut self) -> Result<Option<Message>, Error> {
self.start().await?;
self.flush_resets().await?;
let frame = self.read_frame().await?;
let message = self.handle(frame).await?;
if self.buffer.is_empty() {
self.flush_out().await?;
self.buffer.reclaim(self.limits.idle_capacity as usize);
}
Ok(message)
}
pub async fn handle(&mut self, frame: Frame) -> Result<Option<Message>, Error> {
match frame {
Frame::Settings { ack: false, params } => {
self.idle().await?;
let window_before = self.settings_remote.initial_window_size;
for (id, value) in params {
self.settings_remote.apply(id, value)?;
}
let change = self.settings_remote.initial_window_size as i64 - window_before as i64;
if change != 0 {
for stream in self.streams.values_mut() {
stream.window_remote += change;
}
}
self.hpack_encoder.set_max_capacity(self.settings_remote.header_table_size as usize);
self.queue(&Frame::Settings { ack: true, params: Vec::new() });
Ok(None)
}
Frame::Settings { ack: true, .. } => {
self.idle().await?;
Ok(None)
}
Frame::Ping { ack: false, payload } => {
self.idle().await?;
self.queue(&Frame::Ping { ack: true, payload });
Ok(None)
}
Frame::Ping { ack: true, .. } => {
self.idle().await?;
Ok(None)
}
Frame::WindowUpdate { stream_id, increment } => {
let mut unblocked = false;
if stream_id == StreamID(0) {
let stalled = self.window_remote <= 0;
self.window_remote += increment as i64;
if self.window_remote > Settings::MAXIMUM_WINDOW_SIZE as i64 {
return Err(Error::Protocol("connection send window overflowed".into()));
}
unblocked = stalled && self.window_remote > 0;
} else if let Some(stream) = self.streams.get_mut(&stream_id) {
let stalled = stream.window_remote <= 0;
stream.window_remote += increment as i64;
if stream.window_remote > Settings::MAXIMUM_WINDOW_SIZE as i64 {
let stream_id = stream.id;
self.reset(stream_id, Code::FLOW_CONTROL_ERROR).await?;
return Ok(None);
}
unblocked = stalled && stream.window_remote > 0;
}
if unblocked {
self.idle_frames = 0;
} else {
self.idle().await?;
}
Ok(None)
}
Frame::GoAway { error_code, .. } => {
self.idle().await?;
self.goaway = Some(error_code);
Ok(None)
}
Frame::RstStream { stream_id, .. } => {
self.idle().await?;
let premature = self.streams.remove(&stream_id).is_some_and(|stream| stream.state.sendable());
if premature {
self.premature_resets = self.premature_resets.saturating_add(1);
if self.premature_resets > self.limits.max_premature_resets {
let reason = format!(
"more than {} streams were reset before a response was sent",
self.limits.max_premature_resets
);
return Err(self.overloaded(reason).await);
}
}
Ok(None)
}
Frame::Priority { .. } => {
self.idle().await?;
Ok(None)
}
Frame::PushPromise { .. } => Err(Error::Protocol("PUSH_PROMISE arrived with push disabled".into())),
Frame::Headers { stream_id, end_stream, end_headers, block } => {
self.begin_stream(stream_id)?;
if end_headers {
return self.finish_headers(stream_id, &block, end_stream);
}
let gathered = &mut self.open_stream(stream_id)?.block;
gathered.clear();
gathered.extend_from_slice(&block);
self.continue_headers(stream_id, end_stream).await
}
Frame::Continuation { .. } => {
Err(Error::Protocol("CONTINUATION arrived outside a header block".into()))
}
Frame::Data { stream_id, end_stream, data } => {
match self.streams.get(&stream_id) {
None => return Err(Error::Protocol(format!("DATA on unopened stream {}", stream_id.0))),
Some(stream) if !stream.state.receivable() => {
return Err(Error::Protocol(format!("DATA on closed stream {}", stream_id.0)));
}
Some(_) => {}
}
if data.is_empty() && !end_stream {
self.idle().await?;
} else {
self.idle_frames = 0;
}
let stream = self.open_stream(stream_id)?;
stream.window_local -= data.len() as i64;
if stream.window_local < 0 {
self.reset(stream_id, Code::FLOW_CONTROL_ERROR).await?;
return Ok(None);
}
stream.body.extend_from_slice(&data);
let body = stream.body.len() as u64;
let received = stream.received();
self.buffered_bound = self.buffered_bound.saturating_add(data.len() as u64);
let limit = self.limits.max_message_body_size;
if body > limit {
return Err(Error::Limit(format!("body exceeds {limit} octets")));
}
let limit = self.limits.max_message_size;
if received > limit {
return Err(Error::Limit(format!("message exceeds {limit} octets")));
}
if self.overbuffered() {
let limit = self.limits.max_connection_buffer_size;
return Err(Error::Limit(format!("buffered messages exceed {limit} octets")));
}
self.window_local -= data.len() as i64;
if self.window_local < 0 {
return Err(Error::Protocol("connection receive window overflowed".into()));
}
if !data.is_empty() {
let increment = data.len() as u32;
self.queue(&Frame::WindowUpdate { stream_id: StreamID(0), increment });
self.window_local += increment as i64;
if self.streams.get(&stream_id).is_some_and(|stream| stream.state.receivable()) {
self.queue(&Frame::WindowUpdate { stream_id, increment });
if let Some(stream) = self.streams.get_mut(&stream_id) {
stream.window_local += increment as i64;
}
}
}
if end_stream {
return Ok(self.complete(stream_id, true));
}
Ok(None)
}
}
}
pub async fn continue_headers(&mut self, stream_id: StreamID, end_stream: bool) -> Result<Option<Message>, Error> {
let mut frames = 0u64;
loop {
let size = self.streams.get(&stream_id).map(|stream| stream.block.len()).unwrap_or_default() as u64;
if size > self.limits.max_headers_size {
return Err(Error::Limit(format!("field block exceeds {} octets", self.limits.max_headers_size)));
}
frames += 1;
if frames > self.limits.max_header_count as u64 {
return Err(Error::Limit(format!("field block spans more than {} CONTINUATION frames", self.limits.max_header_count)));
}
match self.read_frame().await? {
Frame::Continuation { stream_id: other, end_headers, block } if other == stream_id => {
self.open_stream(stream_id)?.block.extend_from_slice(&block);
if end_headers {
let gathered = std::mem::take(&mut self.open_stream(stream_id)?.block);
let finished = self.finish_headers(stream_id, &gathered, end_stream);
if let Ok(stream) = self.open_stream(stream_id) {
stream.block = gathered;
stream.block.clear();
}
return finished;
}
}
_ => return Err(Error::Protocol("a field block was interrupted".into())),
}
}
}
pub fn begin_stream(&mut self, stream_id: StreamID) -> Result<(), Error> {
let peer_odd = !self.role.is_client();
if stream_id.0 == 0 || (stream_id.0 % 2 == 1) != peer_odd {
if !self.streams.contains_key(&stream_id) {
return Err(Error::Protocol(format!("stream {} is not the peer's to open", stream_id.0)));
}
return Ok(());
}
if !self.streams.contains_key(&stream_id) {
if stream_id.0 <= self.highest_peer_stream_id {
return Err(Error::Protocol(format!("stream {} does not exceed the last stream the peer opened", stream_id.0)));
}
self.highest_peer_stream_id = stream_id.0;
if let Some(max) = self.settings_local.max_concurrent_streams
&& self.streams.len() as u32 >= max
{
return Err(Error::Protocol(format!("stream {} exceeds the concurrent stream limit", stream_id.0)));
}
let window_local = self.settings_local.initial_window_size as i64;
let window_remote = self.settings_remote.initial_window_size as i64;
self.streams.insert(stream_id, H2Stream::new(stream_id, window_local, window_remote));
}
Ok(())
}
pub fn finish_headers(&mut self, stream_id: StreamID, block: &[u8], end_stream: bool) -> Result<Option<Message>, Error> {
self.idle_frames = 0;
let received = {
let stream = self.open_stream(stream_id)?;
stream.head += block.len() as u64;
stream.received()
};
let limit = self.limits.max_message_size;
if received > limit {
return Err(Error::Limit(format!("message exceeds {limit} octets")));
}
let decoded = self.hpack_decoder.decode(block)?;
if decoded.len() > self.limits.max_header_count as usize {
return Err(Error::Limit(format!("more than {} header fields", self.limits.max_header_count)));
}
let connection_id = self.id.clone();
let security = self.security;
let stream = self.open_stream(stream_id)?;
if let Some(message) = &mut stream.headers {
let mut trailers = Headers::with_capacity(decoded.len());
for field in decoded {
if field.name.starts_with(':') {
return Err(Error::Protocol("trailer section carries a pseudo-header".into()));
}
trailers.append(field.name, field.value);
}
message.trailers = Some(trailers);
} else {
let mut message = common::Fields::into_message(decoded, Version::V2_0)?;
message.stream_id = Some(stream_id);
message.connection_id = Some(connection_id);
security.apply(&mut message);
if message.is_request() {
stream.method = message.method;
}
if message.is_informational() {
stream.state = if stream.state == StreamState::Idle { StreamState::Open } else { stream.state };
return Ok(Some(message));
}
stream.headers = Some(message);
}
stream.state = if stream.state == StreamState::Idle { StreamState::Open } else { stream.state };
let tunneling = stream.headers.as_ref().is_some_and(|message| message.tunneling(stream.method));
if end_stream || tunneling {
return Ok(self.complete(stream_id, end_stream));
}
Ok(None)
}
pub fn complete(&mut self, stream_id: StreamID, end_stream: bool) -> Option<Message> {
let stream = self.streams.get_mut(&stream_id)?;
if end_stream {
stream.state = stream.state.close_remote();
}
let mut message = stream.headers.take()?;
if !stream.body.is_empty() {
message.body = Some(Body::Data(std::mem::take(&mut stream.body).freeze()));
}
self.retire(stream_id);
Some(message)
}
pub fn drain(&mut self, stream_id: StreamID) -> Option<Bytes> {
let stream = self.streams.get_mut(&stream_id)?;
(!stream.body.is_empty()).then(|| std::mem::take(&mut stream.body).freeze())
}
pub async fn flush_resets(&mut self) -> Result<(), Error> {
let pending: Vec<(StreamID, u64)> = self
.streams
.iter_mut()
.filter_map(|(id, stream)| stream.pending_reset.take().map(|code| (*id, code)))
.collect();
for (stream_id, code) in pending {
self.queue(&Frame::RstStream { stream_id, error_code: code as u32 });
self.streams.remove(&stream_id);
}
Ok(())
}
pub async fn write_block(&mut self, stream_id: StreamID, block: &[u8], end_stream: bool) -> Result<(), Error> {
let size = self.settings_remote.max_frame_size as usize;
let mut chunks = block.chunks(size.max(1));
let first = chunks.next().unwrap_or_default();
let mut rest = chunks.peekable();
let end_headers = if rest.peek().is_none() { Flag::END_HEADERS } else { 0 };
let flags = end_headers | if end_stream { Flag::END_STREAM } else { 0 };
FrameHeader::write(&mut self.out, FrameType::Headers, flags, stream_id, first);
while let Some(chunk) = rest.next() {
let flags = if rest.peek().is_none() { Flag::END_HEADERS } else { 0 };
FrameHeader::write(&mut self.out, FrameType::Continuation, flags, stream_id, chunk);
}
Ok(())
}
pub async fn write_data(&mut self, stream_id: StreamID, data: &[u8], end_stream: bool) -> Result<(), Error> {
let mut rest = data;
loop {
let window = self.window_remote.min(self.streams.get(&stream_id).map(|stream| stream.window_remote).unwrap_or_default());
if window <= 0 && !rest.is_empty() {
if let Some(message) = self.pump().await? {
self.ready.push_back(message);
}
continue;
}
let size = rest.len().min(window.max(0) as usize).min(self.settings_remote.max_frame_size as usize);
let (chunk, remaining) = rest.split_at(size);
rest = remaining;
let flags = if end_stream && rest.is_empty() { Flag::END_STREAM } else { 0 };
FrameHeader::write(&mut self.out, FrameType::Data, flags, stream_id, chunk);
if self.out.len() >= self.limits.output_high_water as usize {
self.flush_out().await?;
}
self.window_remote -= size as i64;
if let Some(stream) = self.streams.get_mut(&stream_id) {
stream.window_remote -= size as i64;
}
if rest.is_empty() {
return Ok(());
}
}
}
}
impl<T> H2Connection<T>
where
T: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
pub fn tunnel(self, stream_id: StreamID) -> H2Tunnel {
let (application, internal) = tokio::io::duplex(self.limits.read_chunk_size as usize);
let driver = tokio::spawn(async move { self.drive(stream_id, internal).await });
H2Tunnel { stream: application, driver }
}
pub async fn drive(mut self, stream_id: StreamID, internal: tokio::io::DuplexStream) -> Result<(), Error> {
let (mut reader, mut writer) = tokio::io::split(internal);
let mut scratch = vec![0u8; self.limits.read_chunk_size as usize];
self.start().await?;
loop {
self.flush_out().await?;
tokio::select! {
biased;
frame = self.read_frame() => {
let frame = frame?;
self.handle(frame).await?;
if let Some(data) = self.drain(stream_id) {
writer.write_all(&data).await?;
}
if self.streams.get(&stream_id).is_none_or(|stream| !stream.state.receivable()) {
writer.shutdown().await?;
return Ok(());
}
}
read = reader.read(&mut scratch) => {
match read? {
0 => {
self.write_data(stream_id, &[], true).await?;
return Ok(());
}
read => self.write_data(stream_id, &scratch[..read], false).await?,
}
}
}
}
}
}
pub struct H2Tunnel {
stream: tokio::io::DuplexStream,
driver: tokio::task::JoinHandle<Result<(), Error>>,
}
impl H2Tunnel {
pub fn abort(&self) {
self.driver.abort();
}
pub fn finished(&self) -> bool {
self.driver.is_finished()
}
}
impl AsyncRead for H2Tunnel {
fn poll_read(mut self: std::pin::Pin<&mut Self>, context: &mut std::task::Context<'_>, buffer: &mut tokio::io::ReadBuf<'_>) -> std::task::Poll<std::io::Result<()>> {
std::pin::Pin::new(&mut self.stream).poll_read(context, buffer)
}
}
impl AsyncWrite for H2Tunnel {
fn poll_write(mut self: std::pin::Pin<&mut Self>, context: &mut std::task::Context<'_>, data: &[u8]) -> std::task::Poll<std::io::Result<usize>> {
std::pin::Pin::new(&mut self.stream).poll_write(context, data)
}
fn poll_flush(mut self: std::pin::Pin<&mut Self>, context: &mut std::task::Context<'_>) -> std::task::Poll<std::io::Result<()>> {
std::pin::Pin::new(&mut self.stream).poll_flush(context)
}
fn poll_shutdown(mut self: std::pin::Pin<&mut Self>, context: &mut std::task::Context<'_>) -> std::task::Poll<std::io::Result<()>> {
std::pin::Pin::new(&mut self.stream).poll_shutdown(context)
}
}
impl<T> Connection for H2Connection<T>
where
T: AsyncRead + AsyncWrite + Unpin,
{
fn version(&self) -> Version {
Version::V2_0
}
fn role(&self) -> Role {
self.role
}
fn id(&self) -> ConnectionID {
self.id.clone()
}
fn security(&self) -> Security {
self.security
}
async fn send(&mut self, message: Message) -> Result<(), Error> {
let timeout = self.limits.send_timeout;
let sending = std::pin::pin!(self.send_message(message));
sync::Timeout::within(timeout, sending).await?
}
async fn receive(&mut self) -> Result<Message, Error> {
let timeout = self.limits.receive_timeout;
let receiving = std::pin::pin!(self.receive_message());
sync::Timeout::within(timeout, receiving).await?
}
async fn close(&mut self) {
let last_stream_id = StreamID(self.next_stream_id.saturating_sub(2));
let goaway = Frame::GoAway { last_stream_id, error_code: Code::NO_ERROR, debug_data: Vec::new() };
let _ = self.write(&goaway).await;
let _ = self.transport.shutdown().await;
}
}