use std::collections::{HashMap, VecDeque};
use std::panic::AssertUnwindSafe;
use std::time::Duration;
use futures_util::{
SinkExt, StreamExt,
stream::{SplitSink, SplitStream},
};
use lark_websocket_protobuf::pbbp2::Frame;
use log::{debug, error, trace, warn};
use prost::Message as ProstMessage;
use tokio::net::TcpStream;
use tokio::sync::mpsc;
use tokio::time::{Instant, Interval};
use tokio_tungstenite::tungstenite::protocol::Message as WsMessage;
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, tungstenite::protocol::Message};
use super::client::{
ClientConfig, EventDispatcherHandler, InvalidStateKind, WsClientError, WsClientResult,
WsCloseReason,
};
use super::frame_handler::{
ControlFrameEffect, ControlFrameError, FRAME_METHOD_CONTROL, FRAME_METHOD_DATA, FrameHandler,
};
use super::package::{self, FramePackageBuffer};
const HANDLER_QUEUE_CAP: usize = 64;
const PENDING_OUTBOX_CAP: usize = HANDLER_QUEUE_CAP;
const WORKER_SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(5);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum SessionState {
Active,
Closing,
Closed,
}
impl SessionState {
fn as_str(self) -> &'static str {
match self {
Self::Active => "Active",
Self::Closing => "Closing",
Self::Closed => "Closed",
}
}
}
enum TrySendToWorker {
Full(Box<Frame>),
Closed,
}
#[derive(Debug, Clone)]
enum CloseIntent {
None,
WithoutReason,
WithReason(WsCloseReason),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct SessionOptions {
pub(crate) heartbeat_timeout: Duration,
}
impl Default for SessionOptions {
fn default() -> Self {
Self {
heartbeat_timeout: Duration::from_secs(120),
}
}
}
type HandlerOutcome = WsClientResult<Option<Frame>>;
pub(crate) struct Session {
service_id: i32,
sink: SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, WsMessage>,
stream: SplitStream<WebSocketStream<MaybeTlsStream<TcpStream>>>,
event_handler: EventDispatcherHandler,
package_buffers: HashMap<String, FramePackageBuffer>,
ping_frame_interval: Interval,
heartbeat_timeout: Duration,
state: SessionState,
close_intent: CloseIntent,
inflight_handlers: usize,
pending_outbox: VecDeque<Frame>,
}
impl Session {
pub(crate) fn new(
service_id: i32,
client_config: ClientConfig,
conn: WebSocketStream<MaybeTlsStream<TcpStream>>,
event_handler: EventDispatcherHandler,
options: SessionOptions,
) -> Self {
let (sink, stream) = conn.split();
let ping_secs = (client_config.ping_interval.max(1)) as u64;
Self {
service_id,
sink,
stream,
event_handler,
package_buffers: HashMap::new(),
ping_frame_interval: tokio::time::interval(Duration::from_secs(ping_secs)),
heartbeat_timeout: options.heartbeat_timeout.max(Duration::from_millis(1)),
state: SessionState::Active,
close_intent: CloseIntent::None,
inflight_handlers: 0,
pending_outbox: VecDeque::new(),
}
}
fn ensure_active(&self) -> WsClientResult<()> {
if self.state != SessionState::Active {
return Err(WsClientError::InvalidStateTransition {
kind: InvalidStateKind::ExpectedActive {
actual: self.state.as_str(),
},
});
}
Ok(())
}
fn begin_close(&mut self, reason: Option<WsCloseReason>) -> WsClientResult<()> {
match self.state {
SessionState::Closed => {
return Err(WsClientError::InvalidStateTransition {
kind: InvalidStateKind::AlreadyClosed,
});
}
SessionState::Closing => return Ok(()),
SessionState::Active => {}
}
self.state = SessionState::Closing;
self.close_intent = match reason {
Some(r) => CloseIntent::WithReason(r),
None => CloseIntent::WithoutReason,
};
Ok(())
}
fn drain_close_reason(&mut self) -> Option<WsCloseReason> {
match std::mem::replace(&mut self.close_intent, CloseIntent::None) {
CloseIntent::WithReason(r) => Some(r),
CloseIntent::WithoutReason | CloseIntent::None => None,
}
}
fn take_connection_closed_if_idle(&mut self) -> Option<WsClientResult<()>> {
if self.state == SessionState::Closing
&& self.inflight_handlers == 0
&& self.pending_outbox.is_empty()
{
self.state = SessionState::Closed;
let reason = self.drain_close_reason();
return Some(Err(WsClientError::ConnectionClosed { reason }));
}
None
}
fn note_job_enqueued(&mut self) {
self.inflight_handlers = self.inflight_handlers.saturating_add(1);
}
fn try_send_to_worker(
&mut self,
job_tx: &mpsc::Sender<Frame>,
frame: Frame,
) -> Result<(), TrySendToWorker> {
match job_tx.try_send(frame) {
Ok(()) => {
self.note_job_enqueued();
Ok(())
}
Err(mpsc::error::TrySendError::Full(frame)) => {
Err(TrySendToWorker::Full(Box::new(frame)))
}
Err(mpsc::error::TrySendError::Closed(_)) => Err(TrySendToWorker::Closed),
}
}
fn flush_outbox_try(&mut self, job_tx: &mpsc::Sender<Frame>) -> WsClientResult<()> {
while let Some(frame) = self.pending_outbox.pop_front() {
match self.try_send_to_worker(job_tx, frame) {
Ok(()) => {}
Err(TrySendToWorker::Full(frame)) => {
self.pending_outbox.push_front(*frame);
break;
}
Err(TrySendToWorker::Closed) => {
return Err(WsClientError::InvalidStateTransition {
kind: InvalidStateKind::WorkerGone,
});
}
}
}
Ok(())
}
pub(crate) async fn run(mut self) -> WsClientResult<()> {
let mut last_activity = Instant::now();
let heartbeat_check_period = self
.heartbeat_timeout
.min(Duration::from_secs(1))
.max(Duration::from_millis(50));
let mut heartbeat_check_interval = tokio::time::interval(heartbeat_check_period);
let (job_tx, mut job_rx) = mpsc::channel::<Frame>(HANDLER_QUEUE_CAP);
let (outcome_tx, mut outcome_rx) = mpsc::channel::<HandlerOutcome>(HANDLER_QUEUE_CAP);
let worker_handler = self.event_handler.clone();
let worker = tokio::spawn(async move {
while let Some(frame) = job_rx.recv().await {
let handler = worker_handler.clone();
let outcome = tokio::task::spawn_blocking(move || {
match std::panic::catch_unwind(AssertUnwindSafe(|| {
FrameHandler::handle_data_frame(frame, &handler)
})) {
Ok(opt) => Ok(opt),
Err(_) => Err(WsClientError::HandlerPanicked),
}
})
.await
.unwrap_or_else(|_| Err(WsClientError::HandlerPanicked));
if outcome_tx.send(outcome).await.is_err() {
break;
}
}
});
let result = async {
let mut stream_open = true;
loop {
if let Some(done) = self.take_connection_closed_if_idle() {
return done;
}
self.flush_outbox_try(&job_tx)?;
let need_reserve =
!self.pending_outbox.is_empty() && self.state != SessionState::Closed;
tokio::select! {
permit = job_tx.reserve(), if need_reserve => {
let permit = permit.map_err(|_| {
WsClientError::InvalidStateTransition {
kind: InvalidStateKind::WorkerGone,
}
})?;
if let Some(frame) = self.pending_outbox.pop_front() {
permit.send(frame);
self.note_job_enqueued();
}
}
item = self.stream.next(), if stream_open && self.state != SessionState::Closed => {
match item.transpose() {
Ok(Some(msg)) => {
if msg.is_ping() {
last_activity = Instant::now();
}
if self.state == SessionState::Closing {
if matches!(msg, Message::Binary(_) | Message::Text(_)) {
return Err(WsClientError::InvalidStateTransition {
kind: InvalidStateKind::DataWhileClosing,
});
}
continue;
}
self.handle_message(msg, &job_tx).await?;
}
Ok(None) => {
stream_open = false;
self.begin_close(None)?;
}
Err(e) => {
if let Some(reason) = self.drain_close_reason() {
self.state = SessionState::Closed;
warn!(
"stream error after remote close: {e}; returning close reason"
);
return Err(WsClientError::ConnectionClosed {
reason: Some(reason),
});
}
return Err(e.into());
}
}
}
Some(outcome) = outcome_rx.recv() => {
self.inflight_handlers = self.inflight_handlers.saturating_sub(1);
match outcome {
Ok(Some(response_frame)) => {
if self.state == SessionState::Active {
self.send_frame(response_frame).await?;
}
}
Ok(None) => {}
Err(err) => {
if matches!(self.close_intent, CloseIntent::WithReason(_)) {
self.state = SessionState::Closed;
let reason = self.drain_close_reason();
warn!(
"handler failed during close: {err}; returning close reason"
);
return Err(WsClientError::ConnectionClosed { reason });
}
self.state = SessionState::Closed;
self.close_intent = CloseIntent::None;
return Err(err);
}
}
self.flush_outbox_try(&job_tx)?;
if let Some(done) = self.take_connection_closed_if_idle() {
return done;
}
}
_ = self.ping_frame_interval.tick(), if self.state == SessionState::Active => {
self.send_app_ping().await?;
}
_ = heartbeat_check_interval.tick(), if self.state == SessionState::Active => {
if last_activity.elapsed() > self.heartbeat_timeout {
self.begin_close(None)?;
}
}
}
}
}
.await;
drop(job_tx);
let mut worker = worker;
tokio::select! {
_ = &mut worker => {}
_ = tokio::time::sleep(WORKER_SHUTDOWN_TIMEOUT) => {
warn!(
"handler worker did not finish within {:?}; aborting worker task \
(in-flight spawn_blocking handler may still run until it returns)",
WORKER_SHUTDOWN_TIMEOUT
);
worker.abort();
let _ = worker.await;
}
}
result
}
async fn send_app_ping(&mut self) -> WsClientResult<()> {
self.ensure_active()?;
let frame = FrameHandler::build_ping_frame(self.service_id);
let msg = Message::Binary(frame.encode_to_vec().into());
if let Err(e) = self.sink.send(msg).await {
error!("Failed to send ping message: {e:?}");
return Err(e.into());
}
Ok(())
}
async fn send_frame(&mut self, frame: Frame) -> WsClientResult<()> {
let msg = Message::Binary(frame.encode_to_vec().into());
self.sink.send(msg).await?;
Ok(())
}
async fn handle_message(
&mut self,
msg: WsMessage,
job_tx: &mpsc::Sender<Frame>,
) -> WsClientResult<()> {
self.ensure_active()?;
match msg {
Message::Ping(data) => {
self.sink.send(Message::Pong(data)).await?;
}
Message::Binary(data) => {
let frame = Frame::decode(&*data)?;
trace!("Received frame: {frame:?}");
match frame.method {
FRAME_METHOD_CONTROL => self.apply_control_frame(frame)?,
FRAME_METHOD_DATA => self.enqueue_data_frame(frame, job_tx)?,
method => {
return Err(WsClientError::InvalidFrameMethod { method });
}
}
}
Message::Close(close_frame) => {
let reason = close_frame.map(|frame| WsCloseReason {
code: frame.code,
message: frame.reason.to_string(),
});
self.begin_close(reason)?;
}
_ => return Err(WsClientError::UnexpectedResponse),
}
Ok(())
}
fn apply_control_frame(&mut self, frame: Frame) -> WsClientResult<()> {
match FrameHandler::interpret_control_frame(&frame) {
Ok(ControlFrameEffect::UpdatePingInterval(secs)) => {
self.apply_ping_interval(secs);
Ok(())
}
Ok(ControlFrameEffect::Ignored) => Ok(()),
Err(ControlFrameError::MalformedPong(message)) => {
Err(WsClientError::MalformedControlFrame { message })
}
}
}
fn apply_ping_interval(&mut self, ping_interval: i32) {
let ping_secs = (ping_interval.max(1)) as u64;
self.ping_frame_interval = tokio::time::interval(Duration::from_secs(ping_secs));
self.ping_frame_interval
.reset_after(Duration::from_secs(ping_secs));
debug!("Updated ping interval from pong response: {ping_secs}s");
}
fn enqueue_data_frame(
&mut self,
frame: Frame,
job_tx: &mpsc::Sender<Frame>,
) -> WsClientResult<()> {
self.ensure_active()?;
let Some(frame) = package::assemble_frame(&mut self.package_buffers, frame) else {
return Ok(());
};
match self.try_send_to_worker(job_tx, frame) {
Ok(()) => Ok(()),
Err(TrySendToWorker::Full(frame)) => {
if self.pending_outbox.len() >= PENDING_OUTBOX_CAP {
return Err(WsClientError::BacklogFull {
message: format!("queue {HANDLER_QUEUE_CAP} + outbox {PENDING_OUTBOX_CAP}"),
});
}
self.pending_outbox.push_back(*frame);
Ok(())
}
Err(TrySendToWorker::Closed) => Err(WsClientError::InvalidStateTransition {
kind: InvalidStateKind::WorkerGone,
}),
}
}
}