use std::collections::HashMap;
use std::pin::Pin;
use crate::error;
use crate::result;
use crate::AnySocketAddr;
use crate::solicit::frame::GoawayFrame;
use crate::solicit::frame::HttpFrameType;
use crate::solicit::frame::HttpSettings;
use crate::solicit::frame::RstStreamFrame;
use crate::solicit::frame::WindowUpdateFrame;
use crate::solicit::session::StreamState;
use crate::solicit::session::StreamStateIdleOrClosed;
use crate::solicit::DEFAULT_SETTINGS;
use super::closed_streams::*;
use super::conf::*;
use super::stream::*;
use super::stream_map::*;
use super::types::*;
use super::window_size;
pub use crate::resp::Response;
use crate::client_died_error_holder::ConnDiedType;
use crate::client_died_error_holder::SomethingDiedErrorHolder;
use crate::codec::http_decode_read::HttpDecodeRead;
use crate::codec::queued_write::QueuedWrite;
use crate::common::conn_command_channel::ConnCommandReceiver;
use crate::common::conn_command_channel::ConnCommandSender;
use crate::common::conn_read::ConnReadSideCustom;
use crate::common::conn_write::ConnWriteSideCustom;
use crate::common::init_where::InitWhere;
use crate::hpack;
use crate::solicit::stream_id::StreamId;
use crate::solicit::window_size::NonNegativeWindowSize;
use crate::solicit::window_size::WindowSize;
use crate::ErrorCode;
use futures::channel::oneshot;
use futures::future;
use crate::common::loop_event::LoopEvent;
use crate::log_ndc_future::log_ndc_future;
use futures::future::Future;
use futures::stream::Stream;
use futures::task::Context;
use std::mem;
use std::sync::Arc;
use std::task::Poll;
use tokio::io::split;
use tokio::io::AsyncRead;
use tokio::io::AsyncWrite;
use tokio::io::ReadHalf;
use tokio::io::WriteHalf;
use tokio::runtime::Handle;
pub trait ConnSpecific: Send + 'static {}
pub(crate) struct Conn<T: Types, I: AsyncWrite + AsyncRead + Send + 'static> {
pub peer_addr: AnySocketAddr,
pub conn_died_error_holder: SomethingDiedErrorHolder<ConnDiedType>,
pub specific: T::ConnSpecific,
pub to_write_tx: ConnCommandSender<T>,
pub loop_handle: Handle,
pub streams: StreamMap<T>,
pub peer_closed_streams: ClosedStreams,
pub last_local_stream_id: StreamId,
pub last_peer_stream_id: StreamId,
pub goaway_sent: Option<GoawayFrame>,
pub goaway_received: Option<GoawayFrame>,
pub ping_sent: Option<u64>,
pub out_window_size: WindowSize,
pub in_window_size: NonNegativeWindowSize,
pub pump_out_window_size: window_size::ConnOutWindowSender,
pub framed_read: HttpDecodeRead<ReadHalf<I>>,
pub queued_write: QueuedWrite<WriteHalf<I>>,
pub encoder: hpack::Encoder,
pub write_rx: ConnCommandReceiver<T>,
pub peer_settings: HttpSettings,
pub our_settings_ack: HttpSettings,
pub our_settings_sent: Option<HttpSettings>,
}
impl<T: Types, I: AsyncWrite + AsyncRead + Send + 'static> Drop for Conn<T, I> {
fn drop(&mut self) {
mem::take(&mut self.streams).conn_died(|| self.conn_died_error_holder.error());
}
}
#[derive(Debug, Clone)]
pub struct ConnStateSnapshot {
pub peer_addr: AnySocketAddr,
pub in_window_size: i32,
pub out_window_size: i32,
pub pump_out_window_size: isize,
pub out_buf_bytes: usize,
pub streams: HashMap<StreamId, HttpStreamStateSnapshot>,
}
impl ConnStateSnapshot {
pub fn single_stream(&self) -> (u32, &HttpStreamStateSnapshot) {
let mut iter = self.streams.iter();
let (&id, stream) = iter.next().expect("no streams");
assert!(iter.next().is_none(), "more than one stream");
(id, stream)
}
}
impl<T, I> Conn<T, I>
where
T: Types,
Self: ConnReadSideCustom<Types = T>,
Self: ConnWriteSideCustom<Types = T>,
HttpStreamCommon<T>: HttpStreamData<Types = T>,
I: AsyncWrite + AsyncRead + Send + 'static,
{
pub fn new(
loop_handle: Handle,
specific: T::ConnSpecific,
_conf: CommonConf,
sent_settings: HttpSettings,
to_write_tx: ConnCommandSender<T>,
write_rx: ConnCommandReceiver<T>,
socket: I,
peer_addr: AnySocketAddr,
conn_died_error_holder: SomethingDiedErrorHolder<ConnDiedType>,
) -> Self {
let in_window_size =
NonNegativeWindowSize::new(DEFAULT_SETTINGS.initial_window_size as i32);
let out_window_size = WindowSize::new(DEFAULT_SETTINGS.initial_window_size as i32);
let pump_window_size = window_size::ConnOutWindowSender::new(out_window_size.size() as u32);
let (read, write) = split(socket);
let framed_read = HttpDecodeRead::new(read);
let queued_write = QueuedWrite::new(write);
Conn {
peer_addr,
conn_died_error_holder,
specific,
to_write_tx,
streams: StreamMap::new(),
last_local_stream_id: 0,
last_peer_stream_id: 0,
loop_handle,
goaway_sent: None,
goaway_received: None,
ping_sent: None,
pump_out_window_size: pump_window_size,
peer_closed_streams: ClosedStreams::new(),
framed_read,
queued_write,
write_rx,
encoder: hpack::Encoder::new(),
in_window_size,
out_window_size,
peer_settings: DEFAULT_SETTINGS,
our_settings_ack: DEFAULT_SETTINGS,
our_settings_sent: Some(sent_settings),
}
}
pub fn next_local_stream_id(&mut self) -> StreamId {
let id = match self.last_local_stream_id {
0 => T::CLIENT_OR_SERVER.first_stream_id(),
n => n + 2,
};
self.last_local_stream_id = id;
id
}
pub fn new_stream_data(
&mut self,
stream_id: StreamId,
in_rem_content_length: Option<u64>,
in_message_stage: InMessageStage,
specific: T::HttpStreamSpecific,
) -> (HttpStreamRef<T>, window_size::StreamOutWindowReceiver) {
let (out_window_sender, out_window_receiver) = self
.pump_out_window_size
.new_stream(self.peer_settings.initial_window_size as u32);
let stream = HttpStreamCommon::new(
self.our_settings_sent().initial_window_size,
self.peer_settings.initial_window_size,
out_window_sender,
in_rem_content_length,
in_message_stage,
specific,
);
let stream = self.streams.insert(stream_id, stream);
(stream, out_window_receiver)
}
pub fn dump_state(&self) -> ConnStateSnapshot {
ConnStateSnapshot {
peer_addr: self.peer_addr.clone(),
in_window_size: self.in_window_size.size(),
out_window_size: self.out_window_size.size(),
pump_out_window_size: self.pump_out_window_size.get(),
out_buf_bytes: self.queued_write.queued_bytes_len(),
streams: self.streams.snapshot(),
}
}
pub fn our_settings_sent(&self) -> &HttpSettings {
if let Some(ref sent) = self.our_settings_sent {
&sent
} else {
&self.our_settings_ack
}
}
fn _decrease_out_window(&mut self, size: u32) -> result::Result<()> {
debug_assert!(size < 0x80000000);
self.out_window_size
.try_decrease(size as i32)
.map_err(|_| error::Error::WindowSizeOverflow)
}
pub fn decrease_in_window(&mut self, size: u32) -> result::Result<()> {
debug_assert!(size < 0x80000000);
let old_in_window_size = self.in_window_size.size();
self.in_window_size
.try_decrease_to_non_negative(size as i32)
.map_err(|_| error::Error::WindowSizeOverflow)?;
let new_in_window_size = self.in_window_size.size();
debug!(
"decrease conn window: {} -> {}",
old_in_window_size, new_in_window_size
);
Ok(())
}
pub fn process_dump_state(
&mut self,
sender: oneshot::Sender<ConnStateSnapshot>,
) -> result::Result<()> {
drop(sender.send(self.dump_state()));
Ok(())
}
pub fn send_rst_stream(
&mut self,
stream_id: StreamId,
error_code: ErrorCode,
) -> result::Result<()> {
self.streams.remove_stream(stream_id);
let rst_stream = RstStreamFrame::new(stream_id, error_code);
self.send_frame_and_notify(rst_stream);
Ok(())
}
pub fn send_flow_control_error(&mut self) -> result::Result<()> {
self.send_goaway(ErrorCode::FlowControlError)
}
fn stream_state_idle_or_closed(&self, stream_id: StreamId) -> StreamStateIdleOrClosed {
let last_stream_id = match T::init_where(stream_id) {
InitWhere::Locally => self.last_local_stream_id,
InitWhere::Peer => self.last_peer_stream_id,
};
if stream_id > last_stream_id {
StreamStateIdleOrClosed::Idle
} else {
StreamStateIdleOrClosed::Closed
}
}
fn stream_state(&self, stream_id: StreamId) -> StreamState {
match self.streams.get_stream_state(stream_id) {
Some(state) => state,
None => self.stream_state_idle_or_closed(stream_id).into(),
}
}
pub fn get_stream_maybe_send_error(
&mut self,
stream_id: StreamId,
frame_type: HttpFrameType,
) -> result::Result<Option<HttpStreamRef<T>>> {
let stream_state = self.stream_state(stream_id);
match stream_state {
StreamState::Idle => {
let send_connection_error = match frame_type {
HttpFrameType::Headers
| HttpFrameType::Priority
| HttpFrameType::PushPromise => false,
_ => true,
};
if send_connection_error {
debug!("stream is idle: {}, sending GOAWAY", stream_id);
self.send_goaway(ErrorCode::StreamClosed)?;
}
}
StreamState::Open | StreamState::HalfClosedLocal => {}
StreamState::ReservedLocal | StreamState::ReservedRemote => {}
StreamState::HalfClosedRemote => {
let send_rst = match frame_type {
HttpFrameType::WindowUpdate
| HttpFrameType::Priority
| HttpFrameType::RstStream => false,
_ => true,
};
if send_rst {
debug!(
"stream is half-closed remote: {}, sending RST_STREAM",
stream_id
);
self.send_rst_stream(stream_id, ErrorCode::StreamClosed)?;
}
}
StreamState::Closed => {
let send_stream_closed = match frame_type {
HttpFrameType::RstStream
| HttpFrameType::Priority
| HttpFrameType::WindowUpdate => false,
_ => true,
};
if send_stream_closed {
if self.peer_closed_streams.contains(stream_id) {
debug!("stream is closed by peer: {}, sending GOAWAY", stream_id);
self.send_goaway(ErrorCode::StreamClosed)?;
} else {
debug!("stream is closed by us: {}, sending RST_STREAM", stream_id);
self.send_rst_stream(stream_id, ErrorCode::StreamClosed)?;
}
}
}
}
Ok(self.streams.get_mut(stream_id))
}
pub fn get_stream_for_headers_maybe_send_error(
&mut self,
stream_id: StreamId,
) -> result::Result<Option<HttpStreamRef<T>>> {
self.get_stream_maybe_send_error(stream_id, HttpFrameType::Headers)
}
pub fn increase_in_window(&mut self, stream_id: StreamId, increase: u32) -> result::Result<()> {
if let Some(mut stream) = self.streams.get_mut(stream_id) {
if let Err(_) = stream.stream().in_window_size.try_increase(increase) {
return Err(error::Error::StreamInWindowOverflow(
stream_id,
stream.stream().in_window_size.size(),
increase,
));
}
} else {
return Ok(());
};
let window_update = WindowUpdateFrame::for_stream(stream_id, increase);
self.send_frame_and_notify(window_update);
Ok(())
}
fn poll_next_event(&mut self, cx: &mut Context<'_>) -> Poll<result::Result<LoopEvent<T>>> {
self.poll_flush(cx)?;
if self.queued_write.goaway_queued_and_flushed() {
info!("GOAWAY written and flushed, closing connection");
return Poll::Ready(Ok(LoopEvent::ExitLoop));
}
if self.goaway_received.is_some() && self.streams.is_empty() {
info!("GOAWAY received and streams is empty, closing connection");
return Poll::Ready(Ok(LoopEvent::ExitLoop));
}
match Pin::new(&mut self.write_rx).poll_next(cx) {
Poll::Pending => {}
Poll::Ready(Some(m)) => return Poll::Ready(Ok(LoopEvent::ToWriteMessage(m))),
Poll::Ready(None) => {
return Poll::Ready(Err(error::Error::ClientDied(None)));
}
};
match self.poll_recv_http_frame(cx)? {
Poll::Ready(m) => return Poll::Ready(Ok(LoopEvent::Frame(m))),
Poll::Pending => {}
}
Poll::Pending
}
async fn next_event(&mut self) -> result::Result<LoopEvent<T>> {
future::poll_fn(|cx| self.poll_next_event(cx)).await
}
async fn run_loop(mut self) -> result::Result<()> {
loop {
let event = self.next_event().await?;
match event {
LoopEvent::ToWriteMessage(m) => self.process_message(m)?,
LoopEvent::Frame(f) => self.process_http_frame_of_goaway(f)?,
LoopEvent::ExitLoop => return Ok(()),
}
}
}
pub fn run(self) -> impl Future<Output = result::Result<()>> + Send {
let ndc = Arc::new(format!("{} {}", T::CONN_NDC, self.peer_addr));
log_ndc_future(ndc, self.run_loop())
}
}