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::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
44/// Streams a subscription: handshake, gapless log backfill, then live items.
45///
46/// The broadcast receiver is registered before the basis snapshot is taken,
47/// and live reports at or below the last backfilled `t` are dropped, so no
48/// transaction can fall between backfill and the live stream.
49pub(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                    // The subscriber fell behind the broadcast buffer; end the
126                    // stream so it reconnects and backfills from its basis.
127                    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
139/// Transactor gRPC service over a node.
140pub struct TransactorSvc(pub Arc<TransactorNode>, pub Guard);
141
142/// Authorizes `access` for the request's [`Principal`], mapping the guard's
143/// decision onto a gRPC status. A returned [`ViewFilter`](authz::ViewFilter) is
144/// rejected for now: no transactor/catalog surface applies one yet, so honoring
145/// a filtered decision by returning full data would be an authorization bypass.
146async 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
258/// Catalog gRPC service over a node.
259pub 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
434/// Serves the transactor and catalog services until `shutdown` resolves.
435///
436/// # Errors
437/// Returns an error when the listener cannot be bound or TLS is invalid.
438pub 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    // The interceptor authenticates each request and attaches its `Principal`;
446    // the services carry the same guard to authorize the concrete access.
447    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}