use gnitz_expr::SchemaFacts;
use std::borrow::Cow;
use std::collections::VecDeque;
use std::future::Future;
use std::num::NonZeroU64;
use std::os::fd::{BorrowedFd, RawFd};
use std::pin::Pin;
use std::sync::{Arc, Mutex, MutexGuard, PoisonError};
use std::task::{Context, Poll, Waker};
use std::time::{Duration, Instant};
use crate::error::ClientError;
use crate::protocol::wal_block::decode_wal_block_into;
use crate::{
encode_ddl_txn, encode_frame, encode_push_txn, sys_schema, ClientTransport, ProtocolError, PushFamily, Schema,
ZSetBatch,
};
use gnitz_wire::control::{peek_control_block, ControlHeader, DecodedControl};
use gnitz_wire::txn_frame;
use gnitz_wire::CONNECT_TIMEOUT;
use gnitz_wire::{ClientVerb, WireConflictMode};
use gnitz_wire::{RelClass, RelDescriptorBlob, RelIndex, WireFault, WireStatus};
pub const MAX_IN_FLIGHT: usize = 4096;
pub const MAX_QUEUED_BYTES: usize = 64 << 20;
#[derive(Debug)]
pub struct ScanReply {
pub schema: Arc<Schema>,
pub batch: ZSetBatch,
pub lsn: Option<u64>,
}
#[derive(Debug)]
pub struct RelDescriptor {
pub tid: u64,
pub class: RelClass,
pub pk_repeats: bool,
pub serial: bool,
pub schema: Arc<Schema>,
pub indexes: Vec<RelIndex>,
pub token: u64,
}
pub use gnitz_wire::control::Target;
impl From<&RelDescriptor> for Target {
fn from(rel: &RelDescriptor) -> Self {
Target { tid: rel.tid, token: rel.token }
}
}
#[derive(Debug)]
pub(crate) struct RawBlock {
frame: Vec<u8>,
block: std::ops::Range<usize>,
}
impl RawBlock {
pub(crate) fn block(&self) -> &[u8] {
&self.frame[self.block.clone()]
}
}
#[derive(Copy, Clone, PartialEq, Eq, Debug)]
pub struct DeltaCursor {
pub tag: u64,
pub tick: NonZeroU64,
}
impl DeltaCursor {
pub fn from_pair(tag: u64, tick: u64) -> Option<DeltaCursor> {
NonZeroU64::new(tick).map(|tick| DeltaCursor { tag, tick })
}
pub fn pair(self) -> (u64, u64) {
(self.tag, self.tick.get())
}
pub(crate) fn advanced_to(self, next: DeltaCursor) -> Result<DeltaCursor, ClientError> {
(self.tag == next.tag)
.then_some(next)
.ok_or_else(|| delta_expired("delta cursor's tag names a different boot or relation; bootstrap"))
}
}
fn delta_expired(text: &str) -> ClientError {
ClientError::Refused(WireFault {
status: WireStatus::DeltaExpired,
text: text.into(),
})
}
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub struct Interest {
pub read: bool,
pub write: bool,
}
impl Interest {
pub const NONE: Interest = Interest { read: false, write: false };
pub const READ: Interest = Interest { read: true, write: false };
pub const WRITE: Interest = Interest { read: false, write: true };
pub const BOTH: Interest = Interest { read: true, write: true };
pub fn is_empty(self) -> bool {
self == Interest::NONE
}
pub fn poll_events(self) -> libc::c_short {
(if self.read { libc::POLLIN } else { 0 }) | (if self.write { libc::POLLOUT } else { 0 })
}
pub fn from_revents(revents: libc::c_short) -> Interest {
Interest {
read: revents & (libc::POLLIN | libc::POLLHUP | libc::POLLERR) != 0,
write: revents & libc::POLLOUT != 0,
}
}
}
pub enum Request<'a> {
AllocIds(u64),
AllocSerial { table: Target, count: u64 },
DdlTxn(&'a [(u64, ZSetBatch)]),
PushTxn { families: &'a [PushFamily] },
Push {
target: Target,
schema: &'a Schema,
batch: &'a ZSetBatch,
mode: WireConflictMode,
},
}
impl Request<'_> {
fn encode(self) -> Result<(Vec<u8>, u64), ClientError> {
Ok(match self {
Request::AllocIds(n) => {
let hdr = ControlHeader::naming(ClientVerb::AllocIds, Target::from(0), n);
(encode_frame(hdr, &[], None, None), 0)
}
Request::AllocSerial { table, count } => {
let hdr = ControlHeader::naming(ClientVerb::AllocSerialRange, table, count);
(encode_frame(hdr, &[], None, None), table.tid)
}
Request::DdlTxn(families) => {
for (tid, batch) in families {
if gnitz_wire::sys_family_index(*tid).is_none() {
return Err(ClientError::from(format!("DDL family {tid} is not a system table")));
}
batch.validate(sys_schema(*tid))?;
}
(encode_ddl_txn(families), 0)
}
Request::PushTxn { families } => {
for f in families {
f.batch.validate(&f.schema)?;
}
(encode_push_txn(families), 0)
}
Request::Push { target, schema, batch, mode } => {
batch.validate(schema)?;
let mut hdr = ControlHeader::naming(ClientVerb::Push, target, 0);
hdr.flags.conflict_mode = mode;
(
encode_frame(hdr, &[], Some(&schema.to_block()), Some(batch)),
target.tid,
)
}
})
}
}
enum Slot {
Ack { tid: u64, to: Promise<u64> },
Resolve { to: Promise<Option<Arc<RelDescriptor>>> },
Scan {
tid: u64,
reply_schema: Arc<Schema>,
data: Option<ZSetBatch>,
to: Promise<ScanReply>,
},
Multi {
rels: Vec<(u64, Arc<Schema>)>,
replies: Vec<ScanReply>,
data: Option<ZSetBatch>,
to: Promise<Vec<ScanReply>>,
},
DeltaPoll {
views: Vec<u64>,
at: usize,
poll: u64,
},
}
#[derive(Default)]
enum Awaited<T> {
#[default]
Outstanding,
Arrived(Result<T, ClientError>),
Routed(Box<dyn FnOnce(Result<T, ClientError>) + Send>),
Watched(Waker),
}
struct Cell<T>(Mutex<Awaited<T>>);
impl<T> Cell<T> {
fn lock(&self) -> MutexGuard<'_, Awaited<T>> {
self.0.lock().unwrap_or_else(PoisonError::into_inner)
}
}
pub struct Sent<T>(Arc<Cell<T>>);
pub struct Promise<T>(Option<Arc<Cell<T>>>);
pub fn promise<T>() -> (Promise<T>, Sent<T>) {
let cell = Arc::new(Cell(Mutex::new(Awaited::Outstanding)));
(Promise(Some(Arc::clone(&cell))), Sent(cell))
}
impl<T> Promise<T> {
pub fn fulfil(&mut self, reply: Result<T, ClientError>) {
let Some(cell) = self.0.take() else {
return;
};
let mut state = cell.lock();
match std::mem::take(&mut *state) {
Awaited::Routed(route) => {
drop(state);
route(reply)
}
Awaited::Watched(waiter) => {
*state = Awaited::Arrived(reply);
drop(state);
waiter.wake()
}
Awaited::Outstanding | Awaited::Arrived(_) => *state = Awaited::Arrived(reply),
}
}
}
impl<T> Drop for Promise<T> {
fn drop(&mut self) {
self.fulfil(Err(ClientError::Closed));
}
}
impl<T> Sent<T> {
pub fn ready(value: Result<T, ClientError>) -> Self {
Sent(Arc::new(Cell(Mutex::new(Awaited::Arrived(value)))))
}
pub fn try_take(&mut self) -> Option<Result<T, ClientError>> {
let mut state = self.0.lock();
match std::mem::take(&mut *state) {
Awaited::Arrived(reply) => Some(reply),
other => {
*state = other;
None
}
}
}
pub fn then(self, route: impl FnOnce(Result<T, ClientError>) + Send + 'static) {
let mut state = self.0.lock();
match std::mem::take(&mut *state) {
Awaited::Arrived(reply) => {
drop(state);
route(reply)
}
_ => *state = Awaited::Routed(Box::new(route)),
}
}
}
impl<T> Future for Sent<T> {
type Output = Result<T, ClientError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let mut state = self.0.lock();
match std::mem::take(&mut *state) {
Awaited::Arrived(reply) => Poll::Ready(reply),
_ => {
*state = Awaited::Watched(cx.waker().clone());
Poll::Pending
}
}
}
}
pub(crate) type PollEnd = Result<DeltaCursor, ClientError>;
pub(crate) enum Polled {
Block(RawBlock),
End(PollEnd),
}
#[derive(Default)]
struct Polls {
live: u64,
queue: VecDeque<Polled>,
}
impl Polls {
fn hand(&mut self, poll: u64, polled: Polled) {
if poll == self.live {
self.queue.push_back(polled);
}
}
}
pub struct Session {
transport: ClientTransport,
pending: VecDeque<Slot>,
submitted: u64,
ended: Option<ClientError>,
unread: bool,
polls: Polls,
}
impl Session {
pub fn connect(target: &str) -> Result<Self, ClientError> {
Ok(Self::over(ClientTransport::connect(
target,
Instant::now() + CONNECT_TIMEOUT,
)?))
}
pub(crate) fn over(transport: ClientTransport) -> Self {
Session {
transport,
pending: VecDeque::new(),
submitted: 0,
ended: None,
unread: false,
polls: Polls::default(),
}
}
pub fn requests_sent(&self) -> u64 {
self.submitted
}
pub fn as_raw_fd(&self) -> RawFd {
self.transport.as_raw_fd()
}
pub fn as_fd(&self) -> BorrowedFd<'_> {
self.transport.as_fd()
}
pub fn submit(&mut self, req: Request<'_>) -> Result<Sent<u64>, ClientError> {
let (frame, tid) = req.encode()?;
let (to, sent) = promise();
self.enqueue(frame, Slot::Ack { tid, to })?;
Ok(sent)
}
pub fn submit_resolve(&mut self, qname: &str) -> Result<Sent<Option<Arc<RelDescriptor>>>, ClientError> {
let hdr = ControlHeader::naming(ClientVerb::Resolve, Target::from(0), 0);
let (to, sent) = promise();
self.enqueue(encode_frame(hdr, qname.as_bytes(), None, None), Slot::Resolve { to })?;
Ok(sent)
}
pub fn submit_scan(
&mut self,
target: Target,
spec: &gnitz_wire::ReadSpec,
reply_schema: &Arc<Schema>,
) -> Result<Sent<ScanReply>, ClientError> {
let hdr = ControlHeader::naming(ClientVerb::ScanSpec, target, reply_schema.layout_digest());
let (to, sent) = promise();
let slot = Slot::Scan {
tid: target.tid,
reply_schema: Arc::clone(reply_schema),
data: None,
to,
};
self.enqueue(encode_frame(hdr, &spec.encode(), None, None), slot)?;
Ok(sent)
}
pub fn submit_scan_multi(&mut self, rels: Vec<(u64, Arc<Schema>)>) -> Result<Sent<Vec<ScanReply>>, ClientError> {
if rels.is_empty() {
return Err(ClientError::from("a multi-read names no relation".to_string()));
}
let items: Vec<txn_frame::ScanMultiItem> = rels
.iter()
.map(|(tid, schema)| txn_frame::ScanMultiItem {
tid: *tid,
reply_layout: schema.layout_digest(),
})
.collect();
let (to, sent) = promise();
let slot = Slot::Multi {
replies: Vec::with_capacity(rels.len()),
rels,
data: None,
to,
};
self.enqueue(txn_frame::encode_scan_multi(&items), slot)?;
Ok(sent)
}
pub(crate) fn submit_delta_poll(&mut self, views: &[txn_frame::DeltaPollItem], wait: Duration) {
if views.is_empty() {
return;
}
let slot = Slot::DeltaPoll {
views: views.iter().map(|v| v.view_id).collect(),
at: 0,
poll: self.polls.live,
};
let wait_ms = u64::try_from(wait.as_nanos().div_ceil(1_000_000)).unwrap_or(u64::MAX);
if let Err(why) = self.enqueue(txn_frame::encode_delta_poll(views, wait_ms), slot) {
let ends = views.iter().map(|_| Polled::End(Err(why.clone())));
self.polls.queue.extend(ends);
}
}
fn enqueue(&mut self, frame: Vec<u8>, slot: Slot) -> Result<(), ClientError> {
let (total, limit) = (frame.len(), gnitz_wire::MAX_FRAME_PAYLOAD);
if total > limit {
return Err(ClientError::from(format!(
"request frame is {total} bytes, exceeding the {limit}-byte server ingress cap; \
split the request"
)));
}
if let Some(why) = &self.ended {
return Err(why.clone());
}
if self.at_capacity() {
let queued = self.queued_bytes();
return Err(ClientError::from(if self.pending.len() >= MAX_IN_FLIGHT {
format!("connection has {MAX_IN_FLIGHT} requests in flight")
} else {
format!("connection has {queued} unwritten bytes queued, at the {MAX_QUEUED_BYTES}-byte cap")
}));
}
self.transport.enqueue(frame);
self.submitted += 1;
self.pending.push_back(slot);
Ok(())
}
pub fn step(&mut self, ready: Interest) {
if ready.read {
self.read();
}
if ready.write && self.ended.is_none() {
if let Err(e) = self.transport.flush() {
self.read();
while self.unread {
self.read();
}
self.end(ClientError::ConnectionLost(e));
}
}
}
fn read(&mut self) {
self.unread = false;
if self.ended.is_some() {
return;
}
let Session { transport, pending, polls, .. } = self;
match transport.read(|buf| feed(pending, buf, polls)) {
Ok(more) => self.unread = more,
Err(e) => self.end(ClientError::ConnectionLost(e)),
}
}
pub(crate) fn next_polled(&mut self) -> Option<Polled> {
self.polls.queue.pop_front()
}
pub(crate) fn polled_out(&self) -> bool {
self.polls.queue.is_empty()
}
pub(crate) fn abandon_poll(&mut self) {
self.polls.live += 1;
self.polls.queue.clear();
}
pub(crate) fn unread(&self) -> bool {
self.unread
}
fn queued_bytes(&self) -> usize {
self.transport.queued_bytes()
}
pub fn at_capacity(&self) -> bool {
self.pending.len() >= MAX_IN_FLIGHT || self.queued_bytes() >= MAX_QUEUED_BYTES
}
pub fn interest(&self) -> Interest {
if self.ended.is_some() {
return Interest::NONE;
}
Interest {
read: !self.pending.is_empty(),
write: self.transport.wants_write(),
}
}
pub fn is_closed(&self) -> bool {
self.ended.is_some()
}
fn end(&mut self, why: ClientError) {
if self.ended.is_some() {
return;
}
self.transport.close();
let Session { pending, polls, .. } = self;
for slot in pending.drain(..) {
slot.fail(why.clone(), polls);
}
self.ended = Some(why);
}
pub fn close(&mut self) {
self.end(ClientError::Closed);
}
}
fn feed(pending: &mut VecDeque<Slot>, buf: Cow<'_, [u8]>, polls: &mut Polls) -> Result<(), ProtocolError> {
let Some(head) = pending.front_mut() else {
return Err(ProtocolError::DecodeError("reply frame with no request pending".into()));
};
if let Some(answered) = head.feed(buf, polls)? {
let slot = pending.pop_front().expect("the head was just fed");
if let Err(why) = answered {
slot.fail(why, polls);
}
}
Ok(())
}
impl Slot {
fn feed(
&mut self,
mut buf: Cow<'_, [u8]>,
polls: &mut Polls,
) -> Result<Option<Result<(), ClientError>>, ProtocolError> {
let ctrl = peek_control_block(&buf).map_err(ProtocolError::DecodeError)?;
let named = ctrl.hdr.target_id;
let want = match &*self {
Slot::Ack { tid, .. } | Slot::Scan { tid, .. } => Some(*tid),
Slot::Multi { rels, replies, .. } => Some(rels[replies.len()].0),
Slot::DeltaPoll { views, at, .. } => Some(views[*at]),
Slot::Resolve { .. } => None,
};
if let Some(fault) = ctrl.fault(&buf) {
let refused = ClientError::Refused(fault);
return match (&mut *self, want) {
(Slot::DeltaPoll { views, at, poll }, Some(view)) if named != 0 => {
if named != view {
return Err(out_of_order(view, named));
}
polls.hand(*poll, Polled::End(Err(refused)));
*at += 1;
Ok((*at == views.len()).then_some(Ok(())))
}
_ => Ok(Some(Err(refused))),
};
}
if let Some(want) = want.filter(|&want| want != named) {
return Err(out_of_order(want, named));
}
let frame_schema = match ctrl.schema.clone() {
None => None,
Some(r) if matches!(self, Slot::Resolve { .. }) => Some(Arc::new(
Schema::from_block(&buf[r]).map_err(ProtocolError::DecodeError)?,
)),
Some(_) => {
return Err(ProtocolError::DecodeError(
"a schema block on a reply whose request named its schema".into(),
))
}
};
let decode = |data: &mut Option<ZSetBatch>, schema: &Arc<Schema>, block: &[u8]| {
decode_wal_block_into(data.get_or_insert_with(|| ZSetBatch::new(schema)), block, schema)
};
match (&mut *self, ctrl.data.clone()) {
(_, None) => {}
(Slot::DeltaPoll { poll, .. }, Some(block)) if *poll == polls.live => {
let frame = std::mem::take(&mut buf).into_owned();
polls.hand(*poll, Polled::Block(RawBlock { frame, block }));
}
(Slot::DeltaPoll { .. }, Some(_)) => {}
(Slot::Scan { reply_schema, data, .. }, Some(r)) => decode(data, reply_schema, &buf[r])?,
(Slot::Multi { rels, replies, data, .. }, Some(r)) => decode(data, &rels[replies.len()].1, &buf[r])?,
(Slot::Ack { .. } | Slot::Resolve { .. }, Some(_)) => {
return Err(ProtocolError::DecodeError(
"a data block on a reply that carries no rows".into(),
))
}
}
if ctrl.hdr.flags.continuation {
return Ok(None);
}
let scan_reply = |schema: &Arc<Schema>, data: &mut Option<ZSetBatch>| ScanReply {
batch: data.take().unwrap_or_else(|| ZSetBatch::new(schema)),
schema: Arc::clone(schema),
lsn: Some(ctrl.hdr.arg0),
};
match self {
Slot::Ack { to, .. } => to.fulfil(Ok(ctrl.hdr.arg0)),
Slot::Resolve { to } => to.fulfil(Ok(resolve_descriptor(&ctrl, &buf, frame_schema)?)),
Slot::Scan { reply_schema, data, to, .. } => to.fulfil(Ok(scan_reply(reply_schema, data))),
Slot::Multi { rels, replies, data, to } => {
let reply = scan_reply(&rels[replies.len()].1, data);
replies.push(reply);
if replies.len() < rels.len() {
return Ok(None);
}
to.fulfil(Ok(std::mem::take(replies)))
}
Slot::DeltaPoll { views, at, poll } => {
let cursor = DeltaCursor::from_pair(ctrl.hdr.arg1, ctrl.hdr.arg0)
.ok_or_else(|| ProtocolError::DecodeError("a delta-poll terminal at round 0".into()))?;
polls.hand(*poll, Polled::End(Ok(cursor)));
*at += 1;
if *at < views.len() {
return Ok(None);
}
}
}
Ok(Some(Ok(())))
}
fn fail(self, why: ClientError, polls: &mut Polls) {
match self {
Slot::Ack { mut to, .. } => to.fulfil(Err(why)),
Slot::Resolve { mut to } => to.fulfil(Err(why)),
Slot::Scan { mut to, .. } => to.fulfil(Err(why)),
Slot::Multi { mut to, .. } => to.fulfil(Err(why)),
Slot::DeltaPoll { views, at, poll } => {
for _ in at..views.len() {
polls.hand(poll, Polled::End(Err(why.clone())));
}
}
}
}
}
fn out_of_order(want: u64, got: u64) -> ProtocolError {
ProtocolError::DecodeError(format!("reply out of order: expected target {want}, got {got}"))
}
fn resolve_descriptor(
ctrl: &DecodedControl,
frame: &[u8],
schema: Option<Arc<Schema>>,
) -> Result<Option<Arc<RelDescriptor>>, ProtocolError> {
if ctrl.hdr.target_id == 0 {
return Ok(None);
}
let schema = schema.ok_or_else(|| ProtocolError::DecodeError("RESOLVE reply carries no schema".into()))?;
let desc = RelDescriptorBlob::decode(&frame[ctrl.blob.clone()]).map_err(ProtocolError::DecodeError)?;
Ok(Some(Arc::new(RelDescriptor {
tid: ctrl.hdr.target_id,
class: desc.class,
pk_repeats: desc.pk_repeats,
serial: desc.serial,
schema,
indexes: desc.indexes,
token: ctrl.hdr.arg0,
})))
}
#[cfg(test)]
#[path = "tests/connection.rs"]
mod tests;