use gnitz_expr::SchemaFacts;
use std::borrow::Cow;
use std::collections::VecDeque;
use std::num::NonZeroU64;
use std::os::fd::{OwnedFd, RawFd};
use std::sync::Arc;
use std::time::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, WireFlags};
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,
}
}
}
#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct SlotId(u64);
pub enum Request<'a> {
AllocIds(u64),
AllocSerial { table: Target, count: u64 },
DdlTxn(&'a [(u64, ZSetBatch)]),
PushTxn { families: &'a [PushFamily<'a>] },
Resolve(&'a str),
Push {
target: Target,
schema: &'a Schema,
batch: &'a ZSetBatch,
mode: WireConflictMode,
},
ScanSpec {
target: Target,
spec: &'a gnitz_wire::ReadSpec,
reply_schema: &'a Arc<Schema>,
},
ScanMulti(Vec<(u64, Arc<Schema>)>),
}
pub struct Encoded {
frame: Vec<u8>,
kind: SlotKind,
}
impl Encoded {
fn new(frame: Vec<u8>, kind: SlotKind) -> Result<Self, ClientError> {
let total = frame.len();
let limit = 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"
)));
}
Ok(Encoded { frame, kind })
}
pub(crate) fn delta_poll(views: &[txn_frame::DeltaPollItem]) -> Result<Self, ClientError> {
if views.is_empty() {
return Err(ClientError::from("a delta poll names no view".to_string()));
}
let kind = SlotKind::DeltaPoll {
views: views.iter().map(|v| v.view_id).collect(),
at: 0,
};
Encoded::new(txn_frame::encode_delta_poll(views), kind)
}
}
impl Request<'_> {
pub fn encode(self) -> Result<Encoded, ClientError> {
let (frame, kind) = match self {
Request::AllocIds(n) => {
let hdr = ControlHeader {
flags: WireFlags {
verb: ClientVerb::AllocIds,
..Default::default()
},
arg0: n,
..Default::default()
};
(encode_frame(hdr, &[], None, None), SlotKind::Ack { tid: 0 })
}
Request::AllocSerial { table, count } => {
let hdr = ControlHeader::naming(ClientVerb::AllocSerialRange, table, count);
(encode_frame(hdr, &[], None, None), SlotKind::Ack { tid: 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), SlotKind::Ack { tid: 0 })
}
Request::PushTxn { families } => {
for f in families {
f.batch.validate(f.schema)?;
}
(encode_push_txn(families), SlotKind::Ack { tid: 0 })
}
Request::Resolve(qname) => {
let hdr = ControlHeader {
flags: WireFlags {
verb: ClientVerb::Resolve,
..Default::default()
},
..Default::default()
};
(encode_frame(hdr, qname.as_bytes(), None, None), SlotKind::Resolve)
}
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)),
SlotKind::Ack { tid: target.tid },
)
}
Request::ScanSpec { target, spec, reply_schema } => {
let hdr = ControlHeader::naming(ClientVerb::ScanSpec, target, reply_schema.layout_digest());
(
encode_frame(hdr, &spec.encode(), None, None),
SlotKind::Scan {
tid: target.tid,
reply_schema: Arc::clone(reply_schema),
data: None,
},
)
}
Request::ScanMulti(rels) => {
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 (replies, data) = (Vec::with_capacity(rels.len()), None);
(
txn_frame::encode_scan_multi(&items),
SlotKind::Multi { rels, replies, data },
)
}
};
Encoded::new(frame, kind)
}
}
#[derive(Debug)]
pub enum Reply {
Scan(ScanReply),
Multi(Vec<ScanReply>),
Ack(u64),
Resolve(Option<Arc<RelDescriptor>>),
Polled,
}
impl Reply {
fn kind(&self) -> &'static str {
match self {
Reply::Scan(_) => "Scan",
Reply::Multi(_) => "Multi",
Reply::Ack(_) => "Ack",
Reply::Resolve(_) => "Resolve",
Reply::Polled => "Polled",
}
}
#[inline]
#[track_caller]
pub fn into_ack(self) -> u64 {
match self {
Reply::Ack(value) => value,
other => wrong_shape(other.kind(), "Ack"),
}
}
#[inline]
#[track_caller]
pub fn into_scan(self) -> ScanReply {
match self {
Reply::Scan(r) => r,
other => wrong_shape(other.kind(), "Scan"),
}
}
#[inline]
#[track_caller]
pub fn into_multi(self) -> Vec<ScanReply> {
match self {
Reply::Multi(r) => r,
other => wrong_shape(other.kind(), "Multi"),
}
}
#[inline]
#[track_caller]
pub fn into_resolve(self) -> Option<Arc<RelDescriptor>> {
match self {
Reply::Resolve(d) => d,
other => wrong_shape(other.kind(), "Resolve"),
}
}
}
#[cold]
#[inline(never)]
#[track_caller]
fn wrong_shape(got: &'static str, want: &'static str) -> ! {
panic!(
"the spine resolves a reply against the request that opened its slot; \
wanted Reply::{want}, got Reply::{got}"
)
}
pub type Completions = Vec<(SlotId, Result<Reply, ClientError>)>;
enum SlotKind {
Ack {
tid: u64,
},
Resolve,
Scan {
tid: u64,
reply_schema: Arc<Schema>,
data: Option<ZSetBatch>,
},
Multi {
rels: Vec<(u64, Arc<Schema>)>,
replies: Vec<ScanReply>,
data: Option<ZSetBatch>,
},
DeltaPoll {
views: Vec<u64>,
at: usize,
},
}
struct Slot {
id: SlotId,
kind: SlotKind,
}
pub(crate) type PollEnd = Result<DeltaCursor, ClientError>;
pub(crate) enum Polled {
Block(RawBlock),
End(PollEnd),
}
pub(crate) type PollSink<'a> = dyn FnMut(SlotId, Polled) + 'a;
type Owed = Vec<(SlotId, Polled)>;
fn deliver(sink: Option<&mut PollSink<'_>>, owed: &mut Owed) {
match sink {
Some(sink) => owed.drain(..).for_each(|(slot, p)| sink(slot, p)),
None => owed.clear(),
}
}
pub struct Session {
transport: ClientTransport,
pending: VecDeque<Slot>,
next_slot: u64,
ended: Option<ClientError>,
}
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(),
next_slot: 1,
ended: None,
}
}
pub fn requests_sent(&self) -> u64 {
self.next_slot - 1
}
pub fn as_raw_fd(&self) -> RawFd {
self.transport.as_raw_fd()
}
pub fn try_clone_fd(&self) -> Result<OwnedFd, ClientError> {
Ok(self.transport.try_clone_fd()?)
}
pub fn submit(&mut self, req: Request<'_>) -> Result<SlotId, ClientError> {
self.enqueue(req.encode()?)
}
pub fn enqueue(&mut self, Encoded { frame, kind }: Encoded) -> Result<SlotId, ClientError> {
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);
let id = SlotId(self.next_slot);
self.next_slot += 1;
self.pending.push_back(Slot { id, kind });
Ok(id)
}
pub fn step(&mut self, ready: Interest) -> Completions {
self.step_polling(ready, None)
}
pub(crate) fn step_polling(&mut self, ready: Interest, mut sink: Option<&mut PollSink<'_>>) -> Completions {
let mut done: Completions = Vec::new();
if ready.read {
self.read_frames(sink.as_deref_mut(), &mut done);
}
if ready.write && self.ended.is_none() {
if let Err(e) = self.transport.flush() {
self.read_frames(sink.as_deref_mut(), &mut done);
let mut owed = Owed::new();
self.end(ClientError::ConnectionLost(e), &mut done, &mut owed);
deliver(sink, &mut owed);
}
}
done
}
fn read_frames(&mut self, mut sink: Option<&mut PollSink<'_>>, done: &mut Completions) {
let mut owed = Owed::new();
while self.ended.is_none() {
let Session { transport, pending, .. } = self;
let read = transport.read(|buf| feed(pending, buf, done, &mut owed));
let more = match read {
Ok(more) => more,
Err(e) => {
self.end(ClientError::ConnectionLost(e), done, &mut owed);
false
}
};
deliver(sink.as_deref_mut(), &mut owed);
if !more {
return;
}
}
}
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, done: &mut Completions, owed: &mut Owed) {
if self.ended.is_some() {
return;
}
self.transport.close();
for slot in self.pending.drain(..) {
slot.complete(Err(why.clone()), done, owed);
}
self.ended = Some(why);
}
#[must_use]
pub fn close(&mut self) -> Completions {
let mut done = Vec::new();
self.end(ClientError::Closed, &mut done, &mut Owed::new());
done
}
}
fn feed(
pending: &mut VecDeque<Slot>,
buf: Cow<'_, [u8]>,
done: &mut Completions,
owed: &mut Owed,
) -> Result<(), ProtocolError> {
let Some(head) = pending.front_mut() else {
return Err(ProtocolError::DecodeError("reply frame with no request pending".into()));
};
if let Some(result) = head.feed(buf, owed)? {
let slot = pending.pop_front().expect("the head was just fed");
slot.complete(result, done, owed);
}
Ok(())
}
impl Slot {
fn feed(
&mut self,
mut buf: Cow<'_, [u8]>,
owed: &mut Owed,
) -> Result<Option<Result<Reply, ClientError>>, ProtocolError> {
let ctrl = peek_control_block(&buf).map_err(ProtocolError::DecodeError)?;
let named = ctrl.hdr.target_id;
let want = match &self.kind {
SlotKind::Ack { tid } | SlotKind::Scan { tid, .. } => Some(*tid),
SlotKind::Multi { rels, replies, .. } => Some(rels[replies.len()].0),
SlotKind::DeltaPoll { views, at } => Some(views[*at]),
SlotKind::Resolve => None,
};
if let Some(fault) = ctrl.fault(&buf) {
let refused = ClientError::Refused(fault);
return match (&mut self.kind, want) {
(SlotKind::DeltaPoll { views, at }, Some(view)) if named != 0 => {
if named != view {
return Err(out_of_order(view, named));
}
Ok(end_poll_position(self.id, views, at, Err(refused), owed))
}
_ => 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.kind, SlotKind::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.kind, ctrl.data.clone()) {
(_, None) => {}
(SlotKind::DeltaPoll { .. }, Some(block)) => {
let frame = std::mem::take(&mut buf).into_owned();
owed.push((self.id, Polled::Block(RawBlock { frame, block })));
}
(SlotKind::Scan { reply_schema, data, .. }, Some(r)) => decode(data, reply_schema, &buf[r])?,
(SlotKind::Multi { rels, replies, data }, Some(r)) => decode(data, &rels[replies.len()].1, &buf[r])?,
(SlotKind::Ack { .. } | SlotKind::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),
};
Ok(match &mut self.kind {
SlotKind::Ack { .. } => Some(Ok(Reply::Ack(ctrl.hdr.arg0))),
SlotKind::Resolve => Some(Ok(Reply::Resolve(resolve_descriptor(&ctrl, &buf, frame_schema)?))),
SlotKind::Scan { reply_schema, data, .. } => Some(Ok(Reply::Scan(scan_reply(reply_schema, data)))),
SlotKind::Multi { rels, replies, data } => {
let reply = scan_reply(&rels[replies.len()].1, data);
replies.push(reply);
(replies.len() == rels.len()).then(|| Ok(Reply::Multi(std::mem::take(replies))))
}
SlotKind::DeltaPoll { views, at } => {
let cursor = DeltaCursor::from_pair(ctrl.hdr.arg1, ctrl.hdr.arg0)
.ok_or_else(|| ProtocolError::DecodeError("a delta-poll terminal at round 0".into()))?;
end_poll_position(self.id, views, at, Ok(cursor), owed)
}
})
}
fn complete(self, result: Result<Reply, ClientError>, done: &mut Completions, owed: &mut Owed) {
if let (SlotKind::DeltaPoll { views, at }, Err(why)) = (&self.kind, &result) {
owed.extend(views[*at..].iter().map(|_| (self.id, Polled::End(Err(why.clone())))));
}
done.push((self.id, result));
}
}
fn out_of_order(want: u64, got: u64) -> ProtocolError {
ProtocolError::DecodeError(format!("reply out of order: expected target {want}, got {got}"))
}
fn end_poll_position(
slot: SlotId,
views: &[u64],
at: &mut usize,
end: PollEnd,
owed: &mut Owed,
) -> Option<Result<Reply, ClientError>> {
owed.push((slot, Polled::End(end)));
*at += 1;
(*at == views.len()).then_some(Ok(Reply::Polled))
}
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;