corium_transactor/
server.rs1use std::pin::Pin;
4use std::sync::Arc;
5
6use corium_protocol::auth::{AuthInterceptor, Authenticator};
7use corium_protocol::codec;
8use corium_protocol::pb;
9use corium_protocol::pb::catalog_server::{Catalog, CatalogServer};
10use corium_protocol::pb::transactor_server::{Transactor, TransactorServer};
11use tokio_stream::Stream;
12use tokio_stream::wrappers::ReceiverStream;
13use tonic::{Request, Response, Status};
14
15use crate::node::{DbState, NodeError, TransactorNode};
16
17#[must_use]
19pub fn to_status(error: &NodeError) -> Status {
20 match error {
21 NodeError::UnknownDb(name) => Status::not_found(format!("unknown database {name:?}")),
22 NodeError::InvalidName(_)
23 | NodeError::BadRequest(_)
24 | NodeError::Codec(_)
25 | NodeError::TxForm(_)
26 | NodeError::SchemaForm(_) => Status::invalid_argument(error.to_string()),
27 NodeError::Deposed(_) | NodeError::Standby { .. } | NodeError::UnsupportedFormat { .. } => {
28 Status::failed_precondition(error.to_string())
29 }
30 NodeError::Transact(inner) => match inner {
31 crate::TransactError::Tx(_) => Status::invalid_argument(inner.to_string()),
32 crate::TransactError::Deposed { .. } => Status::failed_precondition(inner.to_string()),
33 _ => Status::internal(inner.to_string()),
34 },
35 NodeError::Store(_) | NodeError::Log(_) | NodeError::Lease(_) => {
36 Status::internal(error.to_string())
37 }
38 }
39}
40
41type ItemStream = Pin<Box<dyn Stream<Item = Result<pb::SubscribeItem, Status>> + Send>>;
42
43pub(crate) fn subscription_stream(
49 state: &Arc<DbState>,
50 from_basis_t: u64,
51 heartbeat_interval_ms: u64,
52) -> ItemStream {
53 let mut live = state.stream_items();
54 let (schema, interner) = state.handshake_snapshot();
55 let basis = state.db().basis_t();
56 let index_basis = state.index_basis();
57 let (tx, rx) = tokio::sync::mpsc::channel::<Result<pb::SubscribeItem, Status>>(64);
58 let state = Arc::clone(state);
59 tokio::spawn(async move {
60 let send = |item: pb::subscribe_item::Item| {
61 let tx = tx.clone();
62 async move {
63 tx.send(Ok(pb::SubscribeItem { item: Some(item) }))
64 .await
65 .is_ok()
66 }
67 };
68 if !send(pb::subscribe_item::Item::Handshake(pb::Handshake {
69 basis_t: basis,
70 index_basis_t: index_basis,
71 schema,
72 heartbeat_interval_ms,
73 }))
74 .await
75 {
76 return;
77 }
78 let mut last_sent = from_basis_t;
79 if from_basis_t < basis {
80 let backfill = {
81 let state = Arc::clone(&state);
82 tokio::task::spawn_blocking(move || {
83 state.tx_range(from_basis_t + 1, Some(basis + 1))
84 })
85 .await
86 };
87 let records = match backfill {
88 Ok(Ok(records)) => records,
89 Ok(Err(error)) => {
90 let _ = tx.send(Err(to_status(&error))).await;
91 return;
92 }
93 Err(error) => {
94 let _ = tx.send(Err(Status::internal(error.to_string()))).await;
95 return;
96 }
97 };
98 for record in records {
99 let datoms = match codec::encode_datoms(&record.datoms, &interner) {
100 Ok(datoms) => datoms,
101 Err(error) => {
102 let _ = tx.send(Err(Status::internal(error.to_string()))).await;
103 return;
104 }
105 };
106 if !send(pb::subscribe_item::Item::Report(pb::TxReport {
107 t: record.t,
108 tx_instant: record.tx_instant,
109 datoms,
110 }))
111 .await
112 {
113 return;
114 }
115 last_sent = record.t;
116 }
117 }
118 loop {
119 match live.recv().await {
120 Ok(pb::subscribe_item::Item::Report(report)) => {
121 if report.t <= last_sent {
122 continue;
123 }
124 last_sent = report.t;
125 if !send(pb::subscribe_item::Item::Report(report)).await {
126 return;
127 }
128 }
129 Ok(item) => {
130 if !send(item).await {
131 return;
132 }
133 }
134 Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
135 let _ = tx
138 .send(Err(Status::data_loss("subscription lagged; resubscribe")))
139 .await;
140 return;
141 }
142 Err(tokio::sync::broadcast::error::RecvError::Closed) => return,
143 }
144 }
145 });
146 Box::pin(ReceiverStream::new(rx))
147}
148
149pub struct TransactorSvc(pub Arc<TransactorNode>);
151
152#[tonic::async_trait]
153impl Transactor for TransactorSvc {
154 async fn transact(
155 &self,
156 request: Request<pb::TransactRequest>,
157 ) -> Result<Response<pb::TransactResponse>, Status> {
158 let request = request.into_inner();
159 check_version(request.protocol_version)?;
160 self.0
161 .transact(&request.db, &request.tx_data)
162 .await
163 .map(Response::new)
164 .map_err(|error| to_status(&error))
165 }
166
167 type SubscribeStream = ItemStream;
168
169 async fn subscribe(
170 &self,
171 request: Request<pb::SubscribeRequest>,
172 ) -> Result<Response<Self::SubscribeStream>, Status> {
173 let request = request.into_inner();
174 check_version(request.protocol_version)?;
175 let state = self
176 .0
177 .db_state(&request.db)
178 .await
179 .map_err(|error| to_status(&error))?;
180 let heartbeat_interval_ms =
181 u64::try_from(self.0.config().heartbeat_interval.as_millis()).unwrap_or(0);
182 Ok(Response::new(subscription_stream(
183 &state,
184 request.from_basis_t,
185 heartbeat_interval_ms,
186 )))
187 }
188
189 async fn sync(
190 &self,
191 request: Request<pb::SyncRequest>,
192 ) -> Result<Response<pb::SyncResponse>, Status> {
193 let request = request.into_inner();
194 let basis_t = self
195 .0
196 .sync(&request.db, request.t)
197 .await
198 .map_err(|error| to_status(&error))?;
199 Ok(Response::new(pb::SyncResponse { basis_t }))
200 }
201
202 async fn status(
203 &self,
204 request: Request<pb::StatusRequest>,
205 ) -> Result<Response<pb::StatusResponse>, Status> {
206 let request = request.into_inner();
207 self.0
208 .status(&request.db)
209 .await
210 .map(Response::new)
211 .map_err(|error| to_status(&error))
212 }
213}
214
215fn check_version(version: u32) -> Result<(), Status> {
216 if version == corium_protocol::PROTOCOL_VERSION {
217 Ok(())
218 } else {
219 Err(Status::failed_precondition(format!(
220 "protocol version {version} is not supported; upgrade to {}",
221 corium_protocol::PROTOCOL_VERSION
222 )))
223 }
224}
225
226pub struct CatalogSvc(pub Arc<TransactorNode>);
228
229#[tonic::async_trait]
230impl Catalog for CatalogSvc {
231 async fn create_database(
232 &self,
233 request: Request<pb::CreateDatabaseRequest>,
234 ) -> Result<Response<pb::CreateDatabaseResponse>, Status> {
235 let request = request.into_inner();
236 let created = self
237 .0
238 .create_db(&request.db, &request.schema)
239 .await
240 .map_err(|error| to_status(&error))?;
241 Ok(Response::new(pb::CreateDatabaseResponse { created }))
242 }
243
244 async fn delete_database(
245 &self,
246 request: Request<pb::DeleteDatabaseRequest>,
247 ) -> Result<Response<pb::DeleteDatabaseResponse>, Status> {
248 let request = request.into_inner();
249 let deleted = self
250 .0
251 .delete_db(&request.db)
252 .await
253 .map_err(|error| to_status(&error))?;
254 Ok(Response::new(pb::DeleteDatabaseResponse { deleted }))
255 }
256
257 async fn list_databases(
258 &self,
259 _request: Request<pb::ListDatabasesRequest>,
260 ) -> Result<Response<pb::ListDatabasesResponse>, Status> {
261 Ok(Response::new(pb::ListDatabasesResponse {
262 dbs: self.0.list_dbs(),
263 }))
264 }
265
266 async fn gc_deleted_databases(
267 &self,
268 request: Request<pb::GcDeletedDatabasesRequest>,
269 ) -> Result<Response<pb::GcDeletedDatabasesResponse>, Status> {
270 let swept = match requested_gc_retention(request.into_inner()) {
271 None => self.0.gc_deleted().await,
272 Some(retention) => self.0.gc_deleted_with_retention(retention).await,
273 };
274 let swept_blobs = swept.map_err(|error| to_status(&error))?;
275 Ok(Response::new(pb::GcDeletedDatabasesResponse {
276 swept_blobs,
277 }))
278 }
279}
280
281fn requested_gc_retention(request: pb::GcDeletedDatabasesRequest) -> Option<std::time::Duration> {
282 request
283 .retention_millis
284 .map(std::time::Duration::from_millis)
285}
286
287pub async fn serve(
292 node: Arc<TransactorNode>,
293 addr: std::net::SocketAddr,
294 authenticator: Arc<dyn Authenticator>,
295 tls: Option<tonic::transport::ServerTlsConfig>,
296 shutdown: impl std::future::Future<Output = ()> + Send,
297) -> Result<(), tonic::transport::Error> {
298 let interceptor = AuthInterceptor::new(authenticator);
299 let mut builder = tonic::transport::Server::builder();
300 if let Some(tls) = tls {
301 builder = builder.tls_config(tls)?;
302 }
303 builder
304 .add_service(TransactorServer::with_interceptor(
305 TransactorSvc(Arc::clone(&node)),
306 interceptor.clone(),
307 ))
308 .add_service(CatalogServer::with_interceptor(
309 CatalogSvc(node),
310 interceptor,
311 ))
312 .serve_with_shutdown(addr, shutdown)
313 .await
314}
315
316#[cfg(test)]
317mod tests {
318 use super::*;
319
320 #[test]
321 fn gc_retention_distinguishes_default_zero_and_subsecond() {
322 let default = pb::GcDeletedDatabasesRequest {
323 retention_millis: None,
324 };
325 let immediate = pb::GcDeletedDatabasesRequest {
326 retention_millis: Some(0),
327 };
328 let subsecond = pb::GcDeletedDatabasesRequest {
329 retention_millis: Some(500),
330 };
331
332 assert_eq!(requested_gc_retention(default), None);
333 assert_eq!(
334 requested_gc_retention(immediate),
335 Some(std::time::Duration::ZERO)
336 );
337 assert_eq!(
338 requested_gc_retention(subsecond),
339 Some(std::time::Duration::from_millis(500))
340 );
341 }
342}