1use 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#[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
43pub(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 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
138pub 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
215pub 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
333pub 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}