use std::pin::Pin;
use std::sync::Arc;
use corium_protocol::auth::{AuthInterceptor, Authenticator};
use corium_protocol::codec;
use corium_protocol::pb;
use corium_protocol::pb::catalog_server::{Catalog, CatalogServer};
use corium_protocol::pb::transactor_server::{Transactor, TransactorServer};
use tokio_stream::Stream;
use tokio_stream::wrappers::ReceiverStream;
use tonic::{Request, Response, Status};
use crate::node::{DbState, NodeError, TransactorNode};
#[must_use]
pub fn to_status(error: &NodeError) -> Status {
match error {
NodeError::UnknownDb(name) => Status::not_found(format!("unknown database {name:?}")),
NodeError::InvalidName(_)
| NodeError::BadRequest(_)
| NodeError::Codec(_)
| NodeError::TxForm(_)
| NodeError::SchemaForm(_) => Status::invalid_argument(error.to_string()),
NodeError::Deposed(_) | NodeError::UnsupportedFormat { .. } => {
Status::failed_precondition(error.to_string())
}
NodeError::Transact(inner) => match inner {
crate::TransactError::Tx(_) => Status::invalid_argument(inner.to_string()),
crate::TransactError::Deposed { .. } => Status::failed_precondition(inner.to_string()),
_ => Status::internal(inner.to_string()),
},
NodeError::Store(_) | NodeError::Log(_) | NodeError::Lease(_) => {
Status::internal(error.to_string())
}
}
}
type ItemStream = Pin<Box<dyn Stream<Item = Result<pb::SubscribeItem, Status>> + Send>>;
pub(crate) fn subscription_stream(state: &Arc<DbState>, from_basis_t: u64) -> ItemStream {
let mut live = state.stream_items();
let (schema, interner) = state.handshake_snapshot();
let basis = state.db().basis_t();
let index_basis = state.index_basis();
let (tx, rx) = tokio::sync::mpsc::channel::<Result<pb::SubscribeItem, Status>>(64);
let state = Arc::clone(state);
tokio::spawn(async move {
let send = |item: pb::subscribe_item::Item| {
let tx = tx.clone();
async move {
tx.send(Ok(pb::SubscribeItem { item: Some(item) }))
.await
.is_ok()
}
};
if !send(pb::subscribe_item::Item::Handshake(pb::Handshake {
basis_t: basis,
index_basis_t: index_basis,
schema,
}))
.await
{
return;
}
let mut last_sent = from_basis_t;
if from_basis_t < basis {
let backfill = {
let state = Arc::clone(&state);
tokio::task::spawn_blocking(move || {
state.tx_range(from_basis_t + 1, Some(basis + 1))
})
.await
};
let records = match backfill {
Ok(Ok(records)) => records,
Ok(Err(error)) => {
let _ = tx.send(Err(to_status(&error))).await;
return;
}
Err(error) => {
let _ = tx.send(Err(Status::internal(error.to_string()))).await;
return;
}
};
for record in records {
let datoms = match codec::encode_datoms(&record.datoms, &interner) {
Ok(datoms) => datoms,
Err(error) => {
let _ = tx.send(Err(Status::internal(error.to_string()))).await;
return;
}
};
if !send(pb::subscribe_item::Item::Report(pb::TxReport {
t: record.t,
tx_instant: record.tx_instant,
datoms,
}))
.await
{
return;
}
last_sent = record.t;
}
}
loop {
match live.recv().await {
Ok(pb::subscribe_item::Item::Report(report)) => {
if report.t <= last_sent {
continue;
}
last_sent = report.t;
if !send(pb::subscribe_item::Item::Report(report)).await {
return;
}
}
Ok(item) => {
if !send(item).await {
return;
}
}
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
let _ = tx
.send(Err(Status::data_loss("subscription lagged; resubscribe")))
.await;
return;
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => return,
}
}
});
Box::pin(ReceiverStream::new(rx))
}
pub struct TransactorSvc(pub Arc<TransactorNode>);
#[tonic::async_trait]
impl Transactor for TransactorSvc {
async fn transact(
&self,
request: Request<pb::TransactRequest>,
) -> Result<Response<pb::TransactResponse>, Status> {
let request = request.into_inner();
check_version(request.protocol_version)?;
self.0
.transact(&request.db, &request.tx_data)
.await
.map(Response::new)
.map_err(|error| to_status(&error))
}
type SubscribeStream = ItemStream;
async fn subscribe(
&self,
request: Request<pb::SubscribeRequest>,
) -> Result<Response<Self::SubscribeStream>, Status> {
let request = request.into_inner();
check_version(request.protocol_version)?;
let state = self
.0
.db_state(&request.db)
.map_err(|error| to_status(&error))?;
Ok(Response::new(subscription_stream(
&state,
request.from_basis_t,
)))
}
async fn sync(
&self,
request: Request<pb::SyncRequest>,
) -> Result<Response<pb::SyncResponse>, Status> {
let request = request.into_inner();
let basis_t = self
.0
.sync(&request.db, request.t)
.await
.map_err(|error| to_status(&error))?;
Ok(Response::new(pb::SyncResponse { basis_t }))
}
async fn status(
&self,
request: Request<pb::StatusRequest>,
) -> Result<Response<pb::StatusResponse>, Status> {
let request = request.into_inner();
self.0
.status(&request.db)
.map(Response::new)
.map_err(|error| to_status(&error))
}
}
fn check_version(version: u32) -> Result<(), Status> {
if version == corium_protocol::PROTOCOL_VERSION {
Ok(())
} else {
Err(Status::failed_precondition(format!(
"protocol version {version} is not supported; upgrade to {}",
corium_protocol::PROTOCOL_VERSION
)))
}
}
pub struct CatalogSvc(pub Arc<TransactorNode>);
#[tonic::async_trait]
impl Catalog for CatalogSvc {
async fn create_database(
&self,
request: Request<pb::CreateDatabaseRequest>,
) -> Result<Response<pb::CreateDatabaseResponse>, Status> {
let request = request.into_inner();
let node = Arc::clone(&self.0);
let created =
tokio::task::spawn_blocking(move || node.create_db(&request.db, &request.schema))
.await
.map_err(|error| Status::internal(error.to_string()))?
.map_err(|error| to_status(&error))?;
Ok(Response::new(pb::CreateDatabaseResponse { created }))
}
async fn delete_database(
&self,
request: Request<pb::DeleteDatabaseRequest>,
) -> Result<Response<pb::DeleteDatabaseResponse>, Status> {
let request = request.into_inner();
let node = Arc::clone(&self.0);
let deleted = tokio::task::spawn_blocking(move || node.delete_db(&request.db))
.await
.map_err(|error| Status::internal(error.to_string()))?
.map_err(|error| to_status(&error))?;
Ok(Response::new(pb::DeleteDatabaseResponse { deleted }))
}
async fn list_databases(
&self,
_request: Request<pb::ListDatabasesRequest>,
) -> Result<Response<pb::ListDatabasesResponse>, Status> {
Ok(Response::new(pb::ListDatabasesResponse {
dbs: self.0.list_dbs(),
}))
}
async fn gc_deleted_databases(
&self,
request: Request<pb::GcDeletedDatabasesRequest>,
) -> Result<Response<pb::GcDeletedDatabasesResponse>, Status> {
let swept = match requested_gc_retention(request.into_inner()) {
None => self.0.gc_deleted().await,
Some(retention) => self.0.gc_deleted_with_retention(retention).await,
};
let swept_blobs = swept.map_err(|error| to_status(&error))?;
Ok(Response::new(pb::GcDeletedDatabasesResponse {
swept_blobs,
}))
}
}
fn requested_gc_retention(request: pb::GcDeletedDatabasesRequest) -> Option<std::time::Duration> {
request
.retention_millis
.map(std::time::Duration::from_millis)
}
pub async fn serve(
node: Arc<TransactorNode>,
addr: std::net::SocketAddr,
authenticator: Arc<dyn Authenticator>,
tls: Option<tonic::transport::ServerTlsConfig>,
shutdown: impl std::future::Future<Output = ()> + Send,
) -> Result<(), tonic::transport::Error> {
let interceptor = AuthInterceptor::new(authenticator);
let mut builder = tonic::transport::Server::builder();
if let Some(tls) = tls {
builder = builder.tls_config(tls)?;
}
builder
.add_service(TransactorServer::with_interceptor(
TransactorSvc(Arc::clone(&node)),
interceptor.clone(),
))
.add_service(CatalogServer::with_interceptor(
CatalogSvc(node),
interceptor,
))
.serve_with_shutdown(addr, shutdown)
.await
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn gc_retention_distinguishes_default_zero_and_subsecond() {
let default = pb::GcDeletedDatabasesRequest {
retention_millis: None,
};
let immediate = pb::GcDeletedDatabasesRequest {
retention_millis: Some(0),
};
let subsecond = pb::GcDeletedDatabasesRequest {
retention_millis: Some(500),
};
assert_eq!(requested_gc_retention(default), None);
assert_eq!(
requested_gc_retention(immediate),
Some(std::time::Duration::ZERO)
);
assert_eq!(
requested_gc_retention(subsecond),
Some(std::time::Duration::from_millis(500))
);
}
}