1use 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#[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
45pub(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 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
141pub struct TransactorSvc(pub Arc<TransactorNode>, pub Guard);
143
144async 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
263pub 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
439pub 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 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}