use crate::connection::{
promise, DeltaCursor, Interest, Polled, RelDescriptor, Request, ScanReply, Sent, Session, Target,
};
use crate::error::ClientError;
use crate::protocol::transport::poll_fd;
use crate::{sys_schema, BatchAppender, PkColumn, ProtocolError, PushFamily, RelName, Schema, ZSetBatch};
use gnitz_expr::{ColumnTable, SchemaFacts};
use gnitz_wire::{ColumnDef, PkBuf, PkKeys, WireConflictMode};
use gnitz_wire::{WireFault, WireStatus};
use std::borrow::Cow;
use std::collections::HashMap;
use std::future::{poll_fn, Future};
use std::os::fd::{AsRawFd, BorrowedFd, RawFd};
use std::pin::Pin;
use std::sync::Arc;
use std::task::{ready, Context, Poll, Waker};
use std::time::Duration;
use gnitz_expr::{LogicalProgram, RowFilter};
use gnitz_wire::sys_rows::{
CircuitRow, ColTabRow, ColTabSlot, FkRef, IdxTabRow, SchemaTabRow, SchemaTabSlot, SysRow, TableTabRow, ViewTabRow,
};
use gnitz_wire::txn_frame::{DeltaPollItem, BLIND};
use gnitz_wire::{payload_bytes, payload_str, payload_u64};
use gnitz_wire::{Circuit, ComputeMap, Cut, KeyRange, ReadBound, ReadSink, ReadSpec};
use gnitz_wire::{
PkColList, PkListRole, TableProps, ViewProps, CIRCUIT_TAB, COL_TAB, IDX_TAB, RELTAB_PAY_NAME, RELTAB_PAY_SCHEMA_ID,
SCHEMA_TAB, TABLE_TAB, VIEW_TAB,
};
pub fn not_found(noun: &'static str, name: &RelName) -> ClientError {
absent(format!("{noun} '{name}' not found"))
}
fn absent(text: String) -> ClientError {
ClientError::Refused(WireFault { status: WireStatus::NotFound, text })
}
pub fn retraction_batch(schema: &Schema, pks: PkColumn) -> ZSetBatch {
let count = pks.len();
ZSetBatch {
pks,
weights: vec![-1; count],
nulls: vec![0; count],
payload: ZSetBatch::filler_columns(schema, count),
blob: vec![],
}
}
pub fn key_reply(schema: &Schema) -> (Arc<Schema>, ReadSink) {
let reply = Schema {
columns: schema.hidden_key_columns().collect(),
pk_cols: (0..schema.pk_cols.len() as u32).collect(),
};
let program = LogicalProgram::copy_cols(&[]).to_blob_bytes();
let map = ComputeMap { program, out_cols: Vec::new() };
(Arc::new(reply), ReadSink { map: Some(map), ..ReadSink::all_rows() })
}
pub const RMW_MAX_ATTEMPTS: usize = 4;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct InlineUniqueIndex {
pub cols: PkColList,
pub name: String,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct IndexRow {
pub owner: u64,
pub name: String,
pub cols: PkColList,
pub is_unique: bool,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum FkTarget {
Table(FkRef),
SelfTable { col: u32 },
}
impl FkTarget {
fn resolve(self, own_id: u64) -> FkRef {
match self {
FkTarget::Table(fk) => fk,
FkTarget::SelfTable { col } => FkRef { table_id: own_id, col },
}
}
}
const SERIAL_RANGE_SIZE: u64 = 64;
#[derive(Default)]
struct DdlBundle(Vec<(u64, ZSetBatch)>);
impl DdlBundle {
fn batch(&mut self, family: u64) -> &mut ZSetBatch {
let at = match self.0.iter().position(|(f, _)| *f == family) {
Some(at) => at,
None => {
self.0.push((family, ZSetBatch::new(sys_schema(family))));
self.0.len() - 1
}
};
&mut self.0[at].1
}
fn put<R: SysRow>(&mut self, row: &R, weight: i64) {
row.write(&mut BatchAppender::new(self.batch(R::FAMILY)), weight);
}
fn same_rows(&self, other: &DdlBundle) -> bool {
fn in_key_order(family: u64, rows: &ZSetBatch) -> ZSetBatch {
let mut live: Vec<usize> = rows.live_rows().collect();
live.sort_unstable_by_key(|&i| rows.pks.get_bytes(i));
let mut out = ZSetBatch::new(sys_schema(family));
for i in live {
out.copy_row_at(rows, i, rows.weights[i]);
}
out
}
self.0.len() == other.0.len()
&& self.0.iter().all(|(family, rows)| {
other
.0
.iter()
.any(|(f, theirs)| f == family && in_key_order(*family, rows) == in_key_order(*family, theirs))
})
}
}
pub const MAX_CHAIN_SEGMENTS: usize = 64;
fn segment_name(vid: u64) -> String {
format!("_seg{vid}")
}
pub fn segment_id(j: u64) -> u64 {
gnitz_wire::CATALOG_ID_CEILING + j
}
fn source_at(source: u64, base: u64) -> u64 {
match source.checked_sub(gnitz_wire::CATALOG_ID_CEILING) {
Some(j) => base + j,
None => source,
}
}
pub struct PlannedView {
pub circuit: Circuit,
pub schema: Arc<Schema>,
pub pk_repeats: bool,
}
pub struct ViewBundle {
pub segments: Vec<PlannedView>,
pub view: PlannedView,
}
impl ViewBundle {
fn put_rows(&self, b: &mut DdlBundle, view_name: &str, props: ViewProps, schema_id: u64, base: u64) {
let owner_vid = base + self.segments.len() as u64;
for (pv, vid) in self.segments.iter().chain([&self.view]).zip(base..) {
let mut circuit = pv.circuit.clone();
for src in circuit.sources_mut() {
*src = source_at(*src, base);
}
let (name, owner_view_id, props) = if vid == owner_vid {
(view_name.to_string(), 0, props)
} else {
(segment_name(vid), owner_vid, ViewProps::default())
};
append_col_rows(b, vid, &pv.schema.columns, &[]);
let circuit = circuit.encode();
b.put(&CircuitRow { view_id: vid, circuit: &circuit }, 1);
let (capacity_bytes, delta_bytes) = props.row_words();
b.put(
&ViewTabRow {
view_id: vid,
schema_id,
name: &name,
pk_col_idx: PkColList::from_slice(&pv.schema.pk_cols).pack(),
capacity_bytes,
delta_bytes,
owner_view_id,
pk_repeats: pv.pk_repeats as u64,
},
1,
);
}
}
}
impl From<PlannedView> for ViewBundle {
fn from(view: PlannedView) -> Self {
ViewBundle { segments: Vec::new(), view }
}
}
pub type ParkHook = Box<dyn FnMut() -> Result<(), Box<dyn std::error::Error + Send + Sync>> + Send>;
pub type Job = Box<dyn FnOnce() + Send>;
pub trait Host: Send {
fn attach(&mut self, fd: BorrowedFd<'_>) -> std::io::Result<()>;
fn poll_io(
&mut self,
want: Interest,
cx: &mut Context<'_>,
io: &mut dyn FnMut(Interest) -> Interest,
) -> Poll<Result<(), ClientError>>;
fn spawn(&mut self, job: Job);
}
#[derive(Default)]
pub struct BlockingHost {
fd: Option<RawFd>,
hook: Option<ParkHook>,
refused: bool,
}
impl BlockingHost {
pub fn with_hook(hook: ParkHook) -> Self {
BlockingHost { hook: Some(hook), ..Default::default() }
}
}
impl Host for BlockingHost {
fn attach(&mut self, fd: BorrowedFd<'_>) -> std::io::Result<()> {
self.fd = Some(fd.as_raw_fd());
self.refused = false;
Ok(())
}
fn poll_io(
&mut self,
want: Interest,
_cx: &mut Context<'_>,
io: &mut dyn FnMut(Interest) -> Interest,
) -> Poll<Result<(), ClientError>> {
let fd = self.fd.expect("a client attaches its host before it waits");
if want.write && !self.refused {
self.refused = io(Interest::WRITE).write;
return Poll::Ready(Ok(()));
}
loop {
match poll_fd(fd, want.poll_events(), None) {
Ok(revents) => {
self.refused = io(Interest::from_revents(revents)).write;
return Poll::Ready(Ok(()));
}
Err(e) if e.kind() == std::io::ErrorKind::Interrupted => {
if let Some(hook) = self.hook.as_mut() {
if let Err(e) = hook() {
return Poll::Ready(Err(ClientError::Interrupted(e.into())));
}
}
}
Err(e) => return Poll::Ready(Err(e.into())),
}
}
}
fn spawn(&mut self, job: Job) {
job()
}
}
pub fn block_on<F: Future>(fut: F) -> F::Output {
let mut fut = std::pin::pin!(fut);
match fut.as_mut().poll(&mut Context::from_waker(Waker::noop())) {
Poll::Ready(out) => out,
Poll::Pending => panic!("block_on drove a client whose host pends; that client belongs to an event loop"),
}
}
fn poll_turn(session: &mut Session, host: &mut dyn Host, cx: &mut Context<'_>) -> Poll<Result<(), ClientError>> {
if session.unread() {
session.step(Interest::READ);
return Poll::Ready(Ok(()));
}
let want = session.interest();
if want.is_empty() {
return Poll::Ready(Err(ClientError::Closed));
}
host.poll_io(want, cx, &mut |ready| {
session.step(ready);
session.interest()
})
}
pub(crate) async fn offload<R: Send + 'static>(
host: &mut dyn Host,
job: impl FnOnce() -> R + Send + 'static,
) -> Result<R, ClientError> {
let (mut ended, job_end) = promise();
host.spawn(Box::new(move || {
ended.fulfil(Ok(std::panic::catch_unwind(std::panic::AssertUnwindSafe(job))))
}));
job_end.await?.map_err(|panic| std::panic::resume_unwind(panic))
}
#[must_use = "a verb does nothing until awaited or detached"]
pub struct Pending<'a, T> {
client: &'a mut GnitzClient,
sent: Sent<T>,
}
impl<T> Unpin for Pending<'_, T> {}
impl<'a, T> Pending<'a, T> {
fn submitted(client: &'a mut GnitzClient, sent: Result<Sent<T>, ClientError>) -> Self {
let sent = sent.unwrap_or_else(|e| Sent::ready(Err(e)));
Pending { client, sent }
}
pub fn detach(self) -> Sent<T> {
self.sent
}
}
impl<T> Future for Pending<'_, T> {
type Output = Result<T, ClientError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
loop {
if let Some(reply) = this.sent.try_take() {
return Poll::Ready(reply);
}
let GnitzClient { session, host, .. } = &mut *this.client;
ready!(poll_turn(session, &mut **host, cx))?;
}
}
}
pub type Op = Box<dyn for<'a> FnOnce(&'a mut GnitzClient) -> BoxFut<'a, ()> + Send>;
pub type BoxFut<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
pub async fn serve(mut client: GnitzClient, mut next: impl FnMut(&mut Context<'_>) -> Poll<Option<Op>>) {
let mut more = true;
loop {
let op = poll_fn(|cx| loop {
if more {
match next(cx) {
Poll::Ready(Some(op)) => return Poll::Ready(Some(op)),
Poll::Ready(None) => more = false,
Poll::Pending => {}
}
}
let GnitzClient { session, host, .. } = &mut client;
if session.interest().is_empty() && !session.unread() {
return match more {
true => Poll::Pending,
false => Poll::Ready(None),
};
}
match poll_turn(session, &mut **host, cx) {
Poll::Ready(Ok(())) => {}
Poll::Ready(Err(_)) => session.close(),
Poll::Pending => return Poll::Pending,
}
})
.await;
let Some(op) = op else { return };
let _ = client.make_room().await;
op(&mut client).await;
}
}
pub struct GnitzClient {
pub(crate) host: Box<dyn Host>,
pub(crate) session: Session,
serial_cache: HashMap<u64, std::ops::Range<u64>>,
txn: Option<TxnBuffer>,
pub(crate) mirror: Option<Box<crate::mirror::MirrorState>>,
kept: HashMap<RelName, Arc<RelDescriptor>>,
}
const _: fn() = || {
fn assert_send<T: Send>() {}
fn assert_send_value<T: Send>(_: T) {}
assert_send::<GnitzClient>();
assert_send::<Sent<ScanReply>>();
let _ = |mut c: GnitzClient, schema: &Arc<Schema>, batch: &ZSetBatch| {
assert_send_value(c.poll_mirror(Duration::ZERO));
assert_send_value(c.create_schema(""));
assert_send_value(c.push(0, schema, batch, WireConflictMode::Update));
assert_send_value(serve(c, |_| Poll::Ready(None)));
};
assert_send::<ZSetBatch>();
assert_send::<ClientError>();
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<Schema>();
};
impl GnitzClient {
pub fn connect(target: &str) -> Result<Self, ClientError> {
block_on(Self::connect_with(target, Box::new(BlockingHost::default())))
}
pub async fn connect_with(target: &str, mut host: Box<dyn Host>) -> Result<Self, ClientError> {
let target = target.to_owned();
let session = offload(&mut *host, move || Session::connect(&target)).await??;
Self::over(session, host)
}
#[cfg(test)]
pub(crate) fn from_session(session: Session) -> GnitzClient {
Self::over(session, Box::new(BlockingHost::default())).expect("a blocking host attaches to any socket")
}
pub fn over(session: Session, mut host: Box<dyn Host>) -> Result<GnitzClient, ClientError> {
host.attach(session.as_fd())?;
Ok(GnitzClient {
host,
session,
serial_cache: HashMap::new(),
txn: None,
mirror: None,
kept: HashMap::new(),
})
}
pub fn requests_sent(&self) -> u64 {
self.session.requests_sent()
}
fn ack(&mut self, req: Request<'_>) -> Pending<'_, u64> {
let sent = self.session.submit(req);
Pending::submitted(self, sent)
}
pub fn wait<T>(&mut self, sent: Sent<T>) -> Pending<'_, T> {
Pending { client: self, sent }
}
pub async fn turn(&mut self) -> Result<(), ClientError> {
let GnitzClient { session, host, .. } = self;
poll_fn(|cx| poll_turn(session, &mut **host, cx)).await
}
pub async fn make_room(&mut self) -> Result<(), ClientError> {
while self.session.at_capacity() {
self.turn().await?;
}
Ok(())
}
pub async fn offload<R: Send + 'static>(
&mut self,
job: impl FnOnce() -> R + Send + 'static,
) -> Result<R, ClientError> {
offload(&mut *self.host, job).await
}
pub async fn reserve_serial_ids(&mut self, table: &RelDescriptor, count: u64) -> Result<u64, ClientError> {
match self.serial_cache.get_mut(&table.tid) {
Some(r) if r.end - r.start >= count => {
let base = r.start;
r.start += count;
Ok(base)
}
_ => {
let want = count.max(SERIAL_RANGE_SIZE);
let base = self
.ack(Request::AllocSerial { table: table.into(), count: want })
.await?;
self.serial_cache.insert(table.tid, base + count..base + want);
Ok(base)
}
}
}
async fn alloc_ids(&mut self, n: u64) -> Result<u64, ClientError> {
self.ack(Request::AllocIds(n)).await
}
pub async fn alloc_id(&mut self) -> Result<u64, ClientError> {
self.alloc_ids(1).await
}
pub fn push<'a>(
&mut self,
target: impl Into<Target>,
schema: &Arc<Schema>,
batch: impl Into<Cow<'a, ZSetBatch>>,
mode: WireConflictMode,
) -> Pending<'_, u64> {
let (target, batch) = (target.into(), batch.into());
if let Some(txn) = &mut self.txn {
let buffered = txn.push(target, schema, batch.into_owned(), mode, BLIND);
let sent = Sent::ready(buffered.map(|()| 0));
return Pending { client: self, sent };
}
let batch = &*batch;
self.ack(Request::Push { target, schema, batch, mode })
}
pub fn scan_spec(
&mut self,
target: impl Into<Target>,
spec: &ReadSpec,
reply_schema: &Arc<Schema>,
) -> Pending<'_, ScanReply> {
let target = target.into();
let sent = self.session.submit_scan(target, spec, reply_schema);
Pending::submitted(self, sent)
}
pub fn mirrored_desc(&self, name: &RelName) -> Option<Arc<RelDescriptor>> {
let m = self.mirror.as_deref()?;
m.views
.iter()
.find(|(&t, v)| v.name == *name && m.store.get().cursor_of(t).is_some())
.map(|(_, v)| Arc::clone(&v.desc))
}
pub fn kept_desc(&self, name: &RelName) -> Option<Arc<RelDescriptor>> {
self.kept.get(name).cloned()
}
pub fn scan_spec_local_first(
&mut self,
target: impl Into<Target>,
spec: ReadSpec,
reply_schema: &Arc<Schema>,
) -> Pending<'_, ScanReply> {
let target = target.into();
let Some(store) = self
.mirror
.as_deref()
.filter(|m| m.cursor_of(target.tid).is_some())
.map(|m| &m.store)
else {
return self.scan_spec(target, &spec, reply_schema);
};
let reply = store
.get()
.scan_spec(target.tid, spec, reply_schema)
.map(|batch| ScanReply {
batch,
schema: Arc::clone(reply_schema),
lsn: None,
});
let sent = Sent::ready(reply.map_err(ClientError::from));
Pending { client: self, sent }
}
pub async fn reconnect(&mut self, target: &str) -> Result<(), ClientError> {
if self.txn_active() {
return Err(ClientError::from(
"reconnect inside a transaction; commit or roll back first".to_string(),
));
}
let target = target.to_owned();
let fresh = self.offload(move || Session::connect(&target)).await??;
let GnitzClient {
host,
session,
serial_cache,
txn: _,
mirror,
kept,
} = self;
host.attach(fresh.as_fd())?;
*session = fresh;
serial_cache.clear();
kept.clear();
if let Some(m) = mirror.as_deref() {
m.store.get().clear_cursors();
}
Ok(())
}
pub async fn delta_bootstrap(
&mut self,
view_id: u64,
view_schema: &Arc<Schema>,
) -> Result<(ScanReply, DeltaCursor), ClientError> {
self.delta_read(view_id, 0, Duration::ZERO, view_schema).await
}
pub async fn delta_poll(
&mut self,
view_id: u64,
cursor: DeltaCursor,
view_schema: &Arc<Schema>,
wait: Duration,
) -> Result<(ScanReply, DeltaCursor), ClientError> {
let (data, next) = self.delta_read(view_id, cursor.tick.get(), wait, view_schema).await?;
Ok((data, cursor.advanced_to(next)?))
}
async fn delta_read(
&mut self,
view_id: u64,
after_tick: u64,
wait: Duration,
reply_schema: &Arc<Schema>,
) -> Result<(ScanReply, DeltaCursor), ClientError> {
let item = DeltaPollItem {
view_id,
after_tick,
reply_layout: reply_schema.layout_digest(),
};
let schema = Arc::clone(reply_schema);
let mut batch = ZSetBatch::new(&schema);
let GnitzClient { session, host, .. } = self;
let mut poll = DeltaPoll::start(session, &[item], wait);
let (mut end, mut undecoded) = (None, None);
while let Some((_, polled)) = poll.next(&mut **host).await? {
match polled {
Polled::Block(b) if undecoded.is_none() => {
undecoded = crate::protocol::wal_block::decode_wal_block_into(&mut batch, b.block(), &schema).err();
}
Polled::Block(_) => {}
Polled::End(e) => end = Some(e),
}
}
if let Some(e) = undecoded {
return Err(e.into());
}
let cursor = end.expect("one item ends once")?;
Ok((ScanReply { schema, batch, lsn: None }, cursor))
}
pub fn scan_many(&mut self, relations: Vec<(u64, Arc<Schema>)>) -> Pending<'_, Vec<ScanReply>> {
let sent = self.session.submit_scan_multi(relations);
Pending::submitted(self, sent)
}
pub async fn create_index(
&mut self,
owner_id: u64,
cols: PkColList,
index_name: &str,
is_unique: bool,
) -> Result<u64, ClientError> {
let index_name = gnitz_wire::canonical_identifier(index_name)?;
let index_id = self.alloc_id().await?;
let mut b = DdlBundle::default();
b.put(
&IdxTabRow {
index_id,
owner_id,
source_col_idx: cols.pack(),
name: &index_name,
is_unique: is_unique as u64,
},
1,
);
self.commit_ddl(b).await?;
Ok(index_id)
}
pub async fn drop_indexes_by_name(&mut self, index_names: &[&str], if_exists: bool) -> Result<(), ClientError> {
self.drop_index_rows(index_names, "index", if_exists, |_| true).await
}
pub async fn drop_unique_constraint(&mut self, tid: u64, name: &str, if_exists: bool) -> Result<(), ClientError> {
self.drop_index_rows(&[name], "constraint", if_exists, |r| r.owner == tid && r.is_unique)
.await
}
async fn drop_index_rows(
&mut self,
names: &[&str],
noun: &'static str,
if_exists: bool,
matches: impl Fn(&IndexRow) -> bool,
) -> Result<(), ClientError> {
let scanned = self.sys_rows(IDX_TAB, ReadBound::None).await?;
let rows = idx_rows(&scanned)?;
let mut b = DdlBundle::default();
let mut retired: Vec<usize> = Vec::with_capacity(names.len());
for name in names {
let name = gnitz_wire::canonical_identifier(name)?;
match rows.iter().find(|(_, r)| r.name == name && matches(r)) {
Some(&(i, _)) => {
if !retired.contains(&i) {
retired.push(i);
b.batch(IDX_TAB).copy_row_at(&scanned, i, -1);
}
}
None if if_exists => {}
None => return Err(absent(format!("{noun} '{name}' not found"))),
}
}
self.commit_ddl(b).await
}
pub async fn index_rows(&mut self) -> Result<Vec<IndexRow>, ClientError> {
let scanned = self.sys_rows(IDX_TAB, ReadBound::None).await?;
Ok(idx_rows(&scanned)?.into_iter().map(|(_, r)| r).collect())
}
pub fn txn_active(&self) -> bool {
self.txn.is_some()
}
pub fn txn_begin(&mut self) -> Result<(), ClientError> {
if self.txn.is_some() {
return Err(ClientError::from("transaction already open".to_string()));
}
self.txn = Some(TxnBuffer::default());
Ok(())
}
pub fn txn_rollback(&mut self) -> Result<(), ClientError> {
self.txn.take().map(|_| ()).ok_or_else(no_transaction)
}
pub async fn txn_commit(&mut self) -> Result<u64, ClientError> {
let buf = self.txn.take().ok_or_else(no_transaction)?;
self.push_txn(&buf).await.map_err(|e| match e {
ClientError::Refused(WireFault { status: WireStatus::StaleCatalog, text }) => {
self.kept.retain(|_, rel| !buf.families_of.contains_key(&rel.tid));
ClientError::Refused(WireFault {
status: WireStatus::TxnConflict,
text: format!("transaction conflict: {text}; retry"),
})
}
e => e,
})
}
async fn push_txn(&mut self, buf: &TxnBuffer) -> Result<u64, ClientError> {
if buf.families.is_empty() {
return Ok(0);
}
self.ack(Request::PushTxn { families: &buf.families }).await
}
pub async fn read_modify_write<E: From<ClientError>>(
&mut self,
target: &RelDescriptor,
bound: ReadBound,
predicate: Vec<u8>,
keys: bool,
mut build: impl FnMut(ZSetBatch) -> Result<ZSetBatch, E>,
) -> Result<usize, E> {
let (tid, schema) = (target.tid, &target.schema);
let (reply, sink) = match keys {
true => key_reply(schema),
false => (Arc::clone(schema), ReadSink::all_rows()),
};
let spec = ReadSpec { bound, predicate, sink };
let mut attempt = 0;
loop {
attempt += 1;
let rest = match (self.txn.as_mut(), &spec.bound) {
(Some(txn), ReadBound::PkSet(set)) => txn.unwritten(tid, set),
_ => None,
};
let (batch, basis) = match rest {
Some(rest) if rest.is_empty() => (ZSetBatch::new(&reply), BLIND),
rest => {
let narrowed = rest.map(|rest| ReadSpec {
bound: ReadBound::PkSet(rest),
predicate: spec.predicate.clone(),
sink: spec.sink.clone(),
});
let ScanReply { batch, lsn, .. } = self
.scan_spec(target, narrowed.as_ref().unwrap_or(&spec), &reply)
.await?;
(batch, lsn.expect("a server read carries its watermark"))
}
};
let mut own = TxnBuffer::default();
let txn = self.txn.as_mut().unwrap_or(&mut own);
let batch = build(txn.overlay(tid, schema, &spec, keys, batch)?)?;
let count = batch.len();
txn.push(target, schema, batch, WireConflictMode::Update, basis)?;
let pushed = match self.txn {
Some(_) => Ok(0),
None => self.push_txn(&own).await,
};
match pushed {
Err(ClientError::Refused(WireFault { status: WireStatus::TxnConflict, .. }))
if attempt < RMW_MAX_ATTEMPTS => {}
r => {
r?;
return Ok(count);
}
}
}
}
async fn commit_ddl(&mut self, bundle: DdlBundle) -> Result<(), ClientError> {
if bundle.0.is_empty() {
return Ok(());
}
if self.txn.is_some() {
return Err(ClientError::from("DDL is not allowed inside a transaction".to_string()));
}
self.ack(Request::DdlTxn(&bundle.0)).await?;
self.kept.clear();
self.after_ddl_commit(&bundle.0).await;
Ok(())
}
async fn after_ddl_commit(&mut self, families: &[(u64, ZSetBatch)]) {
let Some(m) = self.mirror.as_deref() else {
return;
};
let Some((_, b)) = families.iter().find(|(family, _)| *family == VIEW_TAB) else {
return;
};
let vid = |i| b.pks.get(i) as u64;
let mut dropped: Vec<u64> = Vec::new();
let mut renamed: Vec<(RelName, Arc<RelDescriptor>)> = Vec::new();
for v in (0..b.len()).filter(|&i| b.weights[i] < 0).map(vid) {
match (m.views.get(&v), b.live_rows().find(|&j| vid(j) == v)) {
(Some(view), Some(j)) => renamed.push((
payload_str(b, j, RELTAB_PAY_NAME)
.ok()
.and_then(|new| view.name.sibling(new).ok())
.expect("a bundle this client built names its views"),
Arc::clone(&view.desc),
)),
_ => dropped.push(v),
}
}
for v in dropped {
let _ = self.forget_view(v).await;
}
for (name, desc) in renamed {
let tid = desc.tid;
if self.bind(name, desc).await.is_err() {
let _ = self.forget_view(tid).await;
}
}
}
pub async fn create_schema(&mut self, name: &str) -> Result<u64, ClientError> {
let name = gnitz_wire::canonical_identifier(name)?;
let schema_id = self.alloc_id().await?;
let mut b = DdlBundle::default();
b.put(&SchemaTabRow { schema_id, name: &name }, 1);
self.commit_ddl(b).await?;
Ok(schema_id)
}
pub async fn drop_schema(&mut self, name: &str) -> Result<(), ClientError> {
let name = gnitz_wire::canonical_identifier(name)?;
let (schemas, at) = self.lookup_schema(&name).await?;
let schema_id = schemas.pks.get(at) as u64;
let mut b = DdlBundle::default();
for family in [VIEW_TAB, TABLE_TAB] {
let scanned = self.sys_rows(family, ReadBound::None).await?;
for i in scanned.live_rows() {
if payload_u64(&scanned, i, RELTAB_PAY_SCHEMA_ID) == schema_id {
b.batch(family).copy_row_at(&scanned, i, -1);
}
}
}
b.batch(SCHEMA_TAB).copy_row_at(&schemas, at, -1);
self.commit_ddl(b).await
}
pub async fn create_table(
&mut self,
table: &RelName,
schema: &Schema,
fks: &[Option<FkTarget>],
props: TableProps,
unique_indexes: &[InlineUniqueIndex],
) -> Result<u64, ClientError> {
let index_names: Vec<String> = unique_indexes
.iter()
.map(|spec| gnitz_wire::canonical_identifier(&spec.name))
.collect::<Result<Vec<_>, String>>()?;
schema
.validate()
.map_err(|e| ClientError::from(format!("create_table: {e}")))?;
let pk = PkColList::from_slice(&schema.pk_cols);
if !fks.is_empty() && fks.len() != schema.columns.len() {
return Err(ClientError::from(format!(
"create_table: {} foreign-key slots for {} columns",
fks.len(),
schema.columns.len()
)));
}
let (schemas, at) = self.lookup_schema(table.schema()).await?;
let schema_id = schemas.pks.get(at) as u64;
let new_tid = self.alloc_ids(1 + unique_indexes.len() as u64).await?;
let mut b = DdlBundle::default();
append_col_rows(&mut b, new_tid, &schema.columns, fks);
b.put(
&TableTabRow {
table_id: new_tid,
schema_id,
name: table.name(),
pk_col_idx: pk.pack(),
flags: props.pack(),
},
1,
);
for (k, spec) in unique_indexes.iter().enumerate() {
b.put(
&IdxTabRow {
index_id: new_tid + 1 + k as u64,
owner_id: new_tid,
source_col_idx: spec.cols.pack(),
name: &index_names[k],
is_unique: 1,
},
1,
);
}
self.commit_ddl(b).await?;
Ok(new_tid)
}
pub async fn drop_table(&mut self, tables: &[RelName], if_exists: bool) -> Result<(), ClientError> {
self.drop_relations(TABLE_TAB, "table", tables, if_exists).await
}
pub async fn create_view(
&mut self,
view: &RelName,
source: &RelDescriptor,
props: ViewProps,
) -> Result<u64, ClientError> {
let mut circuit = Circuit::default();
let scan = circuit.input_delta(source.tid, ReadBound::None);
circuit.sink(scan);
let planned = PlannedView {
circuit,
schema: Arc::clone(&source.schema),
pk_repeats: source.pk_repeats,
};
self.create_view_chain(view, planned.into(), props, None).await
}
pub async fn create_view_chain(
&mut self,
view: &RelName,
bundle: ViewBundle,
props: ViewProps,
replace: Option<u64>,
) -> Result<u64, ClientError> {
let view_name = view.name();
let n_views = bundle.segments.len() + 1;
if n_views > MAX_CHAIN_SEGMENTS {
return Err(ClientError::from(format!(
"view chain has {n_views} segments, exceeding the {MAX_CHAIN_SEGMENTS}-segment limit",
)));
}
for (k, pv) in bundle.segments.iter().chain([&bundle.view]).enumerate() {
pv.schema
.validate()
.map_err(|e| ClientError::from(format!("View '{view_name}' segment {k}: {e}")))?;
}
let n_segments = bundle.segments.len() as u64;
let mut b = DdlBundle::default();
let schema_id = match replace {
Some(vid) => {
let base = vid.saturating_sub(n_segments);
let members = ReadSpec::all_rows(ReadBound::Range(KeyRange::new(
PkColList::from_slice(&[0]),
&[],
Cut::before(base as u128),
Cut::after(vid as u128),
)));
let reads = [VIEW_TAB, COL_TAB, CIRCUIT_TAB]
.map(|family| (family, self.scan_spec(family, &members, sys_schema(family)).detach()));
let mut standing = DdlBundle::default();
for (family, read) in reads {
standing.0.push((family, self.wait(read).await?.batch));
}
let views = &standing.0[0].1;
let i = views
.live_rows()
.find(|&i| views.pks.get(i) as u64 == vid)
.ok_or_else(|| not_found("view", view))?;
let schema_id = payload_u64(views, i, RELTAB_PAY_SCHEMA_ID);
let mut want = DdlBundle::default();
bundle.put_rows(&mut want, view_name, props, schema_id, base);
if want.same_rows(&standing) {
return Ok(vid);
}
b.batch(VIEW_TAB).copy_row_at(views, i, -1);
schema_id
}
None => {
let (schemas, at) = self.lookup_schema(view.schema()).await?;
schemas.pks.get(at) as u64
}
};
let base = self.alloc_ids(n_views as u64).await?;
bundle.put_rows(&mut b, view_name, props, schema_id, base);
self.commit_ddl(b).await?;
Ok(base + n_segments)
}
pub async fn drop_view(&mut self, views: &[RelName], if_exists: bool) -> Result<(), ClientError> {
self.drop_relations(VIEW_TAB, "view", views, if_exists).await
}
async fn drop_relations(
&mut self,
family: u64,
noun: &'static str,
names: &[RelName],
if_exists: bool,
) -> Result<(), ClientError> {
let mut b = DdlBundle::default();
let mut retired: Vec<u64> = Vec::with_capacity(names.len());
for name in names {
let Some((scanned, i)) = self.relation_retraction(family, noun, name).await? else {
if if_exists {
continue;
}
return Err(not_found(noun, name));
};
let id = scanned.pks.get(i) as u64;
if !retired.contains(&id) {
retired.push(id);
b.batch(family).copy_row_at(&scanned, i, -1);
}
}
self.commit_ddl(b).await
}
async fn relation_retraction(
&mut self,
family: u64,
noun: &'static str,
name: &RelName,
) -> Result<Option<(ZSetBatch, usize)>, ClientError> {
let Some(desc) = self.resolve(name).await? else {
return Ok(None);
};
let (scanned, i) = self
.seek_sys_row(family, &[desc.tid as u128], || not_found(noun, name))
.await?;
if payload_bytes(&scanned, i, RELTAB_PAY_NAME) != name.name().as_bytes() {
return Ok(None);
}
Ok(Some((scanned, i)))
}
pub async fn alter_rename_relation(&mut self, rel: &RelDescriptor, new_name: &str) -> Result<(), ClientError> {
let new_name = gnitz_wire::canonical_identifier(new_name)?;
let tid = rel.tid;
let family = if rel.class.is_view() { VIEW_TAB } else { TABLE_TAB };
self.rewrite_sys_row(
family,
&[tid as u128],
|| absent(format!("relation {tid} not found")),
|b, row| b.set_string_cell(row, RELTAB_PAY_NAME, &new_name),
)
.await
}
pub async fn alter_rename_column(&mut self, tid: u64, col_idx: usize, new_col: &str) -> Result<(), ClientError> {
self.alter_col_pair(tid, col_idx, |b, row| {
b.set_string_cell(row, ColTabSlot::name as usize, new_col)
})
.await
}
pub async fn alter_drop_column(&mut self, tid: u64, col_idx: usize) -> Result<(), ClientError> {
self.alter_col_pair(tid, col_idx, |b, row| {
b.set_u64_cell(row, ColTabSlot::is_hidden as usize, 1)
})
.await
}
pub async fn alter_drop_not_null(&mut self, tid: u64, col_idx: usize) -> Result<(), ClientError> {
self.alter_col_pair(tid, col_idx, |b, row| {
b.set_u64_cell(row, ColTabSlot::is_nullable as usize, 1)
})
.await
}
pub async fn alter_add_column(&mut self, rel: &RelDescriptor, def: &ColumnDef) -> Result<(), ClientError> {
let (tid, col_idx) = (rel.tid, rel.schema.num_columns());
let mut b = DdlBundle::default();
b.put(&ColTabRow::of(tid, col_idx as u64, def, None), 1);
self.commit_ddl(b).await
}
async fn alter_col_pair(
&mut self,
tid: u64,
col_idx: usize,
patch: impl FnOnce(&mut ZSetBatch, usize),
) -> Result<(), ClientError> {
self.rewrite_sys_row(
COL_TAB,
&[tid as u128, col_idx as u128],
|| absent(format!("column index {col_idx} not found on table {tid}")),
patch,
)
.await
}
async fn rewrite_sys_row(
&mut self,
family: u64,
key: &[u128],
missing: impl FnOnce() -> ClientError,
patch: impl FnOnce(&mut ZSetBatch, usize),
) -> Result<(), ClientError> {
let (scanned, i) = self.seek_sys_row(family, key, missing).await?;
let mut b = DdlBundle::default();
let rows = b.batch(family);
rows.copy_row_at(&scanned, i, -1);
rows.copy_row_at(&scanned, i, 1);
patch(rows, 1);
self.commit_ddl(b).await
}
pub async fn resolve_relation(&mut self, name: &RelName) -> Result<Arc<RelDescriptor>, ClientError> {
self.resolve(name).await?.ok_or_else(|| not_found("relation", name))
}
pub async fn resolve(&mut self, name: &RelName) -> Result<Option<Arc<RelDescriptor>>, ClientError> {
let sent = self.session.submit_resolve(name.key());
let found = Pending::submitted(self, sent).await?;
match &found {
Some(desc) => self.kept.insert(name.clone(), Arc::clone(desc)),
None => self.kept.remove(name),
};
Ok(found)
}
async fn lookup_schema(&mut self, schema_name: &str) -> Result<(ZSetBatch, usize), ClientError> {
let batch = self.sys_rows(SCHEMA_TAB, ReadBound::None).await?;
let i = batch
.live_rows()
.find(|&i| payload_bytes(&batch, i, SchemaTabSlot::name as usize) == schema_name.as_bytes())
.ok_or_else(|| absent(format!("schema '{schema_name}' not found")))?;
Ok((batch, i))
}
async fn sys_rows(&mut self, family: u64, bound: ReadBound) -> Result<ZSetBatch, ClientError> {
Ok(self
.scan_spec(family, &ReadSpec::all_rows(bound), sys_schema(family))
.await?
.batch)
}
async fn seek_sys_row(
&mut self,
family: u64,
key: &[u128],
missing: impl FnOnce() -> ClientError,
) -> Result<(ZSetBatch, usize), ClientError> {
let mut keys = PkColumn::empty_for_schema(sys_schema(family));
keys.push_natives(key);
let batch = self.sys_rows(family, ReadBound::PkSet(keys.keys())).await?;
let i = batch.live_rows().next().ok_or_else(missing)?;
Ok((batch, i))
}
}
pub(crate) struct DeltaPoll<'s> {
session: &'s mut Session,
total: usize,
answered: usize,
}
impl Drop for DeltaPoll<'_> {
fn drop(&mut self) {
self.session.abandon_poll();
}
}
impl<'s> DeltaPoll<'s> {
pub(crate) fn start(session: &'s mut Session, items: &[DeltaPollItem], wait: Duration) -> Self {
session.submit_delta_poll(items, wait);
DeltaPoll { session, total: items.len(), answered: 0 }
}
pub(crate) fn drained(&self) -> bool {
self.session.polled_out()
}
pub(crate) async fn next(&mut self, host: &mut dyn Host) -> Result<Option<(usize, Polled)>, ClientError> {
loop {
if let Some(next) = self.session.next_polled() {
let item = self.answered;
self.answered += usize::from(matches!(next, Polled::End(_)));
return Ok(Some((item, next)));
}
if self.answered == self.total {
return Ok(None);
}
poll_fn(|cx| poll_turn(self.session, host, cx)).await?;
}
}
}
fn no_transaction() -> ClientError {
ClientError::from("no transaction open".to_string())
}
#[derive(Default)]
struct TxnBuffer {
families: Vec<PushFamily>,
indexed: Vec<usize>,
families_of: HashMap<u64, Vec<usize>>,
last_op_of: HashMap<u64, HashMap<PkBuf, (usize, usize)>>,
}
impl TxnBuffer {
fn push(
&mut self,
target: impl Into<Target>,
schema: &Arc<Schema>,
batch: ZSetBatch,
mode: WireConflictMode,
basis: u64,
) -> Result<(), ClientError> {
let target = target.into();
let tid = target.tid;
batch
.layout_matches(schema)
.map_err(|e| ClientError::from(format!("relation {tid}: the batch is not in its schema's layout: {e}")))?;
if batch.is_empty() {
return Ok(());
}
self.check_layout(tid, |held| match Arc::ptr_eq(&held.schema, schema) {
true => Ok(()),
false => batch.layout_matches(&held.schema),
})?;
let last = self.families_of.get(&tid).and_then(|v| v.last().copied());
match last.filter(|&i| self.families[i].mode == mode) {
Some(i) => {
let f = &mut self.families[i];
f.batch.extend_from_owned(batch);
f.basis = f.basis.min(basis);
if f.target.token == 0 {
f.target = target;
}
}
None => {
self.families_of.entry(tid).or_default().push(self.families.len());
self.families.push(PushFamily {
target,
schema: Arc::clone(schema),
batch,
mode,
basis,
});
self.indexed.push(0);
}
}
Ok(())
}
fn check_layout(
&self,
tid: u64,
matches: impl FnOnce(&PushFamily) -> Result<(), String>,
) -> Result<(), ClientError> {
let Some(&first) = self.families_of.get(&tid).and_then(|v| v.first()) else {
return Ok(());
};
matches(&self.families[first])
.map_err(|e| ClientError::from(format!("relation {tid} changed layout during this transaction: {e}")))
}
fn index_tid(&mut self, tid: u64) {
let Some(own) = self.families_of.get(&tid) else {
return;
};
let index = self.last_op_of.entry(tid).or_default();
for &fam in own {
let batch = &self.families[fam].batch;
for row in std::mem::replace(&mut self.indexed[fam], batch.len())..batch.len() {
if batch.weights[row] != 0 {
index.insert(PkBuf::from_bytes(batch.pks.get_bytes(row)), (fam, row));
}
}
}
}
fn unwritten(&mut self, tid: u64, set: &PkKeys) -> Option<PkKeys> {
self.index_tid(tid);
let index = self.last_op_of.get(&tid)?;
let first = set.iter().position(|k| index.contains_key(k))?;
let stride = set.stride();
let mut rest = Vec::with_capacity(set.as_bytes().len() - stride);
rest.extend_from_slice(&set.as_bytes()[..first * stride]);
for k in set.iter().skip(first + 1).filter(|k| !index.contains_key(*k)) {
rest.extend_from_slice(k);
}
Some(PkKeys::from_sorted(stride, rest))
}
fn overlay(
&mut self,
tid: u64,
schema: &Schema,
spec: &ReadSpec,
keys: bool,
committed: ZSetBatch,
) -> Result<ZSetBatch, ClientError> {
self.check_layout(tid, |held| held.batch.layout_matches(schema))?;
self.index_tid(tid);
let Some(index) = self.last_op_of.get(&tid) else {
return Ok(committed);
};
let keep: Vec<(usize, i64)> = (0..committed.len())
.filter(|&i| !index.contains_key(committed.pks.get_bytes(i)))
.map(|i| (i, committed.weights[i]))
.collect();
let mut out = committed.gather(&keep);
let mut live = ZSetBatch::new(schema);
let families = &self.families;
let take = |&(fam, row): &(usize, usize)| {
let b = &families[fam].batch;
if b.weights[row] > 0 {
live.copy_row_at(b, row, 1);
}
};
match &spec.bound {
ReadBound::PkSet(set) => set.iter().filter_map(|k| index.get(k)).for_each(take),
_ => index.values().for_each(take),
}
if live.is_empty() {
return Ok(out);
}
let mut ranges = Vec::new();
RowFilter::for_read(&spec.predicate, &spec.bound, schema)
.map_err(|e| ClientError::from(e.to_string()))?
.ranges(&live, &mut ranges);
if keys {
live.payload.clear();
live.blob.clear();
live.retain_ranges(&ranges);
live.nulls.fill(0);
} else {
let kept: Vec<(usize, i64)> = ranges
.iter()
.flat_map(|&(s, e)| s..e)
.map(|r| (r, live.weights[r]))
.collect();
live = live.gather(&kept);
}
out.extend_from_owned(live);
Ok(out)
}
}
fn idx_rows(batch: &ZSetBatch) -> Result<Vec<(usize, IndexRow)>, ClientError> {
batch
.live_rows()
.map(|i| {
let r = IdxTabRow::read(batch, i).map_err(|e| ProtocolError::DecodeError(format!("index row {i}: {e}")))?;
let name = r.name.to_owned();
let cols = PkColList::unpack(r.source_col_idx).map_err(|rule| {
ProtocolError::DecodeError(format!("index '{name}': {}", rule.for_role(PkListRole::ColumnList)))
})?;
let is_unique = gnitz_wire::bool_word(r.is_unique)
.map_err(|e| ProtocolError::DecodeError(format!("index '{name}': {e}")))?;
let owner = r.owner_id;
Ok((i, IndexRow { owner, name, cols, is_unique }))
})
.collect()
}
fn append_col_rows(b: &mut DdlBundle, owner_id: u64, columns: &[ColumnDef], fks: &[Option<FkTarget>]) {
for (i, cd) in columns.iter().enumerate() {
let fk = fks.get(i).copied().flatten().map(|t| t.resolve(owner_id));
b.put(&ColTabRow::of(owner_id, i as u64, cd, fk), 1);
}
}
#[cfg(test)]
#[path = "tests/client.rs"]
mod tests;
#[cfg(test)]
#[path = "benches/client.rs"]
mod bench;