use std::sync::Arc;
use base64::prelude::{BASE64_STANDARD, Engine as _};
use gwk_domain::blob::BLOB_CHUNK_BYTES;
use gwk_domain::ids::{ByteCount, EventCount, EventId, RequestId, Seq, WriterEpoch};
use gwk_domain::port::{BlobError, BlobStore, EventStore, MAX_READ_LIMIT};
use gwk_domain::protocol::{
CONNECTION_EGRESS_BYTES_PER_WINDOW, CONNECTION_INGRESS_BYTES_PER_WINDOW, CONTRACT_VERSION,
FRAME_BODY_MAX_BYTES, FrameKind, KernelErrorCode, KernelRequest, KernelResult,
MAX_SUBSCRIPTIONS_PER_CONNECTION, ProjectionKind, ProjectionRecord, ServerControl,
};
use sqlx::Row;
use tokio::io::{AsyncRead, AsyncWrite};
use tokio::net::UnixStream;
use tokio::sync::{mpsc, watch};
use super::frame::{Budget, Incoming, read_frame, write_frame};
use super::hello::{self, Readiness};
use super::listen::Listener;
use super::{WireError, strict, subscribe};
use crate::blob::store::PgBlobStore;
use crate::epoch::{
GENESIS_EVENT_TYPE, KERNEL_AGGREGATE, KERNEL_SINGLETON, epoch_of, is_public_revision,
};
use crate::error::{KernelError, Result};
use crate::store::PgEventStore;
pub const DRAIN_TIMEOUT_SECS: u64 = 30;
const PAGE_BYTE_BUDGET: usize = (FRAME_BODY_MAX_BYTES as usize) / 2;
const DEFAULT_PAGE_ROWS: u32 = 256;
const BATCH_QUEUE_DEPTH: usize = MAX_SUBSCRIPTIONS_PER_CONNECTION;
const RESPONSE_QUEUE_DEPTH: usize = 1;
pub(crate) struct Outgoing {
pub(crate) control: ServerControl,
pub(crate) delivered: Option<(Arc<std::sync::atomic::AtomicU64>, u64)>,
}
pub(crate) fn fit_page<T: serde::Serialize>(items: Vec<T>) -> Result<(Vec<T>, bool)> {
let mut kept: Vec<T> = Vec::with_capacity(items.len());
let mut bytes = 0usize;
for item in items {
bytes += serde_json::to_vec(&item)
.map_err(|e| KernelError::Config(format!("measure a page item: {e}")))?
.len();
if !kept.is_empty() && bytes > PAGE_BYTE_BUDGET {
return Ok((kept, true));
}
kept.push(item);
}
Ok((kept, false))
}
fn cursor_key(record: &ProjectionRecord, key: &str) -> Result<String> {
let tag = record.kind().as_str();
let json = serde_json::to_value(record)
.map_err(|e| KernelError::Config(format!("serialize a {tag} record: {e}")))?;
json.get(tag)
.and_then(|body| body.get(key))
.and_then(serde_json::Value::as_str)
.map(str::to_owned)
.ok_or_else(|| KernelError::Config(format!("a {tag} record has no {key} to page from")))
}
pub struct Daemon {
store: Arc<PgEventStore>,
public_revision: String,
writer_epoch: WriterEpoch,
wake: Arc<watch::Sender<u64>>,
}
impl Daemon {
pub fn new(store: PgEventStore, public_revision: String) -> Result<Self> {
if !is_public_revision(&public_revision) {
return Err(KernelError::Config(format!(
"public revision {public_revision:?} is not a full 40-hex revision"
)));
}
if store.blobs().is_none() {
return Err(KernelError::Config(
"a serving daemon needs a blob store: every payload too large to inline lives \
there, and half a protocol is not a kernel"
.to_owned(),
));
}
let writer_epoch = WriterEpoch::new(store.boot_epoch().max(0) as u64);
let (wake, _) = watch::channel(0u64);
Ok(Self {
store: Arc::new(store),
public_revision,
writer_epoch,
wake: Arc::new(wake),
})
}
pub fn notify_on_append(&self) {
tokio::spawn(subscribe::watch_events(
self.store.pool().clone(),
Arc::clone(&self.wake),
));
}
fn blobs(&self) -> &PgBlobStore {
self.store
.blobs()
.expect("a daemon cannot be constructed without a blob store")
}
async fn readiness(&self) -> Result<Readiness> {
let mut conn = self.connection().await?;
let epoch = epoch_of(&mut conn)
.await
.map_err(|e| KernelError::Config(format!("read the epoch: {e}")))?;
let watermark = self
.store
.watermark()
.await
.map_err(|e| KernelError::Config(format!("read the watermark: {e}")))?;
Ok(Readiness {
sealed: epoch == crate::epoch::Epoch::Sealed,
watermark,
})
}
async fn connection(&self) -> Result<sqlx::pool::PoolConnection<sqlx::Postgres>> {
self.store
.pool()
.acquire()
.await
.map_err(|e| KernelError::Config(format!("acquire a connection: {e}")))
}
async fn answer(
&self,
request_id: &RequestId,
request: &KernelRequest,
subs: &mut Subscriptions<'_>,
) -> KernelResult {
match self.try_answer(request_id, request, subs).await {
Ok(result) => result,
Err(e) => KernelResult::Error {
code: KernelErrorCode::Storage,
message: e.to_string(),
detail: None,
},
}
}
async fn try_answer(
&self,
request_id: &RequestId,
request: &KernelRequest,
subs: &mut Subscriptions<'_>,
) -> Result<KernelResult> {
let readiness = self.readiness().await?;
Ok(match request {
KernelRequest::Health {} => KernelResult::Health {
ready: true,
sealed: readiness.sealed,
},
KernelRequest::Status {} => KernelResult::Status {
sealed: readiness.sealed,
watermark: readiness.watermark,
writer_epoch: self.writer_epoch,
contract_version: CONTRACT_VERSION,
public_revision: self.public_revision.clone(),
},
KernelRequest::Watermark {} => KernelResult::Watermark {
watermark: readiness.watermark,
},
KernelRequest::VerifySealed {} => self.verify_sealed(readiness.sealed).await?,
KernelRequest::SubmitCommand { envelope } => self.store.submit(envelope).await,
KernelRequest::GetProjection { projection, id } => {
let (mut records, _, _) =
self.projection_page(*projection, None, Some(id), 1).await?;
match records.pop() {
Some(record) => KernelResult::Projection { record },
None => KernelResult::Error {
code: KernelErrorCode::NotFound,
message: format!("no {} with id {id:?}", projection.as_str()),
detail: None,
},
}
}
KernelRequest::ListProjection {
projection,
cursor,
limit,
} => {
let rows = limit
.unwrap_or(DEFAULT_PAGE_ROWS)
.clamp(1, MAX_READ_LIMIT as u32);
let (records, cut, key) = self
.projection_page(*projection, cursor.as_deref(), None, rows)
.await?;
let exhausted = !cut && (records.len() as u32) < rows;
let next_cursor = if exhausted {
None
} else {
records
.last()
.map(|record| cursor_key(record, key))
.transpose()?
};
KernelResult::ProjectionPage {
records,
next_cursor,
}
}
KernelRequest::ReadEvents { cursor, limit } => {
let events = self
.store
.read_from(*cursor, *limit as usize)
.await
.map_err(|e| KernelError::Config(format!("read events: {e}")))?;
let (events, _) = fit_page(events)?;
KernelResult::Events {
cursor: events.last().map(|e| e.global_sequence),
watermark: readiness.watermark,
events,
}
}
KernelRequest::SubscribeEvents { cursor } => subs.admit(request_id, *cursor),
KernelRequest::BlobBegin {
media_type,
byte_size,
} => match self.blobs().begin(media_type.clone(), *byte_size).await {
Ok(upload_id) => KernelResult::BlobBegun { upload_id },
Err(e) => blob_refusal(&e),
},
KernelRequest::BlobChunk {
upload_id,
sequence,
data_base64,
} => self.blob_chunk(upload_id, *sequence, data_base64).await,
KernelRequest::BlobCommit { upload_id, address } => {
match self
.blobs()
.commit(upload_id.clone(), address.clone())
.await
{
Ok((descriptor, deduplicated)) => KernelResult::BlobCommitted {
descriptor,
deduplicated,
},
Err(e) => blob_refusal(&e),
}
}
KernelRequest::BlobAbort { upload_id } => {
match self.blobs().abort(upload_id.clone()).await {
Ok(()) => KernelResult::BlobAborted {
upload_id: upload_id.clone(),
},
Err(e) => blob_refusal(&e),
}
}
KernelRequest::BlobRead {
address,
offset,
length,
} => {
let length = ByteCount::new(length.value().min(BLOB_CHUNK_BYTES as u64));
match self.blobs().read(address, *offset, length).await {
Ok(bytes) => KernelResult::BlobBytes {
address: address.clone(),
offset: *offset,
data_base64: BASE64_STANDARD.encode(&bytes),
},
Err(e) => blob_refusal(&e),
}
}
KernelRequest::BlobStat { address } => match self.blobs().stat(address).await {
Ok(Some(descriptor)) => KernelResult::BlobStat { descriptor },
Ok(None) => KernelResult::Error {
code: KernelErrorCode::NotFound,
message: format!("no blob at {address}"),
detail: None,
},
Err(e) => blob_refusal(&e),
},
})
}
async fn blob_chunk(
&self,
upload_id: &gwk_domain::ids::BlobUploadId,
sequence: u32,
data_base64: &str,
) -> KernelResult {
let chunk = match BASE64_STANDARD.decode(data_base64) {
Ok(chunk) => chunk,
Err(e) => {
return KernelResult::Error {
code: KernelErrorCode::Validation,
message: format!("chunk {sequence} is not valid base64: {e}"),
detail: None,
};
}
};
if chunk.len() > BLOB_CHUNK_BYTES {
return KernelResult::Error {
code: KernelErrorCode::Validation,
message: format!(
"chunk {sequence} carries {} bytes, over the {BLOB_CHUNK_BYTES}-byte chunk",
chunk.len()
),
detail: None,
};
}
match self.blobs().write_chunk(upload_id, sequence, &chunk).await {
Ok(()) => KernelResult::BlobChunkAccepted {
upload_id: upload_id.clone(),
sequence,
},
Err(e) => blob_refusal(&e),
}
}
async fn projection_page(
&self,
kind: ProjectionKind,
cursor: Option<&str>,
exact: Option<&str>,
rows: u32,
) -> Result<(Vec<ProjectionRecord>, bool, &'static str)> {
let (query, key) = crate::checkpoint::read_query(kind).ok_or_else(|| {
KernelError::Config(format!("projection {} has no table", kind.as_str()))
})?;
let mut conn = self.connection().await?;
let raw: Vec<String> = sqlx::query_scalar(query)
.bind(cursor)
.bind(exact)
.bind(i64::from(rows))
.fetch_all(&mut *conn)
.await
.map_err(|e| KernelError::Config(format!("read {} projections: {e}", kind.as_str())))?;
let mut records = Vec::with_capacity(raw.len());
for line in raw {
records.push(
serde_json::from_str::<ProjectionRecord>(&line).map_err(|e| {
KernelError::Config(format!(
"a {} row does not match the contract type: {e}",
kind.as_str()
))
})?,
);
}
let (records, cut) = fit_page(records)?;
Ok((records, cut, key))
}
async fn verify_sealed(&self, sealed: bool) -> Result<KernelResult> {
let mut conn = self.connection().await?;
let row = sqlx::query(
"SELECT event_id, seq::text AS seq_text FROM gwk.event \
WHERE aggregate_type = $1 AND aggregate_id = $2 AND event_type = $3 \
ORDER BY seq LIMIT 1",
)
.bind(KERNEL_AGGREGATE)
.bind(KERNEL_SINGLETON)
.bind(GENESIS_EVENT_TYPE)
.fetch_optional(&mut *conn)
.await
.map_err(|e| KernelError::Config(format!("read genesis: {e}")))?
.ok_or_else(|| KernelError::Config("the log has no genesis event".to_owned()))?;
let event_id: String = row
.try_get("event_id")
.map_err(|e| KernelError::Config(format!("genesis event_id: {e}")))?;
let seq_text: String = row
.try_get("seq_text")
.map_err(|e| KernelError::Config(format!("genesis seq: {e}")))?;
let genesis_watermark = crate::numeric::from_numeric_text(&seq_text)
.map(Seq::new)
.map_err(|e| KernelError::Config(format!("genesis seq: {e}")))?;
let count: i64 = sqlx::query_scalar("SELECT count(*) FROM gwk.event")
.fetch_one(&mut *conn)
.await
.map_err(|e| KernelError::Config(format!("count events: {e}")))?;
Ok(KernelResult::SealedVerification {
sealed,
genesis_event_id: EventId::new(event_id),
genesis_watermark,
event_count: EventCount::new(count.max(0) as u64),
})
}
}
struct Subscriptions<'a> {
daemon: &'a Daemon,
live: tokio::task::JoinSet<()>,
batches: mpsc::Sender<Outgoing>,
responses: mpsc::Sender<ServerControl>,
admitted: Option<(RequestId, Option<Seq>)>,
}
impl<'a> Subscriptions<'a> {
fn new(
daemon: &'a Daemon,
responses: mpsc::Sender<ServerControl>,
batches: mpsc::Sender<Outgoing>,
) -> Self {
Self {
daemon,
live: tokio::task::JoinSet::new(),
batches,
responses,
admitted: None,
}
}
fn admit(&mut self, request_id: &RequestId, cursor: Option<Seq>) -> KernelResult {
while self.live.try_join_next().is_some() {}
if self.live.len() >= MAX_SUBSCRIPTIONS_PER_CONNECTION {
return KernelResult::Error {
code: KernelErrorCode::Overloaded,
message: format!(
"this connection already holds {MAX_SUBSCRIPTIONS_PER_CONNECTION} subscriptions"
),
detail: None,
};
}
self.admitted = Some((request_id.clone(), cursor));
KernelResult::Subscribed { cursor }
}
fn start(&mut self) {
let Some((request_id, cursor)) = self.admitted.take() else {
return;
};
let delivered = Arc::new(std::sync::atomic::AtomicU64::new(
cursor.map_or(0, |seq| seq.value()),
));
self.live.spawn(subscribe::run(
Arc::clone(&self.daemon.store),
request_id,
cursor,
self.batches.clone(),
self.responses.clone(),
self.daemon.wake.subscribe(),
delivered,
));
}
}
fn blob_refusal(error: &BlobError) -> KernelResult {
let code = match error {
BlobError::NotFound => KernelErrorCode::NotFound,
BlobError::Tombstoned => KernelErrorCode::BlobTombstoned,
BlobError::DigestMismatch { .. } | BlobError::Integrity(_) => {
KernelErrorCode::BlobIntegrity
}
BlobError::Pinned => KernelErrorCode::Validation,
BlobError::Storage(_) => KernelErrorCode::Storage,
};
KernelResult::Error {
code,
message: error.to_string(),
detail: None,
}
}
pub async fn serve_connection<R, W>(
daemon: &Daemon,
reader: &mut R,
writer: &mut W,
) -> std::result::Result<(), WireError>
where
R: AsyncRead + Unpin,
W: AsyncWrite + Unpin,
{
let mut ingress = Budget::new(
CONNECTION_INGRESS_BYTES_PER_WINDOW,
CONNECTION_EGRESS_BYTES_PER_WINDOW,
);
let readiness = daemon
.readiness()
.await
.map_err(|e| WireError::new(KernelErrorCode::Storage, format!("readiness: {e}")))?;
hello::negotiate(reader, writer, &mut ingress, readiness).await?;
let mut egress = Budget::new(
CONNECTION_INGRESS_BYTES_PER_WINDOW,
CONNECTION_EGRESS_BYTES_PER_WINDOW,
);
let (responses_tx, responses_rx) = mpsc::channel(RESPONSE_QUEUE_DEPTH);
let (batches_tx, batches_rx) = mpsc::channel(BATCH_QUEUE_DEPTH);
let mut subs = Subscriptions::new(daemon, responses_tx, batches_tx);
tokio::select! {
read = read_requests(daemon, reader, &mut ingress, &mut subs) => read,
written = write_frames(writer, responses_rx, batches_rx, &mut egress) => written,
}
}
async fn read_requests<R>(
daemon: &Daemon,
reader: &mut R,
budget: &mut Budget,
subs: &mut Subscriptions<'_>,
) -> std::result::Result<(), WireError>
where
R: AsyncRead + Unpin,
{
loop {
let frame = match read_frame(reader, FRAME_BODY_MAX_BYTES, budget).await? {
Incoming::Frame(frame) => frame,
Incoming::Closed => return Ok(()),
};
if frame.kind != FrameKind::Json {
return Err(WireError::new(
KernelErrorCode::Handshake,
format!("kind {:?} carries no control value", frame.kind),
));
}
let control: gwk_domain::protocol::ClientControl = strict::decode(&frame.body)?;
let (request_id, request) = match control {
gwk_domain::protocol::ClientControl::Request {
request_id,
request,
} => (request_id, request),
gwk_domain::protocol::ClientControl::Hello { .. } => {
return Err(WireError::new(
KernelErrorCode::Handshake,
"a second hello on an established connection",
));
}
};
let result = daemon.answer(&request_id, &request, subs).await;
let response = ServerControl::Response { request_id, result };
if subs.responses.send(response).await.is_err() {
return Ok(());
}
subs.start();
}
}
async fn write_frames<W>(
writer: &mut W,
mut responses: mpsc::Receiver<ServerControl>,
mut batches: mpsc::Receiver<Outgoing>,
budget: &mut Budget,
) -> std::result::Result<(), WireError>
where
W: AsyncWrite + Unpin,
{
loop {
let outgoing = tokio::select! {
biased;
Some(control) = responses.recv() => Outgoing { control, delivered: None },
Some(outgoing) = batches.recv() => outgoing,
else => return Ok(()),
};
let body = serde_json::to_vec(&outgoing.control).map_err(|e| {
WireError::new(KernelErrorCode::Storage, format!("serialize a frame: {e}"))
})?;
write_frame(writer, FrameKind::Json, &body, budget).await?;
if let Some((cell, seq)) = outgoing.delivered {
cell.store(seq, std::sync::atomic::Ordering::Release);
}
}
}
#[derive(Debug, Default)]
pub struct Stopped {
pub checkpoint: Option<Seq>,
pub checkpoint_error: Option<String>,
}
pub async fn run<S>(listener: Listener, daemon: Arc<Daemon>, shutdown: S) -> Result<Stopped>
where
S: std::future::Future<Output = ()> + Send,
{
daemon.notify_on_append();
let mut connections = tokio::task::JoinSet::new();
let shutdown = std::pin::pin!(shutdown);
let mut shutdown = shutdown;
loop {
tokio::select! {
biased;
() = &mut shutdown => break,
accepted = listener.accept() => {
let (stream, _peer) = accepted?;
let daemon = Arc::clone(&daemon);
connections.spawn(async move {
let (mut reader, mut writer) = tokio::io::split(stream);
let _ = serve_connection(&daemon, &mut reader, &mut writer).await;
});
}
}
}
let drained = tokio::time::timeout(std::time::Duration::from_secs(DRAIN_TIMEOUT_SECS), async {
while connections.join_next().await.is_some() {}
})
.await;
if drained.is_err() {
connections.shutdown().await;
}
let mut stopped = Stopped::default();
match daemon.store.checkpoint_at_watermark().await {
Ok(at) => stopped.checkpoint = at,
Err(e) => stopped.checkpoint_error = Some(e.to_string()),
}
listener.remove();
Ok(stopped)
}
pub async fn serve_stream(
daemon: &Daemon,
stream: UnixStream,
) -> std::result::Result<(), WireError> {
let (mut reader, mut writer) = tokio::io::split(stream);
serve_connection(daemon, &mut reader, &mut writer).await
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_page_is_cut_by_bytes_and_never_to_nothing() {
let items: Vec<String> = (0..8_000).map(|i| format!("{i:0>512}")).collect();
let (kept, cut) = fit_page(items).expect("measure");
assert!(cut, "a page far past the budget was not cut");
let bytes: usize = kept
.iter()
.map(|s| serde_json::to_vec(s).expect("serialize").len())
.sum();
assert!(bytes <= PAGE_BYTE_BUDGET, "the kept page is over budget");
assert!(bytes > PAGE_BYTE_BUDGET / 2, "the cut threw away too much");
let (kept, cut) = fit_page(vec!["x".repeat(PAGE_BYTE_BUDGET * 2)]).expect("measure");
assert_eq!(kept.len(), 1);
assert!(
!cut,
"one item is not a cut page — there is nothing behind it"
);
}
}