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::auth::{AuthInterceptor, Authenticator};
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(_) | NodeError::Log(_) | NodeError::Lease(_) => {
36            Status::internal(error.to_string())
37        }
38    }
39}
40
41type ItemStream = Pin<Box<dyn Stream<Item = Result<pb::SubscribeItem, Status>> + Send>>;
42
43/// Streams a subscription: handshake, gapless log backfill, then live items.
44///
45/// The broadcast receiver is registered before the basis snapshot is taken,
46/// and live reports at or below the last backfilled `t` are dropped, so no
47/// transaction can fall between backfill and the live stream.
48pub(crate) fn subscription_stream(
49    state: &Arc<DbState>,
50    from_basis_t: u64,
51    heartbeat_interval_ms: u64,
52) -> ItemStream {
53    let mut live = state.stream_items();
54    let (schema, interner) = state.handshake_snapshot();
55    let basis = state.db().basis_t();
56    let index_basis = state.index_basis();
57    let (tx, rx) = tokio::sync::mpsc::channel::<Result<pb::SubscribeItem, Status>>(64);
58    let state = Arc::clone(state);
59    tokio::spawn(async move {
60        let send = |item: pb::subscribe_item::Item| {
61            let tx = tx.clone();
62            async move {
63                tx.send(Ok(pb::SubscribeItem { item: Some(item) }))
64                    .await
65                    .is_ok()
66            }
67        };
68        if !send(pb::subscribe_item::Item::Handshake(pb::Handshake {
69            basis_t: basis,
70            index_basis_t: index_basis,
71            schema,
72            heartbeat_interval_ms,
73        }))
74        .await
75        {
76            return;
77        }
78        let mut last_sent = from_basis_t;
79        if from_basis_t < basis {
80            let backfill = {
81                let state = Arc::clone(&state);
82                tokio::task::spawn_blocking(move || {
83                    state.tx_range(from_basis_t + 1, Some(basis + 1))
84                })
85                .await
86            };
87            let records = match backfill {
88                Ok(Ok(records)) => records,
89                Ok(Err(error)) => {
90                    let _ = tx.send(Err(to_status(&error))).await;
91                    return;
92                }
93                Err(error) => {
94                    let _ = tx.send(Err(Status::internal(error.to_string()))).await;
95                    return;
96                }
97            };
98            for record in records {
99                let datoms = match codec::encode_datoms(&record.datoms, &interner) {
100                    Ok(datoms) => datoms,
101                    Err(error) => {
102                        let _ = tx.send(Err(Status::internal(error.to_string()))).await;
103                        return;
104                    }
105                };
106                if !send(pb::subscribe_item::Item::Report(pb::TxReport {
107                    t: record.t,
108                    tx_instant: record.tx_instant,
109                    datoms,
110                }))
111                .await
112                {
113                    return;
114                }
115                last_sent = record.t;
116            }
117        }
118        loop {
119            match live.recv().await {
120                Ok(pb::subscribe_item::Item::Report(report)) => {
121                    if report.t <= last_sent {
122                        continue;
123                    }
124                    last_sent = report.t;
125                    if !send(pb::subscribe_item::Item::Report(report)).await {
126                        return;
127                    }
128                }
129                Ok(item) => {
130                    if !send(item).await {
131                        return;
132                    }
133                }
134                Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
135                    // The subscriber fell behind the broadcast buffer; end the
136                    // stream so it reconnects and backfills from its basis.
137                    let _ = tx
138                        .send(Err(Status::data_loss("subscription lagged; resubscribe")))
139                        .await;
140                    return;
141                }
142                Err(tokio::sync::broadcast::error::RecvError::Closed) => return,
143            }
144        }
145    });
146    Box::pin(ReceiverStream::new(rx))
147}
148
149/// Transactor gRPC service over a node.
150pub struct TransactorSvc(pub Arc<TransactorNode>);
151
152#[tonic::async_trait]
153impl Transactor for TransactorSvc {
154    async fn transact(
155        &self,
156        request: Request<pb::TransactRequest>,
157    ) -> Result<Response<pb::TransactResponse>, Status> {
158        let request = request.into_inner();
159        check_version(request.protocol_version)?;
160        self.0
161            .transact(&request.db, &request.tx_data)
162            .await
163            .map(Response::new)
164            .map_err(|error| to_status(&error))
165    }
166
167    type SubscribeStream = ItemStream;
168
169    async fn subscribe(
170        &self,
171        request: Request<pb::SubscribeRequest>,
172    ) -> Result<Response<Self::SubscribeStream>, Status> {
173        let request = request.into_inner();
174        check_version(request.protocol_version)?;
175        let state = self
176            .0
177            .db_state(&request.db)
178            .await
179            .map_err(|error| to_status(&error))?;
180        let heartbeat_interval_ms =
181            u64::try_from(self.0.config().heartbeat_interval.as_millis()).unwrap_or(0);
182        Ok(Response::new(subscription_stream(
183            &state,
184            request.from_basis_t,
185            heartbeat_interval_ms,
186        )))
187    }
188
189    async fn sync(
190        &self,
191        request: Request<pb::SyncRequest>,
192    ) -> Result<Response<pb::SyncResponse>, Status> {
193        let request = request.into_inner();
194        let basis_t = self
195            .0
196            .sync(&request.db, request.t)
197            .await
198            .map_err(|error| to_status(&error))?;
199        Ok(Response::new(pb::SyncResponse { basis_t }))
200    }
201
202    async fn status(
203        &self,
204        request: Request<pb::StatusRequest>,
205    ) -> Result<Response<pb::StatusResponse>, Status> {
206        let request = request.into_inner();
207        self.0
208            .status(&request.db)
209            .await
210            .map(Response::new)
211            .map_err(|error| to_status(&error))
212    }
213}
214
215fn check_version(version: u32) -> Result<(), Status> {
216    if version == corium_protocol::PROTOCOL_VERSION {
217        Ok(())
218    } else {
219        Err(Status::failed_precondition(format!(
220            "protocol version {version} is not supported; upgrade to {}",
221            corium_protocol::PROTOCOL_VERSION
222        )))
223    }
224}
225
226/// Catalog gRPC service over a node.
227pub struct CatalogSvc(pub Arc<TransactorNode>);
228
229#[tonic::async_trait]
230impl Catalog for CatalogSvc {
231    async fn create_database(
232        &self,
233        request: Request<pb::CreateDatabaseRequest>,
234    ) -> Result<Response<pb::CreateDatabaseResponse>, Status> {
235        let request = request.into_inner();
236        let created = self
237            .0
238            .create_db(&request.db, &request.schema)
239            .await
240            .map_err(|error| to_status(&error))?;
241        Ok(Response::new(pb::CreateDatabaseResponse { created }))
242    }
243
244    async fn delete_database(
245        &self,
246        request: Request<pb::DeleteDatabaseRequest>,
247    ) -> Result<Response<pb::DeleteDatabaseResponse>, Status> {
248        let request = request.into_inner();
249        let deleted = self
250            .0
251            .delete_db(&request.db)
252            .await
253            .map_err(|error| to_status(&error))?;
254        Ok(Response::new(pb::DeleteDatabaseResponse { deleted }))
255    }
256
257    async fn fork_database(
258        &self,
259        request: Request<pb::ForkDatabaseRequest>,
260    ) -> Result<Response<pb::ForkDatabaseResponse>, Status> {
261        let request = request.into_inner();
262        let forked = self
263            .0
264            .fork_db(&request.db, &request.target, request.as_of_t)
265            .await
266            .map_err(|error| to_status(&error))?;
267        Ok(Response::new(pb::ForkDatabaseResponse {
268            created: forked.is_some(),
269            basis_t: forked.unwrap_or(0),
270        }))
271    }
272
273    async fn list_databases(
274        &self,
275        _request: Request<pb::ListDatabasesRequest>,
276    ) -> Result<Response<pb::ListDatabasesResponse>, Status> {
277        Ok(Response::new(pb::ListDatabasesResponse {
278            dbs: self.0.list_dbs(),
279        }))
280    }
281
282    async fn gc_deleted_databases(
283        &self,
284        request: Request<pb::GcDeletedDatabasesRequest>,
285    ) -> Result<Response<pb::GcDeletedDatabasesResponse>, Status> {
286        let swept = match requested_gc_retention(request.into_inner()) {
287            None => self.0.gc_deleted().await,
288            Some(retention) => self.0.gc_deleted_with_retention(retention).await,
289        };
290        let swept_blobs = swept.map_err(|error| to_status(&error))?;
291        Ok(Response::new(pb::GcDeletedDatabasesResponse {
292            swept_blobs,
293        }))
294    }
295
296    async fn request_index(
297        &self,
298        request: Request<pb::RequestIndexRequest>,
299    ) -> Result<Response<pb::RequestIndexResponse>, Status> {
300        let request = request.into_inner();
301        let index_basis_t = self
302            .0
303            .request_index(&request.db)
304            .await
305            .map_err(|error| to_status(&error))?;
306        Ok(Response::new(pb::RequestIndexResponse { index_basis_t }))
307    }
308
309    async fn set_index_policy(
310        &self,
311        request: Request<pb::SetIndexPolicyRequest>,
312    ) -> Result<Response<pb::SetIndexPolicyResponse>, Status> {
313        let request = request.into_inner();
314        let update = IndexPolicyUpdate {
315            interval: request.interval_ms.map(std::time::Duration::from_millis),
316            backoff: request.backoff,
317            tail_threshold: request.tail_threshold,
318            tail_deadline: request
319                .tail_deadline_ms
320                .map(std::time::Duration::from_millis),
321        };
322        let policy = self
323            .0
324            .set_index_policy(&request.db, update)
325            .await
326            .map_err(|error| to_status(&error))?;
327        let millis =
328            |duration: std::time::Duration| u64::try_from(duration.as_millis()).unwrap_or(u64::MAX);
329        Ok(Response::new(pb::SetIndexPolicyResponse {
330            interval_ms: millis(policy.interval),
331            backoff: policy.backoff,
332            tail_threshold: policy.tail_threshold,
333            tail_deadline_ms: millis(policy.tail_deadline),
334        }))
335    }
336}
337
338fn requested_gc_retention(request: pb::GcDeletedDatabasesRequest) -> Option<std::time::Duration> {
339    request
340        .retention_millis
341        .map(std::time::Duration::from_millis)
342}
343
344/// Serves the transactor and catalog services until `shutdown` resolves.
345///
346/// # Errors
347/// Returns an error when the listener cannot be bound or TLS is invalid.
348pub async fn serve(
349    node: Arc<TransactorNode>,
350    addr: std::net::SocketAddr,
351    authenticator: Arc<dyn Authenticator>,
352    tls: Option<tonic::transport::ServerTlsConfig>,
353    shutdown: impl std::future::Future<Output = ()> + Send,
354) -> Result<(), tonic::transport::Error> {
355    let interceptor = AuthInterceptor::new(authenticator);
356    let mut builder = tonic::transport::Server::builder();
357    if let Some(tls) = tls {
358        builder = builder.tls_config(tls)?;
359    }
360    builder
361        .add_service(TransactorServer::with_interceptor(
362            TransactorSvc(Arc::clone(&node)),
363            interceptor.clone(),
364        ))
365        .add_service(CatalogServer::with_interceptor(
366            CatalogSvc(node),
367            interceptor,
368        ))
369        .serve_with_shutdown(addr, shutdown)
370        .await
371}
372
373#[cfg(test)]
374mod tests {
375    use super::*;
376
377    #[test]
378    fn gc_retention_distinguishes_default_zero_and_subsecond() {
379        let default = pb::GcDeletedDatabasesRequest {
380            retention_millis: None,
381        };
382        let immediate = pb::GcDeletedDatabasesRequest {
383            retention_millis: Some(0),
384        };
385        let subsecond = pb::GcDeletedDatabasesRequest {
386            retention_millis: Some(500),
387        };
388
389        assert_eq!(requested_gc_retention(default), None);
390        assert_eq!(
391            requested_gc_retention(immediate),
392            Some(std::time::Duration::ZERO)
393        );
394        assert_eq!(
395            requested_gc_retention(subsecond),
396            Some(std::time::Duration::from_millis(500))
397        );
398    }
399}