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