use std::net::{TcpStream, ToSocketAddrs};
use std::sync::Arc;
use std::time::Duration;
use bytes::Bytes;
use crate::egress::config::ReaderConfig;
use crate::egress::server_event::UpgradeReject;
use crate::egress::tls::build_client_config;
use crate::egress::wire::MsgKind;
use crate::egress::wire::header::{FrameHeader, HEADER_LEN};
use crate::egress::wire::roles;
use crate::egress::ws::client::{Stream, WsClient, WsReadError};
use crate::error::{Error, ErrorCode, Result, fmt};
use crate::ws::handshake::{self, HandshakeError as WsHandshakeError, Headers, HttpReject};
use crate::ws::mask::MaskKeySource;
use crate::ws::nosigpipe::NoSigpipeTcp;
pub(crate) const WRITE_TIMEOUT: Duration = Duration::from_secs(60);
pub(crate) const CLOSE_TIMEOUT: Duration = Duration::from_millis(200);
const MAX_BATCH_WIRE_BYTES: usize = 64 * 1024 * 1024;
const HDR_VERSION: &str = "x-qwp-version";
const HDR_CONTENT_ENCODING: &str = "x-qwp-content-encoding";
const HDR_ROLE: &str = "x-questdb-role";
const HDR_ZONE: &str = "x-questdb-zone";
pub struct WsTransport {
socket: WsClient,
server_version: u8,
}
impl WsTransport {
pub fn connect_to(config: &ReaderConfig, addr_idx: usize) -> Result<Self> {
if addr_idx >= config.addrs.len() {
return Err(fmt!(
ConfigError,
"addr index {} out of range ({} endpoints)",
addr_idx,
config.addrs.len()
));
}
let endpoint = &config.addrs[addr_idx];
let resolved: Vec<_> = (endpoint.host.as_str(), endpoint.port)
.to_socket_addrs()
.map_err(|e| fmt!(CouldNotResolveAddr, "could not resolve {}: {}", endpoint, e))?
.collect();
if resolved.is_empty() {
return Err(fmt!(
CouldNotResolveAddr,
"name resolution returned no addresses for {}",
endpoint
));
}
let tcp = {
let mut last_err: Option<std::io::Error> = None;
let mut connected: Option<TcpStream> = None;
for addr in &resolved {
let res = match config.connect_timeout_ms {
0 => TcpStream::connect(addr),
ms => TcpStream::connect_timeout(addr, Duration::from_millis(ms)),
};
match res {
Ok(s) => {
connected = Some(s);
break;
}
Err(e) => last_err = Some(e),
}
}
match connected {
Some(s) => s,
None => {
let e = last_err.expect("non-empty addrs but no last_err");
let msg = format!(
"could not connect to {} (tried {} address(es)): {}",
endpoint,
resolved.len(),
e
);
return Err(if e.kind() == std::io::ErrorKind::TimedOut {
fmt!(ConnectTimeout, "{}", msg)
} else {
fmt!(SocketError, "{}", msg)
});
}
}
};
let tcp = NoSigpipeTcp::new(tcp).map_err(|e| {
fmt!(
SocketError,
"could not configure SO_NOSIGPIPE on {}: {}",
endpoint,
e
)
})?;
let _ = tcp
.tcp()
.set_read_timeout(Some(Duration::from_millis(config.auth_timeout_ms)));
let mut stream = build_stream(&tcp, endpoint.host.as_str(), config)?;
let host_header = endpoint.to_string();
let path = config.path.clone();
let extra_headers = config.upgrade_headers();
let handshake_result = handshake::upgrade(&mut stream, &host_header, &path, &extra_headers);
let handshake = match handshake_result {
Ok(h) => h,
Err(e) => return Err(map_handshake_error(e)),
};
let server_version = match read_version_header(&handshake.headers)
.and_then(|v| validate_content_encoding(&handshake.headers).map(|_| v))
{
Ok(v) => v,
Err(e) => {
set_tcp_write_timeout(stream.tcp_mut(), Some(CLOSE_TIMEOUT));
stream.shutdown();
return Err(e);
}
};
if server_version > config.max_version {
set_tcp_write_timeout(stream.tcp_mut(), Some(CLOSE_TIMEOUT));
stream.shutdown();
return Err(fmt!(
HandshakeError,
"server negotiated QWP version {} but client advertised max {}",
server_version,
config.max_version
));
}
set_tcp_write_timeout(stream.tcp_mut(), Some(WRITE_TIMEOUT));
set_tcp_read_timeout(stream.tcp_mut(), None);
let mask_keys = MaskKeySource::new().map_err(|e| fmt!(ConfigError, "{}", e.0))?;
let socket = WsClient::new(stream, handshake.leftover, mask_keys, MAX_BATCH_WIRE_BYTES);
Ok(WsTransport {
socket,
server_version,
})
}
pub fn server_version(&self) -> u8 {
self.server_version
}
pub fn write_message(&mut self, payload: Bytes) -> Result<()> {
self.socket
.write_binary_frame(&payload)
.map_err(|e| map_io_error(e, ErrorCode::SocketError))
}
pub fn read_frame(&mut self) -> Result<(FrameHeader, Bytes)> {
let bytes = match self.socket.read_binary_frame() {
Ok(b) => b,
Err(e) => return Err(map_ws_read_error(e)),
};
if bytes.len() < HEADER_LEN {
return Err(fmt!(
ProtocolError,
"WS message too short for frame header: {} bytes",
bytes.len()
));
}
let header = FrameHeader::parse(&bytes[..HEADER_LEN])?;
if header.version != self.server_version {
return Err(fmt!(
ProtocolError,
"frame header version {} != negotiated {}",
header.version,
self.server_version
));
}
if header.payload_length as usize != bytes.len() - HEADER_LEN {
return Err(fmt!(
ProtocolError,
"header payload_length {} != actual {}",
header.payload_length,
bytes.len() - HEADER_LEN
));
}
if bytes.len() > MAX_BATCH_WIRE_BYTES {
return Err(fmt!(
LimitExceeded,
"frame size {} bytes exceeds client cap {} (spec §16: \
RESULT_BATCH max 16 MiB; client allows 4x margin)",
bytes.len(),
MAX_BATCH_WIRE_BYTES
));
}
let payload = bytes.slice(HEADER_LEN..);
Ok((header, payload))
}
pub fn set_read_timeout(&mut self, timeout: Option<Duration>) {
set_tcp_read_timeout(self.socket.stream_mut().tcp_mut(), timeout);
}
pub fn set_write_timeout(&mut self, timeout: Option<Duration>) {
set_tcp_write_timeout(self.socket.stream_mut().tcp_mut(), timeout);
}
pub fn close_in_place(&mut self) {
teardown_inplace(&mut self.socket);
}
pub fn try_write_cancel(&mut self, request_id: i64) {
set_tcp_write_timeout(self.socket.stream_mut().tcp_mut(), Some(CLOSE_TIMEOUT));
let mut payload = Vec::with_capacity(9);
payload.push(MsgKind::Cancel.as_u8());
payload.extend_from_slice(&request_id.to_le_bytes());
let _ = self.socket.write_binary_frame(&payload);
}
}
impl Drop for WsTransport {
fn drop(&mut self) {
teardown_inplace(&mut self.socket);
}
}
fn build_stream(tcp: &NoSigpipeTcp, host: &str, config: &ReaderConfig) -> Result<Stream> {
let owned = tcp
.try_clone()
.map_err(|e| fmt!(SocketError, "could not clone TCP socket: {}", e))?;
if let Some(client_config) = build_client_config(config)? {
let server_name = rustls::pki_types::ServerName::try_from(host.to_string())
.map_err(|e| fmt!(ConfigError, "invalid TLS server name {:?}: {}", host, e))?;
let conn = rustls::ClientConnection::new(Arc::clone(&client_config), server_name)
.map_err(|e| fmt!(TlsError, "rustls handshake setup failed: {}", e))?;
let stream_owned = rustls::StreamOwned::new(conn, owned);
Ok(Stream::Tls(Box::new(stream_owned)))
} else {
Ok(Stream::Plain(owned))
}
}
fn set_tcp_write_timeout(stream: &mut TcpStream, timeout: Option<Duration>) -> bool {
stream.set_write_timeout(timeout).is_ok()
}
fn set_tcp_read_timeout(stream: &mut TcpStream, timeout: Option<Duration>) {
let _ = stream.set_read_timeout(timeout);
}
fn teardown_inplace(socket: &mut WsClient) {
if set_tcp_write_timeout(socket.stream_mut().tcp_mut(), Some(CLOSE_TIMEOUT)) {
let _ = socket.send_close(1000);
}
socket.stream_mut().shutdown();
}
fn read_version_header(headers: &Headers) -> Result<u8> {
let raw = headers.find_ci(HDR_VERSION).ok_or_else(|| {
fmt!(
HandshakeError,
"server response missing X-QWP-Version header"
)
})?;
raw.parse::<u8>()
.map_err(|_| fmt!(HandshakeError, "X-QWP-Version {:?} is not a u8", raw))
}
fn validate_content_encoding(headers: &Headers) -> Result<()> {
let raw = match headers.find_ci(HDR_CONTENT_ENCODING) {
Some(v) => v,
None => return Ok(()),
};
let mut parts = raw.split(';');
let name = parts.next().unwrap_or("").trim();
if name.eq_ignore_ascii_case("raw") || name.eq_ignore_ascii_case("identity") || name.is_empty()
{
Ok(())
} else if name.eq_ignore_ascii_case("zstd") {
#[cfg(feature = "sync-reader-zstd")]
{
Ok(())
}
#[cfg(not(feature = "sync-reader-zstd"))]
{
Err(fmt!(
HandshakeError,
"server selected X-QWP-Content-Encoding {:?} but this client was built \
without the `sync-reader-zstd` feature",
raw
))
}
} else {
Err(fmt!(
HandshakeError,
"server selected X-QWP-Content-Encoding {:?} (unknown codec {:?})",
raw,
name
))
}
}
fn map_io_error(e: std::io::Error, default_code: ErrorCode) -> Error {
let msg = e.to_string();
Error::new(default_code, msg)
}
fn map_ws_read_error(e: WsReadError) -> Error {
match e {
WsReadError::Io(io_err) => Error::new(ErrorCode::SocketError, io_err.to_string()),
WsReadError::Protocol(msg) => Error::new(ErrorCode::ProtocolError, msg),
WsReadError::ServerClose { code } => Error::new(
ErrorCode::SocketError,
match code {
Some(c) => format!("server closed WebSocket (code={})", c),
None => "server closed WebSocket".to_string(),
},
),
}
}
fn map_handshake_error(e: WsHandshakeError) -> Error {
match e {
WsHandshakeError::Io(io_err) => {
let code = if is_tls_io_error(&io_err) {
ErrorCode::TlsError
} else {
ErrorCode::SocketError
};
Error::new(code, format!("WebSocket handshake IO error: {}", io_err))
}
WsHandshakeError::Protocol(msg) => Error::new(
ErrorCode::HandshakeError,
format!("WebSocket handshake protocol error: {}", msg),
),
WsHandshakeError::BadAccept => fmt!(
HandshakeError,
"WebSocket handshake response had invalid Sec-WebSocket-Accept (server not speaking WS \
RFC 6455 or signing with the wrong key)"
),
WsHandshakeError::HttpStatus(reject) => map_http_reject(reject),
}
}
fn map_http_reject(reject: HttpReject) -> Error {
let HttpReject {
status,
headers,
body: _,
} = reject;
if status == 421
&& let Some(upgrade_reject) = parse_upgrade_reject(&headers)
{
return Error::new(
ErrorCode::RoleMismatch,
format!(
"server rejected WebSocket upgrade with 421 + X-QuestDB-Role={} \
(zone={:?}); host is in {} state",
upgrade_reject.role_name,
upgrade_reject.zone,
if upgrade_reject.is_transient() {
"transient (PRIMARY_CATCHUP)"
} else {
"topological"
},
),
)
.with_upgrade_reject(upgrade_reject);
}
let code = if status == 401 || status == 403 {
ErrorCode::AuthError
} else {
ErrorCode::HandshakeError
};
Error::new(
code,
format!("WebSocket handshake failed with HTTP {}", status),
)
}
fn parse_upgrade_reject(headers: &Headers) -> Option<UpgradeReject> {
let role_raw = headers.find_ci(HDR_ROLE)?;
if role_raw.is_empty() {
return None;
}
let role_name = role_raw.to_ascii_uppercase();
let role_byte = roles::byte_for_name(&role_name).unwrap_or(roles::UNKNOWN_NAME);
let zone = headers.find_ci(HDR_ZONE).and_then(|v| {
if v.is_empty() {
None
} else {
Some(v.to_string())
}
});
Some(UpgradeReject::new(role_byte, role_name, zone))
}
fn is_tls_io_error(e: &std::io::Error) -> bool {
if let Some(src) = e.get_ref() {
if src.downcast_ref::<rustls::Error>().is_some() {
return true;
}
let mut cur: Option<&(dyn std::error::Error + 'static)> = src.source();
while let Some(s) = cur {
if s.downcast_ref::<rustls::Error>().is_some() {
return true;
}
cur = s.source();
}
}
false
}
#[cfg(test)]
mod tests {
use super::*;
fn header_map(value: &str) -> Headers {
Headers::from_pairs([("X-QWP-Content-Encoding", value)])
}
#[test]
fn module_is_compilable() {
}
#[test]
fn content_encoding_absent_is_ok() {
validate_content_encoding(&Headers::default()).unwrap();
}
#[test]
fn content_encoding_raw_is_ok() {
validate_content_encoding(&header_map("raw")).unwrap();
validate_content_encoding(&header_map("identity")).unwrap();
}
#[cfg(feature = "sync-reader-zstd")]
#[test]
fn content_encoding_zstd_bare_is_ok() {
validate_content_encoding(&header_map("zstd")).unwrap();
}
#[cfg(feature = "sync-reader-zstd")]
#[test]
fn content_encoding_zstd_with_level_parameter_is_ok() {
validate_content_encoding(&header_map("zstd;level=3")).unwrap();
validate_content_encoding(&header_map("zstd; level=3")).unwrap();
validate_content_encoding(&header_map("zstd;level=1")).unwrap();
validate_content_encoding(&header_map("zstd;level=9")).unwrap();
}
#[cfg(feature = "sync-reader-zstd")]
#[test]
fn content_encoding_zstd_with_unknown_parameter_is_ok() {
validate_content_encoding(&header_map("zstd;dict=42")).unwrap();
validate_content_encoding(&header_map("zstd;foo=bar;baz=qux")).unwrap();
validate_content_encoding(&header_map("zstd;level=99")).unwrap();
}
#[cfg(feature = "sync-reader-zstd")]
#[test]
fn content_encoding_trailing_semicolon_is_ok() {
validate_content_encoding(&header_map("zstd;")).unwrap();
validate_content_encoding(&header_map("zstd; ; ")).unwrap();
}
#[test]
fn content_encoding_unknown_codec_rejected() {
let err = validate_content_encoding(&header_map("brotli")).unwrap_err();
assert_eq!(err.code(), ErrorCode::HandshakeError);
assert!(err.msg().contains("unknown codec"), "got: {}", err.msg());
}
#[test]
fn content_encoding_codec_name_is_case_insensitive() {
validate_content_encoding(&header_map("RAW")).unwrap();
validate_content_encoding(&header_map("Raw")).unwrap();
validate_content_encoding(&header_map("IDENTITY")).unwrap();
validate_content_encoding(&header_map("Identity")).unwrap();
}
#[cfg(feature = "sync-reader-zstd")]
#[test]
fn content_encoding_zstd_case_insensitive() {
validate_content_encoding(&header_map("ZSTD")).unwrap();
validate_content_encoding(&header_map("Zstd")).unwrap();
validate_content_encoding(&header_map("zStd")).unwrap();
validate_content_encoding(&header_map("Zstd;level=3")).unwrap();
}
#[cfg(not(feature = "sync-reader-zstd"))]
#[test]
fn content_encoding_zstd_case_insensitive_rejected_without_feature() {
let err = validate_content_encoding(&header_map("ZSTD")).unwrap_err();
assert_eq!(err.code(), ErrorCode::HandshakeError);
let err = validate_content_encoding(&header_map("Zstd")).unwrap_err();
assert_eq!(err.code(), ErrorCode::HandshakeError);
}
#[test]
fn content_encoding_unknown_codec_with_parameters_still_rejected() {
let err = validate_content_encoding(&header_map("brotli;q=1.0")).unwrap_err();
assert_eq!(err.code(), ErrorCode::HandshakeError);
}
}