1use 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, IndexPolicyUpdate, 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 fork_database(
258 &self,
259 request: Request<pb::ForkDatabaseRequest>,
260 ) -> Result<Response<pb::ForkDatabaseResponse>, Status> {
261 let request = request.into_inner();
262 let forked = self
263 .0
264 .fork_db(&request.db, &request.target, request.as_of_t)
265 .await
266 .map_err(|error| to_status(&error))?;
267 Ok(Response::new(pb::ForkDatabaseResponse {
268 created: forked.is_some(),
269 basis_t: forked.unwrap_or(0),
270 }))
271 }
272
273 async fn list_databases(
274 &self,
275 _request: Request<pb::ListDatabasesRequest>,
276 ) -> Result<Response<pb::ListDatabasesResponse>, Status> {
277 Ok(Response::new(pb::ListDatabasesResponse {
278 dbs: self.0.list_dbs(),
279 }))
280 }
281
282 async fn gc_deleted_databases(
283 &self,
284 request: Request<pb::GcDeletedDatabasesRequest>,
285 ) -> Result<Response<pb::GcDeletedDatabasesResponse>, Status> {
286 let swept = match requested_gc_retention(request.into_inner()) {
287 None => self.0.gc_deleted().await,
288 Some(retention) => self.0.gc_deleted_with_retention(retention).await,
289 };
290 let swept_blobs = swept.map_err(|error| to_status(&error))?;
291 Ok(Response::new(pb::GcDeletedDatabasesResponse {
292 swept_blobs,
293 }))
294 }
295
296 async fn request_index(
297 &self,
298 request: Request<pb::RequestIndexRequest>,
299 ) -> Result<Response<pb::RequestIndexResponse>, Status> {
300 let request = request.into_inner();
301 let index_basis_t = self
302 .0
303 .request_index(&request.db)
304 .await
305 .map_err(|error| to_status(&error))?;
306 Ok(Response::new(pb::RequestIndexResponse { index_basis_t }))
307 }
308
309 async fn set_index_policy(
310 &self,
311 request: Request<pb::SetIndexPolicyRequest>,
312 ) -> Result<Response<pb::SetIndexPolicyResponse>, Status> {
313 let request = request.into_inner();
314 let update = IndexPolicyUpdate {
315 interval: request.interval_ms.map(std::time::Duration::from_millis),
316 backoff: request.backoff,
317 tail_threshold: request.tail_threshold,
318 tail_deadline: request
319 .tail_deadline_ms
320 .map(std::time::Duration::from_millis),
321 };
322 let policy = self
323 .0
324 .set_index_policy(&request.db, update)
325 .await
326 .map_err(|error| to_status(&error))?;
327 let millis =
328 |duration: std::time::Duration| u64::try_from(duration.as_millis()).unwrap_or(u64::MAX);
329 Ok(Response::new(pb::SetIndexPolicyResponse {
330 interval_ms: millis(policy.interval),
331 backoff: policy.backoff,
332 tail_threshold: policy.tail_threshold,
333 tail_deadline_ms: millis(policy.tail_deadline),
334 }))
335 }
336}
337
338fn requested_gc_retention(request: pb::GcDeletedDatabasesRequest) -> Option<std::time::Duration> {
339 request
340 .retention_millis
341 .map(std::time::Duration::from_millis)
342}
343
344pub async fn serve(
349 node: Arc<TransactorNode>,
350 addr: std::net::SocketAddr,
351 authenticator: Arc<dyn Authenticator>,
352 tls: Option<tonic::transport::ServerTlsConfig>,
353 shutdown: impl std::future::Future<Output = ()> + Send,
354) -> Result<(), tonic::transport::Error> {
355 let interceptor = AuthInterceptor::new(authenticator);
356 let mut builder = tonic::transport::Server::builder();
357 if let Some(tls) = tls {
358 builder = builder.tls_config(tls)?;
359 }
360 builder
361 .add_service(TransactorServer::with_interceptor(
362 TransactorSvc(Arc::clone(&node)),
363 interceptor.clone(),
364 ))
365 .add_service(CatalogServer::with_interceptor(
366 CatalogSvc(node),
367 interceptor,
368 ))
369 .serve_with_shutdown(addr, shutdown)
370 .await
371}
372
373#[cfg(test)]
374mod tests {
375 use super::*;
376
377 #[test]
378 fn gc_retention_distinguishes_default_zero_and_subsecond() {
379 let default = pb::GcDeletedDatabasesRequest {
380 retention_millis: None,
381 };
382 let immediate = pb::GcDeletedDatabasesRequest {
383 retention_millis: Some(0),
384 };
385 let subsecond = pb::GcDeletedDatabasesRequest {
386 retention_millis: Some(500),
387 };
388
389 assert_eq!(requested_gc_retention(default), None);
390 assert_eq!(
391 requested_gc_retention(immediate),
392 Some(std::time::Duration::ZERO)
393 );
394 assert_eq!(
395 requested_gc_retention(subsecond),
396 Some(std::time::Duration::from_millis(500))
397 );
398 }
399}