1use std::pin::Pin;
4use std::sync::Arc;
5
6use corium_protocol::authz::{self, Access, Action, Guard, IdentityInterceptor, Principal};
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(_)
36 | NodeError::Log(_)
37 | NodeError::Lease(_)
38 | NodeError::GroupCommit(_) => Status::internal(error.to_string()),
39 }
40}
41
42type ItemStream = Pin<Box<dyn Stream<Item = Result<pb::SubscribeItem, Status>> + Send>>;
43
44pub(crate) fn subscription_stream(
50 state: &Arc<DbState>,
51 from_basis_t: u64,
52 heartbeat_interval_ms: u64,
53) -> ItemStream {
54 let mut live = state.stream_items();
55 let (schema, interner) = state.handshake_snapshot();
56 let basis = state.db().basis_t();
57 let index_basis = state.index_basis();
58 let (tx, rx) = tokio::sync::mpsc::channel::<Result<pb::SubscribeItem, Status>>(64);
59 let state = Arc::clone(state);
60 tokio::spawn(async move {
61 let send = |item: pb::subscribe_item::Item| {
62 let tx = tx.clone();
63 async move {
64 tx.send(Ok(pb::SubscribeItem { item: Some(item) }))
65 .await
66 .is_ok()
67 }
68 };
69 if !send(pb::subscribe_item::Item::Handshake(pb::Handshake {
70 basis_t: basis,
71 index_basis_t: index_basis,
72 schema,
73 heartbeat_interval_ms,
74 }))
75 .await
76 {
77 return;
78 }
79 let mut last_sent = from_basis_t;
80 if from_basis_t < basis {
81 let records = match state.tx_range(from_basis_t + 1, Some(basis + 1)).await {
82 Ok(records) => records,
83 Err(error) => {
84 let _ = tx.send(Err(to_status(&error))).await;
85 return;
86 }
87 };
88 for record in records {
89 let datoms = match codec::encode_datoms(&record.datoms, &interner) {
90 Ok(datoms) => datoms,
91 Err(error) => {
92 let _ = tx.send(Err(Status::internal(error.to_string()))).await;
93 return;
94 }
95 };
96 if !send(pb::subscribe_item::Item::Report(pb::TxReport {
97 t: record.t,
98 tx_instant: record.tx_instant,
99 datoms,
100 }))
101 .await
102 {
103 return;
104 }
105 last_sent = record.t;
106 }
107 }
108 loop {
109 match live.recv().await {
110 Ok(pb::subscribe_item::Item::Report(report)) => {
111 if report.t <= last_sent {
112 continue;
113 }
114 last_sent = report.t;
115 if !send(pb::subscribe_item::Item::Report(report)).await {
116 return;
117 }
118 }
119 Ok(item) => {
120 if !send(item).await {
121 return;
122 }
123 }
124 Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
125 let _ = tx
128 .send(Err(Status::data_loss("subscription lagged; resubscribe")))
129 .await;
130 return;
131 }
132 Err(tokio::sync::broadcast::error::RecvError::Closed) => return,
133 }
134 }
135 });
136 Box::pin(ReceiverStream::new(rx))
137}
138
139pub struct TransactorSvc(pub Arc<TransactorNode>, pub Guard);
141
142async fn authorize(guard: &Guard, principal: &Principal, access: Access) -> Result<(), Status> {
147 match guard.authorize(principal, &access).await {
148 Ok(None) => Ok(()),
149 Ok(Some(_filter)) => Err(Status::unimplemented(
150 "view filtering is not enforced on this surface yet",
151 )),
152 Err(error) => Err(error.to_status()),
153 }
154}
155
156#[tonic::async_trait]
157impl Transactor for TransactorSvc {
158 async fn transact(
159 &self,
160 request: Request<pb::TransactRequest>,
161 ) -> Result<Response<pb::TransactResponse>, Status> {
162 let principal = authz::principal(&request);
163 let request = request.into_inner();
164 check_version(request.protocol_version)?;
165 authorize(
166 &self.1,
167 &principal,
168 Access::on(Action::Transact, &request.db),
169 )
170 .await?;
171 self.0
172 .transact(&request.db, &request.tx_data)
173 .await
174 .map(Response::new)
175 .map_err(|error| to_status(&error))
176 }
177
178 type SubscribeStream = ItemStream;
179
180 async fn subscribe(
181 &self,
182 request: Request<pb::SubscribeRequest>,
183 ) -> Result<Response<Self::SubscribeStream>, Status> {
184 let principal = authz::principal(&request);
185 let request = request.into_inner();
186 check_version(request.protocol_version)?;
187 authorize(
188 &self.1,
189 &principal,
190 Access::on(Action::Subscribe, &request.db),
191 )
192 .await?;
193 let state = self
194 .0
195 .db_state(&request.db)
196 .await
197 .map_err(|error| to_status(&error))?;
198 let heartbeat_interval_ms =
199 u64::try_from(self.0.config().heartbeat_interval.as_millis()).unwrap_or(0);
200 Ok(Response::new(subscription_stream(
201 &state,
202 request.from_basis_t,
203 heartbeat_interval_ms,
204 )))
205 }
206
207 async fn sync(
208 &self,
209 request: Request<pb::SyncRequest>,
210 ) -> Result<Response<pb::SyncResponse>, Status> {
211 let principal = authz::principal(&request);
212 let request = request.into_inner();
213 authorize(
214 &self.1,
215 &principal,
216 Access::on(Action::Inspect, &request.db),
217 )
218 .await?;
219 let basis_t = self
220 .0
221 .sync(&request.db, request.t)
222 .await
223 .map_err(|error| to_status(&error))?;
224 Ok(Response::new(pb::SyncResponse { basis_t }))
225 }
226
227 async fn status(
228 &self,
229 request: Request<pb::StatusRequest>,
230 ) -> Result<Response<pb::StatusResponse>, Status> {
231 let principal = authz::principal(&request);
232 let request = request.into_inner();
233 authorize(
234 &self.1,
235 &principal,
236 Access::on(Action::Inspect, &request.db),
237 )
238 .await?;
239 self.0
240 .status(&request.db)
241 .await
242 .map(Response::new)
243 .map_err(|error| to_status(&error))
244 }
245}
246
247fn check_version(version: u32) -> Result<(), Status> {
248 if version == corium_protocol::PROTOCOL_VERSION {
249 Ok(())
250 } else {
251 Err(Status::failed_precondition(format!(
252 "protocol version {version} is not supported; upgrade to {}",
253 corium_protocol::PROTOCOL_VERSION
254 )))
255 }
256}
257
258pub struct CatalogSvc(pub Arc<TransactorNode>, pub Guard);
260
261#[tonic::async_trait]
262impl Catalog for CatalogSvc {
263 async fn create_database(
264 &self,
265 request: Request<pb::CreateDatabaseRequest>,
266 ) -> Result<Response<pb::CreateDatabaseResponse>, Status> {
267 let principal = authz::principal(&request);
268 let request = request.into_inner();
269 authorize(
270 &self.1,
271 &principal,
272 Access::on(Action::CreateDatabase, &request.db),
273 )
274 .await?;
275 let created = self
276 .0
277 .create_db(&request.db, &request.schema)
278 .await
279 .map_err(|error| to_status(&error))?;
280 Ok(Response::new(pb::CreateDatabaseResponse { created }))
281 }
282
283 async fn delete_database(
284 &self,
285 request: Request<pb::DeleteDatabaseRequest>,
286 ) -> Result<Response<pb::DeleteDatabaseResponse>, Status> {
287 let principal = authz::principal(&request);
288 let request = request.into_inner();
289 authorize(
290 &self.1,
291 &principal,
292 Access::on(Action::DeleteDatabase, &request.db),
293 )
294 .await?;
295 let deleted = self
296 .0
297 .delete_db(&request.db)
298 .await
299 .map_err(|error| to_status(&error))?;
300 Ok(Response::new(pb::DeleteDatabaseResponse { deleted }))
301 }
302
303 async fn fork_database(
304 &self,
305 request: Request<pb::ForkDatabaseRequest>,
306 ) -> Result<Response<pb::ForkDatabaseResponse>, Status> {
307 let principal = authz::principal(&request);
308 let request = request.into_inner();
309 authorize(
310 &self.1,
311 &principal,
312 Access::on(Action::ForkDatabase, &request.target),
313 )
314 .await?;
315 let forked = self
316 .0
317 .fork_db(&request.db, &request.target, request.as_of_t)
318 .await
319 .map_err(|error| to_status(&error))?;
320 Ok(Response::new(pb::ForkDatabaseResponse {
321 created: forked.is_some(),
322 basis_t: forked.unwrap_or(0),
323 }))
324 }
325
326 async fn list_databases(
327 &self,
328 request: Request<pb::ListDatabasesRequest>,
329 ) -> Result<Response<pb::ListDatabasesResponse>, Status> {
330 let principal = authz::principal(&request);
331 authorize(&self.1, &principal, Access::catalog(Action::ListDatabases)).await?;
332 Ok(Response::new(pb::ListDatabasesResponse {
333 dbs: self.0.list_dbs(),
334 }))
335 }
336
337 async fn gc_deleted_databases(
338 &self,
339 request: Request<pb::GcDeletedDatabasesRequest>,
340 ) -> Result<Response<pb::GcDeletedDatabasesResponse>, Status> {
341 let principal = authz::principal(&request);
342 authorize(&self.1, &principal, Access::catalog(Action::GarbageCollect)).await?;
343 let swept = match requested_gc_retention(request.into_inner()) {
344 None => self.0.gc_deleted().await,
345 Some(retention) => self.0.gc_deleted_with_retention(retention).await,
346 };
347 let swept_blobs = swept.map_err(|error| to_status(&error))?;
348 Ok(Response::new(pb::GcDeletedDatabasesResponse {
349 swept_blobs,
350 }))
351 }
352
353 async fn request_index(
354 &self,
355 request: Request<pb::RequestIndexRequest>,
356 ) -> Result<Response<pb::RequestIndexResponse>, Status> {
357 let principal = authz::principal(&request);
358 let request = request.into_inner();
359 authorize(
360 &self.1,
361 &principal,
362 Access::on(Action::ManageIndex, &request.db),
363 )
364 .await?;
365 let index_basis_t = self
366 .0
367 .request_index(&request.db)
368 .await
369 .map_err(|error| to_status(&error))?;
370 Ok(Response::new(pb::RequestIndexResponse { index_basis_t }))
371 }
372
373 async fn set_index_policy(
374 &self,
375 request: Request<pb::SetIndexPolicyRequest>,
376 ) -> Result<Response<pb::SetIndexPolicyResponse>, Status> {
377 let principal = authz::principal(&request);
378 let request = request.into_inner();
379 authorize(
380 &self.1,
381 &principal,
382 Access::on(Action::ManageIndex, &request.db),
383 )
384 .await?;
385 let update = IndexPolicyUpdate {
386 interval: request.interval_ms.map(std::time::Duration::from_millis),
387 backoff: request.backoff,
388 tail_threshold: request.tail_threshold,
389 tail_deadline: request
390 .tail_deadline_ms
391 .map(std::time::Duration::from_millis),
392 };
393 let policy = self
394 .0
395 .set_index_policy(&request.db, update)
396 .await
397 .map_err(|error| to_status(&error))?;
398 let millis =
399 |duration: std::time::Duration| u64::try_from(duration.as_millis()).unwrap_or(u64::MAX);
400 Ok(Response::new(pb::SetIndexPolicyResponse {
401 interval_ms: millis(policy.interval),
402 backoff: policy.backoff,
403 tail_threshold: policy.tail_threshold,
404 tail_deadline_ms: millis(policy.tail_deadline),
405 }))
406 }
407
408 async fn get_storage_info(
409 &self,
410 request: Request<pb::GetStorageInfoRequest>,
411 ) -> Result<Response<pb::GetStorageInfoResponse>, Status> {
412 let principal = authz::principal(&request);
413 let request = request.into_inner();
414 authorize(
415 &self.1,
416 &principal,
417 Access::on(Action::Inspect, &request.db),
418 )
419 .await?;
420 self.0
421 .backup_info(&request.db)
422 .await
423 .map(Response::new)
424 .map_err(|error| to_status(&error))
425 }
426}
427
428fn requested_gc_retention(request: pb::GcDeletedDatabasesRequest) -> Option<std::time::Duration> {
429 request
430 .retention_millis
431 .map(std::time::Duration::from_millis)
432}
433
434pub async fn serve(
439 node: Arc<TransactorNode>,
440 addr: std::net::SocketAddr,
441 guard: Guard,
442 tls: Option<tonic::transport::ServerTlsConfig>,
443 shutdown: impl std::future::Future<Output = ()> + Send,
444) -> Result<(), tonic::transport::Error> {
445 let interceptor = IdentityInterceptor::new(guard.clone());
448 let mut builder = tonic::transport::Server::builder();
449 if let Some(tls) = tls {
450 builder = builder.tls_config(tls)?;
451 }
452 builder
453 .add_service(TransactorServer::with_interceptor(
454 TransactorSvc(Arc::clone(&node), guard.clone()),
455 interceptor.clone(),
456 ))
457 .add_service(CatalogServer::with_interceptor(
458 CatalogSvc(node, guard),
459 interceptor,
460 ))
461 .serve_with_shutdown(addr, shutdown)
462 .await
463}
464
465#[cfg(test)]
466mod tests {
467 use super::*;
468
469 #[test]
470 fn gc_retention_distinguishes_default_zero_and_subsecond() {
471 let default = pb::GcDeletedDatabasesRequest {
472 retention_millis: None,
473 };
474 let immediate = pb::GcDeletedDatabasesRequest {
475 retention_millis: Some(0),
476 };
477 let subsecond = pb::GcDeletedDatabasesRequest {
478 retention_millis: Some(500),
479 };
480
481 assert_eq!(requested_gc_retention(default), None);
482 assert_eq!(
483 requested_gc_retention(immediate),
484 Some(std::time::Duration::ZERO)
485 );
486 assert_eq!(
487 requested_gc_retention(subsecond),
488 Some(std::time::Duration::from_millis(500))
489 );
490 }
491}