use std::net::Ipv4Addr;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use bytes::{Bytes, BytesMut};
use crate::egress::binds::{Bind, SimpleNullKind};
use crate::egress::column::ColumnView;
use crate::egress::config::{Endpoint, ReaderConfig, Target};
use crate::egress::decoder::DecodedBatch;
use crate::egress::decoder::ZstdScratch;
use crate::egress::query_request::{
QUERY_FLAG_RESET_DICT, QueryRequest, QueryRequestBuilder, REQUEST_ID_OFFSET,
};
use crate::egress::schema::Schema;
use crate::egress::server_event::UpgradeReject;
use crate::egress::server_event::{ServerEvent, ServerInfo, ServerRole, decode_frame};
use crate::egress::symbol_dict::SymbolDict;
use crate::egress::tracker::HostHealthTracker;
use crate::egress::transport::{CLOSE_TIMEOUT, WRITE_TIMEOUT, WsTransport};
use crate::egress::wire::capabilities::has_query_flags;
use crate::egress::wire::header::HEADER_LEN;
use crate::egress::wire::msg_kind::MsgKind;
use crate::egress::wire::varint;
use crate::error::{Error, ErrorCode, Result, fmt};
#[derive(Debug, Default)]
pub struct ReaderStats {
pub bytes_received: AtomicU64,
pub credit_granted_total: AtomicU64,
pub read_ns: AtomicU64,
pub decode_ns: AtomicU64,
}
pub struct Reader {
cfg: Arc<ReaderConfig>,
addr_idx: usize,
transport: Option<WsTransport>,
dict: SymbolDict,
query_schema: Option<Schema>,
next_request_id: i64,
cursor_active: bool,
server_info: Option<ServerInfo>,
stats: Arc<ReaderStats>,
zstd_scratch: ZstdScratch,
tracker: HostHealthTracker,
failover_rng: FailoverRng,
}
const _: fn() = || {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<Reader>();
assert_send_sync::<ReaderStats>();
assert_send_sync::<HostHealthTracker>();
};
const _: fn() = || {
fn assert_send<T: Send>() {}
assert_send::<crate::egress::ReaderQuery<'_>>();
assert_send::<crate::egress::Cursor<'_>>();
#[cfg(feature = "arrow-egress")]
assert_send::<crate::egress::arrow::CursorRecordBatchReader<'_, '_>>();
#[cfg(feature = "polars-egress")]
assert_send::<crate::egress::arrow::polars::CursorPolarsIter<'_, '_>>();
};
impl Reader {
pub fn from_conf<T: AsRef<str>>(conf: T) -> Result<Self> {
let cfg = ReaderConfig::from_conf(conf)?;
Self::from_config(&cfg)
}
pub fn from_env() -> Result<Self> {
let conf = std::env::var("QDB_CLIENT_CONF").map_err(|e| match e {
std::env::VarError::NotPresent => {
fmt!(ConfigError, "Environment variable QDB_CLIENT_CONF not set.")
}
std::env::VarError::NotUnicode(_) => fmt!(
InvalidUtf8,
"Environment variable QDB_CLIENT_CONF is set but its value is not valid UTF-8."
),
})?;
Self::from_conf(conf)
}
pub fn from_config(cfg: &ReaderConfig) -> Result<Self> {
cfg.validate()?;
let cfg = Arc::new(cfg.clone());
let mut tracker = HostHealthTracker::new(
cfg.addrs.len(),
cfg.zone.as_deref(),
matches!(cfg.target, Target::Primary),
);
let walk = walk_via_tracker(
&mut tracker,
&cfg,
false,
&[
ErrorCode::ConfigError,
ErrorCode::UnsupportedServer,
ErrorCode::AuthError,
],
)?;
Ok(Reader {
cfg,
addr_idx: walk.session.idx,
transport: Some(walk.session.transport),
dict: SymbolDict::new(),
query_schema: None,
next_request_id: 1,
cursor_active: false,
server_info: walk.session.server_info,
stats: Arc::new(ReaderStats::default()),
zstd_scratch: ZstdScratch::new(),
tracker,
failover_rng: FailoverRng::new(),
})
}
fn connect_endpoint(cfg: &ReaderConfig, idx: usize) -> Result<TransportSession> {
let mut transport = WsTransport::connect_to(cfg, idx).map_err(|e| {
let endpoint = &cfg.addrs[idx];
let mut annotated = Error::new(e.code(), format!("endpoint {}: {}", endpoint, e.msg()));
if let Some(r) = e.upgrade_reject() {
annotated = annotated.with_upgrade_reject(r.clone());
}
if let Some(info) = e.server_info() {
annotated = annotated.with_server_info(info.clone());
}
annotated
})?;
let server_info = if transport.server_version() >= 1 {
Some(read_server_info_frame(
&mut transport,
Duration::from_millis(cfg.server_info_timeout_ms),
)?)
} else {
None
};
if !matches!(cfg.target, Target::Any) {
match server_info.as_ref() {
None => {
return Err(fmt!(
RoleMismatch,
"endpoint {} supplied no SERVER_INFO and cannot match target={:?}",
idx,
cfg.target
));
}
Some(info) if !target_matches(cfg.target, info.role) => {
let role = info.role;
let role_name = role.as_str();
let reject =
UpgradeReject::new(role.as_u8(), role_name.clone(), info.zone_id.clone());
return Err(Error::new(
ErrorCode::RoleMismatch,
format!(
"endpoint {} role={} cluster={:?} does not match target={:?}",
idx, role_name, info.cluster_id, cfg.target,
),
)
.with_upgrade_reject(reject)
.with_server_info(info.clone()));
}
_ => {}
}
}
Ok(TransportSession {
idx,
transport,
server_info,
})
}
fn reconnect_with_failover(
&mut self,
failed_idx: usize,
budget: &mut FailoverBudget,
on_attempt: &mut dyn FnMut(u32),
) -> Result<u32> {
let cfg = Arc::clone(&self.cfg);
let mut last_err: Option<Error> = None;
let mut deadline_exhausted = false;
self.tracker.record_mid_stream_failure(failed_idx);
if let Some(dead) = self.transport.take() {
drop(dead);
}
let mut total_dials: u32 = 0;
let mut attempts_made: u32 = 0;
loop {
match budget.before_reconnect_round(&cfg, &mut self.failover_rng) {
Ok(()) => {}
Err(FailoverBudgetStop::AttemptsExhausted) => break,
Err(FailoverBudgetStop::DeadlineExhausted) => {
deadline_exhausted = true;
break;
}
}
attempts_made = attempts_made.saturating_add(1);
on_attempt(attempts_made);
match walk_via_tracker(
&mut self.tracker,
&cfg,
true,
&[
ErrorCode::ConfigError,
ErrorCode::UnsupportedServer,
ErrorCode::AuthError,
],
) {
Ok(walk) => {
total_dials = total_dials.saturating_add(walk.dials);
self.transport = Some(walk.session.transport);
self.server_info = walk.session.server_info;
self.dict = SymbolDict::new();
self.query_schema = None;
self.addr_idx = walk.session.idx;
return Ok(total_dials);
}
Err(e) => match e.code() {
code if !is_failover_eligible(code) => {
return Err(e);
}
_ => {
warn_on_protocol_error_failover(&e, "reconnect walk");
last_err = Some(e);
}
},
}
}
if deadline_exhausted {
let last_msg = last_err
.as_ref()
.map(|e| e.msg().to_string())
.unwrap_or_else(|| "<no error captured>".to_string());
return Err(fmt!(
SocketError,
"failover wall-clock budget exhausted (failover_max_duration_ms={}) after {} attempt(s); last error: {}",
cfg.failover_max_duration_ms,
attempts_made,
last_msg
));
}
Err(last_err.unwrap_or_else(|| {
fmt!(
SocketError,
"failover exhausted after {} attempts",
attempts_made
)
}))
}
pub fn current_addr(&self) -> &Endpoint {
&self.cfg.addrs[self.addr_idx]
}
fn transport_mut(&mut self) -> Result<&mut WsTransport> {
self.transport.as_mut().ok_or_else(|| {
fmt!(
SocketError,
"Reader connection is closed and cannot be reused: a cursor was dropped before being \
fully read, or a mid-query failover exhausted its retry budget. To keep the \
connection reusable, drain each cursor (call next_batch() until it returns None) or \
call cursor.cancel() before dropping it; otherwise open a fresh Reader."
)
})
}
fn transport_ref(&self) -> Result<&WsTransport> {
self.transport.as_ref().ok_or_else(|| {
fmt!(
SocketError,
"Reader connection is closed and cannot be reused: a cursor was dropped before being \
fully read, or a mid-query failover exhausted its retry budget. To keep the \
connection reusable, drain each cursor (call next_batch() until it returns None) or \
call cursor.cancel() before dropping it; otherwise open a fresh Reader."
)
})
}
fn alloc_request_id(&mut self) -> i64 {
let id = self.next_request_id;
let next = self.next_request_id.wrapping_add(1);
self.next_request_id = if next <= 0 { 1 } else { next };
id
}
pub fn bytes_received(&self) -> u64 {
self.stats.bytes_received.load(Ordering::Relaxed)
}
pub fn transport_torn_down(&self) -> bool {
self.transport.is_none()
}
pub fn credit_granted_total(&self) -> u64 {
self.stats.credit_granted_total.load(Ordering::Relaxed)
}
pub fn read_ns(&self) -> u64 {
self.stats.read_ns.load(Ordering::Relaxed)
}
pub fn decode_ns(&self) -> u64 {
self.stats.decode_ns.load(Ordering::Relaxed)
}
pub fn reset_timing(&self) {
self.stats.read_ns.store(0, Ordering::Relaxed);
self.stats.decode_ns.store(0, Ordering::Relaxed);
}
pub fn stats(&self) -> &Arc<ReaderStats> {
&self.stats
}
pub fn server_info(&self) -> Option<&ServerInfo> {
self.server_info.as_ref()
}
pub fn server_version(&self) -> Result<u8> {
Ok(self.transport_ref()?.server_version())
}
pub fn symbol_dict(&self) -> &SymbolDict {
&self.dict
}
pub fn prepare<S: Into<String>>(&mut self, sql: S) -> ReaderQuery<'_> {
ReaderQuery {
reader: self,
builder: QueryRequest::builder(sql),
reset_symbol_dict: false,
on_failover_reset: None,
on_failover_progress: None,
}
}
pub fn execute<S: Into<String>>(&mut self, sql: S) -> Result<Cursor<'_>> {
self.prepare(sql).execute()
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct FailoverResetEvent {
pub failed_addr: Endpoint,
pub new_addr: Endpoint,
pub new_server_info: Option<ServerInfo>,
pub new_request_id: i64,
pub attempts: u32,
pub trigger: Error,
pub elapsed: std::time::Duration,
}
type FailoverResetCallback<'r> = Box<dyn FnMut(&FailoverResetEvent) + Send + 'r>;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum FailoverPhase {
Disconnected = 0,
Retrying = 1,
Reset = 2,
GaveUp = 3,
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct FailoverProgressEvent {
pub phase: FailoverPhase,
pub failed_addr: Endpoint,
pub new_addr: Option<Endpoint>,
pub new_server_info: Option<ServerInfo>,
pub new_request_id: Option<i64>,
pub attempt: u32,
pub trigger: Error,
pub elapsed: std::time::Duration,
pub final_error: Option<Error>,
}
type FailoverProgressCallback<'r> = Box<dyn FnMut(&FailoverProgressEvent) + Send + 'r>;
#[must_use = "ReaderQuery does nothing until you call .execute(); dropping it discards \
the prepared SQL and any binds without sending a QUERY_REQUEST"]
pub struct ReaderQuery<'r> {
reader: &'r mut Reader,
builder: QueryRequestBuilder,
reset_symbol_dict: bool,
on_failover_reset: Option<FailoverResetCallback<'r>>,
on_failover_progress: Option<FailoverProgressCallback<'r>>,
}
macro_rules! bind_method {
($name:ident, $($arg:ident : $ty:ty),*) => {
pub fn $name(mut self, $($arg : $ty),*) -> Self {
self.builder = self.builder.$name($($arg),*);
self
}
};
}
impl<'r> ReaderQuery<'r> {
pub fn initial_credit(mut self, credit: u64) -> Self {
self.builder = self.builder.initial_credit(credit);
self
}
pub fn reset_symbol_dict(mut self, reset: bool) -> Self {
self.reset_symbol_dict = reset;
self
}
pub fn on_failover_reset<F>(mut self, callback: F) -> Self
where
F: FnMut(&FailoverResetEvent) + Send + 'r,
{
self.on_failover_reset = Some(Box::new(callback));
self
}
pub fn on_failover_progress<F>(mut self, callback: F) -> Self
where
F: FnMut(&FailoverProgressEvent) + Send + 'r,
{
self.on_failover_progress = Some(Box::new(callback));
self
}
pub fn bind(mut self, value: Bind) -> Self {
self.builder = self.builder.bind(value);
self
}
bind_method!(bind_null, kind: SimpleNullKind);
bind_method!(bind_bool, v: bool);
bind_method!(bind_i8, v: i8);
bind_method!(bind_i16, v: i16);
bind_method!(bind_i32, v: i32);
bind_method!(bind_i64, v: i64);
bind_method!(bind_f32, v: f32);
bind_method!(bind_f64, v: f64);
bind_method!(bind_timestamp_micros, v: i64);
bind_method!(bind_timestamp_nanos, v: i64);
bind_method!(bind_date_millis, v: i64);
bind_method!(bind_uuid, v: [u8; 16]);
bind_method!(bind_long256, v: [u8; 32]);
bind_method!(bind_char, v: u16);
bind_method!(bind_ipv4, v: Ipv4Addr);
pub fn bind_varchar<S: Into<String>>(mut self, v: S) -> Self {
self.builder = self.builder.bind_varchar(v);
self
}
pub fn bind_decimal64(mut self, value: i64, scale: i8) -> Self {
self.builder = self.builder.bind_decimal64(value, scale);
self
}
pub fn bind_decimal128(mut self, value: i128, scale: i8) -> Self {
self.builder = self.builder.bind_decimal128(value, scale);
self
}
pub fn bind_decimal256(mut self, bytes: [u8; 32], scale: i8) -> Self {
self.builder = self.builder.bind_decimal256(bytes, scale);
self
}
pub fn bind_geohash(mut self, value: u64, precision_bits: u8) -> Self {
self.builder = self.builder.bind_geohash(value, precision_bits);
self
}
pub fn bind_binary<B: Into<Vec<u8>>>(mut self, v: B) -> Self {
self.builder = self.builder.bind_binary(v);
self
}
pub fn bind_null_varchar(mut self) -> Self {
self.builder = self.builder.bind_null_varchar();
self
}
pub fn bind_null_binary(mut self) -> Self {
self.builder = self.builder.bind_null_binary();
self
}
pub fn bind_null_decimal64(mut self, scale: i8) -> Self {
self.builder = self.builder.bind_null_decimal64(scale);
self
}
pub fn bind_null_decimal128(mut self, scale: i8) -> Self {
self.builder = self.builder.bind_null_decimal128(scale);
self
}
pub fn bind_null_decimal256(mut self, scale: i8) -> Self {
self.builder = self.builder.bind_null_decimal256(scale);
self
}
pub fn bind_null_geohash(mut self, precision_bits: u8) -> Self {
self.builder = self.builder.bind_null_geohash(precision_bits);
self
}
pub fn execute(self) -> Result<Cursor<'r>> {
if self.reader.cursor_active {
return Err(fmt!(
InvalidApiCall,
"another cursor is already in flight on this connection (only one cursor at a time per Reader)"
));
}
let request_id = self.reader.alloc_request_id();
self.reader.query_schema = None;
let server_supports_query_flags = self
.reader
.server_info()
.map(|info| has_query_flags(info.capabilities))
.unwrap_or(false);
let query_flags = if self.reset_symbol_dict && server_supports_query_flags {
QUERY_FLAG_RESET_DICT
} else {
0
};
let req = self
.builder
.request_id(request_id)
.query_flags(query_flags)
.build()?;
let credit_enabled = req.initial_credit() > 0;
let mut encoded_request = Vec::with_capacity(64);
req.encode(&mut encoded_request)?;
if encoded_request.len() < REQUEST_ID_OFFSET + 8
|| encoded_request[0] != MsgKind::QueryRequest.as_u8()
{
return Err(fmt!(
ProtocolError,
"QUERY_REQUEST encoding layout invariant violated (len={}, first={:?})",
encoded_request.len(),
encoded_request.first().copied(),
));
}
debug_assert_eq!(
i64::from_le_bytes(
encoded_request[REQUEST_ID_OFFSET..REQUEST_ID_OFFSET + 8]
.try_into()
.expect("length checked above"),
),
request_id,
"request_id at byte offset {} doesn't match the value just encoded",
REQUEST_ID_OFFSET,
);
let encoded_request: Bytes = encoded_request.into();
self.reader
.transport_mut()?
.write_message(encoded_request.clone())?;
self.reader.cursor_active = true;
let failover_budget = FailoverBudget::new(&self.reader.cfg);
Ok(Cursor {
reader: self.reader,
request_id,
last_batch: None,
terminal: None,
credit_enabled,
cancelling: false,
done: false,
terminal_error: None,
encoded_request,
on_failover_reset: self.on_failover_reset,
on_failover_progress: self.on_failover_progress,
failover_budget,
failover_resets: 0,
decode_failover_rounds: 0,
stale_plan_retries: 0,
data_delivered: false,
#[cfg(feature = "arrow-egress")]
drifted_batch: None,
#[cfg(feature = "arrow-egress")]
sym_values: crate::egress::arrow::SymbolValuesCache::default(),
#[cfg(feature = "arrow-egress")]
sym_scratch: crate::egress::arrow::SymbolBuildScratch::default(),
#[cfg(feature = "polars-egress")]
symbol_registry: None,
#[cfg(feature = "polars-egress")]
symbol_delta_modes: Vec::new(),
})
}
}
fn patch_request_id(buf: Bytes, new_rid: i64) -> Bytes {
let mut buf = match buf.try_into_mut() {
Ok(buf_mut) => buf_mut,
Err(shared) => BytesMut::from(&shared[..]),
};
buf[REQUEST_ID_OFFSET..REQUEST_ID_OFFSET + 8].copy_from_slice(&new_rid.to_le_bytes());
buf.freeze()
}
const CANCEL_DRAIN_READ_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
const MAX_DECODE_FAILOVER_ROUNDS: u32 = 1;
const MAX_STALE_PLAN_RETRIES: u32 = 15;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum StreamFailureKind {
Read,
Decode,
}
impl StreamFailureKind {
fn context(self) -> &'static str {
match self {
StreamFailureKind::Read => "mid-query frame read",
StreamFailureKind::Decode => "mid-query frame decode",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum FailoverBudgetStop {
AttemptsExhausted,
DeadlineExhausted,
}
struct FailoverBudget {
reconnect_rounds_remaining: u32,
next_backoff_ms: u64,
deadline: Option<std::time::Instant>,
}
impl FailoverBudget {
fn new(cfg: &ReaderConfig) -> Self {
let deadline = if cfg.failover_max_duration_ms == 0 {
None
} else {
Some(std::time::Instant::now() + Duration::from_millis(cfg.failover_max_duration_ms))
};
Self {
reconnect_rounds_remaining: cfg.failover_reconnect_rounds(),
next_backoff_ms: cfg.failover_backoff_initial_ms,
deadline,
}
}
fn advance_backoff(&mut self, max_backoff_ms: u64) {
if self.next_backoff_ms > 0 {
self.next_backoff_ms = self.next_backoff_ms.saturating_mul(2).min(max_backoff_ms);
}
}
fn before_reconnect_round(
&mut self,
cfg: &ReaderConfig,
rng: &mut FailoverRng,
) -> std::result::Result<(), FailoverBudgetStop> {
if self.reconnect_rounds_remaining == 0 {
return Err(FailoverBudgetStop::AttemptsExhausted);
}
let jittered_ms = rng.full_jitter_ms(self.next_backoff_ms);
let sleep_dur = match self.deadline {
Some(dl) => match dl.checked_duration_since(std::time::Instant::now()) {
Some(remaining) if !remaining.is_zero() => {
std::cmp::min(Duration::from_millis(jittered_ms), remaining)
}
_ => return Err(FailoverBudgetStop::DeadlineExhausted),
},
None => Duration::from_millis(jittered_ms),
};
std::thread::sleep(sleep_dur);
self.advance_backoff(cfg.failover_backoff_max_ms);
self.reconnect_rounds_remaining = self.reconnect_rounds_remaining.saturating_sub(1);
Ok(())
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum Terminal {
End { final_seq: u64, total_rows: u64 },
ExecDone { op_type: u8, rows_affected: u64 },
}
#[must_use = "Cursor must be drained via next_batch() or cancelled via cancel(); \
dropping mid-stream sends a best-effort CANCEL and closes the WebSocket, \
tearing down the connection for the next query on this Reader"]
pub struct Cursor<'r> {
reader: &'r mut Reader,
request_id: i64,
last_batch: Option<DecodedBatch>,
terminal: Option<Terminal>,
encoded_request: Bytes,
on_failover_reset: Option<FailoverResetCallback<'r>>,
on_failover_progress: Option<FailoverProgressCallback<'r>>,
failover_budget: FailoverBudget,
failover_resets: u32,
decode_failover_rounds: u32,
stale_plan_retries: u32,
data_delivered: bool,
credit_enabled: bool,
cancelling: bool,
done: bool,
terminal_error: Option<Error>,
#[cfg(feature = "arrow-egress")]
drifted_batch: Option<DecodedBatch>,
#[cfg(feature = "arrow-egress")]
sym_values: crate::egress::arrow::SymbolValuesCache,
#[cfg(feature = "arrow-egress")]
sym_scratch: crate::egress::arrow::SymbolBuildScratch,
#[cfg(feature = "polars-egress")]
symbol_registry: Option<crate::egress::arrow::polars::SymbolRegistry>,
#[cfg(feature = "polars-egress")]
symbol_delta_modes: Vec<bool>,
}
enum NextOutcome {
HaveBatch,
Done,
}
impl<'r> Cursor<'r> {
pub fn request_id(&self) -> i64 {
self.request_id
}
pub fn terminal(&self) -> Option<&Terminal> {
self.terminal.as_ref()
}
pub fn connection_reusable(&self) -> bool {
self.done && !self.reader.transport_torn_down()
}
pub fn credit_granted_total(&self) -> u64 {
self.reader
.stats
.credit_granted_total
.load(Ordering::Relaxed)
}
pub fn next_batch(&mut self) -> Result<Option<BatchView<'_>>> {
if self.done {
return match self.terminal_error.as_ref() {
Some(e) => Err(e.clone()),
None => Ok(None),
};
}
match self.next_batch_inner() {
Ok(NextOutcome::HaveBatch) => {
if self.last_batch.is_none() || self.reader.query_schema.is_none() {
let err = fmt!(
ProtocolError,
"internal invariant: next_batch produced a batch without a decoded view or schema"
);
self.terminate_with_close();
if self.done && self.terminal_error.is_none() {
self.terminal_error = Some(err.clone());
}
return Err(err);
}
Ok(Some(BatchView {
decoded: self.last_batch.as_ref().unwrap(),
dict: &self.reader.dict,
schema: self.reader.query_schema.as_ref().unwrap(),
}))
}
Ok(NextOutcome::Done) => Ok(None),
Err(e) => {
if self.done && self.terminal_error.is_none() {
self.terminal_error = Some(e.clone());
}
Err(e)
}
}
}
#[cfg(feature = "arrow-egress")]
pub fn as_arrow_reader<'c>(
&'c mut self,
) -> Result<crate::egress::arrow::CursorRecordBatchReader<'r, 'c>> {
crate::egress::arrow::CursorRecordBatchReader::new(self)
}
#[cfg(feature = "arrow-egress")]
pub fn fetch_all_arrow(
&mut self,
) -> Result<(arrow::datatypes::SchemaRef, Vec<arrow::array::RecordBatch>)> {
self.enable_internal_replay();
let mut reader = self.as_arrow_reader()?;
let mut resets_seen = reader.failover_resets();
let mut batches: Vec<arrow::array::RecordBatch> = Vec::new();
loop {
let Some(item) = reader.next() else { break };
let rb = item.map_err(|e| {
crate::egress::arrow::try_downcast_questdb(&e)
.cloned()
.unwrap_or_else(|| fmt!(ArrowExport, "{}", e))
})?;
let resets_now = reader.failover_resets();
if resets_now != resets_seen {
resets_seen = resets_now;
batches.clear();
}
batches.push(rb);
}
Ok((reader.schema(), batches))
}
#[cfg(feature = "polars-egress")]
pub fn iter_polars<'c>(&'c mut self) -> Result<crate::egress::arrow::CursorPolarsIter<'r, 'c>> {
crate::egress::arrow::CursorPolarsIter::new(self)
}
#[cfg(feature = "arrow-egress")]
pub fn next_arrow_batch(&mut self) -> Result<Option<arrow::array::RecordBatch>> {
self.next_arrow_batch_inner(None, false)
}
#[cfg(feature = "arrow-egress")]
#[doc(hidden)]
pub fn next_arrow_batch_inner(
&mut self,
expected_schema: Option<&arrow::datatypes::SchemaRef>,
compact: bool,
) -> Result<Option<arrow::array::RecordBatch>> {
use crate::egress::arrow::{batch_arrow_schema, batch_to_record_batch_with, schemas_equal};
use std::sync::Arc;
if self.done {
return match self.terminal_error.as_ref() {
Some(e) => Err(e.clone()),
None => Ok(None),
};
}
let decoded = if let Some(stashed) = self.drifted_batch.take() {
stashed
} else {
let outcome = match self.next_batch_inner() {
Ok(o) => o,
Err(e) => {
if self.done && self.terminal_error.is_none() {
self.terminal_error = Some(e.clone());
}
return Err(e);
}
};
match outcome {
NextOutcome::Done => return Ok(None),
NextOutcome::HaveBatch => match self.last_batch.take() {
Some(b) => b,
None => {
let e = fmt!(
ProtocolError,
"internal invariant: next_batch produced a batch without a decoded view"
);
self.stash_arrow_terminal_error(&e);
return Err(e);
}
},
}
};
let egress_schema = match self.reader.query_schema.as_ref() {
Some(s) => s.clone(),
None => {
let e = fmt!(
ProtocolError,
"internal invariant: next_batch produced a batch without a decoded schema"
);
self.stash_arrow_terminal_error(&e);
return Err(e);
}
};
let arrow_schema = match batch_arrow_schema(&egress_schema, &decoded) {
Ok(s) => Arc::new(s),
Err(e) => {
self.stash_arrow_terminal_error(&e);
return Err(e);
}
};
if let Some(expected) = expected_schema
&& !schemas_equal(expected.as_ref(), arrow_schema.as_ref())
{
let e = fmt!(
SchemaDrift,
"mid-stream Arrow schema drift: expected schema differs from batch_seq={}",
decoded.batch_seq
);
self.drifted_batch = Some(decoded);
return Err(e);
}
#[cfg(feature = "polars-egress")]
{
self.symbol_delta_modes.clear();
self.symbol_delta_modes
.extend(decoded.columns.iter().map(|c| {
matches!(
c,
crate::egress::decoder::DecodedColumn::Symbol {
local_dict: None,
..
}
)
}));
}
match batch_to_record_batch_with(
arrow_schema,
&egress_schema,
decoded,
&self.reader.dict,
&mut self.sym_values,
if compact {
Some(&mut self.sym_scratch)
} else {
None
},
) {
Ok(rb) => Ok(Some(rb)),
Err(e) => {
self.stash_arrow_terminal_error(&e);
Err(e)
}
}
}
#[cfg(feature = "polars-egress")]
pub(crate) fn symbol_registry_synced(
&mut self,
) -> Result<&crate::egress::arrow::polars::SymbolRegistry> {
let reg = self
.symbol_registry
.get_or_insert_with(crate::egress::arrow::polars::SymbolRegistry::new);
reg.sync(&self.reader.dict)?;
Ok(reg)
}
#[cfg(feature = "polars-egress")]
pub(crate) fn symbol_delta_modes(&self) -> &[bool] {
&self.symbol_delta_modes
}
#[cfg(feature = "arrow-egress")]
fn stash_arrow_terminal_error(&mut self, err: &Error) {
self.done = true;
if self.terminal_error.is_none() {
self.terminal_error = Some(err.clone());
}
}
fn next_batch_inner(&mut self) -> Result<NextOutcome> {
loop {
let (header, payload) = match self.read_frame_raw() {
Ok(hp) => hp,
Err(e) => {
self.failover_after_stream_failure(e, StreamFailureKind::Read)?;
continue;
}
};
let wire_bytes = HEADER_LEN as u64 + header.payload_length as u64;
let t1 = std::time::Instant::now();
let decode_result = decode_frame(
header,
&payload,
&mut self.reader.dict,
&mut self.reader.query_schema,
&mut self.reader.zstd_scratch,
);
self.reader.stats.decode_ns.fetch_add(
u64::try_from(t1.elapsed().as_nanos()).unwrap_or(u64::MAX),
Ordering::Relaxed,
);
let event = match decode_result {
Ok(ev) => ev,
Err(e) => {
self.failover_after_stream_failure(e, StreamFailureKind::Decode)?;
continue;
}
};
match event {
ServerEvent::Batch(b) => {
if b.request_id != self.request_id {
let err = fmt!(
ProtocolError,
"RESULT_BATCH request_id {} != cursor {}",
b.request_id,
self.request_id
);
self.terminate_with_close();
return Err(err);
}
if self.credit_enabled
&& !self.cancelling
&& let Err(e) = self.send_credit_frame(wire_bytes)
{
self.terminate_with_close();
return Err(e);
}
if self.reader.query_schema.is_none() {
let err = fmt!(ProtocolError, "RESULT_BATCH decoded without a schema");
self.terminate_with_close();
return Err(err);
}
let last = self.last_batch.insert(b);
self.data_delivered = true;
let _ = last;
return Ok(NextOutcome::HaveBatch);
}
ServerEvent::End {
request_id,
final_seq,
total_rows,
} => {
if let Err(e) = self.check_rid(request_id, "RESULT_END") {
self.terminate_with_close();
return Err(e);
}
self.terminal = Some(Terminal::End {
final_seq,
total_rows,
});
self.reader.cursor_active = false;
self.done = true;
return Ok(NextOutcome::Done);
}
ServerEvent::ExecDone {
request_id,
op_type,
rows_affected,
} => {
if let Err(e) = self.check_rid(request_id, "EXEC_DONE") {
self.terminate_with_close();
return Err(e);
}
self.terminal = Some(Terminal::ExecDone {
op_type,
rows_affected,
});
self.reader.cursor_active = false;
self.done = true;
return Ok(NextOutcome::Done);
}
ServerEvent::Error {
request_id,
status,
message,
} => {
if let Err(e) = self.check_rid(request_id, "QUERY_ERROR") {
self.terminate_with_close();
return Err(e);
}
if !self.cancelling
&& !self.data_delivered
&& self.stale_plan_retries < MAX_STALE_PLAN_RETRIES
&& is_stale_plan_error(status, &message)
{
self.stale_plan_retries = self.stale_plan_retries.saturating_add(1);
match self.replay_query_same_connection() {
Ok(()) => continue,
Err(e) => {
self.reader.cursor_active = false;
self.done = true;
return Err(e);
}
}
}
self.reader.cursor_active = false;
self.done = true;
return Err(map_server_status(status, message));
}
ServerEvent::CacheReset { .. } => {
self.reset_symbol_caches();
continue;
}
ServerEvent::ServerInfo(_) => {
continue;
}
}
}
}
pub fn failover_resets(&self) -> u32 {
self.failover_resets
}
pub fn stale_plan_retries(&self) -> u32 {
self.stale_plan_retries
}
#[cfg(feature = "arrow-egress")]
pub(crate) fn enable_internal_replay(&mut self) {
if self.on_failover_reset.is_none() {
self.on_failover_reset = Some(Box::new(|_: &FailoverResetEvent| {}));
}
}
pub fn current_addr(&self) -> &Endpoint {
self.reader.current_addr()
}
pub fn server_version(&self) -> Result<u8> {
self.reader.server_version()
}
pub fn server_info(&self) -> Option<&ServerInfo> {
self.reader.server_info()
}
fn read_frame_raw(
&mut self,
) -> Result<(crate::egress::wire::header::FrameHeader, bytes::Bytes)> {
let t0 = std::time::Instant::now();
let (header, payload) = self.reader.transport_mut()?.read_frame()?;
self.reader.stats.read_ns.fetch_add(
u64::try_from(t0.elapsed().as_nanos()).unwrap_or(u64::MAX),
Ordering::Relaxed,
);
let wire_bytes = HEADER_LEN as u64 + header.payload_length as u64;
self.reader
.stats
.bytes_received
.fetch_add(wire_bytes, Ordering::Relaxed);
Ok((header, payload))
}
fn failover_after_stream_failure(&mut self, e: Error, kind: StreamFailureKind) -> Result<()> {
if self.cancelling || !self.reader.cfg.failover || !is_failover_eligible(e.code()) {
self.terminate_with_close();
return Err(e);
}
if would_silently_duplicate(self.data_delivered, self.on_failover_reset.is_some()) {
let err = fmt!(
FailoverWouldDuplicate,
"mid-query failover would replay rows already delivered to the caller \
(install on_failover_reset to authorize replay); \
cursor terminated. Trigger: {} ({:?})",
e.msg(),
e.code()
);
self.terminate_with_close();
return Err(err);
}
if kind == StreamFailureKind::Decode {
if self.decode_failover_rounds >= MAX_DECODE_FAILOVER_ROUNDS {
self.terminate_with_close();
return Err(e);
}
self.decode_failover_rounds = self.decode_failover_rounds.saturating_add(1);
}
warn_on_protocol_error_failover(&e, kind.context());
self.failover_reconnect_and_replay(e)
}
fn replay_query_same_connection(&mut self) -> Result<()> {
let new_rid = self.reader.alloc_request_id();
self.request_id = new_rid;
self.encoded_request = patch_request_id(std::mem::take(&mut self.encoded_request), new_rid);
self.reader.query_schema = None;
self.last_batch = None;
#[cfg(feature = "arrow-egress")]
{
self.drifted_batch = None;
}
self.reader
.transport_mut()
.and_then(|t| t.write_message(self.encoded_request.clone()))
}
fn reset_symbol_caches(&mut self) {
#[cfg(feature = "arrow-egress")]
{
self.sym_values = crate::egress::arrow::SymbolValuesCache::default();
self.sym_scratch = crate::egress::arrow::SymbolBuildScratch::default();
}
#[cfg(feature = "polars-egress")]
{
self.symbol_registry = None;
}
}
fn failover_reconnect_and_replay(&mut self, trigger: Error) -> Result<()> {
let mut trigger = trigger;
loop {
let started = std::time::Instant::now();
let failed_idx = self.reader.addr_idx;
let failed_addr = self.reader.cfg.addrs[failed_idx].clone();
if let Some(cb) = self.on_failover_progress.as_mut() {
let event = FailoverProgressEvent {
phase: FailoverPhase::Disconnected,
failed_addr: failed_addr.clone(),
new_addr: None,
new_server_info: None,
new_request_id: None,
attempt: 0,
trigger: trigger.clone(),
elapsed: started.elapsed(),
final_error: None,
};
cb(&event);
}
let mut last_attempt: u32 = 0;
let reconnect_result = {
let Self {
reader,
on_failover_progress,
failover_budget,
..
} = self;
let failed_addr_ref = &failed_addr;
let trigger_ref = &trigger;
reader.reconnect_with_failover(failed_idx, failover_budget, &mut |attempt: u32| {
last_attempt = attempt;
if let Some(cb) = on_failover_progress.as_mut() {
let event = FailoverProgressEvent {
phase: FailoverPhase::Retrying,
failed_addr: failed_addr_ref.clone(),
new_addr: None,
new_server_info: None,
new_request_id: None,
attempt,
trigger: trigger_ref.clone(),
elapsed: started.elapsed(),
final_error: None,
};
cb(&event);
}
})
};
let attempts = match reconnect_result {
Ok(n) => n,
Err(e) => {
if let Some(cb) = self.on_failover_progress.as_mut() {
let event = FailoverProgressEvent {
phase: FailoverPhase::GaveUp,
failed_addr: failed_addr.clone(),
new_addr: None,
new_server_info: None,
new_request_id: None,
attempt: last_attempt,
trigger: trigger.clone(),
elapsed: started.elapsed(),
final_error: Some(e.clone()),
};
cb(&event);
}
self.reader.cursor_active = false;
self.done = true;
return Err(if prefer_over_trigger(e.code()) {
e
} else {
trigger
});
}
};
self.last_batch = None;
#[cfg(feature = "arrow-egress")]
{
self.drifted_batch = None;
}
self.reset_symbol_caches();
let new_rid = self.reader.alloc_request_id();
self.request_id = new_rid;
self.encoded_request =
patch_request_id(std::mem::take(&mut self.encoded_request), new_rid);
match self
.reader
.transport_mut()
.and_then(|t| t.write_message(self.encoded_request.clone()))
{
Ok(()) => {
self.failover_resets = self.failover_resets.saturating_add(1);
let new_addr = self.reader.cfg.addrs[self.reader.addr_idx].clone();
let new_server_info = self.reader.server_info.clone();
if let Some(cb) = self.on_failover_progress.as_mut() {
let event = FailoverProgressEvent {
phase: FailoverPhase::Reset,
failed_addr: failed_addr.clone(),
new_addr: Some(new_addr.clone()),
new_server_info: new_server_info.clone(),
new_request_id: Some(new_rid),
attempt: attempts,
trigger: trigger.clone(),
elapsed: started.elapsed(),
final_error: None,
};
cb(&event);
}
if let Some(cb) = self.on_failover_reset.as_mut() {
let event = FailoverResetEvent {
failed_addr,
new_addr,
new_server_info,
new_request_id: new_rid,
attempts,
trigger,
elapsed: started.elapsed(),
};
cb(&event);
}
return Ok(());
}
Err(e) => {
warn_on_protocol_error_failover(&e, "replay query write");
if !self.reader.cfg.failover || !is_failover_eligible(e.code()) {
if let Some(cb) = self.on_failover_progress.as_mut() {
let event = FailoverProgressEvent {
phase: FailoverPhase::GaveUp,
failed_addr: failed_addr.clone(),
new_addr: None,
new_server_info: None,
new_request_id: None,
attempt: attempts,
trigger: trigger.clone(),
elapsed: started.elapsed(),
final_error: Some(e.clone()),
};
cb(&event);
}
if let Some(dead) = self.reader.transport.take() {
drop(dead);
}
self.reader.cursor_active = false;
self.done = true;
return Err(e);
}
trigger = e;
continue;
}
}
}
}
pub fn cancel(&mut self) -> Result<()> {
if self.done {
return Ok(());
}
self.cancelling = true;
let mut payload = Vec::with_capacity(9);
payload.push(MsgKind::Cancel.as_u8());
payload.extend_from_slice(&self.request_id.to_le_bytes());
let write_outcome = match self.reader.transport_mut() {
Ok(t) => t.write_message(Bytes::from(payload)),
Err(e) => Err(e),
};
if let Err(e) = write_outcome {
self.terminate_with_close();
return Err(e);
}
if let Some(t) = self.reader.transport.as_mut() {
t.set_read_timeout(Some(CANCEL_DRAIN_READ_TIMEOUT));
t.set_write_timeout(Some(CLOSE_TIMEOUT));
}
if self.credit_enabled {
let _ = self.write_credit_frame_raw(1);
}
let mut drain_result: Result<()> = Ok(());
while !self.done {
match self.next_batch() {
Ok(Some(_)) => {} Ok(None) => break,
Err(e) => {
if matches!(e.code(), crate::ErrorCode::Cancelled) {
break;
}
drain_result = Err(e);
break;
}
}
}
if let Some(t) = self.reader.transport.as_mut() {
t.set_read_timeout(None);
t.set_write_timeout(Some(WRITE_TIMEOUT));
}
drain_result
}
pub fn add_credit(&mut self, additional_bytes: u64) -> Result<()> {
if self.done {
return Err(match self.terminal_error.as_ref() {
Some(e) => e.clone(),
None => fmt!(InvalidApiCall, "cursor is terminal; add_credit not allowed"),
});
}
let first_err = match self.send_credit_frame(additional_bytes) {
Ok(()) => return Ok(()),
Err(e) => e,
};
if self.cancelling || !self.reader.cfg.failover || !is_failover_eligible(first_err.code()) {
self.terminate_with_close();
return Err(first_err);
}
if would_silently_duplicate(self.data_delivered, self.on_failover_reset.is_some()) {
let err = fmt!(
FailoverWouldDuplicate,
"mid-query failover would replay rows already delivered to the caller \
(install on_failover_reset to authorize replay); \
cursor terminated. Trigger: {} ({:?})",
first_err.msg(),
first_err.code()
);
self.terminate_with_close();
return Err(err);
}
warn_on_protocol_error_failover(&first_err, "add_credit write");
self.failover_reconnect_and_replay(first_err)?;
match self.send_credit_frame(additional_bytes) {
Ok(()) => Ok(()),
Err(e) => {
self.terminate_with_close();
Err(e)
}
}
}
fn send_credit_frame(&mut self, additional_bytes: u64) -> Result<()> {
self.write_credit_frame_raw(additional_bytes)?;
self.reader
.stats
.credit_granted_total
.fetch_add(additional_bytes, Ordering::Relaxed);
Ok(())
}
fn write_credit_frame_raw(&mut self, additional_bytes: u64) -> Result<()> {
let mut payload = Vec::with_capacity(16);
payload.push(MsgKind::Credit.as_u8());
payload.extend_from_slice(&self.request_id.to_le_bytes());
varint::encode_u64(additional_bytes, &mut payload);
self.reader
.transport_mut()?
.write_message(Bytes::from(payload))?;
Ok(())
}
fn check_rid(&self, got: i64, what: &str) -> Result<()> {
if got != self.request_id {
return Err(fmt!(
ProtocolError,
"{} request_id {} != cursor {}",
what,
got,
self.request_id
));
}
Ok(())
}
fn terminate_with_close(&mut self) {
if let Some(mut t) = self.reader.transport.take() {
t.close_in_place();
drop(t);
}
self.reader.cursor_active = false;
self.done = true;
}
}
impl Drop for Cursor<'_> {
fn drop(&mut self) {
if self.reader.cursor_active {
if let Some(mut t) = self.reader.transport.take() {
if !self.cancelling {
t.try_write_cancel(self.request_id);
}
t.close_in_place();
drop(t);
}
self.reader.cursor_active = false;
}
}
}
#[must_use = "BatchView is a borrowed projection; dropping it without iterating \
the rows or calling its accessors throws away the just-decoded batch"]
pub struct BatchView<'c> {
decoded: &'c DecodedBatch,
dict: &'c SymbolDict,
schema: &'c Schema,
}
impl<'c> BatchView<'c> {
pub fn request_id(&self) -> i64 {
self.decoded.request_id
}
pub fn batch_seq(&self) -> u64 {
self.decoded.batch_seq
}
pub fn flags(&self) -> u8 {
self.decoded.flags
}
pub fn schema(&self) -> &'c Schema {
self.schema
}
pub fn row_count(&self) -> usize {
self.decoded.row_count
}
pub fn column_count(&self) -> usize {
self.decoded.columns.len()
}
pub fn column(&self, idx: usize) -> Result<ColumnView<'_>> {
self.decoded.column_view(idx, self.dict)
}
pub fn dict(&self) -> &'c SymbolDict {
self.dict
}
}
fn would_silently_duplicate(data_delivered: bool, has_reset_callback: bool) -> bool {
data_delivered && !has_reset_callback
}
fn is_failover_eligible(code: ErrorCode) -> bool {
matches!(
code,
ErrorCode::SocketError
| ErrorCode::ConnectTimeout
| ErrorCode::HandshakeError
| ErrorCode::TlsError
| ErrorCode::ProtocolError
| ErrorCode::CouldNotResolveAddr
| ErrorCode::RoleMismatch
)
}
fn warn_on_protocol_error_failover(err: &Error, context: &str) {
if err.code() == ErrorCode::ProtocolError {
log::warn!(
"ProtocolError triggered failover ({}): {} — \
reconnecting may mask transient wire-frame corruption \
(truncated frames, malformed varints) or a deterministic \
protocol violation; check server logs if this recurs.",
context,
err.msg()
);
}
}
fn prefer_over_trigger(code: ErrorCode) -> bool {
matches!(
code,
ErrorCode::AuthError
| ErrorCode::RoleMismatch
| ErrorCode::ConfigError
| ErrorCode::UnsupportedServer
| ErrorCode::HandshakeError
| ErrorCode::TlsError
)
}
#[derive(Debug)]
pub(crate) struct FailoverRng {
state: u64,
}
impl FailoverRng {
pub(crate) fn new() -> Self {
use std::sync::atomic::{AtomicU64, Ordering};
static COUNTER: AtomicU64 = AtomicU64::new(0);
let now_ns = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(0);
let bump = COUNTER.fetch_add(1, Ordering::Relaxed);
Self {
state: now_ns ^ bump.wrapping_mul(0x9E37_79B9_7F4A_7C15),
}
}
fn next_u64(&mut self) -> u64 {
self.state = self.state.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = self.state;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
pub(crate) fn full_jitter_ms(&mut self, base: u64) -> u64 {
if base == 0 {
return 0;
}
self.next_u64() % base
}
}
fn target_matches(target: Target, role: ServerRole) -> bool {
match target {
Target::Any => true,
Target::Primary => matches!(
role,
ServerRole::Primary | ServerRole::PrimaryCatchup | ServerRole::Standalone
),
Target::Replica => matches!(role, ServerRole::Replica),
}
}
struct TransportSession {
idx: usize,
transport: WsTransport,
server_info: Option<ServerInfo>,
}
struct WalkOutcome {
session: TransportSession,
dials: u32,
}
fn walk_via_tracker(
tracker: &mut HostHealthTracker,
cfg: &Arc<ReaderConfig>,
allow_reset_pass: bool,
terminal_codes: &[ErrorCode],
) -> Result<WalkOutcome> {
tracker.begin_round(false);
let mut last_role_mismatch: Option<Error> = None;
let mut last_transport_err: Option<Error> = None;
let mut retried_after_reset = false;
let mut dials: u32 = 0;
loop {
let idx = match tracker.pick_next() {
Some(i) => i,
None => {
if allow_reset_pass && !retried_after_reset {
tracker.begin_round(true);
retried_after_reset = true;
continue;
}
break;
}
};
dials = dials.saturating_add(1);
match Reader::connect_endpoint(cfg.as_ref(), idx) {
Ok(session) => {
if let Some(info) = session.server_info.as_ref() {
tracker.record_zone(idx, info.zone_id.as_deref());
}
tracker.record_success(idx);
return Ok(WalkOutcome { session, dials });
}
Err(e) => {
let code = e.code();
if terminal_codes.contains(&code) {
return Err(e);
}
match code {
ErrorCode::RoleMismatch => {
let reject = e.upgrade_reject();
let transient = reject.is_some_and(|r| r.is_transient());
if let Some(r) = reject {
tracker.record_zone(idx, r.zone.as_deref());
}
tracker.record_role_reject(idx, transient);
last_role_mismatch = Some(e);
}
_ => {
tracker.record_transport_error(idx);
last_transport_err = Some(e);
}
}
}
}
}
if let Some(e) = last_role_mismatch {
return Err(e);
}
Err(last_transport_err
.unwrap_or_else(|| fmt!(SocketError, "all {} endpoints unreachable", cfg.addrs.len())))
}
fn read_server_info_frame(transport: &mut WsTransport, timeout: Duration) -> Result<ServerInfo> {
transport.set_read_timeout(Some(timeout));
let result = transport.read_frame();
transport.set_read_timeout(None);
let (header, payload) = result?;
let mut dict = SymbolDict::new();
let mut query_schema: Option<Schema> = None;
let mut zstd_scratch = ZstdScratch::new();
let event = decode_frame(
header,
&payload,
&mut dict,
&mut query_schema,
&mut zstd_scratch,
)?;
match event {
ServerEvent::ServerInfo(info) => Ok(info),
other => Err(fmt!(
ProtocolError,
"expected SERVER_INFO as the first frame, got {:?}",
std::mem::discriminant(&other)
)),
}
}
const STALE_PLAN_PATTERNS: [&str; 2] = [
"cached query plan cannot be used",
"table schema has changed",
];
fn is_stale_plan_error(status: crate::egress::wire::msg_kind::StatusCode, message: &str) -> bool {
use crate::egress::wire::msg_kind::StatusCode as S;
if status != S::InternalError {
return false;
}
let lower = message.to_ascii_lowercase();
STALE_PLAN_PATTERNS.iter().any(|p| lower.contains(p))
}
fn map_server_status(
status: crate::egress::wire::msg_kind::StatusCode,
message: String,
) -> crate::Error {
use crate::ErrorCode as C;
use crate::egress::wire::msg_kind::StatusCode as S;
let code = match status {
S::SchemaMismatch => C::ServerSchemaMismatch,
S::ParseError => C::ServerParseError,
S::InternalError => C::ServerInternalError,
S::SecurityError => C::ServerSecurityError,
S::Cancelled => C::Cancelled,
S::LimitExceeded => C::ServerLimitExceeded,
};
crate::Error::new(code, message)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn reader_stats_arc_clones_share_storage() {
let stats = Arc::new(ReaderStats::default());
let alias = Arc::clone(&stats);
stats.bytes_received.fetch_add(42, Ordering::Relaxed);
stats.credit_granted_total.fetch_add(7, Ordering::Relaxed);
stats.read_ns.fetch_add(1_000, Ordering::Relaxed);
stats.decode_ns.fetch_add(500, Ordering::Relaxed);
assert_eq!(alias.bytes_received.load(Ordering::Relaxed), 42);
assert_eq!(alias.credit_granted_total.load(Ordering::Relaxed), 7);
assert_eq!(alias.read_ns.load(Ordering::Relaxed), 1_000);
assert_eq!(alias.decode_ns.load(Ordering::Relaxed), 500);
alias.read_ns.store(0, Ordering::Relaxed);
alias.decode_ns.store(0, Ordering::Relaxed);
assert_eq!(stats.read_ns.load(Ordering::Relaxed), 0);
assert_eq!(stats.decode_ns.load(Ordering::Relaxed), 0);
}
#[test]
fn request_id_offset_matches_encoder_layout() {
const RID: i64 = 0x0123_4567_89AB_CDEF;
let req = QueryRequest::builder("SELECT 1")
.request_id(RID)
.build()
.expect("build");
let mut buf = Vec::new();
req.encode(&mut buf).expect("encode");
assert!(buf.len() >= REQUEST_ID_OFFSET + 8);
assert_eq!(buf[0], MsgKind::QueryRequest.as_u8());
let mut id_bytes = [0u8; 8];
id_bytes.copy_from_slice(&buf[REQUEST_ID_OFFSET..REQUEST_ID_OFFSET + 8]);
assert_eq!(i64::from_le_bytes(id_bytes), RID);
}
#[test]
fn patch_request_id_preserves_body_and_updates_id() {
const OLD_RID: i64 = 0x1111_2222_3333_4444;
const NEW_RID: i64 = 0x5555_6666_7777_8888;
let req = QueryRequest::builder("SELECT * FROM big_table WHERE x > $1")
.request_id(OLD_RID)
.build()
.expect("build");
let mut original = Vec::with_capacity(64);
req.encode(&mut original).expect("encode");
let original = Bytes::from(original);
let patched = patch_request_id(original.clone(), NEW_RID);
assert_eq!(patched[0], MsgKind::QueryRequest.as_u8());
let mut id_bytes = [0u8; 8];
id_bytes.copy_from_slice(&patched[REQUEST_ID_OFFSET..REQUEST_ID_OFFSET + 8]);
assert_eq!(i64::from_le_bytes(id_bytes), NEW_RID);
assert_eq!(
&patched[..REQUEST_ID_OFFSET],
&original[..REQUEST_ID_OFFSET]
);
assert_eq!(
&patched[REQUEST_ID_OFFSET + 8..],
&original[REQUEST_ID_OFFSET + 8..]
);
let _hold = patched.clone();
let patched_again = patch_request_id(patched, OLD_RID);
let mut id_bytes = [0u8; 8];
id_bytes.copy_from_slice(&patched_again[REQUEST_ID_OFFSET..REQUEST_ID_OFFSET + 8]);
assert_eq!(i64::from_le_bytes(id_bytes), OLD_RID);
}
#[test]
fn is_failover_eligible_matrix() {
use ErrorCode::*;
for code in [
SocketError,
ConnectTimeout,
HandshakeError,
TlsError,
ProtocolError,
CouldNotResolveAddr,
RoleMismatch,
] {
assert!(
is_failover_eligible(code),
"{:?} must be failover-eligible",
code
);
}
for code in [
ConfigError,
InvalidApiCall,
AuthError,
UnsupportedServer,
InvalidUtf8,
InvalidBind,
ServerSchemaMismatch,
ServerParseError,
ServerInternalError,
ServerSecurityError,
LimitExceeded,
ServerLimitExceeded,
Cancelled,
] {
assert!(
!is_failover_eligible(code),
"{:?} must NOT be failover-eligible",
code
);
}
}
#[test]
fn map_server_status_matrix() {
use crate::egress::wire::msg_kind::StatusCode as S;
use ErrorCode as C;
let cases: &[(S, C)] = &[
(S::SchemaMismatch, C::ServerSchemaMismatch),
(S::ParseError, C::ServerParseError),
(S::InternalError, C::ServerInternalError),
(S::SecurityError, C::ServerSecurityError),
(S::Cancelled, C::Cancelled),
(S::LimitExceeded, C::ServerLimitExceeded),
];
for (status, expected_code) in cases {
let err = map_server_status(*status, "msg".to_string());
assert_eq!(
err.code(),
*expected_code,
"status {:?} should map to {:?}",
status,
expected_code
);
assert_eq!(err.msg(), "msg");
}
let mut seen = std::collections::HashSet::new();
for (_, code) in cases {
assert!(
seen.insert(*code),
"ErrorCode {:?} mapped from two distinct StatusCode values",
code
);
}
}
#[test]
fn prefer_over_trigger_matrix() {
use ErrorCode::*;
for code in [
AuthError,
RoleMismatch,
ConfigError,
UnsupportedServer,
HandshakeError,
TlsError,
] {
assert!(
prefer_over_trigger(code),
"{:?} must be preferred over the trigger",
code
);
}
for code in [
SocketError,
ProtocolError,
CouldNotResolveAddr,
InvalidApiCall,
InvalidUtf8,
InvalidBind,
ServerInternalError,
Cancelled,
] {
assert!(
!prefer_over_trigger(code),
"{:?} must NOT be preferred over the trigger",
code
);
}
}
#[test]
fn failover_budget_backoff_base_grows_and_caps() {
let mut budget = FailoverBudget {
reconnect_rounds_remaining: 9,
next_backoff_ms: 10,
deadline: None,
};
let observed: [u64; 9] = std::array::from_fn(|_| {
let base = budget.next_backoff_ms;
budget.advance_backoff(20);
base
});
assert_eq!(observed, [10, 20, 20, 20, 20, 20, 20, 20, 20]);
let mut disabled = FailoverBudget {
reconnect_rounds_remaining: 1,
next_backoff_ms: 0,
deadline: None,
};
disabled.advance_backoff(20);
assert_eq!(disabled.next_backoff_ms, 0);
}
#[test]
fn before_reconnect_round_applies_configured_backoff_cap() {
let cfg = ReaderConfig::from_conf(concat!(
"ws::addr=localhost:9000;",
"failover_max_attempts=4;",
"failover_backoff_initial_ms=1;",
"failover_backoff_max_ms=2"
))
.unwrap();
let mut budget = FailoverBudget::new(&cfg);
let mut rng = FailoverRng { state: 0 };
let observed: [u64; 3] = std::array::from_fn(|_| {
budget.before_reconnect_round(&cfg, &mut rng).unwrap();
budget.next_backoff_ms
});
assert_eq!(observed, [2, 2, 2]);
assert_eq!(budget.reconnect_rounds_remaining, 0);
assert_eq!(
budget.before_reconnect_round(&cfg, &mut rng),
Err(FailoverBudgetStop::AttemptsExhausted)
);
let disabled_cfg = ReaderConfig::from_conf(concat!(
"ws::addr=localhost:9000;",
"failover_max_attempts=2;",
"failover_backoff_initial_ms=0;",
"failover_backoff_max_ms=2"
))
.unwrap();
let mut disabled = FailoverBudget::new(&disabled_cfg);
let mut disabled_rng = FailoverRng { state: 0 };
disabled
.before_reconnect_round(&disabled_cfg, &mut disabled_rng)
.unwrap();
assert_eq!(disabled.next_backoff_ms, 0);
assert_eq!(disabled.reconnect_rounds_remaining, 0);
}
#[test]
fn full_jitter_ms_zero_base_returns_zero() {
let mut rng = FailoverRng::new();
for _ in 0..32 {
assert_eq!(rng.full_jitter_ms(0), 0);
}
}
#[test]
fn full_jitter_ms_draws_are_in_range() {
let mut rng = FailoverRng::new();
for &base in &[1u64, 2, 80, 100, 1_000, 65_537, u32::MAX as u64] {
for _ in 0..10_000 {
let d = rng.full_jitter_ms(base);
assert!(
d < base,
"full_jitter_ms({}) returned {}, which is >= base \
(full-jitter draws must be in [0, base))",
base,
d
);
}
}
}
#[test]
fn full_jitter_ms_distribution_covers_full_range() {
let mut rng = FailoverRng::new();
let mut saw_low = false;
let mut saw_high = false;
for _ in 0..10_000 {
let d = rng.full_jitter_ms(100);
if d < 10 {
saw_low = true;
}
if d >= 90 {
saw_high = true;
}
if saw_low && saw_high {
break;
}
}
assert!(
saw_low,
"expected at least one draw < 10 out of 10k samples"
);
assert!(
saw_high,
"expected at least one draw >= 90 out of 10k samples"
);
}
#[test]
fn would_silently_duplicate_truth_table() {
assert!(!would_silently_duplicate(false, false));
assert!(!would_silently_duplicate(false, true));
assert!(would_silently_duplicate(true, false));
assert!(!would_silently_duplicate(true, true));
}
}