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