use std::collections::{HashMap, VecDeque};
use std::io::{self, Read, Write};
use std::time::{Duration, Instant};
use crate::ingress::RawQwpWsRoundStream;
use crate::ingress::sender::qwp_ws::WsStream;
use crate::ws::frame::{self, FrameError, FrameHeader, Opcode, encode_client_frame};
use crate::ws::mask::{MaskKeySource, apply_mask};
use crate::{Result, error};
use crate::ingress::AckLevel;
use super::qwp_frame_size_error;
pub(crate) const WS_HEADER_RESERVE: usize = 14;
const QWP_STATUS_OK: u8 = 0x00;
const QWP_STATUS_DURABLE_ACK: u8 = 0x02;
const QWP_STATUS_SCHEMA_MISMATCH: u8 = 0x03;
const QWP_STATUS_PARSE_ERROR: u8 = 0x05;
const QWP_STATUS_INTERNAL_ERROR: u8 = 0x06;
const QWP_STATUS_SECURITY_ERROR: u8 = 0x08;
const QWP_STATUS_WRITE_ERROR: u8 = 0x09;
const MAX_INBOUND_FRAME_BYTES: u64 = 256 * 1024 * 1024;
const MAX_IN_FLIGHT: u32 = 128;
const CLOSE_TIMEOUT: Duration = Duration::from_millis(200);
const WS_CLOSE_STATUS_NORMAL: [u8; 2] = 1000u16.to_be_bytes();
struct PendingAck {
fsn: u64,
}
pub(crate) struct ColumnConn {
stream: WsStream,
leftover: Vec<u8>,
write_buf: Vec<u8>,
read_buf: Vec<u8>,
mask_keys: MaskKeySource,
next_fsn: u64,
pending_acks: VecDeque<PendingAck>,
in_flight: u32,
durable_watermarks: HashMap<String, i64>,
pending_durable_targets: HashMap<String, i64>,
must_close: bool,
transport_dead: bool,
spent: bool,
endpoint_idx: usize,
max_buf_size: usize,
request_timeout: Duration,
durable_ack_opt_in: bool,
}
impl ColumnConn {
pub(crate) fn from_round_stream(raw: RawQwpWsRoundStream) -> Result<Self> {
let mask_keys = MaskKeySource::new()
.map_err(|e| error::fmt!(SocketError, "MaskKeySource init failed: {}", e.0))?;
Ok(Self {
stream: raw.stream,
leftover: raw.leftover,
write_buf: Vec::with_capacity(64 * 1024),
read_buf: Vec::with_capacity(4 * 1024),
mask_keys,
next_fsn: 0,
pending_acks: VecDeque::new(),
in_flight: 0,
durable_watermarks: HashMap::new(),
pending_durable_targets: HashMap::new(),
must_close: false,
transport_dead: false,
spent: false,
endpoint_idx: raw.endpoint_idx,
max_buf_size: raw.max_buf_size,
request_timeout: raw.request_timeout,
durable_ack_opt_in: raw.durable_ack_opt_in,
})
}
#[cfg(test)]
pub(crate) fn for_test(stream: WsStream, durable_ack_opt_in: bool) -> Self {
Self {
stream,
leftover: Vec::new(),
write_buf: Vec::new(),
read_buf: Vec::new(),
mask_keys: MaskKeySource::new().expect("test mask key source"),
next_fsn: 0,
pending_acks: VecDeque::new(),
in_flight: 0,
durable_watermarks: HashMap::new(),
pending_durable_targets: HashMap::new(),
must_close: false,
transport_dead: false,
spent: false,
endpoint_idx: 0,
max_buf_size: 1 << 20,
request_timeout: Duration::from_secs(30),
durable_ack_opt_in,
}
}
#[cfg(test)]
pub(crate) fn pending_durable_target_count(&self) -> usize {
self.pending_durable_targets.len()
}
pub(crate) fn must_close(&self) -> bool {
self.must_close || self.spent
}
pub(crate) fn transport_dead(&self) -> bool {
self.transport_dead
}
pub(crate) fn can_drain_in_flight(&self) -> bool {
!self.must_close && !self.transport_dead
}
pub(crate) fn mark_spent(&mut self) {
self.spent = true;
}
pub(crate) fn endpoint_idx(&self) -> usize {
self.endpoint_idx
}
pub(crate) fn mark_must_close(&mut self) {
self.must_close = true;
}
pub(crate) fn publish_qwp<F>(
&mut self,
encode: F,
) -> std::result::Result<PublishedFrame, PublishError>
where
F: FnOnce(&mut Vec<u8>) -> Result<()>,
{
if self.must_close {
return Err(PublishError::BeforeWrite(error::fmt!(
SocketError,
"QWP/WebSocket connection latched as terminal; \
return the sender to the pool and acquire a fresh one."
)));
}
self.write_buf.clear();
self.write_buf.resize(WS_HEADER_RESERVE, 0);
if let Err(e) = encode(&mut self.write_buf) {
self.write_buf.clear();
return Err(PublishError::BeforeWrite(e));
}
let payload_len = self.write_buf.len() - WS_HEADER_RESERVE;
if payload_len > self.max_buf_size {
return Err(PublishError::BeforeWrite(qwp_frame_size_error(
payload_len,
self.max_buf_size,
)));
}
let mask_key = match self.mask_keys.next_key() {
Ok(k) => k,
Err(e) => {
return Err(PublishError::BeforeWrite(self.latch(error::fmt!(
SocketError,
"mask key entropy failed: {}",
e.0
))));
}
};
apply_mask(&mut self.write_buf[WS_HEADER_RESERVE..], mask_key, 0);
let ws_header_len = ws_header_len_for(payload_len);
let header_offset = WS_HEADER_RESERVE - ws_header_len;
write_ws_header(
&mut self.write_buf[header_offset..WS_HEADER_RESERVE],
payload_len,
mask_key,
);
if let Err(e) = self.set_timeouts(Some(self.request_timeout), Some(self.request_timeout)) {
return Err(PublishError::BeforeWrite(e));
}
if let Err(e) = self.stream.write_all(&self.write_buf[header_offset..]) {
return Err(PublishError::DuringWrite(self.latch(error::fmt!(
SocketError,
"QWP/WebSocket socket write failed: {}",
e
))));
}
if let Err(e) = self.stream.flush() {
return Err(PublishError::DuringWrite(self.latch(error::fmt!(
SocketError,
"QWP/WebSocket socket flush failed: {}",
e
))));
}
let fsn = self.next_fsn;
self.next_fsn = self.next_fsn.wrapping_add(1);
Ok(PublishedFrame { fsn })
}
pub(crate) fn push_pending(&mut self, fsn: u64) {
self.pending_acks.push_back(PendingAck { fsn });
self.in_flight += 1;
}
pub(crate) fn in_flight(&self) -> u32 {
self.in_flight
}
pub(crate) fn has_sync_commit_slot(&self) -> bool {
self.in_flight < MAX_IN_FLIGHT - 1
}
pub(crate) fn validate_ack_level(&self, ack_level: AckLevel) -> Result<()> {
if ack_level == AckLevel::Durable && !self.durable_ack_opt_in {
return Err(error::fmt!(
InvalidApiCall,
"AckLevel::Durable requires the pool to be opened with \
`request_durable_ack=on` in the connect string."
));
}
Ok(())
}
pub(crate) fn try_drain_acks(&mut self) -> Result<()> {
while let Some(response) = self.try_recv_qwp_response()? {
self.process_response(response)?;
}
Ok(())
}
pub(crate) fn drain_one_ack_blocking(&mut self) -> Result<()> {
self.set_timeouts(Some(self.request_timeout), Some(self.request_timeout))?;
let target = self.in_flight;
let deadline_anchor = Instant::now();
loop {
let response = self.recv_qwp_response()?;
self.process_response(response)?;
if self.in_flight < target {
return Ok(());
}
if !self.request_timeout.is_zero() && deadline_anchor.elapsed() >= self.request_timeout
{
return Err(self.latch(error::fmt!(
SocketError,
"QWP/WebSocket connection received no slot-freeing ack within {:?}",
self.request_timeout
)));
}
}
}
pub(crate) fn sync_all_acks(&mut self, ack_level: AckLevel) -> Result<()> {
if self.must_close {
return Err(error::fmt!(
SocketError,
"QWP/WebSocket connection latched as terminal."
));
}
self.validate_ack_level(ack_level)?;
self.set_timeouts(Some(self.request_timeout), Some(self.request_timeout))?;
let bounded = !self.request_timeout.is_zero();
let mut deadline_anchor = Instant::now();
let mut last_in_flight = self.in_flight;
while self.in_flight > 0 {
let response = self.recv_qwp_response()?;
self.process_response(response)?;
if self.in_flight < last_in_flight {
last_in_flight = self.in_flight;
deadline_anchor = Instant::now();
} else if bounded && deadline_anchor.elapsed() >= self.request_timeout {
return Err(self.latch(error::fmt!(
SocketError,
"QWP/WebSocket sync stalled: {} frame(s) unacked with no progress for {:?}",
self.in_flight,
self.request_timeout
)));
}
}
if ack_level == AckLevel::Durable {
let mut deadline_anchor = Instant::now();
let mut last_mark = self.durable_progress_mark();
while !self.durability_satisfied() {
let response = self.recv_qwp_response()?;
self.process_response(response)?;
let mark = self.durable_progress_mark();
if mark != last_mark {
last_mark = mark;
deadline_anchor = Instant::now();
} else if bounded && deadline_anchor.elapsed() >= self.request_timeout {
return Err(self.latch(error::fmt!(
SocketError,
"QWP/WebSocket durable sync stalled: no watermark progress for {:?}",
self.request_timeout
)));
}
}
}
self.drop_satisfied_durable_targets();
Ok(())
}
fn durability_satisfied(&self) -> bool {
self.pending_durable_targets.iter().all(|(t, target)| {
self.durable_watermarks.get(t).copied().unwrap_or(i64::MIN) >= *target
})
}
fn durable_progress_mark(&self) -> (usize, i128) {
let sum: i128 = self.durable_watermarks.values().map(|&v| v as i128).sum();
(self.durable_watermarks.len(), sum)
}
fn drop_satisfied_durable_targets(&mut self) {
let watermarks = &self.durable_watermarks;
self.pending_durable_targets
.retain(|t, target| watermarks.get(t).copied().unwrap_or(i64::MIN) < *target);
}
fn process_response(&mut self, response: QwpResponse) -> Result<()> {
match response {
QwpResponse::Ok { sequence, tables } => {
let mut popped = 0u32;
while let Some(front) = self.pending_acks.front() {
if front.fsn > sequence {
break;
}
self.pending_acks.pop_front();
popped += 1;
}
if popped == 0 {
return Ok(());
}
if self.durable_ack_opt_in {
for (t, seq_txn) in tables {
self.pending_durable_targets
.entry(t)
.and_modify(|w| {
if seq_txn > *w {
*w = seq_txn;
}
})
.or_insert(seq_txn);
}
}
self.in_flight = self.in_flight.checked_sub(popped).ok_or_else(|| {
self.must_close = true;
error::fmt!(
SocketError,
"QWP in-flight accounting underflow: {} acked, {} tracked",
popped,
self.in_flight
)
})?;
Ok(())
}
QwpResponse::DurableAck { tables } => {
for (t, seq_txn) in tables {
self.durable_watermarks
.entry(t)
.and_modify(|w| {
if seq_txn > *w {
*w = seq_txn;
}
})
.or_insert(seq_txn);
}
Ok(())
}
QwpResponse::Error {
sequence,
status,
message,
} => {
let err = map_error_status(status, &message);
Err(self.latch(crate::Error::new(
err.code(),
format!(
"QWP server error on fsn {}: status=0x{:02x}, message={:?}",
sequence, status, message
),
)))
}
}
}
pub(crate) fn at_in_flight_cap(&self) -> bool {
self.in_flight >= MAX_IN_FLIGHT
}
fn latch(&mut self, err: crate::Error) -> crate::Error {
self.must_close = true;
if err.code() == crate::ErrorCode::SocketError {
self.transport_dead = true;
}
err
}
fn set_timeouts(&mut self, read: Option<Duration>, write: Option<Duration>) -> Result<()> {
self.stream.set_timeouts(read, write).map_err(|e| {
self.latch(error::fmt!(
SocketError,
"QWP/WebSocket socket set_timeouts failed: {}",
e
))
})
}
fn try_recv_qwp_response(&mut self) -> Result<Option<QwpResponse>> {
loop {
match FrameHeader::parse(&self.leftover) {
Ok(h) => {
if !h.fin {
return Err(self.latch(error::fmt!(
SocketError,
"QWP/WebSocket server sent a fragmented frame; QWP is FIN-only"
)));
}
if h.payload_len > MAX_INBOUND_FRAME_BYTES {
return Err(self.latch(error::fmt!(
SocketError,
"WS frame declared {} payload bytes (max {})",
h.payload_len,
MAX_INBOUND_FRAME_BYTES
)));
}
let payload_len = h.payload_len as usize;
let header_len = h.header_len;
if self.leftover.len() < header_len + payload_len {
if !self.try_fill_leftover()? {
return Ok(None);
}
continue;
}
self.leftover.drain(..header_len);
self.read_buf.clear();
if self.read_buf.try_reserve(payload_len).is_err() {
return Err(self.latch(error::fmt!(
SocketError,
"could not allocate {} bytes for inbound QWP frame",
payload_len
)));
}
self.read_buf
.extend_from_slice(&self.leftover[..payload_len]);
self.leftover.drain(..payload_len);
match h.opcode {
Opcode::Binary => {
return parse_qwp_response(&self.read_buf)
.inspect_err(|_| {
self.must_close = true;
})
.map(Some);
}
Opcode::Ping => {
self.send_pong(payload_len)?;
continue;
}
Opcode::Pong => continue,
Opcode::Close => {
self.must_close = true;
self.transport_dead = true;
return Err(error::fmt!(
SocketError,
"QWP/WebSocket server closed the connection"
));
}
}
}
Err(FrameError::Incomplete) => {
if !self.try_fill_leftover()? {
return Ok(None);
}
}
Err(FrameError::Protocol(msg)) => {
return Err(self.latch(error::fmt!(
SocketError,
"QWP/WebSocket frame parse error: {}",
msg
)));
}
}
}
}
fn recv_qwp_response(&mut self) -> Result<QwpResponse> {
loop {
let header = self.read_ws_frame_header()?;
if !header.fin {
return Err(self.latch(error::fmt!(
SocketError,
"QWP/WebSocket server sent a fragmented frame; QWP is FIN-only"
)));
}
let payload_len = header.payload_len as usize;
if header.payload_len > MAX_INBOUND_FRAME_BYTES {
return Err(self.latch(error::fmt!(
SocketError,
"WS frame declared {} payload bytes (max {})",
header.payload_len,
MAX_INBOUND_FRAME_BYTES
)));
}
self.read_buf.clear();
if self.read_buf.try_reserve(payload_len).is_err() {
return Err(self.latch(error::fmt!(
SocketError,
"could not allocate {} bytes for inbound QWP frame",
payload_len
)));
}
self.read_buf.resize(payload_len, 0);
self.read_exact_into_buf(payload_len)?;
match header.opcode {
Opcode::Binary => {
return parse_qwp_response(&self.read_buf).inspect_err(|_| {
self.must_close = true;
});
}
Opcode::Ping => {
self.send_pong(payload_len)?;
continue;
}
Opcode::Pong => {
continue;
}
Opcode::Close => {
self.must_close = true;
self.transport_dead = true;
return Err(error::fmt!(
SocketError,
"QWP/WebSocket server closed the connection"
));
}
}
}
}
fn read_ws_frame_header(&mut self) -> Result<FrameHeader> {
loop {
match FrameHeader::parse(&self.leftover) {
Ok(h) => {
let header_len = h.header_len;
self.leftover.drain(..header_len);
return Ok(h);
}
Err(FrameError::Incomplete) => {
self.fill_leftover()?;
}
Err(FrameError::Protocol(msg)) => {
return Err(self.latch(error::fmt!(
SocketError,
"QWP/WebSocket frame parse error: {}",
msg
)));
}
}
}
}
fn read_exact_into_buf(&mut self, len: usize) -> Result<()> {
let from_leftover = self.leftover.len().min(len);
self.read_buf[..from_leftover].copy_from_slice(&self.leftover[..from_leftover]);
self.leftover.drain(..from_leftover);
let mut filled = from_leftover;
while filled < len {
let n = self
.stream
.read(&mut self.read_buf[filled..])
.map_err(|e| {
self.latch(error::fmt!(
SocketError,
"QWP/WebSocket socket read failed: {}",
e
))
})?;
if n == 0 {
return Err(self.latch(error::fmt!(
SocketError,
"QWP/WebSocket socket closed unexpectedly during frame read"
)));
}
filled += n;
}
Ok(())
}
fn try_fill_leftover(&mut self) -> Result<bool> {
let mut chunk = [0u8; 4096];
match self.stream.read_nonblocking_once(&mut chunk) {
Ok(0) => Err(self.latch(error::fmt!(
SocketError,
"QWP/WebSocket socket closed unexpectedly"
))),
Ok(n) => {
self.leftover.extend_from_slice(&chunk[..n]);
Ok(true)
}
Err(e) if e.kind() == io::ErrorKind::WouldBlock => Ok(false),
Err(e) => Err(self.latch(error::fmt!(
SocketError,
"QWP/WebSocket non-blocking read failed: {}",
e
))),
}
}
fn fill_leftover(&mut self) -> Result<()> {
let mut chunk = [0u8; 1024];
let n = self.stream.read(&mut chunk).map_err(|e| {
self.latch(error::fmt!(
SocketError,
"QWP/WebSocket socket read failed: {}",
e
))
})?;
if n == 0 {
return Err(self.latch(error::fmt!(
SocketError,
"QWP/WebSocket socket closed unexpectedly while reading frame header"
)));
}
self.leftover.extend_from_slice(&chunk[..n]);
Ok(())
}
fn send_pong(&mut self, payload_len: usize) -> Result<()> {
let mask_key = self.mask_keys.next_key().map_err(|e| {
self.latch(error::fmt!(SocketError, "mask key entropy failed: {}", e.0))
})?;
let mut pong = Vec::with_capacity(WS_HEADER_RESERVE + payload_len);
frame::encode_client_frame(
&mut pong,
Opcode::Pong,
mask_key,
&self.read_buf[..payload_len],
);
self.stream.write_all(&pong).map_err(|e| {
self.latch(error::fmt!(
SocketError,
"QWP/WebSocket pong write failed: {}",
e
))
})?;
self.stream.flush().map_err(|e| {
self.latch(error::fmt!(
SocketError,
"QWP/WebSocket pong flush failed: {}",
e
))
})?;
Ok(())
}
}
impl Drop for ColumnConn {
fn drop(&mut self) {
if self
.stream
.set_timeouts(Some(CLOSE_TIMEOUT), Some(CLOSE_TIMEOUT))
.is_ok()
&& let Ok(mask_key) = self.mask_keys.next_key()
{
self.write_buf.clear();
encode_client_frame(
&mut self.write_buf,
Opcode::Close,
mask_key,
&WS_CLOSE_STATUS_NORMAL,
);
let _ = self.stream.write_all(&self.write_buf);
let _ = self.stream.flush();
}
self.stream.shutdown_tls();
}
}
pub(crate) struct PublishedFrame {
pub(crate) fsn: u64,
}
pub(crate) enum PublishError {
BeforeWrite(crate::Error),
DuringWrite(crate::Error),
}
#[derive(Debug)]
enum QwpResponse {
Ok {
sequence: u64,
tables: Vec<(String, i64)>,
},
DurableAck {
tables: Vec<(String, i64)>,
},
Error {
sequence: u64,
status: u8,
message: String,
},
}
fn parse_qwp_response(payload: &[u8]) -> Result<QwpResponse> {
if payload.is_empty() {
return Err(error::fmt!(SocketError, "Empty QWP response frame"));
}
let status = payload[0];
match status {
QWP_STATUS_OK => {
if payload.len() < 1 + 8 + 2 {
return Err(error::fmt!(SocketError, "QWP OK response truncated"));
}
let sequence = u64::from_le_bytes(payload[1..9].try_into().unwrap());
let tables = parse_table_entries(payload, 9, "QWP OK response")?;
Ok(QwpResponse::Ok { sequence, tables })
}
QWP_STATUS_DURABLE_ACK => {
let tables = parse_table_entries(payload, 1, "QWP durable ACK response")?;
Ok(QwpResponse::DurableAck { tables })
}
_ => {
let (sequence, message) = parse_error_body(payload)?;
Ok(QwpResponse::Error {
sequence,
status,
message,
})
}
}
}
fn parse_table_entries(
payload: &[u8],
table_count_offset: usize,
context: &'static str,
) -> Result<Vec<(String, i64)>> {
let table_count_end = table_count_offset
.checked_add(2)
.ok_or_else(|| error::fmt!(SocketError, "{} table count offset overflow", context))?;
if payload.len() < table_count_end {
return Err(error::fmt!(SocketError, "{} truncated", context));
}
let table_count = u16::from_le_bytes(
payload[table_count_offset..table_count_end]
.try_into()
.unwrap(),
) as usize;
let mut pos = table_count_end;
let max_entries = payload.len().saturating_sub(table_count_end) / 11;
let mut entries: Vec<(String, i64)> = Vec::new();
if entries.try_reserve(table_count.min(max_entries)).is_err() {
return Err(error::fmt!(
SocketError,
"{} could not allocate {} table entries",
context,
table_count
));
}
for _ in 0..table_count {
let name_len_end = pos
.checked_add(2)
.ok_or_else(|| error::fmt!(SocketError, "{} table entry offset overflow", context))?;
if payload.len() < name_len_end {
return Err(error::fmt!(
SocketError,
"{} table entry truncated",
context
));
}
let name_len = u16::from_le_bytes(payload[pos..name_len_end].try_into().unwrap()) as usize;
pos = name_len_end;
if name_len == 0 {
return Err(error::fmt!(SocketError, "{} table name is empty", context));
}
let name_end = pos
.checked_add(name_len)
.ok_or_else(|| error::fmt!(SocketError, "{} table name length overflow", context))?;
let seq_txn_end = name_end
.checked_add(8)
.ok_or_else(|| error::fmt!(SocketError, "{} table entry length overflow", context))?;
if payload.len() < seq_txn_end {
return Err(error::fmt!(
SocketError,
"{} table entry truncated",
context
));
}
let name = std::str::from_utf8(&payload[pos..name_end])
.map_err(|_| error::fmt!(SocketError, "{} table name not UTF-8", context))?
.to_owned();
let seq_txn = i64::from_le_bytes(payload[name_end..seq_txn_end].try_into().unwrap());
entries.push((name, seq_txn));
pos = seq_txn_end;
}
if pos != payload.len() {
return Err(error::fmt!(
SocketError,
"{} has trailing bytes after table entries",
context
));
}
Ok(entries)
}
fn parse_error_body(payload: &[u8]) -> Result<(u64, String)> {
if payload.len() < 1 + 8 + 2 {
return Err(error::fmt!(SocketError, "QWP error response truncated"));
}
let sequence = u64::from_le_bytes(payload[1..9].try_into().unwrap());
let msg_len = u16::from_le_bytes(payload[9..11].try_into().unwrap()) as usize;
if msg_len > 1024 {
return Err(error::fmt!(
SocketError,
"QWP error response message too long (declared {} bytes, max 1024)",
msg_len
));
}
let msg_end = 11usize
.checked_add(msg_len)
.ok_or_else(|| error::fmt!(SocketError, "QWP error response message length overflow"))?;
if payload.len() < msg_end {
return Err(error::fmt!(
SocketError,
"QWP error response truncated (declared {} bytes)",
msg_len
));
}
if payload.len() != msg_end {
return Err(error::fmt!(
SocketError,
"QWP error response has trailing bytes after message"
));
}
let message = std::str::from_utf8(&payload[11..msg_end])
.map_err(|_| error::fmt!(SocketError, "QWP error message not UTF-8"))?
.to_owned();
Ok((sequence, message))
}
fn map_error_status(status: u8, msg: &str) -> crate::Error {
match status {
QWP_STATUS_SCHEMA_MISMATCH => {
error::fmt!(InvalidApiCall, "QWP schema mismatch: {}", msg)
}
QWP_STATUS_PARSE_ERROR => error::fmt!(InvalidApiCall, "QWP parse error: {}", msg),
QWP_STATUS_INTERNAL_ERROR => error::fmt!(ServerFlushError, "QWP internal error: {}", msg),
QWP_STATUS_SECURITY_ERROR => error::fmt!(AuthError, "QWP security error: {}", msg),
QWP_STATUS_WRITE_ERROR => error::fmt!(ServerFlushError, "QWP write error: {}", msg),
_ => error::fmt!(
ServerFlushError,
"QWP unrecognised error status 0x{:02x}: {}",
status,
msg
),
}
}
#[inline]
fn ws_header_len_for(payload_len: usize) -> usize {
if payload_len <= 125 {
2 + 4
} else if payload_len <= 0xFFFF {
4 + 4
} else {
10 + 4
}
}
fn write_ws_header(out: &mut [u8], payload_len: usize, mask_key: [u8; 4]) {
const FIN_BIT: u8 = 0x80;
const BINARY_OPCODE: u8 = 0x2;
const MASK_BIT: u8 = 0x80;
out[0] = FIN_BIT | BINARY_OPCODE;
let len_bytes;
let mask_offset;
if payload_len <= 125 {
out[1] = MASK_BIT | (payload_len as u8);
mask_offset = 2;
len_bytes = 0;
} else if payload_len <= 0xFFFF {
out[1] = MASK_BIT | 126;
out[2..4].copy_from_slice(&(payload_len as u16).to_be_bytes());
mask_offset = 4;
len_bytes = 2;
} else {
out[1] = MASK_BIT | 127;
out[2..10].copy_from_slice(&(payload_len as u64).to_be_bytes());
mask_offset = 10;
len_bytes = 8;
}
let _ = len_bytes;
out[mask_offset..mask_offset + 4].copy_from_slice(&mask_key);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn ws_header_len_matches_payload_length_class() {
assert_eq!(ws_header_len_for(0), 6);
assert_eq!(ws_header_len_for(125), 6);
assert_eq!(ws_header_len_for(126), 8);
assert_eq!(ws_header_len_for(0xFFFF), 8);
assert_eq!(ws_header_len_for(0x1_0000), 14);
assert_eq!(ws_header_len_for(1 << 24), 14);
}
#[test]
fn write_ws_header_short_form() {
let mut buf = [0u8; 6];
write_ws_header(&mut buf, 5, [0xDE, 0xAD, 0xBE, 0xEF]);
assert_eq!(buf[0], 0x82); assert_eq!(buf[1], 0x80 | 5); assert_eq!(&buf[2..6], &[0xDE, 0xAD, 0xBE, 0xEF]);
}
#[test]
fn write_ws_header_16bit_form() {
let mut buf = [0u8; 8];
write_ws_header(&mut buf, 200, [1, 2, 3, 4]);
assert_eq!(buf[0], 0x82);
assert_eq!(buf[1], 0x80 | 126);
assert_eq!(u16::from_be_bytes([buf[2], buf[3]]), 200);
assert_eq!(&buf[4..8], &[1, 2, 3, 4]);
}
#[test]
fn write_ws_header_64bit_form() {
let mut buf = [0u8; 14];
write_ws_header(&mut buf, 0x1_0000, [9, 8, 7, 6]);
assert_eq!(buf[0], 0x82);
assert_eq!(buf[1], 0x80 | 127);
assert_eq!(
u64::from_be_bytes([
buf[2], buf[3], buf[4], buf[5], buf[6], buf[7], buf[8], buf[9]
]),
0x1_0000
);
assert_eq!(&buf[10..14], &[9, 8, 7, 6]);
}
#[test]
fn parse_qwp_ok_with_one_table() {
let mut payload = vec![0u8];
payload.extend_from_slice(&42u64.to_le_bytes());
payload.extend_from_slice(&1u16.to_le_bytes());
payload.extend_from_slice(&2u16.to_le_bytes());
payload.extend_from_slice(b"tx");
payload.extend_from_slice(&7i64.to_le_bytes());
let response = parse_qwp_response(&payload).unwrap();
match response {
QwpResponse::Ok { sequence, tables } => {
assert_eq!(sequence, 42);
assert_eq!(tables, vec![("tx".to_owned(), 7)]);
}
other => panic!("expected Ok, got {other:?}"),
}
}
#[test]
fn parse_qwp_durable_ack_empty() {
let mut payload = vec![QWP_STATUS_DURABLE_ACK];
payload.extend_from_slice(&0u16.to_le_bytes());
let response = parse_qwp_response(&payload).unwrap();
match response {
QwpResponse::DurableAck { tables } => {
assert!(tables.is_empty());
}
other => panic!("expected DurableAck, got {other:?}"),
}
}
#[test]
fn parse_qwp_error_truncated_rejected() {
let err = parse_qwp_response(&[QWP_STATUS_PARSE_ERROR]).unwrap_err();
assert_eq!(err.code(), crate::ErrorCode::SocketError);
}
fn dummy_ws_stream() -> WsStream {
use crate::ws::nosigpipe::NoSigpipeTcp;
use std::net::{TcpListener, TcpStream};
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("local_addr");
let client = TcpStream::connect(addr).expect("connect");
let _server = listener.accept().expect("accept");
WsStream::Plain(NoSigpipeTcp::new(client).expect("nosigpipe"))
}
#[test]
fn process_ok_pops_pending_and_tracks_durable_when_opted_in() {
let mut conn = ColumnConn::for_test(dummy_ws_stream(), true);
conn.push_pending(0);
conn.process_response(QwpResponse::Ok {
sequence: 0,
tables: vec![("trades".to_string(), 7)],
})
.expect("matching ok");
assert_eq!(conn.in_flight(), 0);
assert_eq!(conn.pending_durable_target_count(), 1);
assert!(!conn.must_close());
}
#[test]
fn process_ok_skips_durable_targets_without_opt_in() {
let mut conn = ColumnConn::for_test(dummy_ws_stream(), false);
conn.push_pending(0);
conn.process_response(QwpResponse::Ok {
sequence: 0,
tables: vec![("trades".to_string(), 7), ("quotes".to_string(), 3)],
})
.expect("matching ok");
assert_eq!(conn.in_flight(), 0);
assert_eq!(
conn.pending_durable_target_count(),
0,
"durable targets must not accumulate without request_durable_ack"
);
}
#[test]
fn process_unmatched_ok_is_tolerated_as_noop() {
let mut conn = ColumnConn::for_test(dummy_ws_stream(), false);
conn.process_response(QwpResponse::Ok {
sequence: 0,
tables: vec![],
})
.expect("redundant ok must be tolerated");
assert!(!conn.must_close());
}
#[test]
fn process_stale_ok_below_pending_fsn_is_noop_and_keeps_pending() {
let mut conn = ColumnConn::for_test(dummy_ws_stream(), false);
conn.push_pending(5);
conn.process_response(QwpResponse::Ok {
sequence: 3,
tables: vec![],
})
.expect("stale ok must be tolerated");
assert!(!conn.must_close());
assert_eq!(conn.in_flight(), 1);
conn.process_response(QwpResponse::Ok {
sequence: 5,
tables: vec![],
})
.expect("matching ok must pop");
assert_eq!(conn.in_flight(), 0);
}
#[test]
fn sync_all_acks_fails_fast_on_non_advancing_peer() {
use crate::ws::nosigpipe::NoSigpipeTcp;
use std::net::{TcpListener, TcpStream};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::mpsc;
use std::thread;
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("addr");
let client = TcpStream::connect(addr).expect("connect");
let (server, _) = listener.accept().expect("accept");
let stop = Arc::new(AtomicBool::new(false));
let stop_writer = Arc::clone(&stop);
let flood = thread::spawn(move || {
let mut server = server;
let mut payload = vec![QWP_STATUS_OK];
payload.extend_from_slice(&0u64.to_le_bytes());
payload.extend_from_slice(&0u16.to_le_bytes());
let mut ws_frame = vec![0x82u8, payload.len() as u8];
ws_frame.extend_from_slice(&payload);
while !stop_writer.load(Ordering::Relaxed) {
if server.write_all(&ws_frame).is_err() {
break;
}
}
});
let mut conn = ColumnConn::for_test(
WsStream::Plain(NoSigpipeTcp::new(client).expect("nosigpipe")),
false,
);
conn.request_timeout = Duration::from_millis(150);
conn.push_pending(5);
let (tx, rx) = mpsc::channel();
let worker = thread::spawn(move || {
let code = conn.sync_all_acks(AckLevel::Ok).err().map(|e| e.code());
let _ = tx.send(code);
});
let outcome = rx.recv_timeout(Duration::from_secs(5));
stop.store(true, Ordering::Relaxed);
let _ = worker.join();
let _ = flood.join();
match outcome {
Ok(Some(code)) => assert_eq!(code, crate::ErrorCode::SocketError),
Ok(None) => panic!("sync_all_acks unexpectedly succeeded against a non-advancing peer"),
Err(_) => panic!("sync_all_acks hung: the no-progress deadline did not fire"),
}
}
#[test]
fn redundant_ok_after_durable_prune_does_not_strand_sync() {
let mut conn = ColumnConn::for_test(dummy_ws_stream(), true);
conn.push_pending(0);
conn.process_response(QwpResponse::Ok {
sequence: 0,
tables: vec![("t".to_string(), 7)],
})
.expect("ok ack");
conn.process_response(QwpResponse::DurableAck {
tables: vec![("t".to_string(), 7)],
})
.expect("durable ack");
conn.drop_satisfied_durable_targets();
assert_eq!(
conn.pending_durable_targets.len(),
0,
"satisfied pending targets must be pruned"
);
conn.process_response(QwpResponse::Ok {
sequence: 0,
tables: vec![("t".to_string(), 7)],
})
.expect("redundant ok");
assert_eq!(
conn.pending_durable_targets.len(),
0,
"a redundant OK must not resurrect a satisfied durable target"
);
assert_eq!(
conn.durable_watermarks.get("t").copied(),
Some(7),
"durable watermarks must outlive their satisfied pending target"
);
assert!(
conn.durability_satisfied(),
"no pending targets => durability is trivially satisfied"
);
}
#[test]
fn process_error_frame_maps_status_and_latches() {
for (status, expected) in [
(QWP_STATUS_SCHEMA_MISMATCH, crate::ErrorCode::InvalidApiCall),
(
QWP_STATUS_INTERNAL_ERROR,
crate::ErrorCode::ServerFlushError,
),
(QWP_STATUS_SECURITY_ERROR, crate::ErrorCode::AuthError),
(QWP_STATUS_WRITE_ERROR, crate::ErrorCode::ServerFlushError),
] {
let mut conn = ColumnConn::for_test(dummy_ws_stream(), false);
let err = conn
.process_response(QwpResponse::Error {
sequence: 0,
status,
message: "boom".to_string(),
})
.expect_err("server error must surface");
assert_eq!(err.code(), expected, "status 0x{status:02x}");
assert!(conn.must_close(), "error must latch the connection");
}
}
}