Skip to main content

corium_transactor/
server.rs

1//! Tonic services exposing a [`TransactorNode`] over the network.
2
3use 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/// Maps node errors onto gRPC statuses.
18#[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        // A missing or unresolvable key is an operator misconfiguration, not
32        // a transient fault: the caller must fix the deployment, not retry.
33        | 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
49/// Streams a subscription: handshake, gapless log backfill, then live items.
50///
51/// The broadcast receiver is registered before the basis snapshot is taken,
52/// and live reports at or below the last backfilled `t` are dropped, so no
53/// transaction can fall between backfill and the live stream.
54pub(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                    // The subscriber fell behind the broadcast buffer; end the
132                    // stream so it reconnects and backfills from its basis.
133                    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
145/// Transactor gRPC service over a node.
146pub struct TransactorSvc(pub Arc<TransactorNode>, pub Guard);
147
148/// Authorizes `access` for the request's [`Principal`], mapping the guard's
149/// decision onto a gRPC status. A returned [`ViewFilter`](authz::ViewFilter) is
150/// rejected for now: no transactor/catalog surface applies one yet, so honoring
151/// a filtered decision by returning full data would be an authorization bypass.
152async 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
267/// Catalog gRPC service over a node.
268pub 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
535/// Parses an optional key identity, where empty means "unset".
536fn 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
551/// Serves the transactor and catalog services until `shutdown` resolves.
552///
553/// # Errors
554/// Returns an error when the listener cannot be bound or TLS is invalid.
555pub 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    // The interceptor authenticates each request and attaches its `Principal`;
563    // the services carry the same guard to authorize the concrete access.
564    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}