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::BasisMismatch { .. } => Status::aborted(error.to_string()),
28 NodeError::Deposed(_)
29 | NodeError::Standby { .. }
30 | NodeError::UnsupportedFormat { .. }
31 | NodeError::Keys(_)
34 | NodeError::KeysFenced { .. } => Status::failed_precondition(error.to_string()),
35 NodeError::Transact(inner) => match inner {
36 crate::TransactError::Tx(_) => Status::invalid_argument(inner.to_string()),
37 crate::TransactError::Deposed { .. } => Status::failed_precondition(inner.to_string()),
38 _ => Status::internal(inner.to_string()),
39 },
40 NodeError::Store(_)
41 | NodeError::Log(_)
42 | NodeError::Lease(_)
43 | NodeError::GroupCommit(_) => Status::internal(error.to_string()),
44 }
45}
46
47type ItemStream = Pin<Box<dyn Stream<Item = Result<pb::SubscribeItem, Status>> + Send>>;
48
49pub(crate) fn subscription_stream(
55 state: &Arc<DbState>,
56 from_basis_t: u64,
57 heartbeat_interval_ms: u64,
58) -> ItemStream {
59 let mut live = state.stream_items();
60 let (schema, interner) = state.handshake_snapshot();
61 let basis = state.db().basis_t();
62 let index_basis = state.index_basis();
63 let (tx, rx) = tokio::sync::mpsc::channel::<Result<pb::SubscribeItem, Status>>(64);
64 let state = Arc::clone(state);
65 tokio::spawn(async move {
66 let send = |item: pb::subscribe_item::Item| {
67 let tx = tx.clone();
68 async move {
69 tx.send(Ok(pb::SubscribeItem { item: Some(item) }))
70 .await
71 .is_ok()
72 }
73 };
74 if !send(pb::subscribe_item::Item::Handshake(pb::Handshake {
75 basis_t: basis,
76 index_basis_t: index_basis,
77 schema,
78 heartbeat_interval_ms,
79 }))
80 .await
81 {
82 return;
83 }
84 let mut last_sent = from_basis_t;
85 if from_basis_t < basis {
86 let records = match state.tx_range(from_basis_t + 1, Some(basis + 1)).await {
87 Ok(records) => records,
88 Err(error) => {
89 let _ = tx.send(Err(to_status(&error))).await;
90 return;
91 }
92 };
93 for record in records {
94 let datoms = match codec::encode_datoms(&record.datoms, &interner) {
95 Ok(datoms) => datoms,
96 Err(error) => {
97 let _ = tx.send(Err(Status::internal(error.to_string()))).await;
98 return;
99 }
100 };
101 if !send(pb::subscribe_item::Item::Report(pb::TxReport {
102 t: record.t,
103 tx_instant: corium_db::bootstrap::asserted_instant(record.t, &record.datoms)
104 .unwrap_or(record.tx_instant),
105 datoms,
106 }))
107 .await
108 {
109 return;
110 }
111 last_sent = record.t;
112 }
113 }
114 loop {
115 match live.recv().await {
116 Ok(pb::subscribe_item::Item::Report(report)) => {
117 if report.t <= last_sent {
118 continue;
119 }
120 last_sent = report.t;
121 if !send(pb::subscribe_item::Item::Report(report)).await {
122 return;
123 }
124 }
125 Ok(item) => {
126 if !send(item).await {
127 return;
128 }
129 }
130 Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
131 let _ = tx
134 .send(Err(Status::data_loss("subscription lagged; resubscribe")))
135 .await;
136 return;
137 }
138 Err(tokio::sync::broadcast::error::RecvError::Closed) => return,
139 }
140 }
141 });
142 Box::pin(ReceiverStream::new(rx))
143}
144
145pub struct TransactorSvc(pub Arc<TransactorNode>, pub Guard);
147
148async fn authorize(guard: &Guard, principal: &Principal, access: Access) -> Result<(), Status> {
153 match guard.authorize(principal, &access).await {
154 Ok(None) => Ok(()),
155 Ok(Some(_filter)) => Err(Status::unimplemented(
156 "view filtering is not enforced on this surface yet",
157 )),
158 Err(error) => Err(error.to_status()),
159 }
160}
161
162#[tonic::async_trait]
163impl Transactor for TransactorSvc {
164 async fn transact(
165 &self,
166 request: Request<pb::TransactRequest>,
167 ) -> Result<Response<pb::TransactResponse>, Status> {
168 let principal = authz::principal(&request);
169 let request = request.into_inner();
170 check_version(request.protocol_version)?;
171 authorize(
172 &self.1,
173 &principal,
174 Access::on(Action::Transact, &request.db),
175 )
176 .await?;
177 self.0
178 .transact_at(&request.db, &request.tx_data, request.expected_basis_t)
179 .await
180 .map(Response::new)
181 .map_err(|error| to_status(&error))
182 }
183
184 type SubscribeStream = ItemStream;
185
186 async fn subscribe(
187 &self,
188 request: Request<pb::SubscribeRequest>,
189 ) -> Result<Response<Self::SubscribeStream>, Status> {
190 let principal = authz::principal(&request);
191 let request = request.into_inner();
192 check_version(request.protocol_version)?;
193 authorize(
194 &self.1,
195 &principal,
196 Access::on(Action::Subscribe, &request.db),
197 )
198 .await?;
199 let state = self
200 .0
201 .db_state(&request.db)
202 .await
203 .map_err(|error| to_status(&error))?;
204 let heartbeat_interval_ms =
205 u64::try_from(self.0.config().heartbeat_interval.as_millis()).unwrap_or(0);
206 Ok(Response::new(subscription_stream(
207 &state,
208 request.from_basis_t,
209 heartbeat_interval_ms,
210 )))
211 }
212
213 async fn sync(
214 &self,
215 request: Request<pb::SyncRequest>,
216 ) -> Result<Response<pb::SyncResponse>, Status> {
217 let principal = authz::principal(&request);
218 let request = request.into_inner();
219 authorize(
220 &self.1,
221 &principal,
222 Access::on(Action::Inspect, &request.db),
223 )
224 .await?;
225 let basis_t = self
226 .0
227 .sync(&request.db, request.t)
228 .await
229 .map_err(|error| to_status(&error))?;
230 Ok(Response::new(pb::SyncResponse { basis_t }))
231 }
232
233 async fn status(
234 &self,
235 request: Request<pb::StatusRequest>,
236 ) -> Result<Response<pb::StatusResponse>, Status> {
237 let principal = authz::principal(&request);
238 let request = request.into_inner();
239 authorize(
240 &self.1,
241 &principal,
242 Access::on(Action::Inspect, &request.db),
243 )
244 .await?;
245 self.0
246 .status(&request.db)
247 .await
248 .map(Response::new)
249 .map_err(|error| to_status(&error))
250 }
251}
252
253fn check_version(version: u32) -> Result<(), Status> {
254 if (corium_protocol::MIN_SUPPORTED_PROTOCOL_VERSION..=corium_protocol::PROTOCOL_VERSION)
255 .contains(&version)
256 {
257 Ok(())
258 } else {
259 Err(Status::failed_precondition(format!(
260 "protocol version {version} is not supported; supported range is {}..={}",
261 corium_protocol::MIN_SUPPORTED_PROTOCOL_VERSION,
262 corium_protocol::PROTOCOL_VERSION
263 )))
264 }
265}
266
267pub struct CatalogSvc(pub Arc<TransactorNode>, pub Guard);
269
270#[tonic::async_trait]
271impl Catalog for CatalogSvc {
272 async fn create_database(
273 &self,
274 request: Request<pb::CreateDatabaseRequest>,
275 ) -> Result<Response<pb::CreateDatabaseResponse>, Status> {
276 let principal = authz::principal(&request);
277 let request = request.into_inner();
278 authorize(
279 &self.1,
280 &principal,
281 Access::on(Action::CreateDatabase, &request.db),
282 )
283 .await?;
284 let storage_key = parse_key_id(&request.storage_key)?;
285 let created = self
286 .0
287 .create_db(&request.db, &request.schema, storage_key)
288 .await
289 .map_err(|error| to_status(&error))?;
290 Ok(Response::new(pb::CreateDatabaseResponse { created }))
291 }
292
293 async fn delete_database(
294 &self,
295 request: Request<pb::DeleteDatabaseRequest>,
296 ) -> Result<Response<pb::DeleteDatabaseResponse>, Status> {
297 let principal = authz::principal(&request);
298 let request = request.into_inner();
299 authorize(
300 &self.1,
301 &principal,
302 Access::on(Action::DeleteDatabase, &request.db),
303 )
304 .await?;
305 let deleted = self
306 .0
307 .delete_db(&request.db)
308 .await
309 .map_err(|error| to_status(&error))?;
310 Ok(Response::new(pb::DeleteDatabaseResponse { deleted }))
311 }
312
313 async fn fork_database(
314 &self,
315 request: Request<pb::ForkDatabaseRequest>,
316 ) -> Result<Response<pb::ForkDatabaseResponse>, Status> {
317 let principal = authz::principal(&request);
318 let request = request.into_inner();
319 authorize(
320 &self.1,
321 &principal,
322 Access::on(Action::ForkDatabase, &request.target),
323 )
324 .await?;
325 let forked = self
326 .0
327 .fork_db(&request.db, &request.target, request.as_of_t)
328 .await
329 .map_err(|error| to_status(&error))?;
330 Ok(Response::new(pb::ForkDatabaseResponse {
331 created: forked.is_some(),
332 basis_t: forked.unwrap_or(0),
333 }))
334 }
335
336 async fn list_databases(
337 &self,
338 request: Request<pb::ListDatabasesRequest>,
339 ) -> Result<Response<pb::ListDatabasesResponse>, Status> {
340 let principal = authz::principal(&request);
341 authorize(&self.1, &principal, Access::catalog(Action::ListDatabases)).await?;
342 Ok(Response::new(pb::ListDatabasesResponse {
343 dbs: self.0.list_dbs(),
344 }))
345 }
346
347 async fn gc_deleted_databases(
348 &self,
349 request: Request<pb::GcDeletedDatabasesRequest>,
350 ) -> Result<Response<pb::GcDeletedDatabasesResponse>, Status> {
351 let principal = authz::principal(&request);
352 authorize(&self.1, &principal, Access::catalog(Action::GarbageCollect)).await?;
353 let swept = match requested_gc_retention(request.into_inner()) {
354 None => self.0.gc_deleted().await,
355 Some(retention) => self.0.gc_deleted_with_retention(retention).await,
356 };
357 let swept_blobs = swept.map_err(|error| to_status(&error))?;
358 Ok(Response::new(pb::GcDeletedDatabasesResponse {
359 swept_blobs,
360 }))
361 }
362
363 async fn request_index(
364 &self,
365 request: Request<pb::RequestIndexRequest>,
366 ) -> Result<Response<pb::RequestIndexResponse>, Status> {
367 let principal = authz::principal(&request);
368 let request = request.into_inner();
369 authorize(
370 &self.1,
371 &principal,
372 Access::on(Action::ManageIndex, &request.db),
373 )
374 .await?;
375 let index_basis_t = self
376 .0
377 .request_index(&request.db)
378 .await
379 .map_err(|error| to_status(&error))?;
380 Ok(Response::new(pb::RequestIndexResponse { index_basis_t }))
381 }
382
383 async fn set_index_policy(
384 &self,
385 request: Request<pb::SetIndexPolicyRequest>,
386 ) -> Result<Response<pb::SetIndexPolicyResponse>, Status> {
387 let principal = authz::principal(&request);
388 let request = request.into_inner();
389 authorize(
390 &self.1,
391 &principal,
392 Access::on(Action::ManageIndex, &request.db),
393 )
394 .await?;
395 let update = IndexPolicyUpdate {
396 interval: request.interval_ms.map(std::time::Duration::from_millis),
397 backoff: request.backoff,
398 tail_threshold: request.tail_threshold,
399 tail_deadline: request
400 .tail_deadline_ms
401 .map(std::time::Duration::from_millis),
402 };
403 let policy = self
404 .0
405 .set_index_policy(&request.db, update)
406 .await
407 .map_err(|error| to_status(&error))?;
408 let millis =
409 |duration: std::time::Duration| u64::try_from(duration.as_millis()).unwrap_or(u64::MAX);
410 Ok(Response::new(pb::SetIndexPolicyResponse {
411 interval_ms: millis(policy.interval),
412 backoff: policy.backoff,
413 tail_threshold: policy.tail_threshold,
414 tail_deadline_ms: millis(policy.tail_deadline),
415 }))
416 }
417
418 async fn get_storage_info(
419 &self,
420 request: Request<pb::GetStorageInfoRequest>,
421 ) -> Result<Response<pb::GetStorageInfoResponse>, Status> {
422 let principal = authz::principal(&request);
423 let request = request.into_inner();
424 authorize(
425 &self.1,
426 &principal,
427 Access::on(Action::Inspect, &request.db),
428 )
429 .await?;
430 self.0
431 .backup_info(&request.db)
432 .await
433 .map(Response::new)
434 .map_err(|error| to_status(&error))
435 }
436
437 async fn key_status(
438 &self,
439 request: Request<pb::KeyStatusRequest>,
440 ) -> Result<Response<pb::KeyStatusResponse>, Status> {
441 let principal = authz::principal(&request);
442 let request = request.into_inner();
443 authorize(
444 &self.1,
445 &principal,
446 Access::on(Action::ManageKeys, &request.db),
447 )
448 .await?;
449 let status = self
450 .0
451 .key_status(&request.db)
452 .await
453 .map_err(|error| to_status(&error))?;
454 let basis_t = status.basis_t;
455 let Some(manifest) = status
456 .manifest
457 .filter(|manifest| !manifest.storage_keys.is_empty())
458 else {
459 return Ok(Response::new(pb::KeyStatusResponse {
460 encrypted: false,
461 basis_t,
462 ..pb::KeyStatusResponse::default()
463 }));
464 };
465 let storage_keys = manifest
466 .storage_keys
467 .iter()
468 .map(|key| pb::StorageKeyEpoch {
469 epoch: key.epoch,
470 kek_epoch: key.kek_epoch,
471 algorithm: key.algorithm.to_string(),
472 state: key.state.to_string(),
473 created_at_unix_ms: key.created_at_unix_ms,
474 opened_at_t: key.opened_at_t,
475 live_objects: key.live_objects,
476 records_sealed: manifest
477 .log_records_sealed(key.epoch, basis_t)
478 .unwrap_or_default(),
479 })
480 .collect();
481 Ok(Response::new(pb::KeyStatusResponse {
482 encrypted: true,
483 kek: manifest.kek.to_string(),
484 storage_keys,
485 basis_t,
486 records_per_epoch_limit: corium_store::LOG_RECORDS_PER_EPOCH_LIMIT,
487 rotation_due: manifest.storage_rotation_due(basis_t),
488 keys_unavailable: status.keys_unavailable,
489 keys_fenced: status.keys_fenced,
490 }))
491 }
492
493 async fn rotate_storage_key(
494 &self,
495 request: Request<pb::RotateStorageKeyRequest>,
496 ) -> Result<Response<pb::RotateStorageKeyResponse>, Status> {
497 let principal = authz::principal(&request);
498 let request = request.into_inner();
499 authorize(
500 &self.1,
501 &principal,
502 Access::on(Action::ManageKeys, &request.db),
503 )
504 .await?;
505 let epoch = self
506 .0
507 .rotate_storage_key(&request.db)
508 .await
509 .map_err(|error| to_status(&error))?;
510 Ok(Response::new(pb::RotateStorageKeyResponse { epoch }))
511 }
512
513 async fn rewrap_keys(
514 &self,
515 request: Request<pb::RewrapKeysRequest>,
516 ) -> Result<Response<pb::RewrapKeysResponse>, Status> {
517 let principal = authz::principal(&request);
518 let request = request.into_inner();
519 authorize(
520 &self.1,
521 &principal,
522 Access::on(Action::ManageKeys, &request.db),
523 )
524 .await?;
525 let kek = parse_key_id(&request.kek)?
526 .ok_or_else(|| Status::invalid_argument("a key-encryption key is required"))?;
527 self.0
528 .rewrap_keys(&request.db, kek)
529 .await
530 .map_err(|error| to_status(&error))?;
531 Ok(Response::new(pb::RewrapKeysResponse {}))
532 }
533}
534
535fn parse_key_id(value: &str) -> Result<Option<corium_crypt::KeyId>, Status> {
537 if value.is_empty() {
538 return Ok(None);
539 }
540 corium_crypt::KeyId::new(value)
541 .map(Some)
542 .map_err(|error| Status::invalid_argument(error.to_string()))
543}
544
545fn requested_gc_retention(request: pb::GcDeletedDatabasesRequest) -> Option<std::time::Duration> {
546 request
547 .retention_millis
548 .map(std::time::Duration::from_millis)
549}
550
551pub async fn serve(
556 node: Arc<TransactorNode>,
557 addr: std::net::SocketAddr,
558 guard: Guard,
559 tls: Option<tonic::transport::ServerTlsConfig>,
560 shutdown: impl std::future::Future<Output = ()> + Send,
561) -> Result<(), tonic::transport::Error> {
562 let interceptor = IdentityInterceptor::new(guard.clone());
565 let mut builder = tonic::transport::Server::builder();
566 if let Some(tls) = tls {
567 builder = builder.tls_config(tls)?;
568 }
569 builder
570 .add_service(TransactorServer::with_interceptor(
571 TransactorSvc(Arc::clone(&node), guard.clone()),
572 interceptor.clone(),
573 ))
574 .add_service(CatalogServer::with_interceptor(
575 CatalogSvc(node, guard),
576 interceptor,
577 ))
578 .serve_with_shutdown(addr, shutdown)
579 .await
580}
581
582#[cfg(test)]
583mod tests {
584 use super::*;
585
586 #[test]
587 fn gc_retention_distinguishes_default_zero_and_subsecond() {
588 let default = pb::GcDeletedDatabasesRequest {
589 retention_millis: None,
590 };
591 let immediate = pb::GcDeletedDatabasesRequest {
592 retention_millis: Some(0),
593 };
594 let subsecond = pb::GcDeletedDatabasesRequest {
595 retention_millis: Some(500),
596 };
597
598 assert_eq!(requested_gc_retention(default), None);
599 assert_eq!(
600 requested_gc_retention(immediate),
601 Some(std::time::Duration::ZERO)
602 );
603 assert_eq!(
604 requested_gc_retention(subsecond),
605 Some(std::time::Duration::from_millis(500))
606 );
607 }
608
609 #[test]
610 fn protocol_versions_allow_server_first_rolling_upgrades() {
611 assert!(check_version(corium_protocol::MIN_SUPPORTED_PROTOCOL_VERSION).is_ok());
612 assert!(check_version(corium_protocol::PROTOCOL_VERSION).is_ok());
613 assert_eq!(
614 check_version(corium_protocol::PROTOCOL_VERSION + 1)
615 .unwrap_err()
616 .code(),
617 tonic::Code::FailedPrecondition
618 );
619 }
620}