use alloc::format;
use alloc::string::ToString;
use alloc::sync::Arc;
use alloc::vec::Vec;
use core::time::Duration;
use std::net::TcpStream;
use std::sync::mpsc::{self, Receiver, RecvTimeoutError, Sender};
use std::thread::JoinHandle;
use liminal::protocol::{Frame, SchemaId, decode};
use liminal_protocol::outcome::ReconnectState;
use spin::Mutex;
use crate::SdkError;
use super::binding::{AttemptFateOutcome, OpenRequestDecision, WebSocketAuthorityBinding};
use super::connection_error;
use super::core::{
DriverOutput, FrameCorrelation, ResponseExpectation, SocketCommand, SocketEvent,
WebSocketFrameDriver,
};
use super::liminal_ws_message_bound;
use super::std_socket::{SocketRead, WsSocket};
const CLIENT_MIN_VERSION: liminal::protocol::ProtocolVersion =
liminal::protocol::ProtocolVersion::new(1, 0);
const CLIENT_MAX_VERSION: liminal::protocol::ProtocolVersion =
liminal::protocol::ProtocolVersion::new(1, 0);
const SUBSCRIPTION_STREAM_ID: u32 = 1;
const SUBSCRIBE_MAX_IN_FLIGHT: u32 = 1024;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WebSocketDeliveredMessage {
delivery_seq: u64,
schema_id: SchemaId,
payload: Vec<u8>,
}
impl WebSocketDeliveredMessage {
#[must_use]
pub const fn delivery_seq(&self) -> u64 {
self.delivery_seq
}
#[must_use]
pub const fn schema_id(&self) -> SchemaId {
self.schema_id
}
#[must_use]
pub fn payload(&self) -> &[u8] {
&self.payload
}
#[must_use]
pub fn into_payload(self) -> Vec<u8> {
self.payload
}
}
#[derive(Debug)]
pub struct WebSocketSubscriptionStream {
shutdown: TcpStream,
subscription_id: u64,
inbound: Receiver<WebSocketDeliveredMessage>,
binding: Arc<Mutex<WebSocketAuthorityBinding>>,
reader: Option<JoinHandle<()>>,
}
impl WebSocketSubscriptionStream {
pub fn open(
address: &str,
channel: &str,
accepted_schemas: Vec<SchemaId>,
) -> Result<Self, SdkError> {
let message_bound = liminal_ws_message_bound()?;
let mut binding = WebSocketAuthorityBinding::new();
match binding.request_open() {
OpenRequestDecision::Authorized { .. } => {}
OpenRequestDecision::Refused(refusal) => {
return Err(connection_error(&format!(
"client authority refused the subscription open: {refusal:?}"
)));
}
}
match Self::open_link(address, channel, accepted_schemas, message_bound) {
Ok((socket, driver, subscription_id, pending)) => {
match binding.connection_established() {
AttemptFateOutcome::Recorded { .. } => {}
AttemptFateOutcome::Refused(refusal) => {
return Err(SdkError::Protocol {
description: format!(
"client authority refused the Connected fate for the \
subscription open: {refusal:?}"
),
});
}
}
Self::start(socket, driver, binding, subscription_id, pending)
}
Err(error) => match binding.open_failed() {
AttemptFateOutcome::Recorded { .. } => Err(error),
AttemptFateOutcome::Refused(refusal) => Err(SdkError::Protocol {
description: format!(
"subscription open failed ({error}) and the client authority \
refused the Failed fate: {refusal:?}"
),
}),
},
}
}
pub fn recv_timeout(&self, timeout: Duration) -> Result<WebSocketDeliveredMessage, SdkError> {
self.inbound.recv_timeout(timeout).map_err(|error| {
let detail = match error {
RecvTimeoutError::Timeout => "no delivery arrived within the timeout",
RecvTimeoutError::Disconnected => {
"the subscription reader stopped before a delivery arrived"
}
};
connection_error(&format!("websocket subscription receive failed: {detail}"))
})
}
#[must_use]
pub const fn subscription_id(&self) -> u64 {
self.subscription_id
}
#[must_use]
pub fn reconnect_state(&self) -> ReconnectState {
self.binding.lock().reconnect_state()
}
fn open_link(
address: &str,
channel: &str,
accepted_schemas: Vec<SchemaId>,
message_bound: usize,
) -> Result<
(
WsSocket,
WebSocketFrameDriver,
u64,
Vec<WebSocketDeliveredMessage>,
),
SdkError,
> {
let mut driver = WebSocketFrameDriver::new();
let command = driver
.command_open()
.map_err(|refusal| SdkError::Protocol {
description: format!("subscription driver refused its first open: {refusal:?}"),
})?;
if command != SocketCommand::Open {
return Err(SdkError::Protocol {
description: "subscription driver emitted a non-open first command".to_string(),
});
}
let mut socket = WsSocket::connect(address, message_bound)?;
let step = driver.handle_event(SocketEvent::Opened);
if step.output != DriverOutput::Opened {
return Err(SdkError::Protocol {
description: format!("subscription driver refused the opened socket: {step:?}"),
});
}
let mut pending = Vec::new();
let connect = Frame::Connect {
flags: 0,
min_version: CLIENT_MIN_VERSION,
max_version: CLIENT_MAX_VERSION,
auth_token: Vec::new(),
};
match setup_exchange(&mut socket, &mut driver, &connect, &mut pending)? {
Frame::ConnectAck { .. } => {}
Frame::ConnectError {
reason_code,
message,
..
} => {
return Err(connection_error(&format!(
"server rejected subscription connection (reason {reason_code}): {}",
message.unwrap_or_else(|| "no detail".to_string())
)));
}
other => {
return Err(unexpected_setup_frame("ConnectAck", &other));
}
}
let subscribe = Frame::Subscribe {
flags: 0,
stream_id: SUBSCRIPTION_STREAM_ID,
channel: channel.to_string(),
accepted_schemas,
max_in_flight: SUBSCRIBE_MAX_IN_FLIGHT,
};
let subscription_id =
match setup_exchange(&mut socket, &mut driver, &subscribe, &mut pending)? {
Frame::SubscribeAck {
subscription_id, ..
} => subscription_id,
Frame::SubscribeError {
reason_code,
message,
..
} => {
return Err(SdkError::Protocol {
description: format!(
"server rejected subscribe (reason {reason_code}): {}",
message.unwrap_or_else(|| "no detail".to_string())
),
});
}
other => {
return Err(unexpected_setup_frame("SubscribeAck", &other));
}
};
Ok((socket, driver, subscription_id, pending))
}
fn start(
socket: WsSocket,
driver: WebSocketFrameDriver,
binding: WebSocketAuthorityBinding,
subscription_id: u64,
pending: Vec<WebSocketDeliveredMessage>,
) -> Result<Self, SdkError> {
socket.set_read_timeout(None)?;
let shutdown = socket.try_clone_stream()?;
let binding = Arc::new(Mutex::new(binding));
let reader_binding = Arc::clone(&binding);
let (sender, inbound) = mpsc::channel();
let reader = std::thread::Builder::new()
.name("liminal-ws-subscription-reader".to_string())
.spawn(move || run_reader(socket, driver, &reader_binding, pending, &sender))
.map_err(|source| SdkError::Protocol {
description: format!(
"failed to start websocket subscription reader thread: {source}"
),
})?;
Ok(Self {
shutdown,
subscription_id,
inbound,
binding,
reader: Some(reader),
})
}
}
impl Drop for WebSocketSubscriptionStream {
fn drop(&mut self) {
self.shutdown.shutdown(std::net::Shutdown::Both).ok();
if let Some(reader) = self.reader.take() {
reader.join().ok();
}
}
}
fn setup_exchange(
socket: &mut WsSocket,
driver: &mut WebSocketFrameDriver,
request: &Frame,
pending: &mut Vec<WebSocketDeliveredMessage>,
) -> Result<Frame, SdkError> {
let bytes = super::encode_frame(request)?;
let command = driver
.command_send(bytes, ResponseExpectation::Correlated)
.map_err(|refusal| SdkError::Protocol {
description: format!("subscription driver refused the setup send: {refusal:?}"),
})?;
let SocketCommand::SendBinary(payload) = command else {
return Err(SdkError::Protocol {
description: "subscription driver emitted a non-send command for a send".to_string(),
});
};
if let Err(failure) = socket.send_binary(payload) {
let step = driver.handle_event(SocketEvent::Failed(failure));
if step.command == Some(SocketCommand::Close) {
socket.execute_close();
}
return Err(connection_error(&format!(
"failed to send subscription setup frame: {}",
socket
.last_failure_detail()
.unwrap_or("websocket send failed")
)));
}
loop {
let event = match socket.read_event() {
SocketRead::TimedOut => {
return Err(connection_error(
"subscription connection timed out waiting for a control-frame reply",
));
}
SocketRead::Event(event) => event,
};
let step = driver.handle_event(event);
if step.command == Some(SocketCommand::Close) {
socket.execute_close();
}
match step.output {
DriverOutput::Frame { bytes, correlation } => {
let frame = decode_message(&bytes)?;
match correlation {
FrameCorrelation::UnsolicitedDelivery => {
if let Some(message) = delivered_message(frame) {
pending.push(message);
}
}
FrameCorrelation::CorrelatedResponse | FrameCorrelation::UnsolicitedFrame => {
return Ok(frame);
}
}
}
DriverOutput::Terminal(terminal) => {
return Err(connection_error(&format!(
"subscription connection terminated during setup: {terminal:?}"
)));
}
DriverOutput::Opened
| DriverOutput::PostTerminalIgnored(_)
| DriverOutput::Refused(_) => {
return Err(SdkError::Protocol {
description: format!(
"subscription driver produced an unexpected setup output: {:?}",
step.output
),
});
}
}
}
}
fn run_reader(
mut socket: WsSocket,
mut driver: WebSocketFrameDriver,
binding: &Mutex<WebSocketAuthorityBinding>,
pending: Vec<WebSocketDeliveredMessage>,
sender: &Sender<WebSocketDeliveredMessage>,
) {
for message in pending {
if sender.send(message).is_err() {
close_link(&mut socket, &mut driver);
return;
}
}
loop {
let event = match socket.read_event() {
SocketRead::TimedOut => continue,
SocketRead::Event(event) => event,
};
let step = driver.handle_event(event);
if step.command == Some(SocketCommand::Close) {
socket.execute_close();
}
match step.output {
DriverOutput::Frame { bytes, correlation } => match correlation {
FrameCorrelation::UnsolicitedDelivery => {
let Ok(frame) = decode_message(&bytes) else {
close_link(&mut socket, &mut driver);
continue;
};
if let Some(message) = delivered_message(frame) {
if sender.send(message).is_err() {
close_link(&mut socket, &mut driver);
return;
}
}
}
FrameCorrelation::CorrelatedResponse | FrameCorrelation::UnsolicitedFrame => {
match decode_message(&bytes) {
Ok(Frame::Disconnect { .. }) => {
close_link(&mut socket, &mut driver);
}
Ok(_) => {}
Err(_) => {
close_link(&mut socket, &mut driver);
}
}
}
},
DriverOutput::Terminal(terminal) => {
let _outcome = binding.lock().established_terminal(&terminal);
return;
}
DriverOutput::PostTerminalIgnored(_) => return,
DriverOutput::Opened | DriverOutput::Refused(_) => {}
}
}
}
fn close_link(socket: &mut WsSocket, driver: &mut WebSocketFrameDriver) {
if driver.command_close().is_ok() {
socket.execute_close();
}
}
fn decode_message(bytes: &[u8]) -> Result<Frame, SdkError> {
match decode(bytes) {
Ok((frame, consumed)) if consumed == bytes.len() => Ok(frame),
Ok((_, consumed)) => Err(SdkError::Protocol {
description: format!(
"subscription decode consumed {consumed} of {} message bytes",
bytes.len()
),
}),
Err(error) => Err(SdkError::Protocol {
description: format!("subscription wire codec error: {error}"),
}),
}
}
fn delivered_message(frame: Frame) -> Option<WebSocketDeliveredMessage> {
match frame {
Frame::Deliver {
delivery_seq,
envelope,
..
} => Some(WebSocketDeliveredMessage {
delivery_seq,
schema_id: envelope.schema_id,
payload: envelope.payload,
}),
_ => None,
}
}
fn unexpected_setup_frame(expected: &str, actual: &Frame) -> SdkError {
SdkError::Protocol {
description: format!(
"expected {expected} during subscription setup, received {:?}",
actual.frame_type()
),
}
}